build(deps): bump go-micro to v4.11.4-0.20260929213538-cd412c159a26

Point the go-micro.dev/v4 replace at the latest main-v4 maintenance
branch commit in butonic/go-micro (cd412c15) instead of the v4.11.3
tag. The new base is upstream v4.11.1 plus the registry cache
node-TTL fixes (upstream #2715, #2736, #2740), the concurrent map
access fix (upstream #2794), plus further dependency updates and test
fixes on the maintenance branch.

The branch also modernizes go-micro's own go.mod (go 1.26, lego,
dario.cat/mergo v1, x/deque v2, ...), and MVS pulls the aligned
versions into this module: lego v4.35.2 -> lego/v5 v5.5.2,
go-jose/go-jose/v4 -> v4.1.5, gofrs/flock -> v0.13.1,
mattn/go-isatty -> v0.0.24, go.uber.org/zap -> v1.28.0, x/net ->
0.59.0, miekg/dns -> 1.1.73, moby api/client, ProtonMail/go-crypto
-> 1.5.2, klauspost/compress -> 1.20.1, gorilla/handlers -> 1.5.2,
gobwas/ws -> 1.4.0, nxadm/tail -> 1.4.11, ...

Signed-off-by: Jörn Friedrich Dreyer <jfd@butonic.de>
This commit is contained in:
Jörn Friedrich Dreyer committed 2026-10-01 11:14:15 +02:00
1 parent eb32a80e01
commit 076a3b0212
396 files changed
+46182 -7104

No files matched your search

+38 -37
View File
@@ -13,7 +13,7 @@ require (
github.com/beevik/etree v1.8.0
github.com/blevesearch/bleve/v2 v2.6.1
github.com/cenkalti/backoff v2.2.1+incompatible
github.com/coreos/go-oidc/v3 v3.20.0
github.com/coreos/go-oidc/v3 v3.21.0
github.com/cs3org/go-cs3apis v0.0.0-20260424072047-8d9ef7076ae9
github.com/davidbyttow/govips/v2 v2.18.0
github.com/dhowden/tag v0.0.0-20240417053706-3d75831295e8
@@ -64,7 +64,7 @@ require (
github.com/open-policy-agent/opa v1.19.1
github.com/opencloud-eu/icap-client v0.0.0-20250930132611-28a2afe62d89
github.com/opencloud-eu/libre-graph-api-go v1.0.8-0.20260902170011-45af3945a067
github.com/opencloud-eu/reva/v2 v2.50.1-0.20260924091259-7fa82c610dba
github.com/opencloud-eu/reva/v2 v2.50.1-0.20261001091108-11d87fb6b985
github.com/opensearch-project/opensearch-go/v4 v4.7.3
github.com/orcaman/concurrent-map v1.0.0
github.com/pkg/errors v0.9.1
@@ -106,12 +106,12 @@ require (
golang.org/x/crypto v0.57.0
golang.org/x/exp v0.0.0-20260410095643-746e56fc9e2f
golang.org/x/image v0.46.0
golang.org/x/net v0.58.0
golang.org/x/net v0.59.0
golang.org/x/oauth2 v0.37.0
golang.org/x/sync v0.23.0
golang.org/x/term v0.46.0
golang.org/x/text v0.42.0
google.golang.org/genproto/googleapis/api v0.0.0-20260819154853-08b0e4226688
google.golang.org/genproto/googleapis/api v0.0.0-20260928230214-8a89bd6388cc
google.golang.org/grpc v1.84.0
google.golang.org/protobuf v1.36.12
gopkg.in/yaml.v2 v2.4.0
@@ -123,14 +123,14 @@ require (
require (
contrib.go.opencensus.io/exporter/prometheus v0.4.2 // indirect
filippo.io/edwards25519 v1.2.0 // indirect
github.com/Azure/go-ansiterm v0.0.0-20250102033503-faa5f7b0171c // indirect
github.com/Azure/go-ansiterm v0.0.0-20260917205352-e937bb47801a // indirect
github.com/Azure/go-ntlmssp v0.1.1 // indirect
github.com/BurntSushi/toml v1.6.0 // indirect
github.com/Masterminds/goutils v1.1.1 // indirect
github.com/Masterminds/semver/v3 v3.5.0 // indirect
github.com/Masterminds/sprig v2.22.0+incompatible // indirect
github.com/Microsoft/go-winio v0.6.2 // indirect
github.com/ProtonMail/go-crypto v1.1.6 // indirect
github.com/ProtonMail/go-crypto v1.5.2 // indirect
github.com/RoaringBitmap/roaring/v2 v2.14.5 // indirect
github.com/agnivade/levenshtein v1.2.1 // indirect
github.com/ajg/form v1.5.1 // indirect
@@ -140,7 +140,7 @@ require (
github.com/armon/go-radix v1.0.0 // indirect
github.com/asaskevich/govalidator v0.0.0-20230301143203-a9d515a09cc2 // indirect
github.com/beorn7/perks v1.0.1 // indirect
github.com/bitly/go-simplejson v0.5.0 // indirect
github.com/bitly/go-simplejson v0.5.1 // indirect
github.com/bits-and-blooms/bitset v1.24.2 // indirect
github.com/blevesearch/bleve_index_api v1.4.1 // indirect
github.com/blevesearch/geo v0.2.6 // indirect
@@ -168,10 +168,10 @@ require (
github.com/cevaris/ordered_map v0.0.0-20190319150403-3adeae072e73 // indirect
github.com/clipperhouse/displaywidth v0.11.0 // indirect
github.com/clipperhouse/uax29/v2 v2.7.0 // indirect
github.com/cloudflare/circl v1.6.3 // indirect
github.com/cloudflare/circl v1.6.5 // indirect
github.com/containerd/errdefs v1.0.0 // indirect
github.com/containerd/errdefs/pkg v0.3.0 // indirect
github.com/containerd/log v0.1.0 // indirect
github.com/containerd/log v0.2.0 // indirect
github.com/containerd/platforms v1.0.0-rc.2 // indirect
github.com/coreos/go-semver v0.3.1 // indirect
github.com/coreos/go-systemd/v22 v22.7.0 // indirect
@@ -180,7 +180,7 @@ require (
github.com/cpuguy83/go-md2man/v2 v2.0.7 // indirect
github.com/crewjam/httperr v0.2.0 // indirect
github.com/crewjam/saml v0.4.14 // indirect
github.com/cyphar/filepath-securejoin v0.6.1 // indirect
github.com/cyphar/filepath-securejoin v0.7.0 // indirect
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect
github.com/deckarep/golang-set v1.8.0 // indirect
github.com/decred/dcrd/dcrec/secp256k1/v4 v4.4.1 // indirect
@@ -195,18 +195,18 @@ require (
github.com/egirna/icap v0.0.0-20181108071049-d5ee18bd70bc // indirect
github.com/emirpasic/gods v1.18.1 // indirect
github.com/emvi/iso-639-1 v1.1.1 // indirect
github.com/evanphx/json-patch/v5 v5.5.0 // indirect
github.com/evanphx/json-patch/v5 v5.9.11 // indirect
github.com/fatih/color v1.19.0 // indirect
github.com/felixge/httpsnoop v1.1.0 // indirect
github.com/fsnotify/fsnotify v1.10.1 // indirect
github.com/gdexlab/go-render v1.0.1 // indirect
github.com/go-acme/lego/v4 v4.4.0 // indirect
github.com/go-acme/lego/v5 v5.5.2 // indirect
github.com/go-asn1-ber/asn1-ber v1.5.8 // indirect
github.com/go-git/gcfg v1.5.1-0.20230307220236-3a3c6141e376 // indirect
github.com/go-git/go-billy/v5 v5.9.0 // indirect
github.com/go-git/go-billy/v5 v5.9.1 // indirect
github.com/go-git/go-git/v5 v5.19.2 // indirect
github.com/go-jose/go-jose/v3 v3.0.5 // indirect
github.com/go-jose/go-jose/v4 v4.1.4 // indirect
github.com/go-jose/go-jose/v4 v4.1.5 // indirect
github.com/go-kit/log v0.2.1 // indirect
github.com/go-logfmt/logfmt v0.5.1 // indirect
github.com/go-logr/logr v1.4.4 // indirect
@@ -218,26 +218,26 @@ require (
github.com/go-playground/locales v0.14.1 // indirect
github.com/go-playground/universal-translator v0.18.2 // indirect
github.com/go-redis/redis/v8 v8.11.5 // indirect
github.com/go-sql-driver/mysql v1.10.0 // indirect
github.com/go-sql-driver/mysql v1.10.1 // indirect
github.com/go-task/slim-sprig v0.0.0-20230315185526-52ccab3ef572 // indirect
github.com/go-task/slim-sprig/v3 v3.0.0 // indirect
github.com/go-test/deep v1.1.0 // indirect
github.com/gobwas/glob v0.2.3 // indirect
github.com/gobwas/httphead v0.1.0 // indirect
github.com/gobwas/pool v0.2.1 // indirect
github.com/gobwas/ws v1.2.1 // indirect
github.com/gobwas/ws v1.4.0 // indirect
github.com/goccy/go-json v0.10.6 // indirect
github.com/goccy/go-yaml v1.19.2 // indirect
github.com/gofrs/flock v0.13.0 // indirect
github.com/gofrs/flock v0.13.1 // indirect
github.com/golang-jwt/jwt/v4 v4.5.2 // indirect
github.com/golang/groupcache v0.0.0-20241129210726-2c02b8208cf8 // indirect
github.com/golang/snappy v1.0.0 // indirect
github.com/google/go-querystring v1.1.0 // indirect
github.com/google/go-querystring v1.2.0 // indirect
github.com/google/go-tpm v0.9.8 // indirect
github.com/google/pprof v0.0.0-20260402051712-545e8a4df936 // indirect
github.com/google/renameio/v2 v2.0.2 // indirect
github.com/gookit/goutil v0.8.0 // indirect
github.com/gorilla/handlers v1.5.1 // indirect
github.com/gorilla/handlers v1.5.2 // indirect
github.com/gorilla/schema v1.4.1 // indirect
github.com/grpc-ecosystem/go-grpc-middleware v1.4.0 // indirect
github.com/hashicorp/go-hclog v1.6.3 // indirect
@@ -246,15 +246,15 @@ require (
github.com/hashicorp/yamux v0.1.2 // indirect
github.com/huandu/xstrings v1.5.0 // indirect
github.com/iancoleman/strcase v0.3.0 // indirect
github.com/imdario/mergo v0.3.15 // indirect
github.com/imdario/mergo v0.3.16 // indirect
github.com/inconshreveable/mousetrap v1.1.0 // indirect
github.com/jbenet/go-context v0.0.0-20150711004518-d14ea06fba99 // indirect
github.com/jonboulle/clockwork v0.5.0 // indirect
github.com/json-iterator/go v1.1.12 // indirect
github.com/json-iterator/go v1.1.13-0.20220915233716-71ac16282d12 // indirect
github.com/juliangruber/go-intersect v1.1.0 // indirect
github.com/kevinburke/ssh_config v1.2.0 // indirect
github.com/klauspost/compress v1.20.0 // indirect
github.com/klauspost/cpuid/v2 v2.3.0 // indirect
github.com/kevinburke/ssh_config v1.6.0 // indirect
github.com/klauspost/compress v1.20.1 // indirect
github.com/klauspost/cpuid/v2 v2.4.0 // indirect
github.com/klauspost/crc32 v1.3.0 // indirect
github.com/kovidgoyal/go-parallel v1.1.1 // indirect
github.com/kovidgoyal/go-shm v1.0.0 // indirect
@@ -273,11 +273,11 @@ require (
github.com/magiconair/properties v1.8.10 // indirect
github.com/mattermost/xml-roundtrip-validator v0.1.0 // indirect
github.com/mattn/go-colorable v0.1.15 // indirect
github.com/mattn/go-isatty v0.0.22 // indirect
github.com/mattn/go-isatty v0.0.24 // indirect
github.com/mattn/go-runewidth v0.0.24 // indirect
github.com/mattn/go-sqlite3 v1.14.49 // indirect
github.com/mendsley/gojwk v0.0.0-20141217222730-4d5ec6e58103 // indirect
github.com/miekg/dns v1.1.68 // indirect
github.com/miekg/dns v1.1.73 // indirect
github.com/mileusna/useragent v1.3.5 // indirect
github.com/minio/crc64nvme v1.1.1 // indirect
github.com/minio/highwayhash v1.0.4 // indirect
@@ -286,13 +286,13 @@ require (
github.com/mitchellh/copystructure v1.2.0 // indirect
github.com/mitchellh/reflectwalk v1.0.2 // indirect
github.com/moby/docker-image-spec v1.3.1 // indirect
github.com/moby/go-archive v0.2.0 // indirect
github.com/moby/moby/api v1.55.0 // indirect
github.com/moby/moby/client v0.5.0 // indirect
github.com/moby/go-archive v0.3.3 // indirect
github.com/moby/moby/api v1.56.0 // indirect
github.com/moby/moby/client v0.6.0 // indirect
github.com/moby/patternmatcher v0.6.1 // indirect
github.com/moby/sys/sequential v0.7.0 // indirect
github.com/moby/sys/user v0.4.0 // indirect
github.com/moby/sys/userns v0.1.0 // indirect
github.com/moby/sys/user v0.4.1 // indirect
github.com/moby/sys/userns v0.2.1 // indirect
github.com/moby/term v0.5.2 // indirect
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect
github.com/modern-go/reflect2 v1.0.3-0.20250322232337-35a7c28c31ee // indirect
@@ -301,7 +301,7 @@ require (
github.com/nats-io/jwt/v2 v2.8.2 // indirect
github.com/nats-io/nkeys v0.4.16 // indirect
github.com/nats-io/nuid v1.0.1 // indirect
github.com/nxadm/tail v1.4.8 // indirect
github.com/nxadm/tail v1.4.11 // indirect
github.com/oklog/run v1.2.0 // indirect
github.com/olekukonko/cat v0.0.0-20250911104152-50322a0618f6 // indirect
github.com/olekukonko/ll v0.1.6 // indirect
@@ -339,7 +339,7 @@ require (
github.com/sethvargo/go-diceware v0.6.0 // indirect
github.com/sethvargo/go-password v0.4.0 // indirect
github.com/shirou/gopsutil/v4 v4.26.6 // indirect
github.com/skeema/knownhosts v1.3.1 // indirect
github.com/skeema/knownhosts v1.3.3 // indirect
github.com/sourcegraph/conc v0.3.1-0.20240121214520-5f936abd7ae8 // indirect
github.com/spacewander/go-suffix-tree v0.0.0-20191010040751-0865e368c784 // indirect
github.com/spf13/cast v1.10.0 // indirect
@@ -362,7 +362,7 @@ require (
github.com/xanzy/ssh-agent v0.3.3 // indirect
github.com/xeipuuv/gojsonpointer v0.0.0-20190905194746-02993c407bfb // indirect
github.com/xeipuuv/gojsonreference v0.0.0-20180127040603-bd5ef7bd5415 // indirect
github.com/xrash/smetrics v0.0.0-20240521201337-686a1a2994c1 // indirect
github.com/xrash/smetrics v0.0.0-20250705151800-55b8f293f342 // indirect
github.com/yashtewari/glob-intersection v0.2.0 // indirect
github.com/yusufpapurcu/wmi v1.2.4 // indirect
github.com/zeebo/xxh3 v1.1.0 // indirect
@@ -375,7 +375,7 @@ require (
go.opentelemetry.io/otel/metric v1.46.0 // indirect
go.opentelemetry.io/proto/otlp v1.11.0 // indirect
go.uber.org/multierr v1.11.0 // indirect
go.uber.org/zap v1.27.1 // indirect
go.uber.org/zap v1.28.0 // indirect
go.yaml.in/yaml/v2 v2.4.4 // indirect
go.yaml.in/yaml/v3 v3.0.5 // indirect
golang.org/x/mod v0.41.0 // indirect
@@ -383,7 +383,7 @@ require (
golang.org/x/time v0.16.0 // indirect
golang.org/x/tools v0.49.0 // indirect
google.golang.org/genproto v0.0.0-20260319201613-d00831a3d3e7 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260825221802-da73d73af1c5 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260928230214-8a89bd6388cc // indirect
gopkg.in/cenkalti/backoff.v1 v1.1.0 // indirect
gopkg.in/ini.v1 v1.67.3 // indirect
gopkg.in/tomb.v1 v1.0.0-20141024135613-dd632973f1e7 // indirect
@@ -395,7 +395,8 @@ replace github.com/studio-b12/gowebdav => github.com/kobergj/gowebdav v0.0.0-202
replace github.com/unrolled/secure => github.com/opencloud-eu/secure v0.0.0-20260312082735-b6f5cb2244e4
replace go-micro.dev/v4 => github.com/butonic/go-micro/v4 v4.11.1-0.20241115112658-b5d4de5ed9b3
// TODO we could bump go-micro.dev/v4 to v6 and drop this fork as our v4 changes got merged upstream
replace go-micro.dev/v4 => github.com/butonic/go-micro/v4 v4.11.4-0.20260929213538-cd412c159a26
// exclude the v2 line of go-sqlite3 which was released accidentally and prevents pulling in newer versions of go-sqlite3
// see https://github.com/mattn/go-sqlite3/issues/965 for more details
+82 -380
View File
File diff suppressed because it is too large. Load diff
+3 -3
View File
@@ -26,7 +26,7 @@ func (csiState csiEntryState) Handle(b byte) (s state, e error) {
func (csiState csiEntryState) Transition(s state) error {
csiState.parser.logf("CsiEntry::Transition %s --> %s", csiState.Name(), s.Name())
csiState.baseState.Transition(s)
_ = csiState.baseState.Transition(s)
switch s {
case csiState.parser.ground:
@@ -34,9 +34,9 @@ func (csiState csiEntryState) Transition(s state) error {
case csiState.parser.csiParam:
switch {
case sliceContains(csiParams, csiState.parser.context.currentChar):
csiState.parser.collectParam()
_ = csiState.parser.collectParam()
case sliceContains(intermeds, csiState.parser.context.currentChar):
csiState.parser.collectInter()
_ = csiState.parser.collectInter()
}
}
+2 -3
View File
@@ -16,8 +16,7 @@ func (csiState csiParamState) Handle(b byte) (s state, e error) {
case sliceContains(alphabetics, b):
return csiState.parser.ground, nil
case sliceContains(csiCollectables, b):
csiState.parser.collectParam()
return csiState, nil
return csiState, csiState.parser.collectParam()
case sliceContains(executors, b):
return csiState, csiState.parser.execute()
}
@@ -27,7 +26,7 @@ func (csiState csiParamState) Handle(b byte) (s state, e error) {
func (csiState csiParamState) Transition(s state) error {
csiState.parser.logf("CsiParam::Transition %s --> %s", csiState.Name(), s.Name())
csiState.baseState.Transition(s)
_ = csiState.baseState.Transition(s)
switch s {
case csiState.parser.ground:
+1 -1
View File
@@ -25,7 +25,7 @@ func (escState escapeIntermediateState) Handle(b byte) (s state, e error) {
func (escState escapeIntermediateState) Transition(s state) error {
escState.parser.logf("escapeIntermediateState::Transition %s --> %s", escState.Name(), s.Name())
escState.baseState.Transition(s)
_ = escState.baseState.Transition(s)
switch s {
case escState.parser.ground:
+1 -1
View File
@@ -29,7 +29,7 @@ func (escState escapeState) Handle(b byte) (s state, e error) {
func (escState escapeState) Transition(s state) error {
escState.parser.logf("Escape::Transition %s --> %s", escState.Name(), s.Name())
escState.baseState.Transition(s)
_ = escState.baseState.Transition(s)
switch s {
case escState.parser.ground:
+5 -8
View File
@@ -113,16 +113,13 @@ func (ap *AnsiParser) handle(b byte) error {
if newState == nil {
ap.logf("WARNING: newState is nil")
return errors.New("New state of 'nil' is invalid.")
return errors.New("new state of 'nil' is invalid")
}
if newState != ap.currState {
if err := ap.changeState(newState); err != nil {
return err
}
if newState == ap.currState {
return nil
}
return nil
return ap.changeState(newState)
}
func (ap *AnsiParser) changeState(newState state) error {
@@ -136,7 +133,7 @@ func (ap *AnsiParser) changeState(newState state) error {
// Perform transition action
if err := ap.currState.Transition(newState); err != nil {
ap.logf("Transition from '%s' to '%s' failed with: '%v'", ap.currState.Name(), newState.Name, err)
ap.logf("Transition from '%s' to '%s' failed with: '%v'", ap.currState.Name(), newState.Name(), err)
return err
}
+5 -9
View File
@@ -4,9 +4,9 @@ import (
"strconv"
)
func parseParams(bytes []byte) ([]string, error) {
paramBuff := make([]byte, 0, 0)
params := []string{}
func parseParams(bytes []byte) []string {
var paramBuff []byte
var params []string
for _, v := range bytes {
if v == ';' {
@@ -14,7 +14,7 @@ func parseParams(bytes []byte) ([]string, error) {
// Completed parameter, append it to the list
s := string(paramBuff)
params = append(params, s)
paramBuff = make([]byte, 0, 0)
paramBuff = paramBuff[:0]
}
} else {
paramBuff = append(paramBuff, v)
@@ -27,11 +27,7 @@ func parseParams(bytes []byte) ([]string, error) {
params = append(params, s)
}
return params, nil
}
func parseCmd(context ansiContext) (string, error) {
return string(context.currentChar), nil
return params
}
func getInt(params []string, dflt int) int {
+7 -9
View File
@@ -15,10 +15,9 @@ func (ap *AnsiParser) collectInter() error {
}
func (ap *AnsiParser) escDispatch() error {
cmd, _ := parseCmd(*ap.context)
intermeds := ap.context.interBuffer
cmd := string(ap.context.currentChar)
ap.logf("escDispatch currentChar: %#x", ap.context.currentChar)
ap.logf("escDispatch: %v(%v)", cmd, intermeds)
ap.logf("escDispatch: %s(%q)", cmd, string(ap.context.interBuffer))
switch cmd {
case "D": // IND
@@ -31,14 +30,14 @@ func (ap *AnsiParser) escDispatch() error {
return err
case "M": // RI
return ap.eventHandler.RI()
default:
return nil
}
return nil
}
func (ap *AnsiParser) csiDispatch() error {
cmd, _ := parseCmd(*ap.context)
params, _ := parseParams(ap.context.paramBuffer)
cmd := string(ap.context.currentChar)
params := parseParams(ap.context.paramBuffer)
ap.logf("Parsed params: %v with length: %d", params, len(params))
ap.logf("csiDispatch: %v(%v)", cmd, params)
@@ -109,9 +108,8 @@ func (ap *AnsiParser) print() error {
return ap.eventHandler.Print(ap.context.currentChar)
}
func (ap *AnsiParser) clear() error {
func (ap *AnsiParser) clear() {
ap.context = &ansiContext{}
return nil
}
func (ap *AnsiParser) execute() error {
-2
View File
@@ -1,7 +1,5 @@
package ansiterm
type stateID int
type state interface {
Enter() error
Exit() error
-10
View File
@@ -1,9 +1,5 @@
package ansiterm
import (
"strconv"
)
func sliceContains(bytes []byte, b byte) bool {
for _, v := range bytes {
if v == b {
@@ -13,9 +9,3 @@ func sliceContains(bytes []byte, b byte) bool {
return false
}
func convertBytesToInteger(bytes []byte) int {
s := string(bytes)
i, _ := strconv.Atoi(s)
return i
}
+1 -98
View File
@@ -5,12 +5,9 @@ package winterm
import (
"fmt"
"os"
"strconv"
"strings"
"syscall"
"github.com/Azure/go-ansiterm"
windows "golang.org/x/sys/windows"
"golang.org/x/sys/windows"
)
// Windows keyboard constants
@@ -55,100 +52,6 @@ const (
ENHANCED_KEY = 0x0100
)
type ansiCommand struct {
CommandBytes []byte
Command string
Parameters []string
IsSpecial bool
}
func newAnsiCommand(command []byte) *ansiCommand {
if isCharacterSelectionCmdChar(command[1]) {
// Is Character Set Selection commands
return &ansiCommand{
CommandBytes: command,
Command: string(command),
IsSpecial: true,
}
}
// last char is command character
lastCharIndex := len(command) - 1
ac := &ansiCommand{
CommandBytes: command,
Command: string(command[lastCharIndex]),
IsSpecial: false,
}
// more than a single escape
if lastCharIndex != 0 {
start := 1
// skip if double char escape sequence
if command[0] == ansiterm.ANSI_ESCAPE_PRIMARY && command[1] == ansiterm.ANSI_ESCAPE_SECONDARY {
start++
}
// convert this to GetNextParam method
ac.Parameters = strings.Split(string(command[start:lastCharIndex]), ansiterm.ANSI_PARAMETER_SEP)
}
return ac
}
func (ac *ansiCommand) paramAsSHORT(index int, defaultValue int16) int16 {
if index < 0 || index >= len(ac.Parameters) {
return defaultValue
}
param, err := strconv.ParseInt(ac.Parameters[index], 10, 16)
if err != nil {
return defaultValue
}
return int16(param)
}
func (ac *ansiCommand) String() string {
return fmt.Sprintf("0x%v \"%v\" (\"%v\")",
bytesToHex(ac.CommandBytes),
ac.Command,
strings.Join(ac.Parameters, "\",\""))
}
// isAnsiCommandChar returns true if the passed byte falls within the range of ANSI commands.
// See http://manpages.ubuntu.com/manpages/intrepid/man4/console_codes.4.html.
func isAnsiCommandChar(b byte) bool {
switch {
case ansiterm.ANSI_COMMAND_FIRST <= b && b <= ansiterm.ANSI_COMMAND_LAST && b != ansiterm.ANSI_ESCAPE_SECONDARY:
return true
case b == ansiterm.ANSI_CMD_G1 || b == ansiterm.ANSI_CMD_OSC || b == ansiterm.ANSI_CMD_DECPAM || b == ansiterm.ANSI_CMD_DECPNM:
// non-CSI escape sequence terminator
return true
case b == ansiterm.ANSI_CMD_STR_TERM || b == ansiterm.ANSI_BEL:
// String escape sequence terminator
return true
}
return false
}
func isXtermOscSequence(command []byte, current byte) bool {
return (len(command) >= 2 && command[0] == ansiterm.ANSI_ESCAPE_PRIMARY && command[1] == ansiterm.ANSI_CMD_OSC && current != ansiterm.ANSI_BEL)
}
func isCharacterSelectionCmdChar(b byte) bool {
return (b == ansiterm.ANSI_CMD_G0 || b == ansiterm.ANSI_CMD_G1 || b == ansiterm.ANSI_CMD_G2 || b == ansiterm.ANSI_CMD_G3)
}
// bytesToHex converts a slice of bytes to a human-readable string.
func bytesToHex(b []byte) string {
hex := make([]string, len(b))
for i, ch := range b {
hex[i] = fmt.Sprintf("%X", ch)
}
return strings.Join(hex, "")
}
// ensureInRange adjusts the passed value, if necessary, to ensure it is within
// the passed min / max range.
func ensureInRange(n int16, min int16, max int16) int16 {
+1 -1
View File
@@ -101,7 +101,7 @@ func Wrap(key, plainText []byte) ([]byte, error) {
// Unwrap a key using the RFC 3394 AES Key Wrap Algorithm.
func Unwrap(key, cipherText []byte) ([]byte, error) {
if len(cipherText)%8 != 0 {
if len(cipherText) < 16 || len(cipherText)%8 != 0 {
return nil, ErrUnwrapCiphertext
}
+2
View File
@@ -69,6 +69,8 @@ func (l *lineReader) Read(p []byte) (n int, err error) {
if isPrefix {
return 0, ArmorCorrupt
}
// Trim the line to remove any whitespace
line = bytes.TrimSpace(line)
if bytes.HasPrefix(line, armorEnd) {
l.eof = true
+11 -1
View File
@@ -145,7 +145,14 @@ func Decrypt(priv *PrivateKey, vsG, c, curveOID, fingerprint []byte) (msg []byte
// RFC6637 §8: "m = symm_alg_ID || session key || checksum || pkcs5_padding"
// The last byte should be the length of the padding, as per PKCS5; strip it off.
return m[:len(m)-int(m[len(m)-1])], nil
if len(m) == 0 {
return nil, errors.New("ecdh: invalid padding")
}
padLen := int(m[len(m)-1])
if padLen > len(m) {
return nil, errors.New("ecdh: invalid padding")
}
return m[:len(m)-padLen], nil
}
func buildKey(pub *PublicKey, zb []byte, curveOID, fingerprint []byte, stripLeading, stripTrailing bool) ([]byte, error) {
@@ -196,6 +203,9 @@ func buildKey(pub *PublicKey, zb []byte, curveOID, fingerprint []byte, stripLead
return nil, err
}
mb := h.Sum(nil)
if len(mb) < pub.KDF.Cipher.KeySize() {
return nil, errors.New("ecdh: KDF hash output is shorter than the KDF cipher key size")
}
return mb[:pub.KDF.Cipher.KeySize()], nil // return oBits leftmost bits of MB.
+3
View File
@@ -82,6 +82,9 @@ func Decrypt(priv *PrivateKey, c1, c2 *big.Int) (msg []byte, err error) {
s.Mul(s, c2)
s.Mod(s, priv.P)
em := s.Bytes()
if len(em) == 0 {
return nil, errors.New("elgamal: decryption error")
}
firstByteIsTwo := subtle.ConstantTimeByteEq(em[0], 2)
+10
View File
@@ -180,6 +180,16 @@ func (dke ErrMalformedMessage) Error() string {
return "openpgp: malformed message " + string(dke)
}
type messageTooLargeError int
func (e messageTooLargeError) Error() string {
return "openpgp: decompressed message size exceeds provided limit"
}
// ErrMessageTooLarge is returned if the read data from
// a compressed packet exceeds the provided limit.
var ErrMessageTooLarge error = messageTooLargeError(0)
// ErrEncryptionKeySelection is returned if encryption key selection fails (v2 API).
type ErrEncryptionKeySelection struct {
PrimaryKeyId string
+7 -6
View File
@@ -4,8 +4,11 @@ package algorithm
import (
"crypto/cipher"
"strconv"
"github.com/ProtonMail/go-crypto/eax"
"github.com/ProtonMail/go-crypto/ocb"
"github.com/ProtonMail/go-crypto/openpgp/errors"
)
// AEADMode defines the Authenticated Encryption with Associated Data mode of
@@ -48,8 +51,7 @@ func (mode AEADMode) NonceLength() int {
}
// New returns a fresh instance of the given mode
func (mode AEADMode) New(block cipher.Block) (alg cipher.AEAD) {
var err error
func (mode AEADMode) New(block cipher.Block) (alg cipher.AEAD, err error) {
switch mode {
case AEADModeEAX:
alg, err = eax.NewEAX(block)
@@ -57,9 +59,8 @@ func (mode AEADMode) New(block cipher.Block) (alg cipher.AEAD) {
alg, err = ocb.NewOCB(block)
case AEADModeGCM:
alg, err = cipher.NewGCM(block)
default:
err = errors.UnsupportedError("unknown aead mode: " + strconv.Itoa(int(mode)))
}
if err != nil {
panic(err.Error())
}
return alg
return alg, err
}
+8 -2
View File
@@ -125,7 +125,10 @@ func (c *curve25519) Encaps(rand io.Reader, point []byte) (ephemeral, sharedSecr
// "VB = convert point V to the octet string"
// sharedPoint corresponds to `VB`.
var sharedPoint x25519lib.Key
x25519lib.Shared(&sharedPoint, &ephemeralPrivate, &pubKey)
ok := x25519lib.Shared(&sharedPoint, &ephemeralPrivate, &pubKey)
if !ok {
return nil, nil, errors.KeyInvalidError("ecc: the public key is a low order point")
}
return ephemeralPublic[:], sharedPoint[:], nil
}
@@ -146,7 +149,10 @@ func (c *curve25519) Decaps(vsG, secret []byte) (sharedSecret []byte, err error)
// RFC6637 §8: "Note that the recipient obtains the shared secret by calculating
// S = rV = rvG, where (r,R) is the recipient's key pair."
// sharedPoint corresponds to `S`.
x25519lib.Shared(&sharedPoint, &decodedPrivate, &ephemeralPublic)
ok := x25519lib.Shared(&sharedPoint, &decodedPrivate, &ephemeralPublic)
if !ok {
return nil, errors.KeyInvalidError("ecc: the public key is a low order point")
}
return sharedPoint[:], nil
}
@@ -16,6 +16,8 @@ type ECDSACurve interface {
UnmarshalIntegerPoint([]byte) (x, y *big.Int)
MarshalIntegerSecret(d *big.Int) []byte
UnmarshalIntegerSecret(d []byte) *big.Int
MarshalFieldInteger(d *big.Int) []byte
UnmarshalFieldInteger(d []byte) *big.Int
GenerateECDSA(rand io.Reader) (x, y, secret *big.Int, err error)
Sign(rand io.Reader, x, y, d *big.Int, hash []byte) (r, s *big.Int, err error)
Verify(x, y *big.Int, hash []byte, r, s *big.Int) bool
+13 -1
View File
@@ -56,6 +56,15 @@ func (c *genericCurve) UnmarshalIntegerSecret(d []byte) *big.Int {
return new(big.Int).SetBytes(d)
}
func (c *genericCurve) MarshalFieldInteger(i *big.Int) (b []byte) {
b = make([]byte, (c.Curve.Params().BitSize+7)/8)
return i.FillBytes(b)
}
func (c *genericCurve) UnmarshalFieldInteger(d []byte) *big.Int {
return new(big.Int).SetBytes(d)
}
func (c *genericCurve) GenerateECDH(rand io.Reader) (point, secret []byte, err error) {
secret, x, y, err := elliptic.GenerateKey(c.Curve, rand)
if err != nil {
@@ -78,7 +87,7 @@ func (c *genericCurve) GenerateECDSA(rand io.Reader) (x, y, secret *big.Int, err
func (c *genericCurve) Encaps(rand io.Reader, point []byte) (ephemeral, sharedSecret []byte, err error) {
xP, yP := elliptic.Unmarshal(c.Curve, point)
if xP == nil {
panic("invalid point")
return nil, nil, errors.KeyInvalidError(fmt.Sprintf("ecc (%s): invalid point", c.Curve.Params().Name))
}
d, x, y, err := elliptic.GenerateKey(c.Curve, rand)
@@ -99,6 +108,9 @@ func (c *genericCurve) Encaps(rand io.Reader, point []byte) (ephemeral, sharedSe
func (c *genericCurve) Decaps(ephemeral, secret []byte) (sharedSecret []byte, err error) {
x, y := elliptic.Unmarshal(c.Curve, ephemeral)
if x == nil {
return nil, errors.KeyInvalidError(fmt.Sprintf("ecc (%s): invalid point", c.Curve.Params().Name))
}
zbBig, _ := c.Curve.ScalarMult(x, y, secret)
byteLen := (c.Curve.Params().BitSize + 7) >> 3
zb := make([]byte, byteLen)
+83 -3
View File
@@ -21,7 +21,10 @@ import (
"github.com/ProtonMail/go-crypto/openpgp/errors"
"github.com/ProtonMail/go-crypto/openpgp/internal/algorithm"
"github.com/ProtonMail/go-crypto/openpgp/internal/ecc"
"github.com/ProtonMail/go-crypto/openpgp/mldsa_eddsa"
"github.com/ProtonMail/go-crypto/openpgp/mlkem_ecdh"
"github.com/ProtonMail/go-crypto/openpgp/packet"
"github.com/ProtonMail/go-crypto/openpgp/slhdsa"
"github.com/ProtonMail/go-crypto/openpgp/x25519"
"github.com/ProtonMail/go-crypto/openpgp/x448"
)
@@ -279,7 +282,7 @@ func newSigner(config *packet.Config) (signer interface{}, err error) {
config.RSAPrimes = config.RSAPrimes[2:]
return generateRSAKeyWithPrimes(config.Random(), 2, bits, primes)
}
return rsa.GenerateKey(config.Random(), bits)
return generateRSAKey(config.Random(), bits)
case packet.PubKeyAlgoEdDSA:
if config.V6() {
// Implementations MUST NOT accept or generate v6 key material
@@ -319,6 +322,32 @@ func newSigner(config *packet.Config) (signer interface{}, err error) {
return nil, err
}
return priv, nil
case packet.PubKeyAlgoMldsa65Ed25519, packet.PubKeyAlgoMldsa87Ed448:
if !config.V6() {
return nil, goerrors.New("openpgp: cannot create a non-v6 mldsa_eddsa key")
}
c, err := packet.GetEdDSACurveFromAlgID(config.PublicKeyAlgorithm())
if err != nil {
return nil, err
}
d, err := packet.GetMldsaFromAlgID(config.PublicKeyAlgorithm())
if err != nil {
return nil, err
}
return mldsa_eddsa.GenerateKey(config.Random(), uint8(config.PublicKeyAlgorithm()), c, d)
case packet.PubKeyAlgoSlhdsaShake128s, packet.PubKeyAlgoSlhdsaShake128f, packet.PubKeyAlgoSlhdsaShake256s:
if !config.V6() {
return nil, goerrors.New("openpgp: cannot create a non-v6 SLH-DSH key")
}
scheme, err := packet.GetSlhdsaSchemeFromAlgID(config.PublicKeyAlgorithm())
if err != nil {
return nil, err
}
return slhdsa.GenerateKey(config.Random(), uint8(config.PublicKeyAlgorithm()), scheme)
default:
return nil, errors.InvalidArgumentError("unsupported public key algorithm")
}
@@ -326,7 +355,8 @@ func newSigner(config *packet.Config) (signer interface{}, err error) {
// Generates an encryption/decryption key
func newDecrypter(config *packet.Config) (decrypter interface{}, err error) {
switch config.PublicKeyAlgorithm() {
pubKeyAlgo := config.PublicKeyAlgorithm()
switch pubKeyAlgo {
case packet.PubKeyAlgoRSA:
bits := config.RSAModulusBits()
if bits < 1024 {
@@ -337,7 +367,7 @@ func newDecrypter(config *packet.Config) (decrypter interface{}, err error) {
config.RSAPrimes = config.RSAPrimes[2:]
return generateRSAKeyWithPrimes(config.Random(), 2, bits, primes)
}
return rsa.GenerateKey(config.Random(), bits)
return generateRSAKey(config.Random(), bits)
case packet.PubKeyAlgoEdDSA, packet.PubKeyAlgoECDSA:
fallthrough // When passing EdDSA or ECDSA, we generate an ECDH subkey
case packet.PubKeyAlgoECDH:
@@ -361,6 +391,27 @@ func newDecrypter(config *packet.Config) (decrypter interface{}, err error) {
return x25519.GenerateKey(config.Random())
case packet.PubKeyAlgoEd448, packet.PubKeyAlgoX448: // When passing Ed448, we generate an x448 subkey
return x448.GenerateKey(config.Random())
case packet.PubKeyAlgoMldsa65Ed25519, packet.PubKeyAlgoMldsa87Ed448,
packet.PubKeyAlgoSlhdsaShake128s, packet.PubKeyAlgoSlhdsaShake128f, packet.PubKeyAlgoSlhdsaShake256s:
if pubKeyAlgo, err = packet.GetMatchingMlkem(config.PublicKeyAlgorithm()); err != nil {
return nil, err
}
fallthrough // When passing ML-DSA + EdDSA or ECDSA, we generate a ML-KEM + ECDH subkey
case packet.PubKeyAlgoMlkem768X25519, packet.PubKeyAlgoMlkem1024X448:
if !config.V6() && pubKeyAlgo == packet.PubKeyAlgoMlkem1024X448 {
return nil, goerrors.New("openpgp: cannot create a non-v6 mlkem1024_x448 key")
}
c, err := packet.GetECDHCurveFromAlgID(pubKeyAlgo)
if err != nil {
return nil, err
}
k, err := packet.GetMlkemFromAlgID(pubKeyAlgo)
if err != nil {
return nil, err
}
return mlkem_ecdh.GenerateKey(config.Random(), uint8(pubKeyAlgo), c, k)
default:
return nil, errors.InvalidArgumentError("unsupported public key algorithm")
}
@@ -368,6 +419,29 @@ func newDecrypter(config *packet.Config) (decrypter interface{}, err error) {
var bigOne = big.NewInt(1)
// generateRSAKey generates an RSA keypair and ensures that p < q as required by RFC 9580.
func generateRSAKey(random io.Reader, bits int) (*rsa.PrivateKey, error) {
key, err := rsa.GenerateKey(random, bits)
if err != nil {
return nil, err
}
enforceRSAPrimeOrder(key)
return key, nil
}
// enforceRSAPrimeOrder ensures p < q, as required by RFC 9580 section 5.5.5.1
// for RSA keys. The OpenPGP serialization writes Primes[1] as p and Primes[0]
// as q (so that Go's Qinv is u = p^-1 mod q), hence we need
// Primes[1] < Primes[0]. The Go standard library does not guarantee this, so we
// swap if needed and recompute the precomputed values.
func enforceRSAPrimeOrder(key *rsa.PrivateKey) {
if len(key.Primes) == 2 && key.Primes[0].Cmp(key.Primes[1]) < 0 {
key.Primes[0], key.Primes[1] = key.Primes[1], key.Primes[0]
key.Precomputed = rsa.PrecomputedValues{}
key.Precompute()
}
}
// generateRSAKeyWithPrimes generates a multi-prime RSA keypair of the
// given bit size, using the given random source and pre-populated primes.
func generateRSAKeyWithPrimes(random io.Reader, nprimes int, bits int, prepopulatedPrimes []*big.Int) (*rsa.PrivateKey, error) {
@@ -451,6 +525,12 @@ NextSetOfPrimes:
}
}
// RFC 9580 section 5.5.5.1 requires p < q for RSA keys.
// Primes[1] is serialized as p and Primes[0] as q.
if len(priv.Primes) == 2 && priv.Primes[0].Cmp(priv.Primes[1]) < 0 {
priv.Primes[0], priv.Primes[1] = priv.Primes[1], priv.Primes[0]
}
priv.Precompute()
return priv, nil
}
+15 -2
View File
@@ -106,6 +106,10 @@ func shouldPreferIdentity(existingId, potentialNewId *Identity) bool {
return true
}
if potentialNewId.SelfSignature == nil {
return false
}
if existingId.SelfSignature.IsPrimaryId != nil && *existingId.SelfSignature.IsPrimaryId &&
!(potentialNewId.SelfSignature.IsPrimaryId != nil && *potentialNewId.SelfSignature.IsPrimaryId) {
return false
@@ -142,7 +146,9 @@ func (e *Entity) EncryptionKey(now time.Time) (Key, bool) {
!subkey.PublicKey.KeyExpired(subkey.Sig, now) &&
!subkey.Sig.SigExpired(now) &&
!subkey.Revoked(now) &&
(maxTime.IsZero() || subkey.Sig.CreationTime.After(maxTime)) {
(maxTime.IsZero() ||
subkey.Sig.CreationTime.After(maxTime) ||
(subkey.Sig.CreationTime.Equal(maxTime) && subkey.IsPQ())) {
candidateSubkey = i
maxTime = subkey.Sig.CreationTime
}
@@ -209,7 +215,9 @@ func (e *Entity) signingKeyByIdUsage(now time.Time, id uint64, flags int) (Key,
!subkey.PublicKey.KeyExpired(subkey.Sig, now) &&
!subkey.Sig.SigExpired(now) &&
!subkey.Revoked(now) &&
(maxTime.IsZero() || subkey.Sig.CreationTime.After(maxTime)) &&
(maxTime.IsZero() ||
subkey.Sig.CreationTime.After(maxTime) ||
(subkey.Sig.CreationTime.Equal(maxTime) && subkey.IsPQ())) &&
(id == 0 || subkey.PublicKey.KeyId == id) {
candidateSubkey = idx
maxTime = subkey.Sig.CreationTime
@@ -305,6 +313,11 @@ func (s *Subkey) Revoked(now time.Time) bool {
return revoked(s.Revocations, now)
}
// IsPQ returns true if the algorithm is Post-Quantum safe.
func (s *Subkey) IsPQ() bool {
return s.PublicKey.IsPQ()
}
// Revoked returns whether the key or subkey has been revoked by a self-signature.
// Note that third-party revocation signatures are not supported.
// Note also that Identity revocation should be checked separately.
@@ -0,0 +1,119 @@
// Package mldsa_eddsa implements hybrid ML-DSA + EdDSA encryption for OpenPGP,
// according to https://www.rfc-editor.org/rfc/rfc9980.html#name-composite-signature-schemes.
package mldsa_eddsa
import (
goerrors "errors"
"io"
"github.com/ProtonMail/go-crypto/openpgp/errors"
"github.com/ProtonMail/go-crypto/openpgp/internal/ecc"
"github.com/cloudflare/circl/sign"
"github.com/cloudflare/circl/sign/mldsa/mldsa65"
"github.com/cloudflare/circl/sign/mldsa/mldsa87"
)
const (
MlDsaSeedLen = 32
)
type PublicKey struct {
AlgId uint8
Curve ecc.EdDSACurve
Mldsa sign.Scheme
PublicPoint []byte
PublicMldsa sign.PublicKey
}
type PrivateKey struct {
PublicKey
SecretEc []byte
SecretMldsa sign.PrivateKey
SecretMldsaSeed []byte
}
// GenerateKey generates a ML-DSA + EdDSA composite key as specified in
// https://www.rfc-editor.org/rfc/rfc9980.html#name-key-generation-procedure-2
func GenerateKey(rand io.Reader, algId uint8, c ecc.EdDSACurve, d sign.Scheme) (priv *PrivateKey, err error) {
priv = new(PrivateKey)
priv.PublicKey.AlgId = algId
priv.PublicKey.Curve = c
priv.PublicKey.Mldsa = d
priv.PublicKey.PublicPoint, priv.SecretEc, err = c.GenerateEdDSA(rand)
if err != nil {
return nil, err
}
keySeed := make([]byte, d.SeedSize())
if _, err = rand.Read(keySeed); err != nil {
return nil, err
}
if err := priv.DeriveMlDsaKeys(keySeed, true); err != nil {
return nil, err
}
return priv, nil
}
// DeriveMlDsaKeys derives the ML-DSA keys from the provided seed and stores them inside priv.
func (priv *PrivateKey) DeriveMlDsaKeys(seed []byte, overridePublicKey bool) (err error) {
if len(seed) != MlDsaSeedLen {
return goerrors.New("mldsa_eddsa: ml-dsa secret seed has the wrong length")
}
priv.SecretMldsaSeed = seed
publicKey, privateKey := priv.PublicKey.Mldsa.DeriveKey(priv.SecretMldsaSeed)
if overridePublicKey {
priv.PublicKey.PublicMldsa = publicKey
}
priv.SecretMldsa = privateKey
return nil
}
// Sign generates a ML-DSA + EdDSA composite signature as specified in
// https://www.rfc-editor.org/rfc/rfc9980.html#name-signature-generation
func Sign(priv *PrivateKey, message []byte) (mldsaSig, ecSig []byte, err error) {
ecSig, err = priv.PublicKey.Curve.Sign(priv.PublicKey.PublicPoint, priv.SecretEc, message)
if err != nil {
return nil, nil, err
}
// The default signer interface does not use the hedged variant.
// Thus, we need to use the low level api
switch key := priv.SecretMldsa.(type) {
case *mldsa65.PrivateKey:
mldsaSig = make([]byte, mldsa65.SignatureSize)
err = mldsa65.SignTo(key, message, nil, true, mldsaSig)
case *mldsa87.PrivateKey:
mldsaSig = make([]byte, mldsa87.SignatureSize)
err = mldsa87.SignTo(key, message, nil, true, mldsaSig)
default:
err = goerrors.New("mldsa_eddsa: unsupported ML-DSA private key type")
}
if err != nil {
return nil, nil, err
}
return mldsaSig, ecSig, nil
}
// Verify verifies a ML-DSA + EdDSA composite signature as specified in
// https://www.rfc-editor.org/rfc/rfc9980.html#name-signature-verification
func Verify(pub *PublicKey, message, dSig, ecSig []byte) bool {
return pub.Curve.Verify(pub.PublicPoint, message, ecSig) && pub.Mldsa.Verify(pub.PublicMldsa, message, dSig, nil)
}
// Validate checks that the public key corresponds to the private key
func Validate(priv *PrivateKey) (err error) {
if err = priv.PublicKey.Curve.ValidateEdDSA(priv.PublicKey.PublicPoint, priv.SecretEc); err != nil {
return err
}
if !priv.PublicMldsa.Equal(priv.SecretMldsa.Public()) {
return errors.KeyInvalidError("mldsa_eddsa: invalid public key")
}
return nil
}
+258
View File
@@ -0,0 +1,258 @@
// Package mlkem_ecdh implements hybrid ML-KEM + ECDH encryption, suitable for OpenPGP, experimental.
// It follows the spec https://www.rfc-editor.org/rfc/rfc9980.html#name-composite-kem-schemes
package mlkem_ecdh
import (
goerrors "errors"
"fmt"
"io"
"golang.org/x/crypto/sha3"
"github.com/ProtonMail/go-crypto/openpgp/aes/keywrap"
"github.com/ProtonMail/go-crypto/openpgp/errors"
"github.com/ProtonMail/go-crypto/openpgp/internal/ecc"
"github.com/cloudflare/circl/kem"
)
const (
maxSessionKeyLength = 64
MlKemSeedLen = 64
domSep = "OpenPGPCompositeKDFv1"
)
type PublicKey struct {
AlgId uint8
Curve ecc.ECDHCurve
Mlkem kem.Scheme
PublicMlkem kem.PublicKey
PublicPoint []byte
}
type PrivateKey struct {
PublicKey
SecretEc []byte
SecretMlkem kem.PrivateKey
SecretMlkemSeed []byte
}
// GenerateKey implements ML-KEM + ECC key generation as specified in
// https://www.rfc-editor.org/rfc/rfc9980.html#name-key-generation-procedure
func GenerateKey(rand io.Reader, algId uint8, c ecc.ECDHCurve, k kem.Scheme) (priv *PrivateKey, err error) {
priv = new(PrivateKey)
priv.PublicKey.AlgId = algId
priv.PublicKey.Curve = c
priv.PublicKey.Mlkem = k
priv.PublicKey.PublicPoint, priv.SecretEc, err = c.GenerateECDH(rand)
if err != nil {
return nil, err
}
seed, err := generateRandomSeed(rand, MlKemSeedLen)
if err != nil {
return nil, err
}
if err := priv.DeriveMlKemKeys(seed, true); err != nil {
return nil, err
}
return priv, nil
}
// DeriveMlKemKeys derives the ML-KEM keys from the provided seed and stores them inside priv.
func (priv *PrivateKey) DeriveMlKemKeys(seed []byte, overridePublicKey bool) (err error) {
if len(seed) != MlKemSeedLen {
return goerrors.New("mlkem_ecdh: ml-kem secret seed has the wrong length")
}
priv.SecretMlkemSeed = seed
publicKey, privateKey := priv.PublicKey.Mlkem.DeriveKeyPair(priv.SecretMlkemSeed)
if overridePublicKey {
priv.PublicKey.PublicMlkem = publicKey
}
priv.SecretMlkem = privateKey
return nil
}
// Encrypt implements ML-KEM + ECC encryption as specified in
// https://www.rfc-editor.org/rfc/rfc9980.html#name-encryption-procedure
func Encrypt(rand io.Reader, pub *PublicKey, msg []byte) (kEphemeral, ecEphemeral, ciphertext []byte, err error) {
if len(msg) > maxSessionKeyLength {
return nil, nil, nil, goerrors.New("mlkem_ecdh: session key too long")
}
if len(msg)%8 != 0 {
return nil, nil, nil, goerrors.New("mlkem_ecdh: session key not a multiple of 8")
}
// EC shared secret derivation
ecEphemeral, ecSS, err := pub.Curve.Encaps(rand, pub.PublicPoint)
if err != nil {
return nil, nil, nil, err
}
// ML-KEM shared secret derivation
kyberSeed, err := generateRandomSeed(rand, pub.Mlkem.EncapsulationSeedSize())
if err != nil {
return nil, nil, nil, err
}
kEphemeral, kSS, err := pub.Mlkem.EncapsulateDeterministically(pub.PublicMlkem, kyberSeed)
if err != nil {
return nil, nil, nil, err
}
keyEncryptionKey, err := buildKey(pub, ecSS, ecEphemeral, pub.PublicPoint, kSS)
if err != nil {
return nil, nil, nil, err
}
if ciphertext, err = keywrap.Wrap(keyEncryptionKey, msg); err != nil {
return nil, nil, nil, err
}
return kEphemeral, ecEphemeral, ciphertext, nil
}
// Decrypt implements ML-KEM + ECC decryption as specified in
// https://www.rfc-editor.org/rfc/rfc9980.html#name-decryption-procedure
func Decrypt(priv *PrivateKey, kEphemeral, ecEphemeral, ciphertext []byte) (msg []byte, err error) {
// EC shared secret derivation
ecSS, err := priv.PublicKey.Curve.Decaps(ecEphemeral, priv.SecretEc)
if err != nil {
return nil, err
}
// ML-KEM shared secret derivation
kSS, err := priv.PublicKey.Mlkem.Decapsulate(priv.SecretMlkem, kEphemeral)
if err != nil {
return nil, err
}
kek, err := buildKey(&priv.PublicKey, ecSS, ecEphemeral, priv.PublicPoint, kSS)
if err != nil {
return nil, err
}
return keywrap.Unwrap(kek, ciphertext)
}
// buildKey implements the composite KDF from
// https://www.rfc-editor.org/rfc/rfc9980.html#name-key-combiner
func buildKey(pub *PublicKey, eccSecretPoint, eccCipherText, eccPublicKey, mlkemKeyShare []byte) ([]byte, error) {
/// Set the output `ecdhKeyShare` to `eccSecretPoint`
eccKeyShare := eccSecretPoint
// mlkemKeyShare - the ML-KEM key share encoded as an octet string
// algId - the OpenPGP algorithm ID of the public-key encryption algorithm
// eccKeyShare - the ECDH key share encoded as an octet string
// eccCipherText - the ECDH ciphertext encoded as an octet string
// eccPublicKey - The ECDH public key of the recipient as an octet string
// SHA3-256(mlkemKeyShare || eccKeyShare || eccCipherText || eccPublicKey ||
// algId || "OpenPGPCompositeKDFv1" || len("OpenPGPCompositeKDFv1"))
h := sha3.New256()
_, _ = h.Write(mlkemKeyShare)
_, _ = h.Write(eccKeyShare)
_, _ = h.Write(eccCipherText)
_, _ = h.Write(eccPublicKey)
_, _ = h.Write([]byte{pub.AlgId})
_, _ = h.Write([]byte(domSep))
_, _ = h.Write([]byte{byte(len(domSep))})
return h.Sum(nil), nil
}
// Validate checks that the public key corresponds to the private key
func Validate(priv *PrivateKey) (err error) {
if err = priv.PublicKey.Curve.ValidateECDH(priv.PublicKey.PublicPoint, priv.SecretEc); err != nil {
return err
}
if !priv.PublicKey.PublicMlkem.Equal(priv.SecretMlkem.Public()) {
return errors.KeyInvalidError("mlkem_ecdh: invalid public key")
}
return
}
// EncodeFields encodes an ML-KEM + ECDH session key encryption fields as
// ephemeral ECDH public key | ML-KEM ciphertext | follow byte length | cipherFunction (v3 only) | encryptedSessionKey
// and writes it to writer.
func EncodeFields(w io.Writer, ec, ml, encryptedSessionKey []byte, cipherFunction byte, v6 bool) (err error) {
if _, err = w.Write(ec); err != nil {
return err
}
if _, err = w.Write(ml); err != nil {
return err
}
lenAlgorithm := 0
if !v6 {
lenAlgorithm = 1
}
if _, err = w.Write([]byte{byte(len(encryptedSessionKey) + lenAlgorithm)}); err != nil {
return err
}
if !v6 {
if _, err = w.Write([]byte{cipherFunction}); err != nil {
return err
}
}
if _, err = w.Write(encryptedSessionKey); err != nil {
return err
}
return nil
}
// DecodeFields decodes an ML-KEM + ECDH session key encryption fields as
// ephemeral ECDH public key | ML-KEM ciphertext | follow byte length | cipherFunction (v3 only) | encryptedSessionKey.
func DecodeFields(r io.Reader, lenEcc, lenMlkem int, v6 bool) (ephemeralPublicEcc, ephemeralPublicMlKem, encryptedSessionKey []byte, cipherFunction byte, err error) {
var buf [1]byte
ephemeralPublicEcc = make([]byte, lenEcc)
if _, err = io.ReadFull(r, ephemeralPublicEcc); err != nil {
return
}
ephemeralPublicMlKem = make([]byte, lenMlkem)
if _, err = io.ReadFull(r, ephemeralPublicMlKem); err != nil {
return
}
// A one-octet size of the following fields.
if _, err = io.ReadFull(r, buf[:]); err != nil {
return
}
followingLen := buf[0]
// The one-octet algorithm identifier, if it was passed (in the case of a v3 PKESK packet).
if !v6 {
if _, err = io.ReadFull(r, buf[:]); err != nil {
return
}
cipherFunction = buf[0]
followingLen -= 1
}
// The encrypted session key.
encryptedSessionKey = make([]byte, followingLen)
if _, err = io.ReadFull(r, encryptedSessionKey); err != nil {
return
}
return
}
func generateRandomSeed(rand io.Reader, size int) ([]byte, error) {
randomBytes := make([]byte, size)
if _, err := rand.Read(randomBytes); err != nil {
return nil, fmt.Errorf("failed to generate random bytes: %w", err)
}
return randomBytes, nil
}
+3 -3
View File
@@ -37,7 +37,7 @@ func (conf *AEADConfig) Mode() AEADMode {
// ChunkSizeByte returns the byte indicating the chunk size. The effective
// chunk size is computed with the formula uint64(1) << (chunkSizeByte + 6)
// limit to 16 = 4 MiB
// limit chunkSizeByte to 16 which equals to 2^22 = 4 MiB
// https://www.ietf.org/archive/id/draft-ietf-openpgp-crypto-refresh-07.html#section-5.13.2
func (conf *AEADConfig) ChunkSizeByte() byte {
if conf == nil || conf.ChunkSize == 0 {
@@ -49,8 +49,8 @@ func (conf *AEADConfig) ChunkSizeByte() byte {
switch {
case exponent < 6:
exponent = 6
case exponent > 16:
exponent = 16
case exponent > 22:
exponent = 22
}
return byte(exponent - 6)
+12 -1
View File
@@ -4,6 +4,7 @@ package packet
import (
"io"
"strconv"
"github.com/ProtonMail/go-crypto/openpgp/errors"
"github.com/ProtonMail/go-crypto/openpgp/internal/algorithm"
@@ -50,6 +51,9 @@ func (ae *AEADEncrypted) parse(buf io.Reader) error {
ae.cipher = CipherFunction(c)
ae.mode = mode
ae.chunkSizeByte = headerData[3]
if ae.chunkSizeByte > 16 {
return errors.UnsupportedError("invalid aead chunk size byte: " + strconv.Itoa(int(ae.chunkSizeByte)))
}
return nil
}
@@ -62,8 +66,15 @@ func (ae *AEADEncrypted) Decrypt(ciph CipherFunction, key []byte) (io.ReadCloser
// decrypt prepares an aeadCrypter and returns a ReadCloser from which
// decrypted bytes can be read (see aeadDecrypter.Read()).
func (ae *AEADEncrypted) decrypt(key []byte) (io.ReadCloser, error) {
if ae.cipher.KeySize() != len(key) {
return nil, errors.StructuralError("invalid session key length for cipher: got " + strconv.Itoa(len(key)) + " bytes, but expected " + strconv.Itoa(ae.cipher.KeySize()) + " bytes")
}
blockCipher := ae.cipher.new(key)
aead := ae.mode.new(blockCipher)
aead, err := ae.mode.new(blockCipher)
if err != nil {
return nil, err
}
// Carry the first tagLen bytes
chunkSize := decodeAEADChunkSize(ae.chunkSizeByte)
tagLen := ae.mode.TagLength()
+31
View File
@@ -98,6 +98,16 @@ func (c *Compressed) parse(r io.Reader) error {
return err
}
// LimitedBodyReader wraps the provided body reader with a limiter that restricts
// the number of bytes read to the specified limit.
// If limit is nil, the reader is unbounded.
func (c *Compressed) LimitedBodyReader(limit *int64) io.Reader {
if limit == nil {
return c.Body
}
return &LimitReader{R: c.Body, N: *limit}
}
// compressedWriterCloser represents the serialized compression stream
// header and the compressor. Its Close() method ensures that both the
// compressor and serialized stream header are closed. Its Write()
@@ -159,3 +169,24 @@ func SerializeCompressed(w io.WriteCloser, algo CompressionAlgo, cc *Compression
return
}
// LimitReader is an io.Reader that fails with MessageToLarge if read bytes exceed N.
type LimitReader struct {
R io.Reader // underlying reader
N int64 // max bytes allowed
}
func (l *LimitReader) Read(p []byte) (int, error) {
if l.N <= 0 {
return 0, errors.ErrMessageTooLarge
}
n, err := l.R.Read(p)
l.N -= int64(n)
if err == nil && l.N <= 0 {
err = errors.ErrMessageTooLarge
}
return n, err
}
+40
View File
@@ -47,6 +47,8 @@ var V5Disabled = false
type Config struct {
// Rand provides the source of entropy.
// If nil, the crypto/rand Reader is used.
// Since Go 1.26, standard library calls (e.g., key generation) ignore Rand
// unless GODEBUG=cryptocustomrand=1 is set.
Rand io.Reader
// DefaultHash is the default hash function to be used.
// If zero, SHA-256 is used.
@@ -178,6 +180,23 @@ type Config struct {
// When set to true, a key without flags is treated as if all flags are enabled.
// This behavior is consistent with GPG.
InsecureAllowAllKeyFlagsWhenMissing bool
// InsecureGenerateNonCriticalKeyFlags causes the "Key Flags" signature subpacket
// to be non-critical in newly generated signatures.
// This may be needed for keys to be accepted by older clients who do not recognize
// the subpacket.
// For example, rpm 4.14.3-150400.59.3.1 in OpenSUSE Leap 15.4 does not recognize it.
InsecureGenerateNonCriticalKeyFlags bool
// InsecureGenerateNonCriticalSignatureCreationTime causes the "Signature Creation Time" signature subpacket
// to be non-critical in newly generated signatures.
// This may be needed for keys to be accepted by older clients who do not recognize
// the subpacket.
// For example, yum 3.4.3-168 in CentOS 7 and yum 3.4.3-158 in Amazon Linux 2 do not recognize it.
InsecureGenerateNonCriticalSignatureCreationTime bool
// MaxDecompressedMessageSize specifies the maximum number of bytes that can be
// read from a compressed packet. This serves as an upper limit to prevent
// excessively large decompressed messages.
MaxDecompressedMessageSize *int64
}
func (c *Config) Random() io.Reader {
@@ -415,6 +434,27 @@ func (c *Config) AllowAllKeyFlagsWhenMissing() bool {
return c.InsecureAllowAllKeyFlagsWhenMissing
}
func (c *Config) GenerateNonCriticalKeyFlags() bool {
if c == nil {
return false
}
return c.InsecureGenerateNonCriticalKeyFlags
}
func (c *Config) GenerateNonCriticalSignatureCreationTime() bool {
if c == nil {
return false
}
return c.InsecureGenerateNonCriticalSignatureCreationTime
}
func (c *Config) DecompressedMessageSizeLimit() *int64 {
if c == nil {
return nil
}
return c.MaxDecompressedMessageSize
}
// BoolPointer is a helper function to set a boolean pointer in the Config.
// e.g., config.CheckPacketSequence = BoolPointer(true)
func BoolPointer(value bool) *bool {
+77 -24
View File
@@ -18,6 +18,7 @@ import (
"github.com/ProtonMail/go-crypto/openpgp/elgamal"
"github.com/ProtonMail/go-crypto/openpgp/errors"
"github.com/ProtonMail/go-crypto/openpgp/internal/encoding"
"github.com/ProtonMail/go-crypto/openpgp/mlkem_ecdh"
"github.com/ProtonMail/go-crypto/openpgp/x25519"
"github.com/ProtonMail/go-crypto/openpgp/x448"
)
@@ -33,10 +34,11 @@ type EncryptedKey struct {
CipherFunc CipherFunction // only valid after a successful Decrypt for a v3 packet
Key []byte // only valid after a successful Decrypt
encryptedMPI1, encryptedMPI2 encoding.Field
ephemeralPublicX25519 *x25519.PublicKey // used for x25519
ephemeralPublicX448 *x448.PublicKey // used for x448
encryptedSession []byte // used for x25519 and x448
encryptedMPI1 encoding.Field // used for RSA, Elgamal and ECDH
encryptedMPI2 encoding.Field // used for Elgamal and ECDH
ephemeralPublicEcc []byte // used for X25519, X448 and ML-KEM
ephemeralPublicMlKem []byte // used for ML-KEM
encryptedSession []byte // used for X25519, X448 and ML-KEM
}
func (e *EncryptedKey) parse(r io.Reader) (err error) {
@@ -124,21 +126,29 @@ func (e *EncryptedKey) parse(r io.Reader) (err error) {
return
}
case PubKeyAlgoX25519:
e.ephemeralPublicX25519, e.encryptedSession, cipherFunction, err = x25519.DecodeFields(r, e.Version == 6)
e.ephemeralPublicEcc, e.encryptedSession, cipherFunction, err = x25519.DecodeFields(r, e.Version == 6)
if err != nil {
return
}
case PubKeyAlgoX448:
e.ephemeralPublicX448, e.encryptedSession, cipherFunction, err = x448.DecodeFields(r, e.Version == 6)
e.ephemeralPublicEcc, e.encryptedSession, cipherFunction, err = x448.DecodeFields(r, e.Version == 6)
if err != nil {
return
}
case PubKeyAlgoMlkem768X25519:
if e.ephemeralPublicEcc, e.ephemeralPublicMlKem, e.encryptedSession, cipherFunction, err = mlkem_ecdh.DecodeFields(r, 32, 1088, e.Version == 6); err != nil {
return err
}
case PubKeyAlgoMlkem1024X448:
if e.ephemeralPublicEcc, e.ephemeralPublicMlKem, e.encryptedSession, cipherFunction, err = mlkem_ecdh.DecodeFields(r, 56, 1568, e.Version == 6); err != nil {
return err
}
}
if e.Version < 6 {
switch e.Algo {
case PubKeyAlgoX25519, PubKeyAlgoX448:
case PubKeyAlgoX25519, PubKeyAlgoX448, PubKeyAlgoMlkem768X25519, PubKeyAlgoMlkem1024X448:
e.CipherFunc = CipherFunction(cipherFunction)
// Check for validiy is in the Decrypt method
// Check for validity is in the Decrypt method
}
}
@@ -188,9 +198,11 @@ func (e *EncryptedKey) Decrypt(priv *PrivateKey, config *Config) error {
}
b, err = ecdh.Decrypt(priv.PrivateKey.(*ecdh.PrivateKey), vsG, m, oid, fp)
case PubKeyAlgoX25519:
b, err = x25519.Decrypt(priv.PrivateKey.(*x25519.PrivateKey), e.ephemeralPublicX25519, e.encryptedSession)
b, err = x25519.Decrypt(priv.PrivateKey.(*x25519.PrivateKey), e.ephemeralPublicEcc, e.encryptedSession)
case PubKeyAlgoX448:
b, err = x448.Decrypt(priv.PrivateKey.(*x448.PrivateKey), e.ephemeralPublicX448, e.encryptedSession)
b, err = x448.Decrypt(priv.PrivateKey.(*x448.PrivateKey), e.ephemeralPublicEcc, e.encryptedSession)
case PubKeyAlgoMlkem768X25519, PubKeyAlgoMlkem1024X448:
b, err = mlkem_ecdh.Decrypt(priv.PrivateKey.(*mlkem_ecdh.PrivateKey), e.ephemeralPublicMlKem, e.ephemeralPublicEcc, e.encryptedSession)
default:
err = errors.InvalidArgumentError("cannot decrypt encrypted session key with private key of type " + strconv.Itoa(int(priv.PubKeyAlgo)))
}
@@ -203,6 +215,9 @@ func (e *EncryptedKey) Decrypt(priv *PrivateKey, config *Config) error {
case PubKeyAlgoRSA, PubKeyAlgoRSAEncryptOnly, PubKeyAlgoElGamal, PubKeyAlgoECDH:
keyOffset := 0
if e.Version < 6 {
if len(b) == 0 {
return errors.StructuralError("truncated session key")
}
e.CipherFunc = CipherFunction(b[0])
keyOffset = 1
if !e.CipherFunc.IsSupported() {
@@ -210,22 +225,22 @@ func (e *EncryptedKey) Decrypt(priv *PrivateKey, config *Config) error {
}
}
key, err = decodeChecksumKey(b[keyOffset:])
if err != nil {
return err
}
case PubKeyAlgoX25519, PubKeyAlgoX448:
case PubKeyAlgoX25519, PubKeyAlgoX448, PubKeyAlgoMlkem768X25519, PubKeyAlgoMlkem1024X448:
if e.Version < 6 {
switch e.CipherFunc {
case CipherAES128, CipherAES192, CipherAES256:
break
default:
return errors.StructuralError("v3 PKESK mandates AES as cipher function for x25519 and x448")
return errors.StructuralError("v3 PKESK mandates AES as cipher function for x25519, x448, and PQC")
}
}
key = b[:]
default:
return errors.UnsupportedError("unsupported algorithm for decryption")
}
if err != nil {
return err
}
e.Key = key
return nil
}
@@ -244,6 +259,11 @@ func (e *EncryptedKey) Serialize(w io.Writer) error {
encodedLength = x25519.EncodedFieldsLength(e.encryptedSession, e.Version == 6)
case PubKeyAlgoX448:
encodedLength = x448.EncodedFieldsLength(e.encryptedSession, e.Version == 6)
case PubKeyAlgoMlkem768X25519, PubKeyAlgoMlkem1024X448:
encodedLength = len(e.ephemeralPublicEcc) + len(e.ephemeralPublicMlKem) + len(e.encryptedSession) + 1
if e.Version < 6 {
encodedLength += 1
}
default:
return errors.InvalidArgumentError("don't know how to serialize encrypted key type " + strconv.Itoa(int(e.Algo)))
}
@@ -309,10 +329,13 @@ func (e *EncryptedKey) Serialize(w io.Writer) error {
_, err := w.Write(e.encryptedMPI2.EncodedBytes())
return err
case PubKeyAlgoX25519:
err := x25519.EncodeFields(w, e.ephemeralPublicX25519, e.encryptedSession, byte(e.CipherFunc), e.Version == 6)
err := x25519.EncodeFields(w, e.ephemeralPublicEcc, e.encryptedSession, byte(e.CipherFunc), e.Version == 6)
return err
case PubKeyAlgoX448:
err := x448.EncodeFields(w, e.ephemeralPublicX448, e.encryptedSession, byte(e.CipherFunc), e.Version == 6)
err := x448.EncodeFields(w, e.ephemeralPublicEcc, e.encryptedSession, byte(e.CipherFunc), e.Version == 6)
return err
case PubKeyAlgoMlkem768X25519, PubKeyAlgoMlkem1024X448:
err := mlkem_ecdh.EncodeFields(w, e.ephemeralPublicEcc, e.ephemeralPublicMlKem, e.encryptedSession, byte(e.CipherFunc), e.Version == 6)
return err
default:
panic("internal error")
@@ -346,13 +369,13 @@ func SerializeEncryptedKeyAEADwithHiddenOption(w io.Writer, pub *PublicKey, ciph
if version == 6 && pub.PubKeyAlgo == PubKeyAlgoElGamal {
return errors.InvalidArgumentError("ElGamal v6 PKESK are not allowed")
}
// In v3 PKESKs, for x25519 and x448, mandate using AES
if version == 3 && (pub.PubKeyAlgo == PubKeyAlgoX25519 || pub.PubKeyAlgo == PubKeyAlgoX448) {
switch cipherFunc {
case CipherAES128, CipherAES192, CipherAES256:
break
// In v3 PKESKs, for X25519 and X448, mandate using AES
if version == 3 && cipherFunc != CipherAES128 && cipherFunc != CipherAES192 && cipherFunc != CipherAES256 {
switch pub.PubKeyAlgo {
case PubKeyAlgoX25519, PubKeyAlgoX448, PubKeyAlgoMlkem768X25519, PubKeyAlgoMlkem1024X448:
return errors.InvalidArgumentError("v3 PKESK mandates AES for x25519, x448, and PQC")
default:
return errors.InvalidArgumentError("v3 PKESK mandates AES for x25519 and x448")
break
}
}
@@ -401,7 +424,7 @@ func SerializeEncryptedKeyAEADwithHiddenOption(w io.Writer, pub *PublicKey, ciph
keyOffset = 1
}
encodeChecksumKey(keyBlock[keyOffset:], key)
case PubKeyAlgoX25519, PubKeyAlgoX448:
case PubKeyAlgoX25519, PubKeyAlgoX448, PubKeyAlgoMlkem768X25519, PubKeyAlgoMlkem1024X448:
// algorithm is added in plaintext below
keyBlock = key
}
@@ -417,6 +440,8 @@ func SerializeEncryptedKeyAEADwithHiddenOption(w io.Writer, pub *PublicKey, ciph
return serializeEncryptedKeyX25519(w, config.Random(), buf[:lenHeaderWritten], pub.PublicKey.(*x25519.PublicKey), keyBlock, byte(cipherFunc), version)
case PubKeyAlgoX448:
return serializeEncryptedKeyX448(w, config.Random(), buf[:lenHeaderWritten], pub.PublicKey.(*x448.PublicKey), keyBlock, byte(cipherFunc), version)
case PubKeyAlgoMlkem768X25519, PubKeyAlgoMlkem1024X448:
return serializeEncryptedKeyMlkem(w, config.Random(), buf[:lenHeaderWritten], pub.PublicKey.(*mlkem_ecdh.PublicKey), keyBlock, byte(cipherFunc), version)
case PubKeyAlgoDSA, PubKeyAlgoRSASignOnly:
return errors.InvalidArgumentError("cannot encrypt to public key of type " + strconv.Itoa(int(pub.PubKeyAlgo)))
}
@@ -567,6 +592,9 @@ func checksumKeyMaterial(key []byte) uint16 {
}
func decodeChecksumKey(msg []byte) (key []byte, err error) {
if len(msg) < 2 {
return nil, errors.StructuralError("truncated session key")
}
key = msg[:len(msg)-2]
expectedChecksum := uint16(msg[len(msg)-2])<<8 | uint16(msg[len(msg)-1])
checksum := checksumKeyMaterial(key)
@@ -582,3 +610,28 @@ func encodeChecksumKey(buffer []byte, key []byte) {
buffer[len(key)] = byte(checksum >> 8)
buffer[len(key)+1] = byte(checksum)
}
func serializeEncryptedKeyMlkem(w io.Writer, rand io.Reader, header []byte, pub *mlkem_ecdh.PublicKey, keyBlock []byte, cipherFunc byte, version int) error {
mlE, ecE, c, err := mlkem_ecdh.Encrypt(rand, pub, keyBlock)
if err != nil {
return errors.InvalidArgumentError("ML-KEM + ECDH encryption failed: " + err.Error())
}
packetLen := len(header) /* header length */
packetLen += len(ecE) + len(mlE) + len(c) + 1
if version < 6 {
packetLen += 1
}
err = serializeHeader(w, packetTypeEncryptedKey, packetLen)
if err != nil {
return err
}
_, err = w.Write(header)
if err != nil {
return err
}
return mlkem_ecdh.EncodeFields(w, ecE, mlE, c, cipherFunc, version == 6)
}
+30 -3
View File
@@ -26,6 +26,18 @@ func readFull(r io.Reader, buf []byte) (n int, err error) {
return
}
// readN reads exactly n bytes from r.
func readN(r io.Reader, n uint32) ([]byte, error) {
var buf bytes.Buffer
if _, err := io.CopyN(&buf, r, int64(n)); err != nil {
if err == io.EOF {
err = io.ErrUnexpectedEOF
}
return nil, err
}
return buf.Bytes(), nil
}
// readLength reads an OpenPGP length from r. See RFC 4880, section 4.2.2.
func readLength(r io.Reader) (length int64, isPartial bool, err error) {
var buf [4]byte
@@ -509,13 +521,25 @@ const (
// Deprecated in RFC 4880, Section 13.5. Use key flags instead.
PubKeyAlgoRSAEncryptOnly PublicKeyAlgorithm = 2
PubKeyAlgoRSASignOnly PublicKeyAlgorithm = 3
// PQC DSA algorithms
PubKeyAlgoMldsa65Ed25519 = 30
PubKeyAlgoMldsa87Ed448 = 31
PubKeyAlgoSlhdsaShake128s = 32
PubKeyAlgoSlhdsaShake128f = 33
PubKeyAlgoSlhdsaShake256s = 34
// PQC KEM algorithms
PubKeyAlgoMlkem768X25519 = 35
PubKeyAlgoMlkem1024X448 = 36
)
// CanEncrypt returns true if it's possible to encrypt a message to a public
// key of the given type.
func (pka PublicKeyAlgorithm) CanEncrypt() bool {
switch pka {
case PubKeyAlgoRSA, PubKeyAlgoRSAEncryptOnly, PubKeyAlgoElGamal, PubKeyAlgoECDH, PubKeyAlgoX25519, PubKeyAlgoX448:
case PubKeyAlgoRSA, PubKeyAlgoRSAEncryptOnly, PubKeyAlgoElGamal, PubKeyAlgoECDH, PubKeyAlgoX25519, PubKeyAlgoX448,
PubKeyAlgoMlkem768X25519, PubKeyAlgoMlkem1024X448:
return true
}
return false
@@ -525,7 +549,10 @@ func (pka PublicKeyAlgorithm) CanEncrypt() bool {
// sign a message.
func (pka PublicKeyAlgorithm) CanSign() bool {
switch pka {
case PubKeyAlgoRSA, PubKeyAlgoRSASignOnly, PubKeyAlgoDSA, PubKeyAlgoECDSA, PubKeyAlgoEdDSA, PubKeyAlgoEd25519, PubKeyAlgoEd448:
case PubKeyAlgoRSA, PubKeyAlgoRSASignOnly, PubKeyAlgoDSA, PubKeyAlgoECDSA, PubKeyAlgoEdDSA,
PubKeyAlgoEd25519, PubKeyAlgoEd448,
PubKeyAlgoMldsa65Ed25519, PubKeyAlgoMldsa87Ed448,
PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f, PubKeyAlgoSlhdsaShake256s:
return true
}
return false
@@ -611,7 +638,7 @@ func (mode AEADMode) IsSupported() bool {
}
// new returns a fresh instance of the given mode.
func (mode AEADMode) new(block cipher.Block) cipher.AEAD {
func (mode AEADMode) new(block cipher.Block) (cipher.AEAD, error) {
return algorithm.AEADMode(mode).New(block)
}
+176 -5
View File
@@ -13,6 +13,7 @@ import (
"crypto/sha1"
"crypto/sha256"
"crypto/subtle"
goerrors "errors"
"fmt"
"io"
"math/big"
@@ -27,7 +28,10 @@ import (
"github.com/ProtonMail/go-crypto/openpgp/elgamal"
"github.com/ProtonMail/go-crypto/openpgp/errors"
"github.com/ProtonMail/go-crypto/openpgp/internal/encoding"
"github.com/ProtonMail/go-crypto/openpgp/mldsa_eddsa"
"github.com/ProtonMail/go-crypto/openpgp/mlkem_ecdh"
"github.com/ProtonMail/go-crypto/openpgp/s2k"
"github.com/ProtonMail/go-crypto/openpgp/slhdsa"
"github.com/ProtonMail/go-crypto/openpgp/x25519"
"github.com/ProtonMail/go-crypto/openpgp/x448"
"golang.org/x/crypto/hkdf"
@@ -166,6 +170,10 @@ func NewSignerPrivateKey(creationTime time.Time, signer interface{}) *PrivateKey
pk.PublicKey = *NewEd448PublicKey(creationTime, &pubkey.PublicKey)
case ed448.PrivateKey:
pk.PublicKey = *NewEd448PublicKey(creationTime, &pubkey.PublicKey)
case *mldsa_eddsa.PrivateKey:
pk.PublicKey = *NewMldsaEddsaPublicKey(creationTime, &pubkey.PublicKey)
case *slhdsa.PrivateKey:
pk.PublicKey = *NewSlhdsaPublicKey(creationTime, &pubkey.PublicKey)
default:
panic("openpgp: unknown signer type in NewSignerPrivateKey")
}
@@ -173,7 +181,7 @@ func NewSignerPrivateKey(creationTime time.Time, signer interface{}) *PrivateKey
return pk
}
// NewDecrypterPrivateKey creates a PrivateKey from a *{rsa|elgamal|ecdh|x25519|x448}.PrivateKey.
// NewDecrypterPrivateKey creates a PrivateKey from a *{rsa|elgamal|ecdh|x25519|x448|mlkem_ecdh}.PrivateKey.
func NewDecrypterPrivateKey(creationTime time.Time, decrypter interface{}) *PrivateKey {
pk := new(PrivateKey)
switch priv := decrypter.(type) {
@@ -187,6 +195,8 @@ func NewDecrypterPrivateKey(creationTime time.Time, decrypter interface{}) *Priv
pk.PublicKey = *NewX25519PublicKey(creationTime, &priv.PublicKey)
case *x448.PrivateKey:
pk.PublicKey = *NewX448PublicKey(creationTime, &priv.PublicKey)
case *mlkem_ecdh.PrivateKey:
pk.PublicKey = *NewMlkemEcdhPublicKey(creationTime, &priv.PublicKey)
default:
panic("openpgp: unknown decrypter type in NewDecrypterPrivateKey")
}
@@ -265,6 +275,9 @@ func (pk *PrivateKey) parse(r io.Reader) (err error) {
if pk.s2kParams.Dummy() {
return
}
if !pk.cipher.IsSupported() {
return errors.UnsupportedError("unsupported cipher function in private key")
}
if pk.s2kParams.Mode() == s2k.Argon2S2K && pk.s2kType != S2KAEAD {
return errors.StructuralError("using Argon2 S2K without AEAD is not allowed")
}
@@ -300,6 +313,9 @@ func (pk *PrivateKey) parse(r io.Reader) (err error) {
return
}
if v5 && pk.s2kType == S2KAEAD {
if pk.aead.IvLength() > len(pk.iv) {
return errors.StructuralError("invalid aead IV length for v5 private key")
}
pk.iv = pk.iv[:pk.aead.IvLength()]
}
}
@@ -530,6 +546,40 @@ func serializeEd448PrivateKey(w io.Writer, priv *ed448.PrivateKey) error {
return err
}
// serializeMlkemPrivateKey serializes a ML-KEM + ECC private key according to
// https://www.rfc-editor.org/rfc/rfc9980.html#name-key-material-packets
func serializeMlkemPrivateKey(w io.Writer, priv *mlkem_ecdh.PrivateKey) (err error) {
if _, err = w.Write(priv.SecretEc); err != nil {
return err
}
_, err = w.Write(priv.SecretMlkemSeed)
return err
}
// serializeMldsaEddsaPrivateKey serializes a ML-DSA + EdDSA private key according to
// https://www.rfc-editor.org/rfc/rfc9980.html#name-key-material-packets-2
func serializeMldsaEddsaPrivateKey(w io.Writer, priv *mldsa_eddsa.PrivateKey) error {
if _, err := w.Write(priv.SecretEc); err != nil {
return err
}
if _, err := w.Write(priv.SecretMldsaSeed); err != nil {
return err
}
return nil
}
// serializeSlhDsaPrivateKey serializes a SLH-DSA private key.
func serializeSlhDsaPrivateKey(w io.Writer, priv *slhdsa.PrivateKey) error {
marshalledKey, err := priv.SecretSlhdsa.MarshalBinary()
if err != nil {
return err
}
if _, err := w.Write(marshalledKey); err != nil {
return err
}
return nil
}
// decrypt decrypts an encrypted private key using a decryption key.
func (pk *PrivateKey) decrypt(decryptionKey []byte) error {
if pk.Dummy() {
@@ -542,7 +592,10 @@ func (pk *PrivateKey) decrypt(decryptionKey []byte) error {
var data []byte
switch pk.s2kType {
case S2KAEAD:
aead := pk.aead.new(block)
aead, err := pk.aead.new(block)
if err != nil {
return err
}
additionalData, err := pk.additionalData()
if err != nil {
return err
@@ -694,7 +747,10 @@ func (pk *PrivateKey) encrypt(key []byte, params *s2k.Params, s2kType S2KType, c
if pk.aead == 0 {
return errors.StructuralError("aead mode is not set on key")
}
aead := pk.aead.new(block)
aead, err := pk.aead.new(block)
if err != nil {
return err
}
additionalData, err := pk.additionalData()
if err != nil {
return err
@@ -830,6 +886,12 @@ func (pk *PrivateKey) serializePrivateKey(w io.Writer) (err error) {
err = serializeEd25519PrivateKey(w, priv)
case *ed448.PrivateKey:
err = serializeEd448PrivateKey(w, priv)
case *mlkem_ecdh.PrivateKey:
err = serializeMlkemPrivateKey(w, priv)
case *mldsa_eddsa.PrivateKey:
err = serializeMldsaEddsaPrivateKey(w, priv)
case *slhdsa.PrivateKey:
err = serializeSlhDsaPrivateKey(w, priv)
default:
err = errors.InvalidArgumentError("unknown private key type")
}
@@ -858,6 +920,31 @@ func (pk *PrivateKey) parsePrivateKey(data []byte) (err error) {
return pk.parseEd25519PrivateKey(data)
case PubKeyAlgoEd448:
return pk.parseEd448PrivateKey(data)
case PubKeyAlgoMlkem768X25519:
if !(pk.Version == 4 || pk.Version >= 6) {
return goerrors.New("openpgp: ML-KEM-768+X25519 may only be used with v4 or v6+")
}
return pk.parseMlkemEcdhPrivateKey(data, 32, mlkem_ecdh.MlKemSeedLen)
case PubKeyAlgoMlkem1024X448:
if pk.Version < 6 {
return goerrors.New("openpgp: ML-KEM-1024+X448 may only be used with v6+")
}
return pk.parseMlkemEcdhPrivateKey(data, 56, mlkem_ecdh.MlKemSeedLen)
case PubKeyAlgoMldsa65Ed25519:
if pk.Version < 6 {
return goerrors.New("openpgp: ML-DSA-65+Ed25519 may only be used with v6+")
}
return pk.parseMldsaEddsaPrivateKey(data, 32, mldsa_eddsa.MlDsaSeedLen)
case PubKeyAlgoMldsa87Ed448:
if pk.Version < 6 {
return goerrors.New("openpgp: ML-DSA-87+Ed448 may only be used with v6+")
}
return pk.parseMldsaEddsaPrivateKey(data, 57, mldsa_eddsa.MlDsaSeedLen)
case PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f, PubKeyAlgoSlhdsaShake256s:
if pk.Version < 6 {
return goerrors.New("openpgp: SLH-DSA may only be used with v6+")
}
return pk.parseSlhdsaPrivateKey(data)
default:
err = errors.StructuralError("unknown private key type")
return
@@ -887,8 +974,10 @@ func (pk *PrivateKey) parseRSAPrivateKey(data []byte) (err error) {
rsaPriv.D = new(big.Int).SetBytes(d.Bytes())
rsaPriv.Primes = make([]*big.Int, 2)
rsaPriv.Primes[0] = new(big.Int).SetBytes(p.Bytes())
rsaPriv.Primes[1] = new(big.Int).SetBytes(q.Bytes())
// Mirror serializeRSAPrivateKey: Primes[1] is p and Primes[0] is q, so that
// Go's Qinv matches u = p^-1 mod q.
rsaPriv.Primes[0] = new(big.Int).SetBytes(q.Bytes())
rsaPriv.Primes[1] = new(big.Int).SetBytes(p.Bytes())
if err := rsaPriv.Validate(); err != nil {
return errors.KeyInvalidError(err.Error())
}
@@ -1121,6 +1210,88 @@ func (pk *PrivateKey) applyHKDF(inputKey []byte) []byte {
return encryptionKey
}
// parseMldsaEddsaPrivateKey parses a ML-DSA + EdDSA private key as specified in
// https://www.rfc-editor.org/rfc/rfc9980.html#name-key-material-packets-2
func (pk *PrivateKey) parseMldsaEddsaPrivateKey(data []byte, ecLen, seedLen int) (err error) {
if pk.Version != 6 {
return goerrors.New("openpgp: cannot parse non-v6 ML-DSA + EdDSA key")
}
pub := pk.PublicKey.PublicKey.(*mldsa_eddsa.PublicKey)
priv := new(mldsa_eddsa.PrivateKey)
priv.PublicKey = *pub
if len(data) != ecLen + seedLen {
return errors.StructuralError("wrong ML-DSA+EdDSA key size")
}
ecKey := make([]byte, ecLen)
copy(ecKey, data[:ecLen])
priv.SecretEc = ecKey
seed := make([]byte, seedLen)
copy(seed, data[ecLen:ecLen+seedLen])
if err = priv.DeriveMlDsaKeys(seed, false); err != nil {
return err
}
if err := mldsa_eddsa.Validate(priv); err != nil {
return err
}
pk.PrivateKey = priv
return nil
}
// parseMlkemEcdhPrivateKey parses a ML-KEM + ECC private key as specified in
// https://www.rfc-editor.org/rfc/rfc9980.html#name-key-material-packets
func (pk *PrivateKey) parseMlkemEcdhPrivateKey(data []byte, ecLen, seedLen int) (err error) {
pub := pk.PublicKey.PublicKey.(*mlkem_ecdh.PublicKey)
priv := new(mlkem_ecdh.PrivateKey)
priv.PublicKey = *pub
if len(data) != ecLen + seedLen {
return errors.StructuralError("wrong ML-KEM+ECDH key size")
}
ecKey := make([]byte, ecLen)
copy(ecKey, data[:ecLen])
priv.SecretEc = ecKey
seed := make([]byte, seedLen)
copy(seed, data[ecLen:ecLen+seedLen])
if err = priv.DeriveMlKemKeys(seed, false); err != nil {
return err
}
if err := mlkem_ecdh.Validate(priv); err != nil {
return err
}
pk.PrivateKey = priv
return nil
}
// parseSlhdsaPrivateKey parses a SLH-DSA private key.
func (pk *PrivateKey) parseSlhdsaPrivateKey(data []byte) (err error) {
if pk.Version != 6 {
return goerrors.New("openpgp: cannot parse non-v6 SLH-DSA key")
}
parsedPublicKey := pk.PublicKey.PublicKey.(*slhdsa.PublicKey)
parsedPrivateKey := new(slhdsa.PrivateKey)
parsedPrivateKey.PublicKey = *parsedPublicKey
parsedPrivateKey.SecretSlhdsa, err = parsedPrivateKey.Slhdsa.UnmarshalBinaryPrivateKey(data)
if err != nil {
return goerrors.New("openpgp: failed to unmarshal SLH-DSA key")
}
if err := slhdsa.Validate(parsedPrivateKey); err != nil {
return err
}
pk.PrivateKey = parsedPrivateKey
return nil
}
func validateDSAParameters(priv *dsa.PrivateKey) error {
p := priv.P // group prime
q := priv.Q // subgroup order
+323 -3
View File
@@ -11,6 +11,7 @@ import (
"crypto/sha256"
_ "crypto/sha512"
"encoding/binary"
goerrors "errors"
"fmt"
"hash"
"io"
@@ -28,8 +29,18 @@ import (
"github.com/ProtonMail/go-crypto/openpgp/internal/algorithm"
"github.com/ProtonMail/go-crypto/openpgp/internal/ecc"
"github.com/ProtonMail/go-crypto/openpgp/internal/encoding"
"github.com/ProtonMail/go-crypto/openpgp/mldsa_eddsa"
"github.com/ProtonMail/go-crypto/openpgp/mlkem_ecdh"
"github.com/ProtonMail/go-crypto/openpgp/slhdsa"
"github.com/ProtonMail/go-crypto/openpgp/x25519"
"github.com/ProtonMail/go-crypto/openpgp/x448"
"github.com/cloudflare/circl/kem"
"github.com/cloudflare/circl/kem/mlkem/mlkem1024"
"github.com/cloudflare/circl/kem/mlkem/mlkem768"
"github.com/cloudflare/circl/sign"
"github.com/cloudflare/circl/sign/mldsa/mldsa65"
"github.com/cloudflare/circl/sign/mldsa/mldsa87"
slhdsaCircl "github.com/cloudflare/circl/sign/slhdsa"
)
// PublicKey represents an OpenPGP public key. See RFC 4880, section 5.5.2.
@@ -37,7 +48,7 @@ type PublicKey struct {
Version int
CreationTime time.Time
PubKeyAlgo PublicKeyAlgorithm
PublicKey interface{} // *rsa.PublicKey, *dsa.PublicKey, *ecdsa.PublicKey or *eddsa.PublicKey, *x25519.PublicKey, *x448.PublicKey, *ed25519.PublicKey, *ed448.PublicKey
PublicKey interface{} // *rsa.PublicKey, *dsa.PublicKey, *ecdsa.PublicKey or *eddsa.PublicKey, *x25519.PublicKey, *x448.PublicKey, *ed25519.PublicKey, *ed448.PublicKey, or *mlkem_ecdh.PublicKey
Fingerprint []byte
KeyId uint64
IsSubkey bool
@@ -230,6 +241,42 @@ func NewEd448PublicKey(creationTime time.Time, pub *ed448.PublicKey) *PublicKey
return pk
}
func NewMlkemEcdhPublicKey(creationTime time.Time, pub *mlkem_ecdh.PublicKey) *PublicKey {
pk := &PublicKey{
Version: 4,
CreationTime: creationTime,
PubKeyAlgo: PublicKeyAlgorithm(pub.AlgId),
PublicKey: pub,
}
pk.setFingerprintAndKeyId()
return pk
}
func NewMldsaEddsaPublicKey(creationTime time.Time, pub *mldsa_eddsa.PublicKey) *PublicKey {
pk := &PublicKey{
Version: 6,
CreationTime: creationTime,
PubKeyAlgo: PublicKeyAlgorithm(pub.AlgId),
PublicKey: pub,
}
pk.setFingerprintAndKeyId()
return pk
}
func NewSlhdsaPublicKey(creationTime time.Time, pub *slhdsa.PublicKey) *PublicKey {
pk := &PublicKey{
Version: 6,
CreationTime: creationTime,
PubKeyAlgo: PublicKeyAlgorithm(pub.AlgId),
PublicKey: pub,
}
pk.setFingerprintAndKeyId()
return pk
}
func (pk *PublicKey) parse(r io.Reader) (err error) {
// RFC 4880, section 5.5.2
var buf [6]byte
@@ -258,7 +305,7 @@ func (pk *PublicKey) parse(r io.Reader) (err error) {
}
pk.CreationTime = time.Unix(int64(uint32(buf[1])<<24|uint32(buf[2])<<16|uint32(buf[3])<<8|uint32(buf[4])), 0)
pk.PubKeyAlgo = PublicKeyAlgorithm(buf[5])
// Ignore four-ocet length
// Ignore four-octet length
switch pk.PubKeyAlgo {
case PubKeyAlgoRSA, PubKeyAlgoRSAEncryptOnly, PubKeyAlgoRSASignOnly:
err = pk.parseRSA(r)
@@ -280,6 +327,31 @@ func (pk *PublicKey) parse(r io.Reader) (err error) {
err = pk.parseEd25519(r)
case PubKeyAlgoEd448:
err = pk.parseEd448(r)
case PubKeyAlgoMlkem768X25519:
if !(pk.Version == 4 || pk.Version >= 6) {
return goerrors.New("openpgp: ML-KEM-768+X25519 may only be used with v4 or v6+")
}
err = pk.parseMlkemEcdh(r, 32, mlkem768.PublicKeySize)
case PubKeyAlgoMlkem1024X448:
if pk.Version < 6 {
return goerrors.New("openpgp: ML-KEM-1024+X448 may only be used with v6+")
}
err = pk.parseMlkemEcdh(r, 56, mlkem1024.PublicKeySize)
case PubKeyAlgoMldsa65Ed25519:
if pk.Version < 6 {
return goerrors.New("openpgp: ML-DSA-65+Ed25519 may only be used with v6+")
}
err = pk.parseMldsaEddsa(r, 32, mldsa65.PublicKeySize)
case PubKeyAlgoMldsa87Ed448:
if pk.Version < 6 {
return goerrors.New("openpgp: ML-DSA-87+Ed448 may only be used with v6+")
}
err = pk.parseMldsaEddsa(r, 57, mldsa87.PublicKeySize)
case PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f, PubKeyAlgoSlhdsaShake256s:
if pk.Version < 6 {
return goerrors.New("openpgp: SLH-DSA may only be used with v6+")
}
err = pk.parseSlhDsa(r)
default:
err = errors.UnsupportedError("public key type: " + strconv.Itoa(int(pk.PubKeyAlgo)))
}
@@ -488,6 +560,10 @@ func (pk *PublicKey) parseECDH(r io.Reader) (err error) {
if !ok {
return errors.UnsupportedError("unsupported ECDH KDF cipher: " + strconv.Itoa(int(pk.kdf.Bytes()[2])))
}
// RFC 6637, section 7: the KDF hash must be at least as long as the KEK.
if kdfHash.HashFunc().Size() < kdfCipher.KeySize() {
return errors.StructuralError("ECDH KDF hash output is shorter than the KDF cipher key size")
}
ecdhKey := ecdh.NewPublicKey(c, kdfHash, kdfCipher)
err = ecdhKey.UnmarshalPoint(pk.p.Bytes())
@@ -594,6 +670,98 @@ func (pk *PublicKey) parseEd448(r io.Reader) (err error) {
return
}
// parseMlkemEcdh parses a ML-KEM + ECC public key as specified in
// https://www.rfc-editor.org/rfc/rfc9980.html#name-key-material-packets
func (pk *PublicKey) parseMlkemEcdh(r io.Reader, ecLen, kLen int) (err error) {
ecKey := make([]byte, ecLen)
if _, err = io.ReadFull(r, ecKey); err != nil {
return
}
mlkemKey := make([]byte, kLen)
if _, err = io.ReadFull(r, mlkemKey); err != nil {
return
}
pub := &mlkem_ecdh.PublicKey{
AlgId: uint8(pk.PubKeyAlgo),
PublicPoint: ecKey,
}
if pub.Curve, err = GetECDHCurveFromAlgID(pk.PubKeyAlgo); err != nil {
return err
}
if pub.Mlkem, err = GetMlkemFromAlgID(pk.PubKeyAlgo); err != nil {
return err
}
if pub.PublicMlkem, err = pub.Mlkem.UnmarshalBinaryPublicKey(mlkemKey); err != nil {
return err
}
pk.PublicKey = pub
return
}
// parseMldsaEddsa parses a ML-DSA + EdDSA public key as specified in
// https://www.rfc-editor.org/rfc/rfc9980.html#name-key-material-packets-2
func (pk *PublicKey) parseMldsaEddsa(r io.Reader, ecLen, dLen int) (err error) {
ecKey := make([]byte, ecLen)
if _, err = io.ReadFull(r, ecKey); err != nil {
return
}
mldsaKey := make([]byte, dLen)
if _, err = io.ReadFull(r, mldsaKey); err != nil {
return
}
pub := &mldsa_eddsa.PublicKey{
AlgId: uint8(pk.PubKeyAlgo),
PublicPoint: ecKey,
}
if pub.Curve, err = GetEdDSACurveFromAlgID(pk.PubKeyAlgo); err != nil {
return err
}
if pub.Mldsa, err = GetMldsaFromAlgID(pk.PubKeyAlgo); err != nil {
return err
}
if pub.PublicMldsa, err = pub.Mldsa.UnmarshalBinaryPublicKey(mldsaKey); err != nil {
return err
}
pk.PublicKey = pub
return
}
func (pk *PublicKey) parseSlhDsa(r io.Reader) (err error) {
parsedPublicKey := &slhdsa.PublicKey{
AlgId: uint8(pk.PubKeyAlgo),
}
if parsedPublicKey.Slhdsa, err = GetSlhdsaSchemeFromAlgID(pk.PubKeyAlgo); err != nil {
return err
}
keyLen := parsedPublicKey.Slhdsa.PublicKeySize()
key := make([]byte, keyLen)
if _, err = io.ReadFull(r, key); err != nil {
return err
}
if parsedPublicKey.PublicSlhdsa, err = parsedPublicKey.Slhdsa.UnmarshalBinaryPublicKey(key); err != nil {
return err
}
pk.PublicKey = parsedPublicKey
return nil
}
// SerializeForHash serializes the PublicKey to w with the special packet
// header format needed for hashing.
func (pk *PublicKey) SerializeForHash(w io.Writer) error {
@@ -681,6 +849,21 @@ func (pk *PublicKey) algorithmSpecificByteCount() uint32 {
length += ed25519.PublicKeySize
case PubKeyAlgoEd448:
length += ed448.PublicKeySize
case PubKeyAlgoMlkem768X25519:
length += x25519.KeySize
length += mlkem768.PublicKeySize
case PubKeyAlgoMlkem1024X448:
length += x448.KeySize
length += mlkem1024.PublicKeySize
case PubKeyAlgoMldsa65Ed25519:
length += ed25519.PublicKeySize
length += mldsa65.PublicKeySize
case PubKeyAlgoMldsa87Ed448:
length += ed448.PublicKeySize
length += mldsa87.PublicKeySize
case PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f, PubKeyAlgoSlhdsaShake256s:
publicKey := pk.PublicKey.(*slhdsa.PublicKey)
length += uint32(publicKey.Slhdsa.PublicKeySize())
default:
panic("unknown public key algorithm")
}
@@ -773,13 +956,46 @@ func (pk *PublicKey) serializeWithoutHeaders(w io.Writer) (err error) {
publicKey := pk.PublicKey.(*ed448.PublicKey)
_, err = w.Write(publicKey.Point)
return
case PubKeyAlgoMlkem768X25519, PubKeyAlgoMlkem1024X448:
publicKey := pk.PublicKey.(*mlkem_ecdh.PublicKey)
if _, err = w.Write(publicKey.PublicPoint); err != nil {
return
}
var mlkemBin []byte
mlkemBin, err = publicKey.PublicMlkem.MarshalBinary()
if err != nil {
return
}
_, err = w.Write(mlkemBin)
return
case PubKeyAlgoMldsa65Ed25519, PubKeyAlgoMldsa87Ed448:
publicKey := pk.PublicKey.(*mldsa_eddsa.PublicKey)
if _, err = w.Write(publicKey.PublicPoint); err != nil {
return
}
var mldsaBin []byte
mldsaBin, err = publicKey.PublicMldsa.MarshalBinary()
if err != nil {
return
}
_, err = w.Write(mldsaBin)
return
case PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f, PubKeyAlgoSlhdsaShake256s:
publicKey := pk.PublicKey.(*slhdsa.PublicKey)
var slhdsaBin []byte
slhdsaBin, err = publicKey.PublicSlhdsa.MarshalBinary()
if err != nil {
return
}
_, err = w.Write(slhdsaBin)
return
}
return errors.InvalidArgumentError("bad public-key algorithm")
}
// CanSign returns true iff this public key can generate signatures
func (pk *PublicKey) CanSign() bool {
return pk.PubKeyAlgo != PubKeyAlgoRSAEncryptOnly && pk.PubKeyAlgo != PubKeyAlgoElGamal && pk.PubKeyAlgo != PubKeyAlgoECDH
return pk.PubKeyAlgo.CanSign()
}
// VerifyHashTag returns nil iff sig appears to be a plausible signature of the data
@@ -859,6 +1075,18 @@ func (pk *PublicKey) VerifySignature(signed hash.Hash, sig *Signature) (err erro
return errors.SignatureError("ed448 verification failure")
}
return nil
case PubKeyAlgoMldsa65Ed25519, PubKeyAlgoMldsa87Ed448:
mldsaEddsaPublicKey := pk.PublicKey.(*mldsa_eddsa.PublicKey)
if !mldsa_eddsa.Verify(mldsaEddsaPublicKey, hashBytes, sig.MldsaSig, sig.EdSig) {
return errors.SignatureError("MldsaEddsa verification failure")
}
return nil
case PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f, PubKeyAlgoSlhdsaShake256s:
slhDsaPublicKey := pk.PublicKey.(*slhdsa.PublicKey)
if !slhdsa.Verify(slhDsaPublicKey, hashBytes, sig.SlhdsaSig) {
return errors.SignatureError("Slhdsa verification failure")
}
return nil
default:
return errors.SignatureError("Unsupported public key algorithm used in signature")
}
@@ -1085,6 +1313,15 @@ func (pk *PublicKey) BitLength() (bitLength uint16, err error) {
bitLength = ed25519.PublicKeySize * 8
case PubKeyAlgoEd448:
bitLength = ed448.PublicKeySize * 8
case PubKeyAlgoMlkem768X25519, PubKeyAlgoMlkem1024X448:
publicKey := pk.PublicKey.(*mlkem_ecdh.PublicKey)
bitLength = uint16(publicKey.Mlkem.PublicKeySize() * 8)
case PubKeyAlgoMldsa65Ed25519, PubKeyAlgoMldsa87Ed448:
publicKey := pk.PublicKey.(*mldsa_eddsa.PublicKey)
bitLength = uint16(publicKey.Mldsa.PublicKeySize() * 8)
case PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f, PubKeyAlgoSlhdsaShake256s:
publicKey := pk.PublicKey.(*slhdsa.PublicKey)
bitLength = uint16(publicKey.Slhdsa.PublicKeySize() * 8)
default:
err = errors.InvalidArgumentError("bad public-key algorithm")
}
@@ -1123,3 +1360,86 @@ func (pk *PublicKey) KeyExpired(sig *Signature, currentTime time.Time) bool {
expiry := pk.CreationTime.Add(time.Duration(*sig.KeyLifetimeSecs) * time.Second)
return currentTime.Unix() > expiry.Unix()
}
// IsPQ returns true if the algorithm of this public key is Post-Quantum safe.
func (pg *PublicKey) IsPQ() bool {
switch pg.PubKeyAlgo {
case PubKeyAlgoMlkem768X25519, PubKeyAlgoMlkem1024X448,
PubKeyAlgoMldsa65Ed25519, PubKeyAlgoMldsa87Ed448,
PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f, PubKeyAlgoSlhdsaShake256s:
return true
default:
return false
}
}
func GetMatchingMlkem(algId PublicKeyAlgorithm) (PublicKeyAlgorithm, error) {
switch algId {
case PubKeyAlgoMldsa65Ed25519, PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f:
return PubKeyAlgoMlkem768X25519, nil
case PubKeyAlgoMldsa87Ed448, PubKeyAlgoSlhdsaShake256s:
return PubKeyAlgoMlkem1024X448, nil
default:
return 0, goerrors.New("packet: unsupported pq public key algorithm")
}
}
// GetMlkemFromAlgID returns the ML-KEM instance from the matching KEM
func GetMlkemFromAlgID(algId PublicKeyAlgorithm) (kem.Scheme, error) {
switch algId {
case PubKeyAlgoMlkem768X25519:
return mlkem768.Scheme(), nil
case PubKeyAlgoMlkem1024X448:
return mlkem1024.Scheme(), nil
default:
return nil, goerrors.New("packet: unsupported ML-KEM public key algorithm")
}
}
// GetSlhdsaSchemeFromAlgID returns the SLH-DSA instance from the matching KEM
func GetSlhdsaSchemeFromAlgID(algId PublicKeyAlgorithm) (sign.Scheme, error) {
switch algId {
case PubKeyAlgoSlhdsaShake128s:
return slhdsaCircl.SHAKE_128s.Scheme(), nil
case PubKeyAlgoSlhdsaShake128f:
return slhdsaCircl.SHAKE_128f.Scheme(), nil
case PubKeyAlgoSlhdsaShake256s:
return slhdsaCircl.SHAKE_256s.Scheme(), nil
default:
return nil, goerrors.New("packet: unsupported SLH-DSA public key algorithm")
}
}
// GetECDHCurveFromAlgID returns the ECDH curve instance from the matching KEM
func GetECDHCurveFromAlgID(algId PublicKeyAlgorithm) (ecc.ECDHCurve, error) {
switch algId {
case PubKeyAlgoMlkem768X25519:
return ecc.NewCurve25519(), nil
case PubKeyAlgoMlkem1024X448:
return ecc.NewX448(), nil
default:
return nil, goerrors.New("packet: unsupported ECDH public key algorithm")
}
}
func GetEdDSACurveFromAlgID(algId PublicKeyAlgorithm) (ecc.EdDSACurve, error) {
switch algId {
case PubKeyAlgoMldsa65Ed25519:
return ecc.NewEd25519(), nil
case PubKeyAlgoMldsa87Ed448:
return ecc.NewEd448(), nil
default:
return nil, goerrors.New("packet: unsupported EdDSA public key algorithm")
}
}
func GetMldsaFromAlgID(algId PublicKeyAlgorithm) (sign.Scheme, error) {
switch algId {
case PubKeyAlgoMldsa65Ed25519:
return mldsa65.Scheme(), nil
case PubKeyAlgoMldsa87Ed448:
return mldsa87.Scheme(), nil
default:
return nil, goerrors.New("packet: unsupported ML-DSA public key algorithm")
}
}
+111 -21
View File
@@ -23,6 +23,10 @@ import (
"github.com/ProtonMail/go-crypto/openpgp/errors"
"github.com/ProtonMail/go-crypto/openpgp/internal/algorithm"
"github.com/ProtonMail/go-crypto/openpgp/internal/encoding"
"github.com/ProtonMail/go-crypto/openpgp/mldsa_eddsa"
"github.com/ProtonMail/go-crypto/openpgp/slhdsa"
"github.com/cloudflare/circl/sign/mldsa/mldsa65"
"github.com/cloudflare/circl/sign/mldsa/mldsa87"
)
const (
@@ -81,17 +85,21 @@ type Signature struct {
ECDSASigR, ECDSASigS encoding.Field
EdDSASigR, EdDSASigS encoding.Field
EdSig []byte
MldsaSig []byte
SlhdsaSig []byte
// rawSubpackets contains the unparsed subpackets, in order.
rawSubpackets []outputSubpacket
// The following are optional so are nil when not included in the
// signature.
// The exception is IssuerKeyVersion, which defaults to 0.
SigLifetimeSecs, KeyLifetimeSecs *uint32
PreferredSymmetric, PreferredHash, PreferredCompression []uint8
PreferredCipherSuites [][2]uint8
IssuerKeyId *uint64
IssuerKeyVersion uint8
IssuerFingerprint []byte
SignerUserId *string
IsPrimaryId *bool
@@ -198,7 +206,11 @@ func (sig *Signature) parse(r io.Reader) (err error) {
sig.SigType = SignatureType(buf[0])
sig.PubKeyAlgo = PublicKeyAlgorithm(buf[1])
switch sig.PubKeyAlgo {
case PubKeyAlgoRSA, PubKeyAlgoRSASignOnly, PubKeyAlgoDSA, PubKeyAlgoECDSA, PubKeyAlgoEdDSA, PubKeyAlgoEd25519, PubKeyAlgoEd448:
case PubKeyAlgoRSA, PubKeyAlgoRSASignOnly, PubKeyAlgoDSA,
PubKeyAlgoECDSA, PubKeyAlgoEdDSA,
PubKeyAlgoEd25519, PubKeyAlgoEd448,
PubKeyAlgoMldsa65Ed25519, PubKeyAlgoMldsa87Ed448,
PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f, PubKeyAlgoSlhdsaShake256s:
default:
err = errors.UnsupportedError("public key algorithm " + strconv.Itoa(int(sig.PubKeyAlgo)))
return
@@ -216,19 +228,18 @@ func (sig *Signature) parse(r io.Reader) (err error) {
return errors.UnsupportedError("hash function " + strconv.Itoa(int(buf[2])))
}
var hashedSubpacketsLength int
var hashedSubpacketsLength uint32
if sig.Version == 6 {
// For a v6 signature, a four-octet length is used.
hashedSubpacketsLength =
int(buf[3])<<24 |
int(buf[4])<<16 |
int(buf[5])<<8 |
int(buf[6])
uint32(buf[3])<<24 |
uint32(buf[4])<<16 |
uint32(buf[5])<<8 |
uint32(buf[6])
} else {
hashedSubpacketsLength = int(buf[3])<<8 | int(buf[4])
hashedSubpacketsLength = uint32(buf[3])<<8 | uint32(buf[4])
}
hashedSubpackets := make([]byte, hashedSubpacketsLength)
_, err = readFull(r, hashedSubpackets)
hashedSubpackets, err := readN(r, hashedSubpacketsLength)
if err != nil {
return
}
@@ -257,8 +268,7 @@ func (sig *Signature) parse(r io.Reader) (err error) {
} else {
unhashedSubpacketsLength = uint32(buf[0])<<8 | uint32(buf[1])
}
unhashedSubpackets := make([]byte, unhashedSubpacketsLength)
_, err = readFull(r, unhashedSubpackets)
unhashedSubpackets, err := readN(r, unhashedSubpacketsLength)
if err != nil {
return
}
@@ -336,12 +346,48 @@ func (sig *Signature) parse(r io.Reader) (err error) {
if err != nil {
return
}
case PubKeyAlgoMldsa65Ed25519:
if err = sig.parseMldsaEddsaSignature(r, 64, mldsa65.SignatureSize); err != nil {
return
}
case PubKeyAlgoMldsa87Ed448:
if err = sig.parseMldsaEddsaSignature(r, 114, mldsa87.SignatureSize); err != nil {
return
}
case PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f, PubKeyAlgoSlhdsaShake256s:
if err = sig.parseSlhdsaSignature(r, sig.PubKeyAlgo); err != nil {
return
}
default:
panic("unreachable")
}
return
}
// parseMldsaEddsaSignature parses an ML-DSA + EdDSA signature as specified in
// https://www.rfc-editor.org/rfc/rfc9980.html#name-signature-packet-packet-typ
func (sig *Signature) parseMldsaEddsaSignature(r io.Reader, ecLen, dLen int) (err error) {
sig.EdSig = make([]byte, ecLen)
if _, err = io.ReadFull(r, sig.EdSig); err != nil {
return
}
sig.MldsaSig = make([]byte, dLen)
_, err = io.ReadFull(r, sig.MldsaSig)
return
}
// parseSlhdsaSignature parses an SLH-DSA signature as specified in
func (sig *Signature) parseSlhdsaSignature(r io.Reader, algID PublicKeyAlgorithm) (err error) {
scheme, err := GetSlhdsaSchemeFromAlgID(algID)
if err != nil {
return err
}
sig.SlhdsaSig = make([]byte, scheme.SignatureSize())
_, err = io.ReadFull(r, sig.SlhdsaSig)
return
}
// parseSignatureSubpackets parses subpackets of the main signature packet. See
// RFC 9580, section 5.2.3.1.
func parseSignatureSubpackets(sig *Signature, subpackets []byte, isHashed bool) (err error) {
@@ -455,6 +501,10 @@ func parseSignatureSubpacket(sig *Signature, subpacket []byte, isHashed bool) (r
sig.SigLifetimeSecs = new(uint32)
*sig.SigLifetimeSecs = binary.BigEndian.Uint32(subpacket)
case exportableCertSubpacket:
if len(subpacket) < 1 {
err = errors.StructuralError("exportable certification subpacket with a bad length")
return
}
if subpacket[0] == 0 {
err = errors.UnsupportedError("signature with non-exportable certification")
return
@@ -618,6 +668,11 @@ func parseSignatureSubpacket(sig *Signature, subpacket []byte, isHashed bool) (r
err = errors.StructuralError("Cannot have multiple embedded signatures")
return
}
// A Primary Key Binding signature must not itself carry an embedded
// signature, otherwise the structure can recurse without bound.
if sig.SigType == SigTypePrimaryKeyBinding {
return nil, errors.StructuralError("embedded signature within a primary key binding signature")
}
sig.EmbeddedSignature = new(Signature)
if err := sig.EmbeddedSignature.parse(bytes.NewBuffer(subpacket)); err != nil {
return nil, err
@@ -637,6 +692,7 @@ func parseSignatureSubpacket(sig *Signature, subpacket []byte, isHashed bool) (r
if v >= 5 && l != 32 || v < 5 && l != 20 {
return nil, errors.StructuralError("bad fingerprint length")
}
sig.IssuerKeyVersion = subpacket[0]
sig.IssuerFingerprint = make([]byte, l)
copy(sig.IssuerFingerprint, subpacket[1:])
sig.IssuerKeyId = new(uint64)
@@ -694,14 +750,14 @@ func subpacketLengthLength(length int) int {
}
func (sig *Signature) CheckKeyIdOrFingerprint(pk *PublicKey) bool {
if sig.IssuerFingerprint != nil && len(sig.IssuerFingerprint) >= 20 {
return bytes.Equal(sig.IssuerFingerprint, pk.Fingerprint)
if sig.IssuerKeyVersion != 0 && len(sig.IssuerFingerprint) >= 20 {
return int(sig.IssuerKeyVersion) == pk.Version && bytes.Equal(sig.IssuerFingerprint, pk.Fingerprint)
}
return sig.IssuerKeyId != nil && *sig.IssuerKeyId == pk.KeyId
}
func (sig *Signature) CheckKeyIdOrFingerprintExplicit(fingerprint []byte, keyId uint64) bool {
if sig.IssuerFingerprint != nil && len(sig.IssuerFingerprint) >= 20 && fingerprint != nil {
if sig.IssuerKeyVersion != 0 && len(sig.IssuerFingerprint) >= 20 && fingerprint != nil {
return bytes.Equal(sig.IssuerFingerprint, fingerprint)
}
return sig.IssuerKeyId != nil && *sig.IssuerKeyId == keyId
@@ -918,6 +974,7 @@ func (sig *Signature) Sign(h hash.Hash, priv *PrivateKey, config *Config) (err e
return errors.ErrDummyPrivateKey("dummy key found")
}
sig.Version = priv.PublicKey.Version
sig.IssuerKeyVersion = uint8(priv.PublicKey.Version)
sig.IssuerFingerprint = priv.PublicKey.Fingerprint
if sig.Version < 6 && config.RandomizeSignaturesViaNotation() {
sig.removeNotationsWithName(SaltNotationName)
@@ -933,7 +990,7 @@ func (sig *Signature) Sign(h hash.Hash, priv *PrivateKey, config *Config) (err e
}
sig.Notations = append(sig.Notations, &notation)
}
sig.outSubpackets, err = sig.buildSubpackets(priv.PublicKey)
sig.outSubpackets, err = sig.buildSubpackets(config)
if err != nil {
return err
}
@@ -996,6 +1053,27 @@ func (sig *Signature) Sign(h hash.Hash, priv *PrivateKey, config *Config) (err e
if err == nil {
sig.EdSig = signature
}
case PubKeyAlgoMldsa65Ed25519, PubKeyAlgoMldsa87Ed448:
if sig.Version != 6 {
return errors.StructuralError("cannot use MldsaEdDsa on a non-v6 signature")
}
sk := priv.PrivateKey.(*mldsa_eddsa.PrivateKey)
dSig, ecSig, err := mldsa_eddsa.Sign(sk, digest)
if err == nil {
sig.MldsaSig = dSig
sig.EdSig = ecSig
}
case PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f, PubKeyAlgoSlhdsaShake256s:
if sig.Version != 6 {
return errors.StructuralError("cannot use SLH-DSA on a non-v6 signature")
}
sk := priv.PrivateKey.(*slhdsa.PrivateKey)
dSig, err := slhdsa.Sign(sk, digest)
if err == nil {
sig.SlhdsaSig = dSig
}
default:
err = errors.UnsupportedError("public key algorithm: " + strconv.Itoa(int(sig.PubKeyAlgo)))
}
@@ -1113,7 +1191,7 @@ func (sig *Signature) Serialize(w io.Writer) (err error) {
if len(sig.outSubpackets) == 0 {
sig.outSubpackets = sig.rawSubpackets
}
if sig.RSASignature == nil && sig.DSASigR == nil && sig.ECDSASigR == nil && sig.EdDSASigR == nil && sig.EdSig == nil {
if sig.RSASignature == nil && sig.DSASigR == nil && sig.ECDSASigR == nil && sig.EdDSASigR == nil && sig.EdSig == nil && sig.SlhdsaSig == nil {
return errors.InvalidArgumentError("Signature: need to call Sign, SignUserId or SignKey before Serialize")
}
@@ -1134,6 +1212,11 @@ func (sig *Signature) Serialize(w io.Writer) (err error) {
sigLength = ed25519.SignatureSize
case PubKeyAlgoEd448:
sigLength = ed448.SignatureSize
case PubKeyAlgoMldsa65Ed25519, PubKeyAlgoMldsa87Ed448:
sigLength = len(sig.EdSig)
sigLength += len(sig.MldsaSig)
case PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f, PubKeyAlgoSlhdsaShake256s:
sigLength += len(sig.SlhdsaSig)
default:
panic("impossible")
}
@@ -1240,6 +1323,13 @@ func (sig *Signature) serializeBody(w io.Writer) (err error) {
err = ed25519.WriteSignature(w, sig.EdSig)
case PubKeyAlgoEd448:
err = ed448.WriteSignature(w, sig.EdSig)
case PubKeyAlgoMldsa65Ed25519, PubKeyAlgoMldsa87Ed448:
if _, err = w.Write(sig.EdSig); err != nil {
return
}
_, err = w.Write(sig.MldsaSig)
case PubKeyAlgoSlhdsaShake128s, PubKeyAlgoSlhdsaShake128f, PubKeyAlgoSlhdsaShake256s:
_, err = w.Write(sig.SlhdsaSig)
default:
panic("impossible")
}
@@ -1254,11 +1344,11 @@ type outputSubpacket struct {
contents []byte
}
func (sig *Signature) buildSubpackets(issuer PublicKey) (subpackets []outputSubpacket, err error) {
func (sig *Signature) buildSubpackets(config *Config) (subpackets []outputSubpacket, err error) {
creationTime := make([]byte, 4)
binary.BigEndian.PutUint32(creationTime, uint32(sig.CreationTime.Unix()))
// Signature Creation Time
subpackets = append(subpackets, outputSubpacket{true, creationTimeSubpacket, true, creationTime})
subpackets = append(subpackets, outputSubpacket{true, creationTimeSubpacket, !config.GenerateNonCriticalSignatureCreationTime(), creationTime})
// Signature Expiration Time
if sig.SigLifetimeSecs != nil && *sig.SigLifetimeSecs != 0 {
sigLifetime := make([]byte, 4)
@@ -1357,7 +1447,7 @@ func (sig *Signature) buildSubpackets(issuer PublicKey) (subpackets []outputSubp
if sig.FlagGroupKey {
flags |= KeyFlagGroupKey
}
subpackets = append(subpackets, outputSubpacket{true, keyFlagsSubpacket, true, []byte{flags}})
subpackets = append(subpackets, outputSubpacket{true, keyFlagsSubpacket, !config.GenerateNonCriticalKeyFlags(), []byte{flags}})
}
// Signer's User ID
if sig.SignerUserId != nil {
@@ -1391,8 +1481,8 @@ func (sig *Signature) buildSubpackets(issuer PublicKey) (subpackets []outputSubp
subpackets = append(subpackets, outputSubpacket{true, embeddedSignatureSubpacket, true, buf.Bytes()})
}
// Issuer Fingerprint
if sig.IssuerFingerprint != nil {
contents := append([]uint8{uint8(issuer.Version)}, sig.IssuerFingerprint...)
if sig.IssuerKeyVersion != 0 && sig.IssuerFingerprint != nil {
contents := append([]uint8{sig.IssuerKeyVersion}, sig.IssuerFingerprint...)
subpackets = append(subpackets, outputSubpacket{true, issuerFingerprintSubpacket, sig.Version >= 5, contents})
}
// Intended Recipient Fingerprint
@@ -158,7 +158,10 @@ func (ske *SymmetricKeyEncrypted) decryptV4(key []byte) ([]byte, CipherFunction,
func (ske *SymmetricKeyEncrypted) aeadDecrypt(version int, key []byte) ([]byte, error) {
adata := []byte{0xc3, byte(version), byte(ske.CipherFunc), byte(ske.Mode)}
aead := getEncryptedKeyAeadInstance(ske.CipherFunc, ske.Mode, key, adata, version)
aead, err := getEncryptedKeyAeadInstance(ske.CipherFunc, ske.Mode, key, adata, version)
if err != nil {
return nil, err
}
plaintextKey, err := aead.Open(nil, ske.iv, ske.encryptedKey, adata)
if err != nil {
@@ -291,7 +294,11 @@ func SerializeSymmetricKeyEncryptedAEADReuseKey(w io.Writer, sessionKey []byte,
case 5, 6:
mode := config.AEAD().Mode()
adata := []byte{0xc3, byte(version), byte(cipherFunc), byte(mode)}
aead := getEncryptedKeyAeadInstance(cipherFunc, mode, keyEncryptingKey, adata, version)
var aead cipher.AEAD
aead, err = getEncryptedKeyAeadInstance(cipherFunc, mode, keyEncryptingKey, adata, version)
if err != nil {
return
}
// Sample iv using random reader
iv := make([]byte, config.AEAD().Mode().IvLength())
@@ -315,7 +322,7 @@ func SerializeSymmetricKeyEncryptedAEADReuseKey(w io.Writer, sessionKey []byte,
return
}
func getEncryptedKeyAeadInstance(c CipherFunction, mode AEADMode, inputKey, associatedData []byte, version int) (aead cipher.AEAD) {
func getEncryptedKeyAeadInstance(c CipherFunction, mode AEADMode, inputKey, associatedData []byte, version int) (aead cipher.AEAD, err error) {
var blockCipher cipher.Block
if version > 5 {
hkdfReader := hkdf.New(sha256.New, inputKey, []byte{}, associatedData)
@@ -68,7 +68,11 @@ func (se *SymmetricallyEncrypted) decryptAead(inputKey []byte) (io.ReadCloser, e
return nil, errors.StructuralError(fmt.Sprintf("invalid session key length for cipher: got %d bytes, but expected %d bytes", len(inputKey), se.Cipher.KeySize()))
}
aead, nonce := getSymmetricallyEncryptedAeadInstance(se.Cipher, se.Mode, inputKey, se.Salt[:], se.associatedData())
aead, nonce, err := getSymmetricallyEncryptedAeadInstance(se.Cipher, se.Mode, inputKey, se.Salt[:], se.associatedData())
if err != nil {
return nil, err
}
// Carry the first tagLen bytes
chunkSize := decodeAEADChunkSize(se.ChunkSizeByte)
tagLen := se.Mode.TagLength()
@@ -131,7 +135,10 @@ func serializeSymmetricallyEncryptedAead(ciphertext io.WriteCloser, cipherSuite
return nil, err
}
aead, nonce := getSymmetricallyEncryptedAeadInstance(cipherSuite.Cipher, cipherSuite.Mode, inputKey, salt, prefix)
aead, nonce, err := getSymmetricallyEncryptedAeadInstance(cipherSuite.Cipher, cipherSuite.Mode, inputKey, salt, prefix)
if err != nil {
return nil, err
}
chunkSize := decodeAEADChunkSize(chunkSizeByte)
tagLen := aead.Overhead()
@@ -150,19 +157,22 @@ func serializeSymmetricallyEncryptedAead(ciphertext io.WriteCloser, cipherSuite
}, nil
}
func getSymmetricallyEncryptedAeadInstance(c CipherFunction, mode AEADMode, inputKey, salt, associatedData []byte) (aead cipher.AEAD, nonce []byte) {
func getSymmetricallyEncryptedAeadInstance(c CipherFunction, mode AEADMode, inputKey, salt, associatedData []byte) (aead cipher.AEAD, nonce []byte, err error) {
hkdfReader := hkdf.New(sha256.New, inputKey, salt, associatedData)
encryptionKey := make([]byte, c.KeySize())
_, _ = readFull(hkdfReader, encryptionKey)
blockCipher := c.new(encryptionKey)
aead, err = mode.new(blockCipher)
if err != nil {
return nil, nil, err
}
nonce = make([]byte, mode.IvLength())
// Last 64 bits of nonce are the counter
_, _ = readFull(hkdfReader, nonce[:len(nonce)-8])
blockCipher := c.new(encryptionKey)
aead = mode.new(blockCipher)
return
}
+8 -3
View File
@@ -23,6 +23,9 @@ import (
// SignatureType is the armor type for a PGP signature.
var SignatureType = "PGP SIGNATURE"
// MessageType is the armor type for a PGP message.
var MessageType = "PGP MESSAGE"
// readArmored reads an armored block with the given type.
func readArmored(r io.Reader, expectedType string) (body io.Reader, err error) {
block, err := armor.Decode(r)
@@ -118,7 +121,9 @@ ParsePackets:
// This packet contains the decryption key encrypted to a public key.
md.EncryptedToKeyIds = append(md.EncryptedToKeyIds, p.KeyId)
switch p.Algo {
case packet.PubKeyAlgoRSA, packet.PubKeyAlgoRSAEncryptOnly, packet.PubKeyAlgoElGamal, packet.PubKeyAlgoECDH, packet.PubKeyAlgoX25519, packet.PubKeyAlgoX448:
case packet.PubKeyAlgoRSA, packet.PubKeyAlgoRSAEncryptOnly, packet.PubKeyAlgoElGamal, packet.PubKeyAlgoECDH,
packet.PubKeyAlgoX25519, packet.PubKeyAlgoX448,
packet.PubKeyAlgoMlkem768X25519, packet.PubKeyAlgoMlkem1024X448:
break
default:
continue
@@ -259,7 +264,7 @@ FindLiteralData:
}
switch p := p.(type) {
case *packet.Compressed:
if err := packets.Push(p.Body); err != nil {
if err := packets.Push(p.LimitedBodyReader(config.DecompressedMessageSizeLimit())); err != nil {
return nil, err
}
case *packet.OnePassSignature:
@@ -407,7 +412,7 @@ func (scr *signatureCheckReader) Read(buf []byte) (int, error) {
}
// If signature KeyID matches
if scr.md.SignedBy != nil && *sig.IssuerKeyId == scr.md.SignedByKeyId {
if scr.md.SignedBy != nil && sig.IssuerKeyId != nil && *sig.IssuerKeyId == scr.md.SignedByKeyId {
key := scr.md.SignedBy
signatureError := key.PublicKey.VerifySignature(scr.h, sig)
if signatureError == nil {
File diff suppressed because one or more lines are too long.
+66
View File
@@ -0,0 +1,66 @@
// Package slhdsa implements SLH-DSA-SHAKE for OpenPGP,
// according to https://www.rfc-editor.org/rfc/rfc9980.html#name-slh-dsa-2.
package slhdsa
import (
goerrors "errors"
"fmt"
"io"
"github.com/ProtonMail/go-crypto/openpgp/errors"
"github.com/cloudflare/circl/sign"
)
type PublicKey struct {
AlgId uint8
Slhdsa sign.Scheme
PublicSlhdsa sign.PublicKey
}
type PrivateKey struct {
PublicKey
SecretSlhdsa sign.PrivateKey
}
// GenerateKey generates a SLH-DSA key.
func GenerateKey(rand io.Reader, algId uint8, scheme sign.Scheme) (priv *PrivateKey, err error) {
priv = new(PrivateKey)
priv.PublicKey.AlgId = algId
priv.PublicKey.Slhdsa = scheme
keySeed := make([]byte, scheme.SeedSize())
if _, err = rand.Read(keySeed); err != nil {
return nil, err
}
priv.PublicKey.PublicSlhdsa, priv.SecretSlhdsa = priv.PublicKey.Slhdsa.DeriveKey(keySeed)
return priv, nil
}
// Sign generates a SLH-DSA signature.
func Sign(priv *PrivateKey, message []byte) (signature []byte, err error) {
signature, err = priv.SecretSlhdsa.Sign(nil, message, nil)
if err != nil {
return nil, fmt.Errorf("slhdsa: unable to sign with SLH-DSA: %s", err)
}
if signature == nil {
return nil, goerrors.New("slhdsa: unable to sign with SLH-DSA")
}
return signature, nil
}
// Verify verifies the SLH-DSA signature.
func Verify(pub *PublicKey, message, dSig []byte) bool {
return pub.Slhdsa.Verify(pub.PublicSlhdsa, message, dSig, nil)
}
// Validate checks that the public key corresponds to the private key
func Validate(priv *PrivateKey) (err error) {
if !priv.PublicSlhdsa.Equal(priv.SecretSlhdsa.Public()) {
return errors.KeyInvalidError("slhdsa: invalid public key")
}
return nil
}
+109 -29
View File
@@ -253,34 +253,12 @@ func writeAndSign(payload io.WriteCloser, candidateHashes []uint8, signed *Entit
}
var hash crypto.Hash
for _, hashId := range candidateHashes {
if h, ok := algorithm.HashIdToHash(hashId); ok && h.Available() {
hash = h
break
}
}
// If the hash specified by config is a candidate, we'll use that.
if configuredHash := config.Hash(); configuredHash.Available() {
for _, hashId := range candidateHashes {
if h, ok := algorithm.HashIdToHash(hashId); ok && h == configuredHash {
hash = h
break
}
}
}
if hash == 0 {
hashId := candidateHashes[0]
name, ok := algorithm.HashIdToString(hashId)
if !ok {
name = "#" + strconv.Itoa(int(hashId))
}
return nil, errors.InvalidArgumentError("cannot encrypt because no candidate hash functions are compiled in. (Wanted " + name + " in this case.)")
}
var salt []byte
if signer != nil {
if hash, err = selectHash(candidateHashes, config.Hash(), signer); err != nil {
return nil, err
}
var opsVersion = 3
if signer.Version == 6 {
opsVersion = signer.Version
@@ -391,6 +369,7 @@ func encrypt(keyWriter io.Writer, dataWriter io.Writer, to []*Entity, signed *En
// AEAD is used only if config enables it and every key supports it
aeadSupported := config.AEAD() != nil
allPQ := len(to) > 0
for i := range to {
var ok bool
encryptKeys[i], ok = to[i].EncryptionKey(config.Now())
@@ -398,6 +377,10 @@ func encrypt(keyWriter io.Writer, dataWriter io.Writer, to []*Entity, signed *En
return nil, errors.InvalidArgumentError("cannot encrypt a message to key id " + strconv.FormatUint(to[i].PrimaryKey.KeyId, 16) + " because it has no valid encryption keys")
}
if !encryptKeys[i].PublicKey.IsPQ() {
allPQ = false
}
primarySelfSignature, _ := to[i].PrimarySelfSignature()
if primarySelfSignature == nil {
return nil, errors.InvalidArgumentError("entity without a self-signature")
@@ -424,8 +407,12 @@ func encrypt(keyWriter io.Writer, dataWriter io.Writer, to []*Entity, signed *En
candidateHashes = []uint8{hashToHashId(crypto.SHA256)}
}
if len(candidateCipherSuites) == 0 {
// https://www.ietf.org/archive/id/draft-ietf-openpgp-crypto-refresh-07.html#section-9.6
candidateCipherSuites = [][2]uint8{{uint8(packet.CipherAES128), uint8(packet.AEADModeOCB)}}
if allPQ {
candidateCipherSuites = [][2]uint8{{uint8(packet.CipherAES256), uint8(packet.AEADModeOCB)}}
} else {
// https://www.ietf.org/archive/id/draft-ietf-openpgp-crypto-refresh-07.html#section-9.6
candidateCipherSuites = [][2]uint8{{uint8(packet.CipherAES128), uint8(packet.AEADModeOCB)}}
}
}
cipher := packet.CipherFunction(candidateCiphers[0])
@@ -558,15 +545,37 @@ func (s signatureWriter) Close() error {
return s.encryptedData.Close()
}
func selectHashForSigningKey(config *packet.Config, signer *packet.PublicKey) crypto.Hash {
acceptableHashes := acceptableHashesToWrite(signer)
hash, ok := algorithm.HashToHashId(config.Hash())
if !ok {
return config.Hash()
}
for _, acceptableHashes := range acceptableHashes {
if acceptableHashes == hash {
return config.Hash()
}
}
if len(acceptableHashes) > 0 {
defaultAcceptedHash, ok := algorithm.HashIdToHash(acceptableHashes[0])
if ok {
return defaultAcceptedHash
}
}
return config.Hash()
}
func createSignaturePacket(signer *packet.PublicKey, sigType packet.SignatureType, config *packet.Config) *packet.Signature {
sigLifetimeSecs := config.SigLifetime()
hash := selectHashForSigningKey(config, signer)
return &packet.Signature{
Version: signer.Version,
SigType: sigType,
PubKeyAlgo: signer.PubKeyAlgo,
Hash: config.Hash(),
Hash: hash,
CreationTime: config.Now(),
IssuerKeyId: &signer.KeyId,
IssuerKeyVersion: uint8(signer.Version),
IssuerFingerprint: signer.Fingerprint,
Notations: config.Notations(),
SigLifetimeSecs: &sigLifetimeSecs,
@@ -618,3 +627,74 @@ func handleCompression(compressed io.WriteCloser, candidateCompression []uint8,
}
return data, nil
}
// selectHash selects the preferred hash given the candidateHashes and the configuredHash
func selectHash(candidateHashes []byte, configuredHash crypto.Hash, signer *packet.PrivateKey) (hash crypto.Hash, err error) {
acceptableHashes := acceptableHashesToWrite(&signer.PublicKey)
candidateHashes = intersectPreferences(acceptableHashes, candidateHashes)
for _, hashId := range candidateHashes {
if h, ok := algorithm.HashIdToHash(hashId); ok && h.Available() {
hash = h
break
}
}
// If the hash specified by config is a candidate, we'll use that.
if configuredHash.Available() {
for _, hashId := range candidateHashes {
if h, ok := algorithm.HashIdToHash(hashId); ok && h == configuredHash {
hash = h
break
}
}
}
if hash == 0 {
if len(acceptableHashes) > 0 {
if h, ok := algorithm.HashIdToHash(acceptableHashes[0]); ok {
hash = h
} else {
return 0, errors.UnsupportedError("no candidate hash functions are compiled in.")
}
} else {
return 0, errors.UnsupportedError("no candidate hash functions are compiled in.")
}
}
return
}
func acceptableHashesToWrite(singingKey *packet.PublicKey) []uint8 {
switch singingKey.PubKeyAlgo {
case packet.PubKeyAlgoEd448, packet.PubKeyAlgoMldsa87Ed448, packet.PubKeyAlgoSlhdsaShake256s:
return []uint8{
hashToHashId(crypto.SHA512),
hashToHashId(crypto.SHA3_512),
}
case packet.PubKeyAlgoECDSA, packet.PubKeyAlgoEdDSA:
if curve, err := singingKey.Curve(); err == nil {
if curve == packet.Curve448 ||
curve == packet.CurveNistP521 ||
curve == packet.CurveBrainpoolP512 {
return []uint8{
hashToHashId(crypto.SHA512),
hashToHashId(crypto.SHA3_512),
}
} else if curve == packet.CurveBrainpoolP384 ||
curve == packet.CurveNistP384 {
return []uint8{
hashToHashId(crypto.SHA384),
hashToHashId(crypto.SHA512),
hashToHashId(crypto.SHA3_512),
}
}
}
}
return []uint8{
hashToHashId(crypto.SHA256),
hashToHashId(crypto.SHA384),
hashToHashId(crypto.SHA512),
hashToHashId(crypto.SHA3_256),
hashToHashId(crypto.SHA3_512),
}
}
+12 -16
View File
@@ -80,7 +80,7 @@ func generateKey(rand io.Reader, privateKey *x25519lib.Key, publicKey *x25519lib
// Encrypt encrypts a sessionKey with x25519 according to
// the OpenPGP crypto refresh specification section 5.1.6. The function assumes that the
// sessionKey has the correct format and padding according to the specification.
func Encrypt(rand io.Reader, publicKey *PublicKey, sessionKey []byte) (ephemeralPublicKey *PublicKey, encryptedSessionKey []byte, err error) {
func Encrypt(rand io.Reader, publicKey *PublicKey, sessionKey []byte) (ephemeralPublicKey, encryptedSessionKey []byte, err error) {
var ephemeralPrivate, ephemeralPublic, staticPublic, shared x25519lib.Key
// Check that the input static public key has 32 bytes
if len(publicKey.Point) != KeySize {
@@ -99,11 +99,9 @@ func Encrypt(rand io.Reader, publicKey *PublicKey, sessionKey []byte) (ephemeral
err = errors.KeyInvalidError("x25519: the public key is a low order point")
return
}
ephemeralPublicKey = ephemeralPublic[:]
// Derive the encryption key from the shared secret
encryptionKey := applyHKDF(ephemeralPublic[:], publicKey.Point[:], shared[:])
ephemeralPublicKey = &PublicKey{
Point: ephemeralPublic[:],
}
encryptionKey := applyHKDF(ephemeralPublicKey, publicKey.Point[:], shared[:])
// Encrypt the sessionKey with aes key wrapping
encryptedSessionKey, err = keywrap.Wrap(encryptionKey, sessionKey)
return
@@ -111,14 +109,14 @@ func Encrypt(rand io.Reader, publicKey *PublicKey, sessionKey []byte) (ephemeral
// Decrypt decrypts a session key stored in ciphertext with the provided x25519
// private key and ephemeral public key.
func Decrypt(privateKey *PrivateKey, ephemeralPublicKey *PublicKey, ciphertext []byte) (encodedSessionKey []byte, err error) {
func Decrypt(privateKey *PrivateKey, ephemeralPublicKey, ciphertext []byte) (encodedSessionKey []byte, err error) {
var ephemeralPublic, staticPrivate, shared x25519lib.Key
// Check that the input ephemeral public key has 32 bytes
if len(ephemeralPublicKey.Point) != KeySize {
if len(ephemeralPublicKey) != KeySize {
err = errors.KeyInvalidError("x25519: the public key has the wrong size")
return
}
copy(ephemeralPublic[:], ephemeralPublicKey.Point)
copy(ephemeralPublic[:], ephemeralPublicKey)
subtle.ConstantTimeCopy(1, staticPrivate[:], privateKey.Secret)
// Compute shared key
ok := x25519lib.Shared(&shared, &staticPrivate, &ephemeralPublic)
@@ -127,7 +125,7 @@ func Decrypt(privateKey *PrivateKey, ephemeralPublicKey *PublicKey, ciphertext [
return
}
// Derive the encryption key from the shared secret
encryptionKey := applyHKDF(ephemeralPublicKey.Point[:], privateKey.PublicKey.Point[:], shared[:])
encryptionKey := applyHKDF(ephemeralPublicKey, privateKey.PublicKey.Point[:], shared[:])
// Decrypt the session key with aes key wrapping
encodedSessionKey, err = keywrap.Unwrap(encryptionKey, ciphertext)
return
@@ -168,12 +166,12 @@ func EncodedFieldsLength(encryptedSessionKey []byte, v6 bool) int {
// EncodeField encodes x25519 session key encryption fields as
// ephemeral x25519 public key | follow byte length | cipherFunction (v3 only) | encryptedSessionKey
// and writes it to writer.
func EncodeFields(writer io.Writer, ephemeralPublicKey *PublicKey, encryptedSessionKey []byte, cipherFunction byte, v6 bool) (err error) {
func EncodeFields(writer io.Writer, ephemeralPublicKey, encryptedSessionKey []byte, cipherFunction byte, v6 bool) (err error) {
lenAlgorithm := 0
if !v6 {
lenAlgorithm = 1
}
if _, err = writer.Write(ephemeralPublicKey.Point); err != nil {
if _, err = writer.Write(ephemeralPublicKey); err != nil {
return err
}
if _, err = writer.Write([]byte{byte(len(encryptedSessionKey) + lenAlgorithm)}); err != nil {
@@ -190,13 +188,11 @@ func EncodeFields(writer io.Writer, ephemeralPublicKey *PublicKey, encryptedSess
// DecodeField decodes a x25519 session key encryption as
// ephemeral x25519 public key | follow byte length | cipherFunction (v3 only) | encryptedSessionKey.
func DecodeFields(reader io.Reader, v6 bool) (ephemeralPublicKey *PublicKey, encryptedSessionKey []byte, cipherFunction byte, err error) {
func DecodeFields(reader io.Reader, v6 bool) (ephemeralPublicKey, encryptedSessionKey []byte, cipherFunction byte, err error) {
var buf [1]byte
ephemeralPublicKey = &PublicKey{
Point: make([]byte, KeySize),
}
ephemeralPublicKey = make([]byte, KeySize)
// 32 octets representing an ephemeral x25519 public key.
if _, err = io.ReadFull(reader, ephemeralPublicKey.Point); err != nil {
if _, err = io.ReadFull(reader, ephemeralPublicKey); err != nil {
return nil, nil, 0, err
}
// A one-octet size of the following fields.
+12 -16
View File
@@ -81,7 +81,7 @@ func generateKey(rand io.Reader, privateKey *x448lib.Key, publicKey *x448lib.Key
// Encrypt encrypts a sessionKey with x448 according to
// the OpenPGP crypto refresh specification section 5.1.7. The function assumes that the
// sessionKey has the correct format and padding according to the specification.
func Encrypt(rand io.Reader, publicKey *PublicKey, sessionKey []byte) (ephemeralPublicKey *PublicKey, encryptedSessionKey []byte, err error) {
func Encrypt(rand io.Reader, publicKey *PublicKey, sessionKey []byte) (ephemeralPublicKey, encryptedSessionKey []byte, err error) {
var ephemeralPrivate, ephemeralPublic, staticPublic, shared x448lib.Key
// Check that the input static public key has 56 bytes.
if len(publicKey.Point) != KeySize {
@@ -99,11 +99,9 @@ func Encrypt(rand io.Reader, publicKey *PublicKey, sessionKey []byte) (ephemeral
err = errors.KeyInvalidError("x448: the public key is a low order point")
return nil, nil, err
}
ephemeralPublicKey = ephemeralPublic[:]
// Derive the encryption key from the shared secret.
encryptionKey := applyHKDF(ephemeralPublic[:], publicKey.Point[:], shared[:])
ephemeralPublicKey = &PublicKey{
Point: ephemeralPublic[:],
}
encryptionKey := applyHKDF(ephemeralPublicKey, publicKey.Point[:], shared[:])
// Encrypt the sessionKey with aes key wrapping.
encryptedSessionKey, err = keywrap.Wrap(encryptionKey, sessionKey)
if err != nil {
@@ -114,14 +112,14 @@ func Encrypt(rand io.Reader, publicKey *PublicKey, sessionKey []byte) (ephemeral
// Decrypt decrypts a session key stored in ciphertext with the provided x448
// private key and ephemeral public key.
func Decrypt(privateKey *PrivateKey, ephemeralPublicKey *PublicKey, ciphertext []byte) (encodedSessionKey []byte, err error) {
func Decrypt(privateKey *PrivateKey, ephemeralPublicKey, ciphertext []byte) (encodedSessionKey []byte, err error) {
var ephemeralPublic, staticPrivate, shared x448lib.Key
// Check that the input ephemeral public key has 56 bytes.
if len(ephemeralPublicKey.Point) != KeySize {
if len(ephemeralPublicKey) != KeySize {
err = errors.KeyInvalidError("x448: the public key has the wrong size")
return nil, err
}
copy(ephemeralPublic[:], ephemeralPublicKey.Point)
copy(ephemeralPublic[:], ephemeralPublicKey)
subtle.ConstantTimeCopy(1, staticPrivate[:], privateKey.Secret)
// Compute shared key.
ok := x448lib.Shared(&shared, &staticPrivate, &ephemeralPublic)
@@ -130,7 +128,7 @@ func Decrypt(privateKey *PrivateKey, ephemeralPublicKey *PublicKey, ciphertext [
return nil, err
}
// Derive the encryption key from the shared secret.
encryptionKey := applyHKDF(ephemeralPublicKey.Point[:], privateKey.PublicKey.Point[:], shared[:])
encryptionKey := applyHKDF(ephemeralPublicKey, privateKey.PublicKey.Point[:], shared[:])
// Decrypt the session key with aes key wrapping.
encodedSessionKey, err = keywrap.Unwrap(encryptionKey, ciphertext)
if err != nil {
@@ -174,12 +172,12 @@ func EncodedFieldsLength(encryptedSessionKey []byte, v6 bool) int {
// EncodeField encodes x448 session key encryption fields as
// ephemeral x448 public key | follow byte length | cipherFunction (v3 only) | encryptedSessionKey
// and writes it to writer.
func EncodeFields(writer io.Writer, ephemeralPublicKey *PublicKey, encryptedSessionKey []byte, cipherFunction byte, v6 bool) (err error) {
func EncodeFields(writer io.Writer, ephemeralPublicKey, encryptedSessionKey []byte, cipherFunction byte, v6 bool) (err error) {
lenAlgorithm := 0
if !v6 {
lenAlgorithm = 1
}
if _, err = writer.Write(ephemeralPublicKey.Point); err != nil {
if _, err = writer.Write(ephemeralPublicKey); err != nil {
return err
}
if _, err = writer.Write([]byte{byte(len(encryptedSessionKey) + lenAlgorithm)}); err != nil {
@@ -198,13 +196,11 @@ func EncodeFields(writer io.Writer, ephemeralPublicKey *PublicKey, encryptedSess
// DecodeField decodes a x448 session key encryption as
// ephemeral x448 public key | follow byte length | cipherFunction (v3 only) | encryptedSessionKey.
func DecodeFields(reader io.Reader, v6 bool) (ephemeralPublicKey *PublicKey, encryptedSessionKey []byte, cipherFunction byte, err error) {
func DecodeFields(reader io.Reader, v6 bool) (ephemeralPublicKey, encryptedSessionKey []byte, cipherFunction byte, err error) {
var buf [1]byte
ephemeralPublicKey = &PublicKey{
Point: make([]byte, KeySize),
}
ephemeralPublicKey = make([]byte, KeySize)
// 56 octets representing an ephemeral x448 public key.
if _, err = io.ReadFull(reader, ephemeralPublicKey.Point); err != nil {
if _, err = io.ReadFull(reader, ephemeralPublicKey); err != nil {
return nil, nil, 0, err
}
// A one-octet size of the following fields.
-10
View File
@@ -1,10 +0,0 @@
language: go
go:
- 1.0.3
- 1.1.2
- 1.2
- tip
install:
- go get github.com/bmizerany/assert
notifications:
email: false
+5 -2
View File
@@ -2,7 +2,10 @@
a Go package to interact with arbitrary JSON
[![Build Status](https://secure.travis-ci.org/bitly/go-simplejson.png)](http://travis-ci.org/bitly/go-simplejson)
[![Build Status](https://github.com/bitly/go-simplejson/actions/workflows/ci.yaml/badge.svg)](https://github.com/bitly/go-simplejson/actions)
[![GoDoc](https://pkg.go.dev/badge/github.com/bitly/go-simplejson)](https://pkg.go.dev/github.com/bitly/go-simplejson)
[![GitHub release](https://img.shields.io/github/release/bitly/go-simplejson.svg)](https://github.com/bitly/go-simplejson/releases/latest)
### Importing
@@ -10,4 +13,4 @@ a Go package to interact with arbitrary JSON
### Documentation
Visit the docs on [gopkgdoc](http://godoc.org/github.com/bitly/go-simplejson)
Visit the docs on [Go package discovery & docs](https://pkg.go.dev/github.com/bitly/go-simplejson)
+35 -23
View File
@@ -8,7 +8,7 @@ import (
// returns the current implementation version
func Version() string {
return "0.5.0"
return "0.5.1"
}
type Json struct {
@@ -115,7 +115,8 @@ func (j *Json) Del(key string) {
// for `key` in its `map` representation
//
// useful for chaining operations (to traverse a nested JSON):
// js.Get("top_level").Get("dict").Get("value").Int()
//
// js.Get("top_level").Get("dict").Get("value").Int()
func (j *Json) Get(key string) *Json {
m, err := j.Map()
if err == nil {
@@ -129,7 +130,7 @@ func (j *Json) Get(key string) *Json {
// GetPath searches for the item as specified by the branch
// without the need to deep dive using Get()'s.
//
// js.GetPath("top_level", "dict")
// js.GetPath("top_level", "dict")
func (j *Json) GetPath(branch ...string) *Json {
jin := j
for _, p := range branch {
@@ -143,7 +144,8 @@ func (j *Json) GetPath(branch ...string) *Json {
//
// this is the analog to Get when accessing elements of
// a json array instead of a json object:
// js.Get("top_level").Get("array").GetIndex(1).Get("key").Int()
//
// js.Get("top_level").Get("array").GetIndex(1).Get("key").Int()
func (j *Json) GetIndex(index int) *Json {
a, err := j.Array()
if err == nil {
@@ -158,9 +160,10 @@ func (j *Json) GetIndex(index int) *Json {
// a `bool` identifying success or failure
//
// useful for chained operations when success is important:
// if data, ok := js.Get("top_level").CheckGet("inner"); ok {
// log.Println(data)
// }
//
// if data, ok := js.Get("top_level").CheckGet("inner"); ok {
// log.Println(data)
// }
func (j *Json) CheckGet(key string) (*Json, bool) {
m, err := j.Map()
if err == nil {
@@ -225,7 +228,7 @@ func (j *Json) StringArray() ([]string, error) {
}
s, ok := a.(string)
if !ok {
return nil, err
return nil, errors.New("type assertion to []string failed")
}
retArr = append(retArr, s)
}
@@ -235,9 +238,10 @@ func (j *Json) StringArray() ([]string, error) {
// MustArray guarantees the return of a `[]interface{}` (with optional default)
//
// useful when you want to interate over array values in a succinct manner:
// for i, v := range js.Get("results").MustArray() {
// fmt.Println(i, v)
// }
//
// for i, v := range js.Get("results").MustArray() {
// fmt.Println(i, v)
// }
func (j *Json) MustArray(args ...[]interface{}) []interface{} {
var def []interface{}
@@ -260,9 +264,10 @@ func (j *Json) MustArray(args ...[]interface{}) []interface{} {
// MustMap guarantees the return of a `map[string]interface{}` (with optional default)
//
// useful when you want to interate over map values in a succinct manner:
// for k, v := range js.Get("dictionary").MustMap() {
// fmt.Println(k, v)
// }
//
// for k, v := range js.Get("dictionary").MustMap() {
// fmt.Println(k, v)
// }
func (j *Json) MustMap(args ...map[string]interface{}) map[string]interface{} {
var def map[string]interface{}
@@ -285,7 +290,8 @@ func (j *Json) MustMap(args ...map[string]interface{}) map[string]interface{} {
// MustString guarantees the return of a `string` (with optional default)
//
// useful when you explicitly want a `string` in a single value return context:
// myFunc(js.Get("param1").MustString(), js.Get("optional_param").MustString("my_default"))
//
// myFunc(js.Get("param1").MustString(), js.Get("optional_param").MustString("my_default"))
func (j *Json) MustString(args ...string) string {
var def string
@@ -308,9 +314,10 @@ func (j *Json) MustString(args ...string) string {
// MustStringArray guarantees the return of a `[]string` (with optional default)
//
// useful when you want to interate over array values in a succinct manner:
// for i, s := range js.Get("results").MustStringArray() {
// fmt.Println(i, s)
// }
//
// for i, s := range js.Get("results").MustStringArray() {
// fmt.Println(i, s)
// }
func (j *Json) MustStringArray(args ...[]string) []string {
var def []string
@@ -333,7 +340,8 @@ func (j *Json) MustStringArray(args ...[]string) []string {
// MustInt guarantees the return of an `int` (with optional default)
//
// useful when you explicitly want an `int` in a single value return context:
// myFunc(js.Get("param1").MustInt(), js.Get("optional_param").MustInt(5150))
//
// myFunc(js.Get("param1").MustInt(), js.Get("optional_param").MustInt(5150))
func (j *Json) MustInt(args ...int) int {
var def int
@@ -356,7 +364,8 @@ func (j *Json) MustInt(args ...int) int {
// MustFloat64 guarantees the return of a `float64` (with optional default)
//
// useful when you explicitly want a `float64` in a single value return context:
// myFunc(js.Get("param1").MustFloat64(), js.Get("optional_param").MustFloat64(5.150))
//
// myFunc(js.Get("param1").MustFloat64(), js.Get("optional_param").MustFloat64(5.150))
func (j *Json) MustFloat64(args ...float64) float64 {
var def float64
@@ -379,7 +388,8 @@ func (j *Json) MustFloat64(args ...float64) float64 {
// MustBool guarantees the return of a `bool` (with optional default)
//
// useful when you explicitly want a `bool` in a single value return context:
// myFunc(js.Get("param1").MustBool(), js.Get("optional_param").MustBool(true))
//
// myFunc(js.Get("param1").MustBool(), js.Get("optional_param").MustBool(true))
func (j *Json) MustBool(args ...bool) bool {
var def bool
@@ -402,7 +412,8 @@ func (j *Json) MustBool(args ...bool) bool {
// MustInt64 guarantees the return of an `int64` (with optional default)
//
// useful when you explicitly want an `int64` in a single value return context:
// myFunc(js.Get("param1").MustInt64(), js.Get("optional_param").MustInt64(5150))
//
// myFunc(js.Get("param1").MustInt64(), js.Get("optional_param").MustInt64(5150))
func (j *Json) MustInt64(args ...int64) int64 {
var def int64
@@ -425,7 +436,8 @@ func (j *Json) MustInt64(args ...int64) int64 {
// MustUInt64 guarantees the return of an `uint64` (with optional default)
//
// useful when you explicitly want an `uint64` in a single value return context:
// myFunc(js.Get("param1").MustUint64(), js.Get("optional_param").MustUint64(5150))
//
// myFunc(js.Get("param1").MustUint64(), js.Get("optional_param").MustUint64(5150))
func (j *Json) MustUint64(args ...uint64) uint64 {
var def uint64
-75
View File
@@ -1,75 +0,0 @@
// +build !go1.1
package simplejson
import (
"encoding/json"
"errors"
"io"
"reflect"
)
// NewFromReader returns a *Json by decoding from an io.Reader
func NewFromReader(r io.Reader) (*Json, error) {
j := new(Json)
dec := json.NewDecoder(r)
err := dec.Decode(&j.data)
return j, err
}
// Implements the json.Unmarshaler interface.
func (j *Json) UnmarshalJSON(p []byte) error {
return json.Unmarshal(p, &j.data)
}
// Float64 coerces into a float64
func (j *Json) Float64() (float64, error) {
switch j.data.(type) {
case float32, float64:
return reflect.ValueOf(j.data).Float(), nil
case int, int8, int16, int32, int64:
return float64(reflect.ValueOf(j.data).Int()), nil
case uint, uint8, uint16, uint32, uint64:
return float64(reflect.ValueOf(j.data).Uint()), nil
}
return 0, errors.New("invalid value type")
}
// Int coerces into an int
func (j *Json) Int() (int, error) {
switch j.data.(type) {
case float32, float64:
return int(reflect.ValueOf(j.data).Float()), nil
case int, int8, int16, int32, int64:
return int(reflect.ValueOf(j.data).Int()), nil
case uint, uint8, uint16, uint32, uint64:
return int(reflect.ValueOf(j.data).Uint()), nil
}
return 0, errors.New("invalid value type")
}
// Int64 coerces into an int64
func (j *Json) Int64() (int64, error) {
switch j.data.(type) {
case float32, float64:
return int64(reflect.ValueOf(j.data).Float()), nil
case int, int8, int16, int32, int64:
return reflect.ValueOf(j.data).Int(), nil
case uint, uint8, uint16, uint32, uint64:
return int64(reflect.ValueOf(j.data).Uint()), nil
}
return 0, errors.New("invalid value type")
}
// Uint64 coerces into an uint64
func (j *Json) Uint64() (uint64, error) {
switch j.data.(type) {
case float32, float64:
return uint64(reflect.ValueOf(j.data).Float()), nil
case int, int8, int16, int32, int64:
return uint64(reflect.ValueOf(j.data).Int()), nil
case uint, uint8, uint16, uint32, uint64:
return reflect.ValueOf(j.data).Uint(), nil
}
return 0, errors.New("invalid value type")
}
@@ -1,5 +1,3 @@
// +build go1.1
package simplejson
import (
+6
View File
@@ -61,6 +61,9 @@ func (Curve) Double(P *Point) *Point { R := *P; R.Double(); return &R }
func (Curve) Add(P, Q *Point) *Point { R := *P; R.Add(Q); return &R }
// ScalarMult returns kP. This function runs in constant time.
//
// The result equals [k]P only when P is in the prime-order subgroup; any
// torsion component of P is dropped.
func (e Curve) ScalarMult(k *Scalar, P *Point) *Point {
k4 := &Scalar{}
k4.divBy4(k)
@@ -75,6 +78,9 @@ func (e Curve) ScalarBaseMult(k *Scalar) *Point {
}
// CombinedMult returns mG+nP, where G is the generator point. This function is non-constant time.
//
// Like ScalarMult, the result equals [m]G+[n]P only when P is in the
// prime-order subgroup; any torsion component of P is dropped.
func (e Curve) CombinedMult(m, n *Scalar, P *Point) *Point {
m4 := &Scalar{}
n4 := &Scalar{}
+11 -1
View File
@@ -46,6 +46,9 @@ func FromBytes(in []byte) (*Point, error) {
}
err := errors.New("invalid decoding")
P := &Point{}
if in[fp.Size]&0x7F != 0 {
return nil, err
}
signX := in[fp.Size] >> 7
copy(P.y[:], in[:fp.Size])
p := fp.P()
@@ -135,7 +138,14 @@ func (P *Point) MarshalBinary() (data []byte, err error) {
}
// UnmarshalBinary must be able to decode the form generated by MarshalBinary.
func (P *Point) UnmarshalBinary(data []byte) error { Q, err := FromBytes(data); *P = *Q; return err }
func (P *Point) UnmarshalBinary(data []byte) error {
Q, err := FromBytes(data)
if err != nil {
return err
}
*P = *Q
return nil
}
// Double sets P = 2Q.
func (P *Point) Double() { P.Add(P) }
+121
View File
@@ -0,0 +1,121 @@
// Package kem provides a unified interface for KEM schemes.
//
// A register of schemes is available in the package
//
// github.com/cloudflare/circl/kem/schemes
package kem
import (
"encoding"
"errors"
)
// A KEM public key
type PublicKey interface {
// Returns the scheme for this public key
Scheme() Scheme
encoding.BinaryMarshaler
Equal(PublicKey) bool
}
// A KEM private key
type PrivateKey interface {
// Returns the scheme for this private key
Scheme() Scheme
encoding.BinaryMarshaler
Equal(PrivateKey) bool
Public() PublicKey
}
// A Scheme represents a specific instance of a KEM.
type Scheme interface {
// Name of the scheme
Name() string
// GenerateKeyPair creates a new key pair.
GenerateKeyPair() (PublicKey, PrivateKey, error)
// Encapsulate generates a shared key ss for the public key and
// encapsulates it into a ciphertext ct.
Encapsulate(pk PublicKey) (ct, ss []byte, err error)
// Returns the shared key encapsulated in ciphertext ct for the
// private key sk.
Decapsulate(sk PrivateKey, ct []byte) ([]byte, error)
// Unmarshals a PublicKey from the provided buffer.
UnmarshalBinaryPublicKey([]byte) (PublicKey, error)
// Unmarshals a PrivateKey from the provided buffer.
UnmarshalBinaryPrivateKey([]byte) (PrivateKey, error)
// Size of encapsulated keys.
CiphertextSize() int
// Size of established shared keys.
SharedKeySize() int
// Size of packed private keys.
PrivateKeySize() int
// Size of packed public keys.
PublicKeySize() int
// DeriveKeyPair deterministically derives a pair of keys from a seed.
// Panics if the length of seed is not equal to the value returned by
// SeedSize.
DeriveKeyPair(seed []byte) (PublicKey, PrivateKey)
// Size of seed used in DeriveKey
SeedSize() int
// EncapsulateDeterministically generates a shared key ss for the public
// key deterministically from the given seed and encapsulates it into
// a ciphertext ct. If unsure, you're better off using Encapsulate().
EncapsulateDeterministically(pk PublicKey, seed []byte) (
ct, ss []byte, err error)
// Size of seed used in EncapsulateDeterministically().
EncapsulationSeedSize() int
}
// AuthScheme represents a KEM that supports authenticated key encapsulation.
type AuthScheme interface {
Scheme
AuthEncapsulate(pkr PublicKey, sks PrivateKey) (ct, ss []byte, err error)
AuthEncapsulateDeterministically(pkr PublicKey, sks PrivateKey, seed []byte) (ct, ss []byte, err error)
AuthDecapsulate(skr PrivateKey, ct []byte, pks PublicKey) ([]byte, error)
}
var (
// ErrTypeMismatch is the error used if types of, for instance, private
// and public keys don't match
ErrTypeMismatch = errors.New("types mismatch")
// ErrSeedSize is the error used if the provided seed is of the wrong
// size.
ErrSeedSize = errors.New("wrong seed size")
// ErrPubKeySize is the error used if the provided public key is of
// the wrong size.
ErrPubKeySize = errors.New("wrong size for public key")
// ErrCiphertextSize is the error used if the provided ciphertext
// is of the wrong size.
ErrCiphertextSize = errors.New("wrong size for ciphertext")
// ErrPrivKeySize is the error used if the provided private key is of
// the wrong size.
ErrPrivKeySize = errors.New("wrong size for private key")
// ErrPubKey is the error used if the provided public key is invalid.
ErrPubKey = errors.New("invalid public key")
// ErrPrivKey is the error used if the provided private key is invalid.
ErrPrivKey = errors.New("invalid private key")
// ErrCipherText is the error used if the provided ciphertext is invalid.
ErrCipherText = errors.New("invalid ciphertext")
)
+407
View File
@@ -0,0 +1,407 @@
// Code generated from pkg.templ.go. DO NOT EDIT.
// Package mlkem1024 implements the IND-CCA2 secure key encapsulation mechanism
// ML-KEM-1024 as defined in FIPS203.
package mlkem1024
import (
"bytes"
"crypto/subtle"
"io"
cryptoRand "crypto/rand"
"github.com/cloudflare/circl/internal/sha3"
"github.com/cloudflare/circl/kem"
cpapke "github.com/cloudflare/circl/pke/kyber/kyber1024"
)
const (
// Size of seed for NewKeyFromSeed
KeySeedSize = cpapke.KeySeedSize + 32
// Size of seed for EncapsulateTo.
EncapsulationSeedSize = 32
// Size of the established shared key.
SharedKeySize = 32
// Size of the encapsulated shared key.
CiphertextSize = cpapke.CiphertextSize
// Size of a packed public key.
PublicKeySize = cpapke.PublicKeySize
// Size of a packed private key.
PrivateKeySize = cpapke.PrivateKeySize + cpapke.PublicKeySize + 64
)
// Type of a ML-KEM-1024 public key
type PublicKey struct {
pk *cpapke.PublicKey
hpk [32]byte // H(pk)
}
// Type of a ML-KEM-1024 private key
type PrivateKey struct {
sk *cpapke.PrivateKey
pk *cpapke.PublicKey
hpk [32]byte // H(pk)
z [32]byte
}
// NewKeyFromSeed derives a public/private keypair deterministically
// from the given seed.
//
// Panics if seed is not of length KeySeedSize.
func NewKeyFromSeed(seed []byte) (*PublicKey, *PrivateKey) {
var sk PrivateKey
var pk PublicKey
if len(seed) != KeySeedSize {
panic("seed must be of length KeySeedSize")
}
pk.pk, sk.sk = cpapke.NewKeyFromSeedMLKEM(seed[:cpapke.KeySeedSize])
sk.pk = pk.pk
copy(sk.z[:], seed[cpapke.KeySeedSize:])
// Compute H(pk)
var ppk [cpapke.PublicKeySize]byte
sk.pk.Pack(ppk[:])
h := sha3.New256()
h.Write(ppk[:])
h.Read(sk.hpk[:])
copy(pk.hpk[:], sk.hpk[:])
return &pk, &sk
}
// GenerateKeyPair generates public and private keys using entropy from rand.
// If rand is nil, crypto/rand.Reader will be used.
func GenerateKeyPair(rand io.Reader) (*PublicKey, *PrivateKey, error) {
var seed [KeySeedSize]byte
if rand == nil {
rand = cryptoRand.Reader
}
_, err := io.ReadFull(rand, seed[:])
if err != nil {
return nil, nil, err
}
pk, sk := NewKeyFromSeed(seed[:])
return pk, sk, nil
}
// EncapsulateTo generates a shared key and ciphertext that contains it
// for the public key using randomness from seed and writes the shared key
// to ss and ciphertext to ct.
//
// Panics if ss, ct or seed are not of length SharedKeySize, CiphertextSize
// and EncapsulationSeedSize respectively.
//
// seed may be nil, in which case crypto/rand.Reader is used to generate one.
func (pk *PublicKey) EncapsulateTo(ct, ss []byte, seed []byte) {
if seed == nil {
seed = make([]byte, EncapsulationSeedSize)
if _, err := cryptoRand.Read(seed[:]); err != nil {
panic(err)
}
} else {
if len(seed) != EncapsulationSeedSize {
panic("seed must be of length EncapsulationSeedSize")
}
}
if len(ct) != CiphertextSize {
panic("ct must be of length CiphertextSize")
}
if len(ss) != SharedKeySize {
panic("ss must be of length SharedKeySize")
}
var m [32]byte
copy(m[:], seed)
// (K', r) = G(m ‖ H(pk))
var kr [64]byte
g := sha3.New512()
g.Write(m[:])
g.Write(pk.hpk[:])
g.Read(kr[:])
// c = Kyber.CPAPKE.Enc(pk, m, r)
pk.pk.EncryptTo(ct, m[:], kr[32:])
copy(ss, kr[:SharedKeySize])
}
// DecapsulateTo computes the shared key which is encapsulated in ct
// for the private key.
//
// Panics if ct or ss are not of length CiphertextSize and SharedKeySize
// respectively.
func (sk *PrivateKey) DecapsulateTo(ss, ct []byte) {
if len(ct) != CiphertextSize {
panic("ct must be of length CiphertextSize")
}
if len(ss) != SharedKeySize {
panic("ss must be of length SharedKeySize")
}
// m' = Kyber.CPAPKE.Dec(sk, ct)
var m2 [32]byte
sk.sk.DecryptTo(m2[:], ct)
// (K'', r') = G(m' ‖ H(pk))
var kr2 [64]byte
g := sha3.New512()
g.Write(m2[:])
g.Write(sk.hpk[:])
g.Read(kr2[:])
// c' = Kyber.CPAPKE.Enc(pk, m', r')
var ct2 [CiphertextSize]byte
sk.pk.EncryptTo(ct2[:], m2[:], kr2[32:])
var ss2 [SharedKeySize]byte
// Compute shared secret in case of rejection: ss₂ = PRF(z ‖ c)
prf := sha3.NewShake256()
prf.Write(sk.z[:])
prf.Write(ct[:CiphertextSize])
prf.Read(ss2[:])
// Set ss2 to the real shared secret if c = c'.
subtle.ConstantTimeCopy(
subtle.ConstantTimeCompare(ct, ct2[:]),
ss2[:],
kr2[:SharedKeySize],
)
copy(ss, ss2[:])
}
// Packs sk to buf.
//
// Panics if buf is not of size PrivateKeySize.
func (sk *PrivateKey) Pack(buf []byte) {
if len(buf) != PrivateKeySize {
panic("buf must be of length PrivateKeySize")
}
sk.sk.Pack(buf[:cpapke.PrivateKeySize])
buf = buf[cpapke.PrivateKeySize:]
sk.pk.Pack(buf[:cpapke.PublicKeySize])
buf = buf[cpapke.PublicKeySize:]
copy(buf, sk.hpk[:])
buf = buf[32:]
copy(buf, sk.z[:])
}
// Unpacks sk from buf.
//
// Panics if buf is not of size PrivateKeySize.
//
// Returns an error if buf is not of size PrivateKeySize, or private key
// doesn't pass the ML-KEM decapsulation key check.
func (sk *PrivateKey) Unpack(buf []byte) error {
if len(buf) != PrivateKeySize {
return kem.ErrPrivKeySize
}
sk.sk = new(cpapke.PrivateKey)
sk.sk.Unpack(buf[:cpapke.PrivateKeySize])
buf = buf[cpapke.PrivateKeySize:]
sk.pk = new(cpapke.PublicKey)
sk.pk.Unpack(buf[:cpapke.PublicKeySize])
var hpk [32]byte
h := sha3.New256()
h.Write(buf[:cpapke.PublicKeySize])
h.Read(hpk[:])
buf = buf[cpapke.PublicKeySize:]
copy(sk.hpk[:], buf[:32])
copy(sk.z[:], buf[32:])
if !bytes.Equal(hpk[:], sk.hpk[:]) {
return kem.ErrPrivKey
}
return nil
}
// Packs pk to buf.
//
// Panics if buf is not of size PublicKeySize.
func (pk *PublicKey) Pack(buf []byte) {
if len(buf) != PublicKeySize {
panic("buf must be of length PublicKeySize")
}
pk.pk.Pack(buf)
}
// Unpacks pk from buf.
//
// Returns an error if buf is not of size PublicKeySize, or the public key
// is not normalized.
func (pk *PublicKey) Unpack(buf []byte) error {
if len(buf) != PublicKeySize {
return kem.ErrPubKeySize
}
pk.pk = new(cpapke.PublicKey)
if err := pk.pk.UnpackMLKEM(buf); err != nil {
return err
}
// Compute cached H(pk)
h := sha3.New256()
h.Write(buf)
h.Read(pk.hpk[:])
return nil
}
// Boilerplate down below for the KEM scheme API.
type scheme struct{}
var sch kem.Scheme = &scheme{}
// Scheme returns a KEM interface.
func Scheme() kem.Scheme { return sch }
func (*scheme) Name() string { return "ML-KEM-1024" }
func (*scheme) PublicKeySize() int { return PublicKeySize }
func (*scheme) PrivateKeySize() int { return PrivateKeySize }
func (*scheme) SeedSize() int { return KeySeedSize }
func (*scheme) SharedKeySize() int { return SharedKeySize }
func (*scheme) CiphertextSize() int { return CiphertextSize }
func (*scheme) EncapsulationSeedSize() int { return EncapsulationSeedSize }
func (sk *PrivateKey) Scheme() kem.Scheme { return sch }
func (pk *PublicKey) Scheme() kem.Scheme { return sch }
func (sk *PrivateKey) MarshalBinary() ([]byte, error) {
var ret [PrivateKeySize]byte
sk.Pack(ret[:])
return ret[:], nil
}
func (sk *PrivateKey) Equal(other kem.PrivateKey) bool {
oth, ok := other.(*PrivateKey)
if !ok {
return false
}
if sk.pk == nil && oth.pk == nil {
return true
}
if sk.pk == nil || oth.pk == nil {
return false
}
if !bytes.Equal(sk.hpk[:], oth.hpk[:]) ||
subtle.ConstantTimeCompare(sk.z[:], oth.z[:]) != 1 {
return false
}
return sk.sk.Equal(oth.sk)
}
func (pk *PublicKey) Equal(other kem.PublicKey) bool {
oth, ok := other.(*PublicKey)
if !ok {
return false
}
if pk.pk == nil && oth.pk == nil {
return true
}
if pk.pk == nil || oth.pk == nil {
return false
}
return bytes.Equal(pk.hpk[:], oth.hpk[:])
}
func (sk *PrivateKey) Public() kem.PublicKey {
pk := new(PublicKey)
pk.pk = sk.pk
copy(pk.hpk[:], sk.hpk[:])
return pk
}
func (pk *PublicKey) MarshalBinary() ([]byte, error) {
var ret [PublicKeySize]byte
pk.Pack(ret[:])
return ret[:], nil
}
func (*scheme) GenerateKeyPair() (kem.PublicKey, kem.PrivateKey, error) {
return GenerateKeyPair(cryptoRand.Reader)
}
func (*scheme) DeriveKeyPair(seed []byte) (kem.PublicKey, kem.PrivateKey) {
if len(seed) != KeySeedSize {
panic(kem.ErrSeedSize)
}
return NewKeyFromSeed(seed[:])
}
func (*scheme) Encapsulate(pk kem.PublicKey) (ct, ss []byte, err error) {
ct = make([]byte, CiphertextSize)
ss = make([]byte, SharedKeySize)
pub, ok := pk.(*PublicKey)
if !ok {
return nil, nil, kem.ErrTypeMismatch
}
pub.EncapsulateTo(ct, ss, nil)
return
}
func (*scheme) EncapsulateDeterministically(pk kem.PublicKey, seed []byte) (
ct, ss []byte, err error) {
if len(seed) != EncapsulationSeedSize {
return nil, nil, kem.ErrSeedSize
}
ct = make([]byte, CiphertextSize)
ss = make([]byte, SharedKeySize)
pub, ok := pk.(*PublicKey)
if !ok {
return nil, nil, kem.ErrTypeMismatch
}
pub.EncapsulateTo(ct, ss, seed)
return
}
func (*scheme) Decapsulate(sk kem.PrivateKey, ct []byte) ([]byte, error) {
if len(ct) != CiphertextSize {
return nil, kem.ErrCiphertextSize
}
priv, ok := sk.(*PrivateKey)
if !ok {
return nil, kem.ErrTypeMismatch
}
ss := make([]byte, SharedKeySize)
priv.DecapsulateTo(ss, ct)
return ss, nil
}
func (*scheme) UnmarshalBinaryPublicKey(buf []byte) (kem.PublicKey, error) {
var ret PublicKey
if err := ret.Unpack(buf); err != nil {
return nil, err
}
return &ret, nil
}
func (*scheme) UnmarshalBinaryPrivateKey(buf []byte) (kem.PrivateKey, error) {
if len(buf) != PrivateKeySize {
return nil, kem.ErrPrivKeySize
}
var ret PrivateKey
if err := ret.Unpack(buf); err != nil {
return nil, err
}
return &ret, nil
}
+407
View File
@@ -0,0 +1,407 @@
// Code generated from pkg.templ.go. DO NOT EDIT.
// Package mlkem768 implements the IND-CCA2 secure key encapsulation mechanism
// ML-KEM-768 as defined in FIPS203.
package mlkem768
import (
"bytes"
"crypto/subtle"
"io"
cryptoRand "crypto/rand"
"github.com/cloudflare/circl/internal/sha3"
"github.com/cloudflare/circl/kem"
cpapke "github.com/cloudflare/circl/pke/kyber/kyber768"
)
const (
// Size of seed for NewKeyFromSeed
KeySeedSize = cpapke.KeySeedSize + 32
// Size of seed for EncapsulateTo.
EncapsulationSeedSize = 32
// Size of the established shared key.
SharedKeySize = 32
// Size of the encapsulated shared key.
CiphertextSize = cpapke.CiphertextSize
// Size of a packed public key.
PublicKeySize = cpapke.PublicKeySize
// Size of a packed private key.
PrivateKeySize = cpapke.PrivateKeySize + cpapke.PublicKeySize + 64
)
// Type of a ML-KEM-768 public key
type PublicKey struct {
pk *cpapke.PublicKey
hpk [32]byte // H(pk)
}
// Type of a ML-KEM-768 private key
type PrivateKey struct {
sk *cpapke.PrivateKey
pk *cpapke.PublicKey
hpk [32]byte // H(pk)
z [32]byte
}
// NewKeyFromSeed derives a public/private keypair deterministically
// from the given seed.
//
// Panics if seed is not of length KeySeedSize.
func NewKeyFromSeed(seed []byte) (*PublicKey, *PrivateKey) {
var sk PrivateKey
var pk PublicKey
if len(seed) != KeySeedSize {
panic("seed must be of length KeySeedSize")
}
pk.pk, sk.sk = cpapke.NewKeyFromSeedMLKEM(seed[:cpapke.KeySeedSize])
sk.pk = pk.pk
copy(sk.z[:], seed[cpapke.KeySeedSize:])
// Compute H(pk)
var ppk [cpapke.PublicKeySize]byte
sk.pk.Pack(ppk[:])
h := sha3.New256()
h.Write(ppk[:])
h.Read(sk.hpk[:])
copy(pk.hpk[:], sk.hpk[:])
return &pk, &sk
}
// GenerateKeyPair generates public and private keys using entropy from rand.
// If rand is nil, crypto/rand.Reader will be used.
func GenerateKeyPair(rand io.Reader) (*PublicKey, *PrivateKey, error) {
var seed [KeySeedSize]byte
if rand == nil {
rand = cryptoRand.Reader
}
_, err := io.ReadFull(rand, seed[:])
if err != nil {
return nil, nil, err
}
pk, sk := NewKeyFromSeed(seed[:])
return pk, sk, nil
}
// EncapsulateTo generates a shared key and ciphertext that contains it
// for the public key using randomness from seed and writes the shared key
// to ss and ciphertext to ct.
//
// Panics if ss, ct or seed are not of length SharedKeySize, CiphertextSize
// and EncapsulationSeedSize respectively.
//
// seed may be nil, in which case crypto/rand.Reader is used to generate one.
func (pk *PublicKey) EncapsulateTo(ct, ss []byte, seed []byte) {
if seed == nil {
seed = make([]byte, EncapsulationSeedSize)
if _, err := cryptoRand.Read(seed[:]); err != nil {
panic(err)
}
} else {
if len(seed) != EncapsulationSeedSize {
panic("seed must be of length EncapsulationSeedSize")
}
}
if len(ct) != CiphertextSize {
panic("ct must be of length CiphertextSize")
}
if len(ss) != SharedKeySize {
panic("ss must be of length SharedKeySize")
}
var m [32]byte
copy(m[:], seed)
// (K', r) = G(m ‖ H(pk))
var kr [64]byte
g := sha3.New512()
g.Write(m[:])
g.Write(pk.hpk[:])
g.Read(kr[:])
// c = Kyber.CPAPKE.Enc(pk, m, r)
pk.pk.EncryptTo(ct, m[:], kr[32:])
copy(ss, kr[:SharedKeySize])
}
// DecapsulateTo computes the shared key which is encapsulated in ct
// for the private key.
//
// Panics if ct or ss are not of length CiphertextSize and SharedKeySize
// respectively.
func (sk *PrivateKey) DecapsulateTo(ss, ct []byte) {
if len(ct) != CiphertextSize {
panic("ct must be of length CiphertextSize")
}
if len(ss) != SharedKeySize {
panic("ss must be of length SharedKeySize")
}
// m' = Kyber.CPAPKE.Dec(sk, ct)
var m2 [32]byte
sk.sk.DecryptTo(m2[:], ct)
// (K'', r') = G(m' ‖ H(pk))
var kr2 [64]byte
g := sha3.New512()
g.Write(m2[:])
g.Write(sk.hpk[:])
g.Read(kr2[:])
// c' = Kyber.CPAPKE.Enc(pk, m', r')
var ct2 [CiphertextSize]byte
sk.pk.EncryptTo(ct2[:], m2[:], kr2[32:])
var ss2 [SharedKeySize]byte
// Compute shared secret in case of rejection: ss₂ = PRF(z ‖ c)
prf := sha3.NewShake256()
prf.Write(sk.z[:])
prf.Write(ct[:CiphertextSize])
prf.Read(ss2[:])
// Set ss2 to the real shared secret if c = c'.
subtle.ConstantTimeCopy(
subtle.ConstantTimeCompare(ct, ct2[:]),
ss2[:],
kr2[:SharedKeySize],
)
copy(ss, ss2[:])
}
// Packs sk to buf.
//
// Panics if buf is not of size PrivateKeySize.
func (sk *PrivateKey) Pack(buf []byte) {
if len(buf) != PrivateKeySize {
panic("buf must be of length PrivateKeySize")
}
sk.sk.Pack(buf[:cpapke.PrivateKeySize])
buf = buf[cpapke.PrivateKeySize:]
sk.pk.Pack(buf[:cpapke.PublicKeySize])
buf = buf[cpapke.PublicKeySize:]
copy(buf, sk.hpk[:])
buf = buf[32:]
copy(buf, sk.z[:])
}
// Unpacks sk from buf.
//
// Panics if buf is not of size PrivateKeySize.
//
// Returns an error if buf is not of size PrivateKeySize, or private key
// doesn't pass the ML-KEM decapsulation key check.
func (sk *PrivateKey) Unpack(buf []byte) error {
if len(buf) != PrivateKeySize {
return kem.ErrPrivKeySize
}
sk.sk = new(cpapke.PrivateKey)
sk.sk.Unpack(buf[:cpapke.PrivateKeySize])
buf = buf[cpapke.PrivateKeySize:]
sk.pk = new(cpapke.PublicKey)
sk.pk.Unpack(buf[:cpapke.PublicKeySize])
var hpk [32]byte
h := sha3.New256()
h.Write(buf[:cpapke.PublicKeySize])
h.Read(hpk[:])
buf = buf[cpapke.PublicKeySize:]
copy(sk.hpk[:], buf[:32])
copy(sk.z[:], buf[32:])
if !bytes.Equal(hpk[:], sk.hpk[:]) {
return kem.ErrPrivKey
}
return nil
}
// Packs pk to buf.
//
// Panics if buf is not of size PublicKeySize.
func (pk *PublicKey) Pack(buf []byte) {
if len(buf) != PublicKeySize {
panic("buf must be of length PublicKeySize")
}
pk.pk.Pack(buf)
}
// Unpacks pk from buf.
//
// Returns an error if buf is not of size PublicKeySize, or the public key
// is not normalized.
func (pk *PublicKey) Unpack(buf []byte) error {
if len(buf) != PublicKeySize {
return kem.ErrPubKeySize
}
pk.pk = new(cpapke.PublicKey)
if err := pk.pk.UnpackMLKEM(buf); err != nil {
return err
}
// Compute cached H(pk)
h := sha3.New256()
h.Write(buf)
h.Read(pk.hpk[:])
return nil
}
// Boilerplate down below for the KEM scheme API.
type scheme struct{}
var sch kem.Scheme = &scheme{}
// Scheme returns a KEM interface.
func Scheme() kem.Scheme { return sch }
func (*scheme) Name() string { return "ML-KEM-768" }
func (*scheme) PublicKeySize() int { return PublicKeySize }
func (*scheme) PrivateKeySize() int { return PrivateKeySize }
func (*scheme) SeedSize() int { return KeySeedSize }
func (*scheme) SharedKeySize() int { return SharedKeySize }
func (*scheme) CiphertextSize() int { return CiphertextSize }
func (*scheme) EncapsulationSeedSize() int { return EncapsulationSeedSize }
func (sk *PrivateKey) Scheme() kem.Scheme { return sch }
func (pk *PublicKey) Scheme() kem.Scheme { return sch }
func (sk *PrivateKey) MarshalBinary() ([]byte, error) {
var ret [PrivateKeySize]byte
sk.Pack(ret[:])
return ret[:], nil
}
func (sk *PrivateKey) Equal(other kem.PrivateKey) bool {
oth, ok := other.(*PrivateKey)
if !ok {
return false
}
if sk.pk == nil && oth.pk == nil {
return true
}
if sk.pk == nil || oth.pk == nil {
return false
}
if !bytes.Equal(sk.hpk[:], oth.hpk[:]) ||
subtle.ConstantTimeCompare(sk.z[:], oth.z[:]) != 1 {
return false
}
return sk.sk.Equal(oth.sk)
}
func (pk *PublicKey) Equal(other kem.PublicKey) bool {
oth, ok := other.(*PublicKey)
if !ok {
return false
}
if pk.pk == nil && oth.pk == nil {
return true
}
if pk.pk == nil || oth.pk == nil {
return false
}
return bytes.Equal(pk.hpk[:], oth.hpk[:])
}
func (sk *PrivateKey) Public() kem.PublicKey {
pk := new(PublicKey)
pk.pk = sk.pk
copy(pk.hpk[:], sk.hpk[:])
return pk
}
func (pk *PublicKey) MarshalBinary() ([]byte, error) {
var ret [PublicKeySize]byte
pk.Pack(ret[:])
return ret[:], nil
}
func (*scheme) GenerateKeyPair() (kem.PublicKey, kem.PrivateKey, error) {
return GenerateKeyPair(cryptoRand.Reader)
}
func (*scheme) DeriveKeyPair(seed []byte) (kem.PublicKey, kem.PrivateKey) {
if len(seed) != KeySeedSize {
panic(kem.ErrSeedSize)
}
return NewKeyFromSeed(seed[:])
}
func (*scheme) Encapsulate(pk kem.PublicKey) (ct, ss []byte, err error) {
ct = make([]byte, CiphertextSize)
ss = make([]byte, SharedKeySize)
pub, ok := pk.(*PublicKey)
if !ok {
return nil, nil, kem.ErrTypeMismatch
}
pub.EncapsulateTo(ct, ss, nil)
return
}
func (*scheme) EncapsulateDeterministically(pk kem.PublicKey, seed []byte) (
ct, ss []byte, err error) {
if len(seed) != EncapsulationSeedSize {
return nil, nil, kem.ErrSeedSize
}
ct = make([]byte, CiphertextSize)
ss = make([]byte, SharedKeySize)
pub, ok := pk.(*PublicKey)
if !ok {
return nil, nil, kem.ErrTypeMismatch
}
pub.EncapsulateTo(ct, ss, seed)
return
}
func (*scheme) Decapsulate(sk kem.PrivateKey, ct []byte) ([]byte, error) {
if len(ct) != CiphertextSize {
return nil, kem.ErrCiphertextSize
}
priv, ok := sk.(*PrivateKey)
if !ok {
return nil, kem.ErrTypeMismatch
}
ss := make([]byte, SharedKeySize)
priv.DecapsulateTo(ss, ct)
return ss, nil
}
func (*scheme) UnmarshalBinaryPublicKey(buf []byte) (kem.PublicKey, error) {
var ret PublicKey
if err := ret.Unpack(buf); err != nil {
return nil, err
}
return &ret, nil
}
func (*scheme) UnmarshalBinaryPrivateKey(buf []byte) (kem.PrivateKey, error) {
if len(buf) != PrivateKeySize {
return nil, kem.ErrPrivKeySize
}
var ret PrivateKey
if err := ret.Unpack(buf); err != nil {
return nil, err
}
return &ret, nil
}
+35 -12
View File
@@ -9,9 +9,6 @@ package mlsbset
import (
"errors"
"fmt"
"math/big"
"github.com/cloudflare/circl/internal/conv"
)
// EltG is a group element.
@@ -69,17 +66,43 @@ func (m Encoder) Encode(k []byte) (*Power, error) {
k = append(k, make([]byte, ap)...)
s := m.signs(k)
b := make([]int32, m.p.L-m.p.D)
c := conv.BytesLe2BigInt(k)
c.Rsh(c, m.p.D)
var bi big.Int
// Original algorithm starts with c := k >> D, then computes
//
// b_(i-D) = s_(i%D) * lsb(c)
// c = [ (c>>1)+1 if b_(i-D) = -1
// [ c>>1 otherwise
//
// To prevent keeping a large k around, we note that at any step i we have
//
// c = (k >> i) + t for t in {0,1}
//
// Base case is obvious. For induction, write kbit for the i-th bit of k.
// Note lsb(c) = kbit ^ t. From that we can compute b_(i-D). Now we need
// to compute the next t.
//
// Consider c >> 1 = (k >> i) + t) >> 1. This equals k >> (i+1)
// unless t = 1 = kbit.
//
// If b_(i-D) is negative, then we must have had 1=lsb(c)=kbit^t, and so
// c >> 1 = k >> (i+1), as desired with new t equal to 1.
// For the other case assume b_(i-D) isn't negative. Now it is possible
// that t = 1 = kbit, and only in that case the new t is equal to 1.
var t uint64
for i := m.p.D; i < m.p.L; i++ {
c0 := int32(c.Bit(0))
b[i-m.p.D] = s[i%m.p.D] * c0
bi.SetInt64(int64(b[i-m.p.D] >> 1))
c.Rsh(c, 1)
c.Sub(c, &bi)
si := s[i%m.p.D]
kbit := uint64(k[i>>3]>>(i&7)) & 1
lsbc := kbit ^ t
neg := uint64(si>>31) & 1 // 1 iff si == -1
b[i-m.p.D] = si * int32(lsbc)
t = (kbit & t) | (lsbc & neg)
}
// carry = (k >> L) + t: any bits of k at positions >= L (present when L is
// not a multiple of 8) plus the final carry, matching the original.
carry := int(t)
for pos := m.p.L; int(pos>>3) < len(k); pos++ {
carry += int((k[pos>>3]>>(pos&7))&1) << (pos - m.p.L)
}
carry := int(c.Int64())
return &Power{m, s, b, carry}, nil
}
+302
View File
@@ -0,0 +1,302 @@
//go:build amd64 && !purego
// +build amd64,!purego
package common
import (
"golang.org/x/sys/cpu"
)
// ZetasAVX2 contains all ζ used in NTT (like the Zetas array), but also
// the values int16(zeta * 62209) for each zeta, which is used in
// Montgomery reduction. There is some duplication and reordering as
// compared to Zetas to make it more convenient for use with AVX2.
var ZetasAVX2 = [...]int16{
// level 1: int16(Zetas[1]*62209) and Zetas[1]
31499, 2571,
// level 2
//
// int16(Zetas[2]*62209), Zetas[2], int16(Zetas[3]*62209), Zetas[3]
14746, 2970, 788, 1812,
// level 3, like level 2.
13525, 1493, -12402, 1422, 28191, 287, -16694, 202,
0, 0, // padding
// layer 4. offset: 1*16
//
// The precomputed multiplication and zetas are grouped by 16 at a
// time as used in the set of butterflies, etc.
-20906, -20906, -20906, -20906, -20906, -20906, -20906, -20906,
27758, 27758, 27758, 27758, 27758, 27758, 27758, 27758,
3158, 3158, 3158, 3158, 3158, 3158, 3158, 3158,
622, 622, 622, 622, 622, 622, 622, 622,
-3799, -3799, -3799, -3799, -3799, -3799, -3799, -3799,
-15690, -15690, -15690, -15690, -15690, -15690, -15690, -15690,
1577, 1577, 1577, 1577, 1577, 1577, 1577, 1577,
182, 182, 182, 182, 182, 182, 182, 182,
10690, 10690, 10690, 10690, 10690, 10690, 10690, 10690,
1359, 1359, 1359, 1359, 1359, 1359, 1359, 1359,
962, 962, 962, 962, 962, 962, 962, 962,
2127, 2127, 2127, 2127, 2127, 2127, 2127, 2127,
-11201, -11201, -11201, -11201, -11201, -11201, -11201, -11201,
31164, 31164, 31164, 31164, 31164, 31164, 31164, 31164,
1855, 1855, 1855, 1855, 1855, 1855, 1855, 1855,
1468, 1468, 1468, 1468, 1468, 1468, 1468, 1468,
// layer 5. offset: 9*16
-5827, -5827, -5827, -5827, 17364, 17364, 17364, 17364,
-26360, -26360, -26360, -26360, -29057, -29057, -29057, -29057,
573, 573, 573, 573, 2004, 2004, 2004, 2004,
264, 264, 264, 264, 383, 383, 383, 383,
5572, 5572, 5572, 5572, -1102, -1102, -1102, -1102,
21439, 21439, 21439, 21439, -26241, -26241, -26241, -26241,
2500, 2500, 2500, 2500, 1458, 1458, 1458, 1458,
1727, 1727, 1727, 1727, 3199, 3199, 3199, 3199,
-28072, -28072, -28072, -28072, 24313, 24313, 24313, 24313,
-10532, -10532, -10532, -10532, 8800, 8800, 8800, 8800,
2648, 2648, 2648, 2648, 1017, 1017, 1017, 1017,
732, 732, 732, 732, 608, 608, 608, 608,
18427, 18427, 18427, 18427, 8859, 8859, 8859, 8859,
26676, 26676, 26676, 26676, -16162, -16162, -16162, -16162,
1787, 1787, 1787, 1787, 411, 411, 411, 411,
3124, 3124, 3124, 3124, 1758, 1758, 1758, 1758,
// layer 6. offset: 17*16
-5689, -5689, -6516, -6516, 1497, 1497, 30967, 30967,
-23564, -23564, 20179, 20179, 20711, 20711, 25081, 25081,
1223, 1223, 652, 652, 2777, 2777, 1015, 1015,
2036, 2036, 1491, 1491, 3047, 3047, 1785, 1785,
-12796, -12796, 26617, 26617, 16065, 16065, -12441, -12441,
9135, 9135, -649, -649, -25986, -25986, 27837, 27837,
516, 516, 3321, 3321, 3009, 3009, 2663, 2663,
1711, 1711, 2167, 2167, 126, 126, 1469, 1469,
19884, 19884, -28249, -28249, -15886, -15886, -8898, -8898,
-28309, -28309, 9076, 9076, -30198, -30198, 18250, 18250,
2476, 2476, 3239, 3239, 3058, 3058, 830, 830,
107, 107, 1908, 1908, 3082, 3082, 2378, 2378,
13427, 13427, 14017, 14017, -29155, -29155, -12756, -12756,
16832, 16832, 4312, 4312, -24155, -24155, -17914, -17914,
2931, 2931, 961, 961, 1821, 1821, 2604, 2604,
448, 448, 2264, 2264, 677, 677, 2054, 2054,
// layer 7. offset: 25*16
-334, 11182, -11477, 13387, -32226, -14233, 20494, -21655,
-27738, 13131, 945, -4586, -14882, 23093, 6182, 5493,
2226, 430, 555, 843, 2078, 871, 1550, 105,
422, 587, 177, 3094, 3038, 2869, 1574, 1653,
32011, -32502, 10631, 30318, 29176, -18741, -28761, 12639,
-18485, 20100, 17561, 18525, -14430, 19529, -5275, -12618,
3083, 778, 1159, 3182, 2552, 1483, 2727, 1119,
1739, 644, 2457, 349, 418, 329, 3173, 3254,
-31183, 20297, 25435, 2146, -7382, 15356, 24392, -32384,
-20926, -6279, 10946, -14902, 24215, -11044, 16990, 14470,
817, 1097, 603, 610, 1322, 2044, 1864, 384,
2114, 3193, 1218, 1994, 2455, 220, 2142, 1670,
10336, -21497, -7933, -20198, -22501, 23211, 10907, -17442,
31637, -23859, 28644, -20257, 23998, 7757, -17422, 23132,
2144, 1799, 2051, 794, 1819, 2475, 2459, 478,
3221, 3021, 996, 991, 958, 1869, 1522, 1628,
// layer 1 inverse
23132, -17422, 7757, 23998, -20257, 28644, -23859, 31637,
-17442, 10907, 23211, -22501, -20198, -7933, -21497, 10336,
1628, 1522, 1869, 958, 991, 996, 3021, 3221,
478, 2459, 2475, 1819, 794, 2051, 1799, 2144,
14470, 16990, -11044, 24215, -14902, 10946, -6279, -20926,
-32384, 24392, 15356, -7382, 2146, 25435, 20297, -31183,
1670, 2142, 220, 2455, 1994, 1218, 3193, 2114,
384, 1864, 2044, 1322, 610, 603, 1097, 817,
-12618, -5275, 19529, -14430, 18525, 17561, 20100, -18485,
12639, -28761, -18741, 29176, 30318, 10631, -32502, 32011,
3254, 3173, 329, 418, 349, 2457, 644, 1739,
1119, 2727, 1483, 2552, 3182, 1159, 778, 3083,
5493, 6182, 23093, -14882, -4586, 945, 13131, -27738,
-21655, 20494, -14233, -32226, 13387, -11477, 11182, -334,
1653, 1574, 2869, 3038, 3094, 177, 587, 422,
105, 1550, 871, 2078, 843, 555, 430, 2226,
// layer 2 inverse
-17914, -17914, -24155, -24155, 4312, 4312, 16832, 16832,
-12756, -12756, -29155, -29155, 14017, 14017, 13427, 13427,
2054, 2054, 677, 677, 2264, 2264, 448, 448,
2604, 2604, 1821, 1821, 961, 961, 2931, 2931,
18250, 18250, -30198, -30198, 9076, 9076, -28309, -28309,
-8898, -8898, -15886, -15886, -28249, -28249, 19884, 19884,
2378, 2378, 3082, 3082, 1908, 1908, 107, 107,
830, 830, 3058, 3058, 3239, 3239, 2476, 2476,
27837, 27837, -25986, -25986, -649, -649, 9135, 9135,
-12441, -12441, 16065, 16065, 26617, 26617, -12796, -12796,
1469, 1469, 126, 126, 2167, 2167, 1711, 1711,
2663, 2663, 3009, 3009, 3321, 3321, 516, 516,
25081, 25081, 20711, 20711, 20179, 20179, -23564, -23564,
30967, 30967, 1497, 1497, -6516, -6516, -5689, -5689,
1785, 1785, 3047, 3047, 1491, 1491, 2036, 2036,
1015, 1015, 2777, 2777, 652, 652, 1223, 1223,
// layer 3 inverse
-16162, -16162, -16162, -16162, 26676, 26676, 26676, 26676,
8859, 8859, 8859, 8859, 18427, 18427, 18427, 18427,
1758, 1758, 1758, 1758, 3124, 3124, 3124, 3124,
411, 411, 411, 411, 1787, 1787, 1787, 1787,
8800, 8800, 8800, 8800, -10532, -10532, -10532, -10532,
24313, 24313, 24313, 24313, -28072, -28072, -28072, -28072,
608, 608, 608, 608, 732, 732, 732, 732,
1017, 1017, 1017, 1017, 2648, 2648, 2648, 2648,
-26241, -26241, -26241, -26241, 21439, 21439, 21439, 21439,
-1102, -1102, -1102, -1102, 5572, 5572, 5572, 5572,
3199, 3199, 3199, 3199, 1727, 1727, 1727, 1727,
1458, 1458, 1458, 1458, 2500, 2500, 2500, 2500,
-29057, -29057, -29057, -29057, -26360, -26360, -26360, -26360,
17364, 17364, 17364, 17364, -5827, -5827, -5827, -5827,
383, 383, 383, 383, 264, 264, 264, 264,
2004, 2004, 2004, 2004, 573, 573, 573, 573,
// layer 4 inverse
31164, 31164, 31164, 31164, 31164, 31164, 31164, 31164,
-11201, -11201, -11201, -11201, -11201, -11201, -11201, -11201,
1468, 1468, 1468, 1468, 1468, 1468, 1468, 1468,
1855, 1855, 1855, 1855, 1855, 1855, 1855, 1855,
1359, 1359, 1359, 1359, 1359, 1359, 1359, 1359,
10690, 10690, 10690, 10690, 10690, 10690, 10690, 10690,
2127, 2127, 2127, 2127, 2127, 2127, 2127, 2127,
962, 962, 962, 962, 962, 962, 962, 962,
-15690, -15690, -15690, -15690, -15690, -15690, -15690, -15690,
-3799, -3799, -3799, -3799, -3799, -3799, -3799, -3799,
182, 182, 182, 182, 182, 182, 182, 182,
1577, 1577, 1577, 1577, 1577, 1577, 1577, 1577,
27758, 27758, 27758, 27758, 27758, 27758, 27758, 27758,
-20906, -20906, -20906, -20906, -20906, -20906, -20906, -20906,
622, 622, 622, 622, 622, 622, 622, 622,
3158, 3158, 3158, 3158, 3158, 3158, 3158, 3158,
// layer 5 inverse
-16694, 202, 28191, 287, -12402, 1422, 13525, 1493,
// layer 6 inverse
788, 1812, 14746, 2970,
// layer 7 inverse
31499, 2571,
}
// Sets p to a + b. Does not normalize coefficients.
func (p *Poly) Add(a, b *Poly) {
if cpu.X86.HasAVX2 {
addAVX2(
(*[N]int16)(p),
(*[N]int16)(a),
(*[N]int16)(b),
)
} else {
p.addGeneric(a, b)
}
}
// Sets p to a - b. Does not normalize coefficients.
func (p *Poly) Sub(a, b *Poly) {
if cpu.X86.HasAVX2 {
subAVX2(
(*[N]int16)(p),
(*[N]int16)(a),
(*[N]int16)(b),
)
} else {
p.subGeneric(a, b)
}
}
// Executes an in-place forward "NTT" on p.
//
// Assumes the coefficients are in absolute value ≤q. The resulting
// coefficients are in absolute value ≤7q. If the input is in Montgomery
// form, then the result is in Montgomery form and so (by linearity of the NTT)
// if the input is in regular form, then the result is also in regular form.
// The order of coefficients will be "tangled". These can be put back into
// their proper order by calling Detangle().
func (p *Poly) NTT() {
if cpu.X86.HasAVX2 {
nttAVX2((*[N]int16)(p))
} else {
p.nttGeneric()
}
}
// Executes an in-place inverse "NTT" on p and multiply by the Montgomery
// factor R.
//
// Requires coefficients to be in "tangled" order, see Tangle().
// Assumes the coefficients are in absolute value ≤q. The resulting
// coefficients are in absolute value ≤q. If the input is in Montgomery
// form, then the result is in Montgomery form and so (by linearity)
// if the input is in regular form, then the result is also in regular form.
func (p *Poly) InvNTT() {
if cpu.X86.HasAVX2 {
invNttAVX2((*[N]int16)(p))
} else {
p.invNTTGeneric()
}
}
// Sets p to the "pointwise" multiplication of a and b.
//
// That is: InvNTT(p) = InvNTT(a) * InvNTT(b). Assumes a and b are in
// Montgomery form. Products between coefficients of a and b must be strictly
// bounded in absolute value by 2¹⁵q. p will be in Montgomery form and
// bounded in absolute value by 2q.
//
// Requires a and b to be in "tangled" order, see Tangle(). p will be in
// tangled order as well.
func (p *Poly) MulHat(a, b *Poly) {
if cpu.X86.HasAVX2 {
mulHatAVX2(
(*[N]int16)(p),
(*[N]int16)(a),
(*[N]int16)(b),
)
} else {
p.mulHatGeneric(a, b)
}
}
// Puts p into the right form to be used with (among others) InvNTT().
func (p *Poly) Tangle() {
if cpu.X86.HasAVX2 {
tangleAVX2((*[N]int16)(p))
}
// When AVX2 is not available, we use the standard order.
}
// Puts p back into standard form.
func (p *Poly) Detangle() {
if cpu.X86.HasAVX2 {
detangleAVX2((*[N]int16)(p))
}
// When AVX2 is not available, we use the standard order.
}
// Almost normalizes coefficients.
//
// Ensures each coefficient is in {0, …, q}.
func (p *Poly) BarrettReduce() {
if cpu.X86.HasAVX2 {
barrettReduceAVX2((*[N]int16)(p))
} else {
p.barrettReduceGeneric()
}
}
// Normalizes coefficients.
//
// Ensures each coefficient is in {0, …, q-1}.
func (p *Poly) Normalize() {
if cpu.X86.HasAVX2 {
normalizeAVX2((*[N]int16)(p))
} else {
p.normalizeGeneric()
}
}
File diff suppressed because it is too large. Load diff
+83
View File
@@ -0,0 +1,83 @@
//go:build arm64 && !purego
// +build arm64,!purego
package common
// Sets p to a + b. Does not normalize coefficients.
func (p *Poly) Add(a, b *Poly) {
polyAddARM64(p, a, b)
}
// Sets p to a - b. Does not normalize coefficients.
func (p *Poly) Sub(a, b *Poly) {
polySubARM64(p, a, b)
}
// Executes an in-place forward "NTT" on p.
//
// Assumes the coefficients are in absolute value ≤q. The resulting
// coefficients are in absolute value ≤7q. If the input is in Montgomery
// form, then the result is in Montgomery form and so (by linearity of the NTT)
// if the input is in regular form, then the result is also in regular form.
// The order of coefficients will be "tangled". These can be put back into
// their proper order by calling Detangle().
func (p *Poly) NTT() {
p.nttGeneric()
}
// Executes an in-place inverse "NTT" on p and multiply by the Montgomery
// factor R.
//
// Requires coefficients to be in "tangled" order, see Tangle().
// Assumes the coefficients are in absolute value ≤q. The resulting
// coefficients are in absolute value ≤q. If the input is in Montgomery
// form, then the result is in Montgomery form and so (by linearity)
// if the input is in regular form, then the result is also in regular form.
func (p *Poly) InvNTT() {
p.invNTTGeneric()
}
// Sets p to the "pointwise" multiplication of a and b.
//
// That is: InvNTT(p) = InvNTT(a) * InvNTT(b). Assumes a and b are in
// Montgomery form. Products between coefficients of a and b must be strictly
// bounded in absolute value by 2¹⁵q. p will be in Montgomery form and
// bounded in absolute value by 2q.
//
// Requires a and b to be in "tangled" order, see Tangle(). p will be in
// tangled order as well.
func (p *Poly) MulHat(a, b *Poly) {
p.mulHatGeneric(a, b)
}
// Puts p into the right form to be used with (among others) InvNTT().
func (p *Poly) Tangle() {
// In the generic implementation there is no advantage to using a
// different order, so we use the standard order everywhere.
}
// Puts p back into standard form.
func (p *Poly) Detangle() {
// In the generic implementation there is no advantage to using a
// different order, so we use the standard order everywhere.
}
// Almost normalizes coefficients.
//
// Ensures each coefficient is in {0, …, q}.
func (p *Poly) BarrettReduce() {
p.barrettReduceGeneric()
}
// Normalizes coefficients.
//
// Ensures each coefficient is in {0, …, q-1}.
func (p *Poly) Normalize() {
p.normalizeGeneric()
}
//go:noescape
func polyAddARM64(p, a, b *Poly)
//go:noescape
func polySubARM64(p, a, b *Poly)
+53
View File
@@ -0,0 +1,53 @@
//go:build arm64 && !purego
#include "go_asm.h"
#include "textflag.h"
// func polyAddARM64(p, a, b *Poly)
TEXT ·polyAddARM64(SB), NOSPLIT|NOFRAME, $0-24
MOVD p+0(FP), R0
MOVD a+8(FP), R1
MOVD b+16(FP), R2
MOVW $(const_N / 32), R3
loop:
VLD1.P (64)(R1), [V0.H8, V1.H8, V2.H8, V3.H8]
VLD1.P (64)(R2), [V4.H8, V5.H8, V6.H8, V7.H8]
VADD V4.H8, V0.H8, V0.H8
VADD V5.H8, V1.H8, V1.H8
VADD V6.H8, V2.H8, V2.H8
VADD V7.H8, V3.H8, V3.H8
VST1.P [V0.H8, V1.H8, V2.H8, V3.H8], (64)(R0)
SUBS $1, R3, R3
BGT loop
RET
// func polySubARM64(p, a, b *Poly)
TEXT ·polySubARM64(SB), NOSPLIT|NOFRAME, $0-24
MOVD p+0(FP), R0
MOVD a+8(FP), R1
MOVD b+16(FP), R2
MOVW $(const_N / 32), R3
loop:
VLD1.P (64)(R1), [V0.H8, V1.H8, V2.H8, V3.H8]
VLD1.P (64)(R2), [V4.H8, V5.H8, V6.H8, V7.H8]
VSUB V4.H8, V0.H8, V0.H8
VSUB V5.H8, V1.H8, V1.H8
VSUB V6.H8, V2.H8, V2.H8
VSUB V7.H8, V3.H8, V3.H8
VST1.P [V0.H8, V1.H8, V2.H8, V3.H8], (64)(R0)
SUBS $1, R3, R3
BGT loop
RET
+74
View File
@@ -0,0 +1,74 @@
package common
// Given -2¹⁵ q ≤ x < 2¹⁵ q, returns -q < y < q with x 2⁻¹⁶ = y (mod q).
func montReduce(x int32) int16 {
// This is Montgomery reduction with R=2¹⁶.
//
// Note gcd(2¹⁶, q) = 1 as q is prime. Write q' := 62209 = q⁻¹ mod R.
// First we compute
//
// m := ((x mod R) q') mod R
// = x q' mod R
// = int16(x q')
// = int16(int32(x) * int32(q'))
//
// Note that x q' might be as big as 2³² and could overflow the int32
// multiplication in the last line. However for any int32s a and b,
// we have int32(int64(a)*int64(b)) = int32(a*b) and so the result is ok.
m := int16(x * 62209)
// Note that x - m q is divisible by R; indeed modulo R we have
//
// x - m q ≡ x - x q' q ≡ x - x q⁻¹ q ≡ x - x = 0.
//
// We return y := (x - m q) / R. Note that y is indeed correct as
// modulo q we have
//
// y ≡ x R⁻¹ - m q R⁻¹ = x R⁻¹
//
// and as both 2¹⁵ q ≤ m q, x < 2¹⁵ q, we have
// 2¹⁶ q ≤ x - m q < 2¹⁶ and so q ≤ (x - m q) / R < q as desired.
return int16(uint32(x-int32(m)*int32(Q)) >> 16)
}
// Given any x, returns x R mod q where R=2¹⁶.
func toMont(x int16) int16 {
// Note |1353 x| ≤ 1353 2¹⁵ ≤ 13318 q ≤ 2¹⁵ q and so we're within
// the bounds of montReduce.
return montReduce(int32(x) * 1353) // 1353 = R² mod q.
}
// Given any x, compute 0 ≤ y ≤ q with x = y (mod q).
//
// Beware: we might have barrettReduce(x) = q ≠ 0 for some x. In fact,
// this happens if and only if x = -nq for some positive integer n.
func barrettReduce(x int16) int16 {
// This is standard Barrett reduction.
//
// For any x we have x mod q = x - ⌊x/q⌋ q. We will use 20159/2²⁶ as
// an approximation of 1/q. Note that 0 ≤ 20159/2²⁶ - 1/q ≤ 0.135/2²⁶
// and so | x 20156/2²⁶ - x/q | ≤ 2⁻¹⁰ for |x| ≤ 2¹⁶. For all x
// not a multiple of q, the number x/q is further than 1/q from any integer
// and so ⌊x 20156/2²⁶⌋ = ⌊x/q⌋. If x is a multiple of q and x is positive,
// then x 20156/2²⁶ is larger than x/q so ⌊x 20156/2²⁶⌋ = ⌊x/q⌋ as well.
// Finally, if x is negative multiple of q, then ⌊x 20156/2²⁶⌋ = ⌊x/q⌋-1.
// Thus
// [ q if x=-nq for pos. integer n
// x - ⌊x 20156/2²⁶⌋ q = [
// [ x mod q otherwise
//
// To compute actually compute this, note that
//
// ⌊x 20156/2²⁶⌋ = (20159 x) >> 26.
return x - int16((int32(x)*20159)>>26)*Q
}
// Returns x if x < q and x - q otherwise. Assumes x ≥ -29439.
func csubq(x int16) int16 {
x -= Q // no overflow due to assumption x ≥ -29439.
// If x is positive, then x >> 15 = 0. If x is negative,
// then uint16(x >> 15) = 2¹⁶-1. So this will add back in q
// if x was smaller than q.
x += (x >> 15) & Q
return x
}
@@ -0,0 +1,77 @@
//go:build (!amd64 && !arm64) || purego
// +build !amd64,!arm64 purego
package common
// Sets p to a + b. Does not normalize coefficients.
func (p *Poly) Add(a, b *Poly) {
p.addGeneric(a, b)
}
// Sets p to a - b. Does not normalize coefficients.
func (p *Poly) Sub(a, b *Poly) {
p.subGeneric(a, b)
}
// Executes an in-place forward "NTT" on p.
//
// Assumes the coefficients are in absolute value ≤q. The resulting
// coefficients are in absolute value ≤7q. If the input is in Montgomery
// form, then the result is in Montgomery form and so (by linearity of the NTT)
// if the input is in regular form, then the result is also in regular form.
// The order of coefficients will be "tangled". These can be put back into
// their proper order by calling Detangle().
func (p *Poly) NTT() {
p.nttGeneric()
}
// Executes an in-place inverse "NTT" on p and multiply by the Montgomery
// factor R.
//
// Requires coefficients to be in "tangled" order, see Tangle().
// Assumes the coefficients are in absolute value ≤q. The resulting
// coefficients are in absolute value ≤q. If the input is in Montgomery
// form, then the result is in Montgomery form and so (by linearity)
// if the input is in regular form, then the result is also in regular form.
func (p *Poly) InvNTT() {
p.invNTTGeneric()
}
// Sets p to the "pointwise" multiplication of a and b.
//
// That is: InvNTT(p) = InvNTT(a) * InvNTT(b). Assumes a and b are in
// Montgomery form. Products between coefficients of a and b must be strictly
// bounded in absolute value by 2¹⁵q. p will be in Montgomery form and
// bounded in absolute value by 2q.
//
// Requires a and b to be in "tangled" order, see Tangle(). p will be in
// tangled order as well.
func (p *Poly) MulHat(a, b *Poly) {
p.mulHatGeneric(a, b)
}
// Puts p into the right form to be used with (among others) InvNTT().
func (p *Poly) Tangle() {
// In the generic implementation there is no advantage to using a
// different order, so we use the standard order everywhere.
}
// Puts p back into standard form.
func (p *Poly) Detangle() {
// In the generic implementation there is no advantage to using a
// different order, so we use the standard order everywhere.
}
// Almost normalizes coefficients.
//
// Ensures each coefficient is in {0, …, q}.
func (p *Poly) BarrettReduce() {
p.barrettReduceGeneric()
}
// Normalizes coefficients.
//
// Ensures each coefficient is in {0, …, q-1}.
func (p *Poly) Normalize() {
p.normalizeGeneric()
}
+193
View File
@@ -0,0 +1,193 @@
package common
// Zetas lists precomputed powers of the primitive root of unity in
// Montgomery representation used for the NTT:
//
// Zetas[i] = ζᵇʳᵛ⁽ⁱ⁾ R mod q
//
// where ζ = 17, brv(i) is the bitreversal of a 7-bit number and R=2¹⁶ mod q.
//
// The following Python code generates the Zetas arrays:
//
// q = 13*2**8 + 1; zeta = 17
// R = 2**16 % q # Montgomery const.
// def brv(x): return int(''.join(reversed(bin(x)[2:].zfill(7))),2)
// print([(pow(zeta, brv(i), q)*R)%q for i in range(128)])
var Zetas = [128]int16{
2285, 2571, 2970, 1812, 1493, 1422, 287, 202, 3158, 622, 1577, 182,
962, 2127, 1855, 1468, 573, 2004, 264, 383, 2500, 1458, 1727, 3199,
2648, 1017, 732, 608, 1787, 411, 3124, 1758, 1223, 652, 2777, 1015,
2036, 1491, 3047, 1785, 516, 3321, 3009, 2663, 1711, 2167, 126,
1469, 2476, 3239, 3058, 830, 107, 1908, 3082, 2378, 2931, 961, 1821,
2604, 448, 2264, 677, 2054, 2226, 430, 555, 843, 2078, 871, 1550,
105, 422, 587, 177, 3094, 3038, 2869, 1574, 1653, 3083, 778, 1159,
3182, 2552, 1483, 2727, 1119, 1739, 644, 2457, 349, 418, 329, 3173,
3254, 817, 1097, 603, 610, 1322, 2044, 1864, 384, 2114, 3193, 1218,
1994, 2455, 220, 2142, 1670, 2144, 1799, 2051, 794, 1819, 2475,
2459, 478, 3221, 3021, 996, 991, 958, 1869, 1522, 1628,
}
// InvNTTReductions keeps track of which coefficients to apply Barrett
// reduction to in Poly.InvNTT().
//
// Generated in a lazily: once a butterfly is computed which is about to
// overflow the int16, the largest coefficient is reduced. If that is
// not enough, the other coefficient is reduced as well.
//
// This is actually optimal, as proven in https://eprint.iacr.org/2020/1377.pdf
var InvNTTReductions = [...]int{
-1, // after layer 1
-1, // after layer 2
16, 17, 48, 49, 80, 81, 112, 113, 144, 145, 176, 177, 208, 209, 240,
241, -1, // after layer 3
0, 1, 32, 33, 34, 35, 64, 65, 96, 97, 98, 99, 128, 129, 160, 161, 162, 163,
192, 193, 224, 225, 226, 227, -1, // after layer 4
2, 3, 66, 67, 68, 69, 70, 71, 130, 131, 194, 195, 196, 197, 198,
199, -1, // after layer 5
4, 5, 6, 7, 132, 133, 134, 135, 136, 137, 138, 139, 140, 141, 142,
143, -1, // after layer 6
-1, // after layer 7
}
// Executes an in-place forward "NTT" on p.
//
// Assumes the coefficients are in absolute value ≤q. The resulting
// coefficients are in absolute value ≤7q. If the input is in Montgomery
// form, then the result is in Montgomery form and so (by linearity of the NTT)
// if the input is in regular form, then the result is also in regular form.
// The order of coefficients will be "tangled". These can be put back into
// their proper order by calling Detangle().
func (p *Poly) nttGeneric() {
// Note that ℤ_q does not have a primitive 512ᵗʰ root of unity (as 512
// does not divide into q-1) and so we cannot do a regular NTT. ℤ_q
// does have a primitive 256ᵗʰ root of unity, the smallest of which
// is ζ := 17.
//
// Recall that our base ring R := ℤ_q[x] / (x²⁵⁶ + 1). The polynomial
// x²⁵⁶+1 will not split completely (as its roots would be 512ᵗʰ roots
// of unity.) However, it does split almost (using ζ¹²⁸ = -1):
//
// x²⁵⁶ + 1 = (x²)¹²⁸ - ζ¹²⁸
// = ((x²)⁶⁴ - ζ⁶⁴)((x²)⁶⁴ + ζ⁶⁴)
// = ((x²)³² - ζ³²)((x²)³² + ζ³²)((x²)³² - ζ⁹⁶)((x²)³² + ζ⁹⁶)
// ⋮
// = (x² - ζ)(x² + ζ)(x² - ζ⁶⁵)(x² + ζ⁶⁵) … (x² + ζ¹²⁷)
//
// Note that the powers of ζ that appear (from the second line down) are
// in binary
//
// 0100000 1100000
// 0010000 1010000 0110000 1110000
// 0001000 1001000 0101000 1101000 0011000 1011000 0111000 1111000
// …
//
// That is: brv(2), brv(3), brv(4), …, where brv(x) denotes the 7-bit
// bitreversal of x. These powers of ζ are given by the Zetas array.
//
// The polynomials x² ± ζⁱ are irreducible and coprime, hence by
// the Chinese Remainder Theorem we know
//
// ℤ_q[x]/(x²⁵⁶+1) → ℤ_q[x]/(x²-ζ) x … x ℤ_q[x]/(x²+ζ¹²⁷)
//
// given by a ↦ ( a mod x²-ζ, …, a mod x²+ζ¹²⁷ )
// is an isomorphism, which is the "NTT". It can be efficiently computed by
//
//
// a ↦ ( a mod (x²)⁶⁴ - ζ⁶⁴, a mod (x²)⁶⁴ + ζ⁶⁴ )
// ↦ ( a mod (x²)³² - ζ³², a mod (x²)³² + ζ³²,
// a mod (x²)⁹⁶ - ζ⁹⁶, a mod (x²)⁹⁶ + ζ⁹⁶ )
//
// et cetera
//
// If N was 8 then this can be pictured in the following diagram:
//
// https://cnx.org/resources/17ee4dfe517a6adda05377b25a00bf6e6c93c334/File0026.png
//
// Each cross is a Cooley-Tukey butterfly: it's the map
//
// (a, b) ↦ (a + ζb, a - ζb)
//
// for the appropriate power ζ for that column and row group.
k := 0 // Index into Zetas
// l runs effectively over the columns in the diagram above; it is half the
// height of a row group, i.e. the number of butterflies in each row group.
// In the diagram above it would be 4, 2, 1.
for l := N / 2; l > 1; l >>= 1 {
// On the nᵗʰ iteration of the l-loop, the absolute value of the
// coefficients are bounded by nq.
// offset effectively loops over the row groups in this column; it is
// the first row in the row group.
for offset := 0; offset < N-l; offset += 2 * l {
k++
zeta := int32(Zetas[k])
// j loops over each butterfly in the row group.
for j := offset; j < offset+l; j++ {
t := montReduce(zeta * int32(p[j+l]))
p[j+l] = p[j] - t
p[j] += t
}
}
}
}
// Executes an in-place inverse "NTT" on p and multiply by the Montgomery
// factor R.
//
// Requires coefficients to be in "tangled" order, see Tangle().
// Assumes the coefficients are in absolute value ≤q. The resulting
// coefficients are in absolute value ≤q. If the input is in Montgomery
// form, then the result is in Montgomery form and so (by linearity)
// if the input is in regular form, then the result is also in regular form.
func (p *Poly) invNTTGeneric() {
k := 127 // Index into Zetas
r := -1 // Index into InvNTTReductions.
// We basically do the opposite of NTT, but postpone dividing by 2 in the
// inverse of the Cooley-Tukey butterfly and accumulate that into a big
// division by 2⁷ at the end. See the comments in the NTT() function.
for l := 2; l < N; l <<= 1 {
for offset := 0; offset < N-l; offset += 2 * l {
// As we're inverting, we need powers of ζ⁻¹ (instead of ζ).
// To be precise, we need ζᵇʳᵛ⁽ᵏ⁾⁻¹²⁸. However, as ζ⁻¹²⁸ = -1,
// we can use the existing Zetas table instead of
// keeping a separate InvZetas table as in Dilithium.
minZeta := int32(Zetas[k])
k--
for j := offset; j < offset+l; j++ {
// Gentleman-Sande butterfly: (a, b) ↦ (a + b, ζ(a-b))
t := p[j+l] - p[j]
p[j] += p[j+l]
p[j+l] = montReduce(minZeta * int32(t))
// Note that if we had |a| < αq and |b| < βq before the
// butterfly, then now we have |a| < (α+β)q and |b| < q.
}
}
// We let the InvNTTReductions instruct us which coefficients to
// Barrett reduce. See TestInvNTTReductions, which tests whether
// there is an overflow.
for {
r++
i := InvNTTReductions[r]
if i < 0 {
break
}
p[i] = barrettReduce(p[i])
}
}
for j := 0; j < N; j++ {
// Note 1441 = (128)⁻¹ R². The coefficients are bounded by 9q, so
// as 1441 * 9 ≈ 2¹⁴ < 2¹⁵, we're within the required bounds
// for montReduce().
p[j] = montReduce(1441 * int32(p[j]))
}
}
+22
View File
@@ -0,0 +1,22 @@
package common
import (
"github.com/cloudflare/circl/pke/kyber/internal/common/params"
)
const (
// Q is the parameter q ≡ 3329 = 2¹¹ + 2¹⁰ + 2⁸ + 1.
Q = params.Q
// N is the parameter N: the length of the polynomials
N = params.N
// PolySize is the size of a packed polynomial.
PolySize = params.PolySize
// PlaintextSize is the size of the plaintext
PlaintextSize = params.PlaintextSize
// Eta2 is the parameter η₂
Eta2 = params.Eta2
)
@@ -0,0 +1,21 @@
package params
// We put these parameters in a separate package so that the Go code,
// such as asm/src.go, that generates assembler can import it.
const (
// Q is the parameter q ≡ 3329 = 2¹¹ + 2¹⁰ + 2⁸ + 1.
Q int16 = 3329
// N is the parameter N: the length of the polynomials
N = 256
// PolySize is the size of a packed polynomial.
PolySize = 384
// PlaintextSize is the size of the plaintext
PlaintextSize = 32
// Eta2 is the parameter η₂
Eta2 = 2
)
+332
View File
@@ -0,0 +1,332 @@
package common
// An element of our base ring R which are polynomials over ℤ_q
// modulo the equation Xᴺ = -1, where q=3329 and N=256.
//
// This type is also used to store NTT-transformed polynomials,
// see Poly.NTT().
//
// Coefficients aren't always reduced. See Normalize().
type Poly [N]int16
// Sets p to a + b. Does not normalize coefficients.
func (p *Poly) addGeneric(a, b *Poly) {
for i := 0; i < N; i++ {
p[i] = a[i] + b[i]
}
}
// Sets p to a - b. Does not normalize coefficients.
func (p *Poly) subGeneric(a, b *Poly) {
for i := 0; i < N; i++ {
p[i] = a[i] - b[i]
}
}
// Almost normalizes coefficients.
//
// Ensures each coefficient is in {0, …, q}.
func (p *Poly) barrettReduceGeneric() {
for i := 0; i < N; i++ {
p[i] = barrettReduce(p[i])
}
}
// Normalizes coefficients.
//
// Ensures each coefficient is in {0, …, q-1}.
func (p *Poly) normalizeGeneric() {
for i := 0; i < N; i++ {
p[i] = csubq(barrettReduce(p[i]))
}
}
// Multiplies p in-place by the Montgomery factor 2¹⁶.
//
// Coefficients of p can be arbitrary. Resulting coefficients are bounded
// in absolute value by q.
func (p *Poly) ToMont() {
for i := 0; i < N; i++ {
p[i] = toMont(p[i])
}
}
// Sets p to the "pointwise" multiplication of a and b.
//
// That is: InvNTT(p) = InvNTT(a) * InvNTT(b). Assumes a and b are in
// Montgomery form. Products between coefficients of a and b must be strictly
// bounded in absolute value by 2¹⁵q. p will be in Montgomery form and
// bounded in absolute value by 2q.
//
// Requires a and b to be in "tangled" order, see Tangle(). p will be in
// tangled order as well.
func (p *Poly) mulHatGeneric(a, b *Poly) {
// Recall from the discussion in NTT(), that a transformed polynomial is
// an element of ℤ_q[x]/(x²-ζ) x … x ℤ_q[x]/(x²+ζ¹²⁷);
// that is: 128 degree-one polynomials instead of simply 256 elements
// from ℤ_q as in the regular NTT. So instead of pointwise multiplication,
// we multiply the 128 pairs of degree-one polynomials modulo the
// right equation:
//
// (a₁ + a₂x)(b₁ + b₂x) = a₁b₁ + a₂b₂ζ' + (a₁b₂ + a₂b₁)x,
//
// where ζ' is the appropriate power of ζ.
k := 64
for i := 0; i < N; i += 4 {
zeta := int32(Zetas[k])
k++
p0 := montReduce(int32(a[i+1]) * int32(b[i+1]))
p0 = montReduce(int32(p0) * zeta)
p0 += montReduce(int32(a[i]) * int32(b[i]))
p1 := montReduce(int32(a[i]) * int32(b[i+1]))
p1 += montReduce(int32(a[i+1]) * int32(b[i]))
p[i] = p0
p[i+1] = p1
p2 := montReduce(int32(a[i+3]) * int32(b[i+3]))
p2 = -montReduce(int32(p2) * zeta)
p2 += montReduce(int32(a[i+2]) * int32(b[i+2]))
p3 := montReduce(int32(a[i+2]) * int32(b[i+3]))
p3 += montReduce(int32(a[i+3]) * int32(b[i+2]))
p[i+2] = p2
p[i+3] = p3
}
}
// Packs p into buf. buf should be of length PolySize.
//
// Assumes p is normalized (and not just Barrett reduced) and "tangled",
// see Tangle().
func (p *Poly) Pack(buf []byte) {
q := *p
q.Detangle()
for i := 0; i < 128; i++ {
t0 := q[2*i]
t1 := q[2*i+1]
buf[3*i] = byte(t0)
buf[3*i+1] = byte(t0>>8) | byte(t1<<4)
buf[3*i+2] = byte(t1 >> 4)
}
}
// Unpacks p from buf.
//
// buf should be of length PolySize. p will be "tangled", see Detangle().
//
// p will not be normalized; instead 0 ≤ p[i] < 4096.
func (p *Poly) Unpack(buf []byte) {
for i := 0; i < 128; i++ {
p[2*i] = int16(buf[3*i]) | ((int16(buf[3*i+1]) << 8) & 0xfff)
p[2*i+1] = int16(buf[3*i+1]>>4) | (int16(buf[3*i+2]) << 4)
}
p.Tangle()
}
// Set p to Decompress_q(m, 1).
//
// p will be normalized. m has to be of PlaintextSize.
func (p *Poly) DecompressMessage(m []byte) {
// Decompress_q(x, 1) = ⌈xq/2⌋ = ⌊xq/2+½⌋ = (xq+1) >> 1 and so
// Decompress_q(0, 1) = 0 and Decompress_q(1, 1) = (q+1)/2.
for i := 0; i < 32; i++ {
for j := 0; j < 8; j++ {
bit := (m[i] >> uint(j)) & 1
// Set coefficient to either 0 or (q+1)/2 depending on the bit.
p[8*i+j] = -int16(bit) & ((Q + 1) / 2)
}
}
}
// Writes Compress_q(p, 1) to m.
//
// Assumes p is normalized. m has to be of length at least PlaintextSize.
func (p *Poly) CompressMessageTo(m []byte) {
// Compress_q(x, 1) is 1 on {833, …, 2496} and zero elsewhere.
for i := 0; i < 32; i++ {
m[i] = 0
for j := 0; j < 8; j++ {
x := 1664 - p[8*i+j]
// With the previous substitution, we want to return 1 if
// and only if x is in {831, …, -832}.
x = (x >> 15) ^ x
// Note (x >> 15)ˣ if x≥0 and -x-1 otherwise. Thus now we want
// to return 1 iff x ≤ 831, ie. x - 832 < 0.
x -= 832
m[i] |= ((byte(x >> 15)) & 1) << uint(j)
}
}
}
// Set p to Decompress_q(m, 1).
//
// Assumes d is in {4, 5, 10, 11}. p will be normalized.
func (p *Poly) Decompress(m []byte, d int) {
// Decompress_q(x, d) = ⌈(q/2ᵈ)x⌋
// = ⌊(q/2ᵈ)x+½⌋
// = ⌊(qx + 2ᵈ⁻¹)/2ᵈ⌋
// = (qx + (1<<(d-1))) >> d
switch d {
case 4:
for i := 0; i < N/2; i++ {
p[2*i] = int16(((1 << 3) +
uint32(m[i]&15)*uint32(Q)) >> 4)
p[2*i+1] = int16(((1 << 3) +
uint32(m[i]>>4)*uint32(Q)) >> 4)
}
case 5:
var t [8]uint16
idx := 0
for i := 0; i < N/8; i++ {
t[0] = uint16(m[idx])
t[1] = (uint16(m[idx]) >> 5) | (uint16(m[idx+1] << 3))
t[2] = uint16(m[idx+1]) >> 2
t[3] = (uint16(m[idx+1]) >> 7) | (uint16(m[idx+2] << 1))
t[4] = (uint16(m[idx+2]) >> 4) | (uint16(m[idx+3] << 4))
t[5] = uint16(m[idx+3]) >> 1
t[6] = (uint16(m[idx+3]) >> 6) | (uint16(m[idx+4] << 2))
t[7] = uint16(m[idx+4]) >> 3
for j := 0; j < 8; j++ {
p[8*i+j] = int16(((1 << 4) +
uint32(t[j]&((1<<5)-1))*uint32(Q)) >> 5)
}
idx += 5
}
case 10:
var t [4]uint16
idx := 0
for i := 0; i < N/4; i++ {
t[0] = uint16(m[idx]) | (uint16(m[idx+1]) << 8)
t[1] = (uint16(m[idx+1]) >> 2) | (uint16(m[idx+2]) << 6)
t[2] = (uint16(m[idx+2]) >> 4) | (uint16(m[idx+3]) << 4)
t[3] = (uint16(m[idx+3]) >> 6) | (uint16(m[idx+4]) << 2)
for j := 0; j < 4; j++ {
p[4*i+j] = int16(((1 << 9) +
uint32(t[j]&((1<<10)-1))*uint32(Q)) >> 10)
}
idx += 5
}
case 11:
var t [8]uint16
idx := 0
for i := 0; i < N/8; i++ {
t[0] = uint16(m[idx]) | (uint16(m[idx+1]) << 8)
t[1] = (uint16(m[idx+1]) >> 3) | (uint16(m[idx+2]) << 5)
t[2] = (uint16(m[idx+2]) >> 6) | (uint16(m[idx+3]) << 2) | (uint16(m[idx+4]) << 10)
t[3] = (uint16(m[idx+4]) >> 1) | (uint16(m[idx+5]) << 7)
t[4] = (uint16(m[idx+5]) >> 4) | (uint16(m[idx+6]) << 4)
t[5] = (uint16(m[idx+6]) >> 7) | (uint16(m[idx+7]) << 1) | (uint16(m[idx+8]) << 9)
t[6] = (uint16(m[idx+8]) >> 2) | (uint16(m[idx+9]) << 6)
t[7] = (uint16(m[idx+9]) >> 5) | (uint16(m[idx+10]) << 3)
for j := 0; j < 8; j++ {
p[8*i+j] = int16(((1 << 10) +
uint32(t[j]&((1<<11)-1))*uint32(Q)) >> 11)
}
idx += 11
}
default:
panic("unsupported d")
}
}
// Writes Compress_q(p, d) to m.
//
// Assumes p is normalized and d is in {4, 5, 10, 11}.
func (p *Poly) CompressTo(m []byte, d int) {
// Compress_q(x, d) = ⌈(2ᵈ/q)x⌋ mod⁺ 2ᵈ
// = ⌊(2ᵈ/q)x+½⌋ mod⁺ 2ᵈ
// = ⌊((x << d) + q/2) / q⌋ mod⁺ 2ᵈ
// = DIV((x << d) + q/2, q) & ((1<<d) - 1)
//
// We approximate DIV(x, q) by computing (x*a)>>e, where a/(2^e) ≈ 1/q.
// For d in {10,11} we use 20,642,679/2^36, which computes division by x/q
// correctly for 0 ≤ x < 41,522,616, which fits (q << 11) + q/2 comfortably.
// For d in {4,5} we use 315/2^20, which doesn't compute division by x/q
// correctly for all inputs, but it's close enough that the end result
// of the compression is correct. The advantage is that we do not need
// to use a 64-bit intermediate value.
switch d {
case 4:
var t [8]uint16
idx := 0
for i := 0; i < N/8; i++ {
for j := 0; j < 8; j++ {
t[j] = uint16((((uint32(p[8*i+j])<<4)+uint32(Q)/2)*315)>>
20) & ((1 << 4) - 1)
}
m[idx] = byte(t[0]) | byte(t[1]<<4)
m[idx+1] = byte(t[2]) | byte(t[3]<<4)
m[idx+2] = byte(t[4]) | byte(t[5]<<4)
m[idx+3] = byte(t[6]) | byte(t[7]<<4)
idx += 4
}
case 5:
var t [8]uint16
idx := 0
for i := 0; i < N/8; i++ {
for j := 0; j < 8; j++ {
t[j] = uint16((((uint32(p[8*i+j])<<5)+uint32(Q)/2)*315)>>
20) & ((1 << 5) - 1)
}
m[idx] = byte(t[0]) | byte(t[1]<<5)
m[idx+1] = byte(t[1]>>3) | byte(t[2]<<2) | byte(t[3]<<7)
m[idx+2] = byte(t[3]>>1) | byte(t[4]<<4)
m[idx+3] = byte(t[4]>>4) | byte(t[5]<<1) | byte(t[6]<<6)
m[idx+4] = byte(t[6]>>2) | byte(t[7]<<3)
idx += 5
}
case 10:
var t [4]uint16
idx := 0
for i := 0; i < N/4; i++ {
for j := 0; j < 4; j++ {
t[j] = uint16((uint64((uint32(p[4*i+j])<<10)+uint32(Q)/2)*
20642679)>>36) & ((1 << 10) - 1)
}
m[idx] = byte(t[0])
m[idx+1] = byte(t[0]>>8) | byte(t[1]<<2)
m[idx+2] = byte(t[1]>>6) | byte(t[2]<<4)
m[idx+3] = byte(t[2]>>4) | byte(t[3]<<6)
m[idx+4] = byte(t[3] >> 2)
idx += 5
}
case 11:
var t [8]uint16
idx := 0
for i := 0; i < N/8; i++ {
for j := 0; j < 8; j++ {
t[j] = uint16((uint64((uint32(p[8*i+j])<<11)+uint32(Q)/2)*
20642679)>>36) & ((1 << 11) - 1)
}
m[idx] = byte(t[0])
m[idx+1] = byte(t[0]>>8) | byte(t[1]<<3)
m[idx+2] = byte(t[1]>>5) | byte(t[2]<<6)
m[idx+3] = byte(t[2] >> 2)
m[idx+4] = byte(t[2]>>10) | byte(t[3]<<1)
m[idx+5] = byte(t[3]>>7) | byte(t[4]<<4)
m[idx+6] = byte(t[4]>>4) | byte(t[5]<<7)
m[idx+7] = byte(t[5] >> 1)
m[idx+8] = byte(t[5]>>9) | byte(t[6]<<2)
m[idx+9] = byte(t[6]>>6) | byte(t[7]<<5)
m[idx+10] = byte(t[7] >> 3)
idx += 11
}
default:
panic("unsupported d")
}
}
+236
View File
@@ -0,0 +1,236 @@
package common
import (
"encoding/binary"
"github.com/cloudflare/circl/internal/sha3"
"github.com/cloudflare/circl/simd/keccakf1600"
)
// DeriveX4Available indicates whether the system supports the quick fourway
// sampling variants like PolyDeriveUniformX4.
var DeriveX4Available = keccakf1600.IsEnabledX4()
// Samples p from a centered binomial distribution with given η.
//
// Essentially CBD_η(PRF(seed, nonce)) from the specification.
func (p *Poly) DeriveNoise(seed []byte, nonce uint8, eta int) {
switch eta {
case 2:
p.DeriveNoise2(seed, nonce)
case 3:
p.DeriveNoise3(seed, nonce)
default:
panic("unsupported eta")
}
}
// Sample p from a centered binomial distribution with n=6 and p=½ - that is:
// coefficients are in {-3, -2, -1, 0, 1, 2, 3} with probabilities {1/64, 3/32,
// 15/64, 5/16, 16/64, 3/32, 1/64}.
func (p *Poly) DeriveNoise3(seed []byte, nonce uint8) {
keySuffix := [1]byte{nonce}
h := sha3.NewShake256()
_, _ = h.Write(seed[:])
_, _ = h.Write(keySuffix[:])
// The distribution at hand is exactly the same as that
// of (a₁ + a₂ + a₃) - (b₁ + b₂+b₃) where a_i,b_i~U(1). Thus we need
// 6 bits per coefficients, thus 192 bytes of input entropy.
// We add two extra zero bytes in the buffer to be able to read 8 bytes
// at the same time (while using only 6.)
var buf [192 + 2]byte
_, _ = h.Read(buf[:192])
for i := 0; i < 32; i++ {
// t is interpreted as a₁ + 2a₂ + 4a₃ + 8b₁ + 16b₂ + ….
t := binary.LittleEndian.Uint64(buf[6*i:])
d := t & 0x249249249249 // a₁ + 8b₁ + …
d += (t >> 1) & 0x249249249249 // a₁ + a₂ + 8(b₁ + b₂) + …
d += (t >> 2) & 0x249249249249 // a₁ + a₂ + a₃ + 4(b₁ + b₂ + b₃) + …
for j := 0; j < 8; j++ {
a := int16(d) & 0x7 // a₁ + a₂ + a₃
d >>= 3
b := int16(d) & 0x7 // b₁ + b₂ + b₃
d >>= 3
p[8*i+j] = a - b
}
}
}
// Sample p from a centered binomial distribution with n=4 and p=½ - that is:
// coefficients are in {-2, -1, 0, 1, 2} with probabilities {1/16, 1/4,
// 3/8, 1/4, 1/16}.
func (p *Poly) DeriveNoise2(seed []byte, nonce uint8) {
keySuffix := [1]byte{nonce}
h := sha3.NewShake256()
_, _ = h.Write(seed[:])
_, _ = h.Write(keySuffix[:])
// The distribution at hand is exactly the same as that
// of (a + a') - (b + b') where a,a',b,b'~U(1). Thus we need 4 bits per
// coefficients, thus 128 bytes of input entropy.
var buf [128]byte
_, _ = h.Read(buf[:])
for i := 0; i < 16; i++ {
// t is interpreted as a + 2a' + 4b + 8b' + ….
t := binary.LittleEndian.Uint64(buf[8*i:])
d := t & 0x5555555555555555 // a + 4b + …
d += (t >> 1) & 0x5555555555555555 // a+a' + 4(b + b') + …
for j := 0; j < 16; j++ {
a := int16(d) & 0x3
d >>= 2
b := int16(d) & 0x3
d >>= 2
p[16*i+j] = a - b
}
}
}
// For each i, sample ps[i] uniformly from the given seed for coordinates
// xs[i] and ys[i]. ps[i] may be nil and is ignored in that case.
//
// Can only be called when DeriveX4Available is true.
func PolyDeriveUniformX4(ps [4]*Poly, seed *[32]byte, xs, ys [4]uint8) {
var perm keccakf1600.StateX4
state := perm.Initialize(false)
// Absorb the seed in the four states
for i := 0; i < 4; i++ {
v := binary.LittleEndian.Uint64(seed[8*i : 8*(i+1)])
for j := 0; j < 4; j++ {
state[i*4+j] = v
}
}
// Absorb the coordinates, the SHAKE128 domain separator (0b1111), the
// start of the padding (0b…001) and the end of the padding 0b100….
// Recall that the rate of SHAKE128 is 168; ie. 21 uint64s.
for j := 0; j < 4; j++ {
state[4*4+j] = uint64(xs[j]) | (uint64(ys[j]) << 8) | (0x1f << 16)
state[20*4+j] = 0x80 << 56
}
var idx [4]int // indices into ps
for j := 0; j < 4; j++ {
if ps[j] == nil {
idx[j] = N // mark nil polynomials as completed
}
}
done := false
for !done {
// Applies KeccaK-f[1600] to state to get the next 21 uint64s of each of
// the four SHAKE128 streams.
perm.Permute()
done = true
PolyLoop:
for j := 0; j < 4; j++ {
if idx[j] == N {
continue
}
for i := 0; i < 7; i++ {
var t [16]uint16
v1 := state[i*3*4+j]
v2 := state[(i*3+1)*4+j]
v3 := state[(i*3+2)*4+j]
t[0] = uint16(v1) & 0xfff
t[1] = uint16(v1>>12) & 0xfff
t[2] = uint16(v1>>24) & 0xfff
t[3] = uint16(v1>>36) & 0xfff
t[4] = uint16(v1>>48) & 0xfff
t[5] = uint16((v1>>60)|(v2<<4)) & 0xfff
t[6] = uint16(v2>>8) & 0xfff
t[7] = uint16(v2>>20) & 0xfff
t[8] = uint16(v2>>32) & 0xfff
t[9] = uint16(v2>>44) & 0xfff
t[10] = uint16((v2>>56)|(v3<<8)) & 0xfff
t[11] = uint16(v3>>4) & 0xfff
t[12] = uint16(v3>>16) & 0xfff
t[13] = uint16(v3>>28) & 0xfff
t[14] = uint16(v3>>40) & 0xfff
t[15] = uint16(v3>>52) & 0xfff
for k := 0; k < 16; k++ {
if t[k] < uint16(Q) {
ps[j][idx[j]] = int16(t[k])
idx[j]++
if idx[j] == N {
continue PolyLoop
}
}
}
}
done = false
}
}
for i := 0; i < 4; i++ {
if ps[i] != nil {
ps[i].Tangle()
}
}
}
// Sample p uniformly from the given seed and x and y coordinates.
//
// Coefficients are reduced and will be in "tangled" order. See Tangle().
func (p *Poly) DeriveUniform(seed *[32]byte, x, y uint8) {
var seedSuffix [2]byte
var buf [168]byte // rate of SHAKE-128
seedSuffix[0] = x
seedSuffix[1] = y
h := sha3.NewShake128()
_, _ = h.Write(seed[:])
_, _ = h.Write(seedSuffix[:])
i := 0
for {
_, _ = h.Read(buf[:])
for j := 0; j < 168; j += 3 {
t1 := (uint16(buf[j]) | (uint16(buf[j+1]) << 8)) & 0xfff //#nosec G602 -- buf has fixed length 168
t2 := (uint16(buf[j+1]>>4) | (uint16(buf[j+2]) << 4)) & 0xfff //#nosec G602 -- buf has fixed length 168
if t1 < uint16(Q) {
p[i] = int16(t1)
i++
if i == N {
break
}
}
if t2 < uint16(Q) {
p[i] = int16(t2)
i++
if i == N {
break
}
}
}
if i == N {
break
}
}
p.Tangle()
}
@@ -0,0 +1,32 @@
// Code generated by command: go run src.go -out ../amd64.s -stubs ../stubs_amd64.go -pkg common. DO NOT EDIT.
//go:build amd64 && !purego
package common
//go:noescape
func addAVX2(p *[256]int16, a *[256]int16, b *[256]int16)
//go:noescape
func subAVX2(p *[256]int16, a *[256]int16, b *[256]int16)
//go:noescape
func nttAVX2(p *[256]int16)
//go:noescape
func invNttAVX2(p *[256]int16)
//go:noescape
func mulHatAVX2(p *[256]int16, a *[256]int16, b *[256]int16)
//go:noescape
func detangleAVX2(p *[256]int16)
//go:noescape
func tangleAVX2(p *[256]int16)
//go:noescape
func barrettReduceAVX2(p *[256]int16)
//go:noescape
func normalizeAVX2(p *[256]int16)
@@ -0,0 +1,192 @@
// Code generated from kyber512/internal/cpapke.go by gen.go
package internal
import (
"bytes"
"github.com/cloudflare/circl/internal/sha3"
"github.com/cloudflare/circl/kem"
"github.com/cloudflare/circl/pke/kyber/internal/common"
)
// A Kyber.CPAPKE private key.
type PrivateKey struct {
sh Vec // NTT(s), normalized
}
// A Kyber.CPAPKE public key.
type PublicKey struct {
rho [32]byte // ρ, the seed for the matrix A
th Vec // NTT(t), normalized
// cached values
aT Mat // the matrix Aᵀ
}
// Packs the private key to buf.
func (sk *PrivateKey) Pack(buf []byte) {
sk.sh.Pack(buf)
}
// Unpacks the private key from buf.
func (sk *PrivateKey) Unpack(buf []byte) {
sk.sh.Unpack(buf)
sk.sh.Normalize()
}
// Packs the public key to buf.
func (pk *PublicKey) Pack(buf []byte) {
pk.th.Pack(buf)
copy(buf[K*common.PolySize:], pk.rho[:])
}
// Unpacks the public key from buf. Checks if the public key is normalized.
func (pk *PublicKey) UnpackMLKEM(buf []byte) error {
pk.Unpack(buf)
// FIPS 203 §7.2 "encapsulation key check" (2).
var buf2 [K * common.PolySize]byte
pk.th.Pack(buf2[:])
if !bytes.Equal(buf[:len(buf2)], buf2[:]) {
return kem.ErrPubKey
}
return nil
}
// Unpacks the public key from buf.
func (pk *PublicKey) Unpack(buf []byte) {
pk.th.Unpack(buf)
pk.th.Normalize()
copy(pk.rho[:], buf[K*common.PolySize:])
pk.aT.Derive(&pk.rho, true)
}
// Derives a new Kyber.CPAPKE keypair from the given seed.
func NewKeyFromSeed(seed []byte) (*PublicKey, *PrivateKey) {
var pk PublicKey
var sk PrivateKey
var expandedSeed [64]byte
h := sha3.New512()
_, _ = h.Write(seed)
// This writes hash into expandedSeed. Yes, this is idiomatic Go.
_, _ = h.Read(expandedSeed[:])
copy(pk.rho[:], expandedSeed[:32])
sigma := expandedSeed[32:] // σ, the noise seed
pk.aT.Derive(&pk.rho, false) // Expand ρ to matrix A; we'll transpose later
var eh Vec
sk.sh.DeriveNoise(sigma, 0, Eta1) // Sample secret vector s
sk.sh.NTT()
sk.sh.Normalize()
eh.DeriveNoise(sigma, K, Eta1) // Sample blind e
eh.NTT()
// Next, we compute t = A s + e.
for i := 0; i < K; i++ {
// Note that coefficients of s are bounded by q and those of A
// are bounded by 4.5q and so their product is bounded by 2¹⁵q
// as required for multiplication.
PolyDotHat(&pk.th[i], &pk.aT[i], &sk.sh)
// A and s were not in Montgomery form, so the Montgomery
// multiplications in the inner product added a factor R⁻¹ which
// we'll cancel out now. This will also ensure the coefficients of
// t are bounded in absolute value by q.
pk.th[i].ToMont()
}
pk.th.Add(&pk.th, &eh) // bounded by 8q.
pk.th.Normalize()
pk.aT.Transpose()
return &pk, &sk
}
// Decrypts ciphertext ct meant for private key sk to plaintext pt.
func (sk *PrivateKey) DecryptTo(pt, ct []byte) {
var u Vec
var v, m common.Poly
u.Decompress(ct, DU)
v.Decompress(ct[K*compressedPolySize(DU):], DV)
// Compute m = v - <s, u>
u.NTT()
PolyDotHat(&m, &sk.sh, &u)
m.BarrettReduce()
m.InvNTT()
m.Sub(&v, &m)
m.Normalize()
// Compress polynomial m to original message
m.CompressMessageTo(pt)
}
// Encrypts message pt for the public key to ciphertext ct using randomness
// from seed.
//
// seed has to be of length SeedSize, pt of PlaintextSize and ct of
// CiphertextSize.
func (pk *PublicKey) EncryptTo(ct, pt, seed []byte) {
var rh, e1, u Vec
var e2, v, m common.Poly
// Sample r, e₁ and e₂ from B_η
rh.DeriveNoise(seed, 0, Eta1)
rh.NTT()
rh.BarrettReduce()
e1.DeriveNoise(seed, K, common.Eta2)
e2.DeriveNoise(seed, 2*K, common.Eta2)
// Next we compute u = Aᵀ r + e₁. First Aᵀ.
for i := 0; i < K; i++ {
// Note that coefficients of r are bounded by q and those of Aᵀ
// are bounded by 4.5q and so their product is bounded by 2¹⁵q
// as required for multiplication.
PolyDotHat(&u[i], &pk.aT[i], &rh)
}
u.BarrettReduce()
// Aᵀ and r were not in Montgomery form, so the Montgomery
// multiplications in the inner product added a factor R⁻¹ which
// the InvNTT cancels out.
u.InvNTT()
u.Add(&u, &e1) // u = Aᵀ r + e₁
// Next compute v = <t, r> + e₂ + Decompress_q(m, 1).
PolyDotHat(&v, &pk.th, &rh)
v.BarrettReduce()
v.InvNTT()
m.DecompressMessage(pt)
v.Add(&v, &m)
v.Add(&v, &e2) // v = <t, r> + e₂ + Decompress_q(m, 1)
// Pack ciphertext
u.Normalize()
v.Normalize()
u.CompressTo(ct, DU)
v.CompressTo(ct[K*compressedPolySize(DU):], DV)
}
// Returns whether sk equals other.
func (sk *PrivateKey) Equal(other *PrivateKey) bool {
ret := int16(0)
for i := 0; i < K; i++ {
for j := 0; j < common.N; j++ {
ret |= sk.sh[i][j] ^ other.sh[i][j]
}
}
return ret == 0
}
+85
View File
@@ -0,0 +1,85 @@
// Code generated from kyber512/internal/mat.go by gen.go
package internal
import (
"github.com/cloudflare/circl/pke/kyber/internal/common"
)
// A k by k matrix of polynomials.
type Mat [K]Vec
// Expands the given seed to the corresponding matrix A or its transpose Aᵀ.
func (m *Mat) Derive(seed *[32]byte, transpose bool) {
if !common.DeriveX4Available {
if transpose {
for i := 0; i < K; i++ {
for j := 0; j < K; j++ {
m[i][j].DeriveUniform(seed, uint8(i), uint8(j))
}
}
} else {
for i := 0; i < K; i++ {
for j := 0; j < K; j++ {
m[i][j].DeriveUniform(seed, uint8(j), uint8(i))
}
}
}
return
}
var ps [4]*common.Poly
var xs [4]uint8
var ys [4]uint8
x := uint8(0)
y := uint8(0)
for x != K {
idx := 0
for ; idx < 4; idx++ {
ps[idx] = &m[x][y]
if transpose {
xs[idx] = x
ys[idx] = y
} else {
xs[idx] = y
ys[idx] = x
}
y++
if y == K {
x++
y = 0
if x == K {
if idx == 0 {
// If there is just one left, then a plain DeriveUniform
// is quicker than the X4 variant.
ps[0].DeriveUniform(seed, xs[0], ys[0])
return
}
for idx++; idx < 4; idx++ {
ps[idx] = nil
}
break
}
}
}
common.PolyDeriveUniformX4(ps, seed, xs, ys)
}
}
// Transposes A in place.
func (m *Mat) Transpose() {
for i := 0; i < K-1; i++ {
for j := i + 1; j < K; j++ {
t := m[i][j]
m[i][j] = m[j][i]
m[j][i] = t
}
}
}
@@ -0,0 +1,21 @@
// Code generated from params.templ.go. DO NOT EDIT.
package internal
import (
"github.com/cloudflare/circl/pke/kyber/internal/common"
)
const (
K = 4
Eta1 = 2
DU = 11
DV = 5
PublicKeySize = 32 + K*common.PolySize
PrivateKeySize = K * common.PolySize
PlaintextSize = common.PlaintextSize
SeedSize = 32
CiphertextSize = 1568
)
+125
View File
@@ -0,0 +1,125 @@
// Code generated from kyber512/internal/vec.go by gen.go
package internal
import (
"github.com/cloudflare/circl/pke/kyber/internal/common"
)
// A vector of K polynomials
type Vec [K]common.Poly
// Samples v[i] from a centered binomial distribution with given η,
// seed and nonce+i.
//
// Essentially CBD_η(PRF(seed, nonce+i)) from the specification.
func (v *Vec) DeriveNoise(seed []byte, nonce uint8, eta int) {
for i := 0; i < K; i++ {
v[i].DeriveNoise(seed, nonce+uint8(i), eta)
}
}
// Sets p to the inner product of a and b using "pointwise" multiplication.
//
// See MulHat() and NTT() for a description of the multiplication.
// Assumes a and b are in Montgomery form. p will be in Montgomery form,
// and its coefficients will be bounded in absolute value by 2kq.
// If a and b are not in Montgomery form, then the action is the same
// as "pointwise" multiplication followed by multiplying by R⁻¹, the inverse
// of the Montgomery factor.
func PolyDotHat(p *common.Poly, a, b *Vec) {
var t common.Poly
*p = common.Poly{} // set p to zero
for i := 0; i < K; i++ {
t.MulHat(&a[i], &b[i])
p.Add(&t, p)
}
}
// Almost normalizes coefficients in-place.
//
// Ensures each coefficient is in {0, …, q}.
func (v *Vec) BarrettReduce() {
for i := 0; i < K; i++ {
v[i].BarrettReduce()
}
}
// Normalizes coefficients in-place.
//
// Ensures each coefficient is in {0, …, q-1}.
func (v *Vec) Normalize() {
for i := 0; i < K; i++ {
v[i].Normalize()
}
}
// Applies in-place inverse NTT(). See Poly.InvNTT() for assumptions.
func (v *Vec) InvNTT() {
for i := 0; i < K; i++ {
v[i].InvNTT()
}
}
// Applies in-place forward NTT(). See Poly.NTT() for assumptions.
func (v *Vec) NTT() {
for i := 0; i < K; i++ {
v[i].NTT()
}
}
// Sets v to a + b.
func (v *Vec) Add(a, b *Vec) {
for i := 0; i < K; i++ {
v[i].Add(&a[i], &b[i])
}
}
// Packs v into buf, which must be of length K*PolySize.
func (v *Vec) Pack(buf []byte) {
for i := 0; i < K; i++ {
v[i].Pack(buf[common.PolySize*i:])
}
}
// Unpacks v from buf which must be of length K*PolySize.
func (v *Vec) Unpack(buf []byte) {
for i := 0; i < K; i++ {
v[i].Unpack(buf[common.PolySize*i:])
}
}
// Writes Compress_q(v, d) to m.
//
// Assumes v is normalized and d is in {3, 4, 5, 10, 11}.
func (v *Vec) CompressTo(m []byte, d int) {
size := compressedPolySize(d)
for i := 0; i < K; i++ {
v[i].CompressTo(m[size*i:], d)
}
}
// Set v to Decompress_q(m, 1).
//
// Assumes d is in {3, 4, 5, 10, 11}. v will be normalized.
func (v *Vec) Decompress(m []byte, d int) {
size := compressedPolySize(d)
for i := 0; i < K; i++ {
v[i].Decompress(m[size*i:], d)
}
}
// ⌈(256 d)/8⌉
func compressedPolySize(d int) int {
switch d {
case 4:
return 128
case 5:
return 160
case 10:
return 320
case 11:
return 352
}
panic("unsupported d")
}
+175
View File
@@ -0,0 +1,175 @@
// Code generated from pkg.templ.go. DO NOT EDIT.
// kyber1024 implements the IND-CPA-secure Public Key Encryption
// scheme Kyber1024.CPAPKE as submitted to round 3 of the NIST PQC competition
// and described in
//
// https://pq-crystals.org/kyber/data/kyber-specification-round3.pdf
package kyber1024
import (
cryptoRand "crypto/rand"
"io"
"github.com/cloudflare/circl/kem"
"github.com/cloudflare/circl/pke/kyber/kyber1024/internal"
)
const (
// Size of seed for NewKeyFromSeed
KeySeedSize = internal.SeedSize
// Size of seed for EncryptTo
EncryptionSeedSize = internal.SeedSize
// Size of a packed PublicKey
PublicKeySize = internal.PublicKeySize
// Size of a packed PrivateKey
PrivateKeySize = internal.PrivateKeySize
// Size of a ciphertext
CiphertextSize = internal.CiphertextSize
// Size of a plaintext
PlaintextSize = internal.PlaintextSize
)
// PublicKey is the type of Kyber1024.CPAPKE public key
type PublicKey internal.PublicKey
// PrivateKey is the type of Kyber1024.CPAPKE private key
type PrivateKey internal.PrivateKey
// GenerateKey generates a public/private key pair using entropy from rand.
// If rand is nil, crypto/rand.Reader will be used.
func GenerateKey(rand io.Reader) (*PublicKey, *PrivateKey, error) {
var seed [KeySeedSize]byte
if rand == nil {
rand = cryptoRand.Reader
}
_, err := io.ReadFull(rand, seed[:])
if err != nil {
return nil, nil, err
}
pk, sk := internal.NewKeyFromSeed(seed[:])
return (*PublicKey)(pk), (*PrivateKey)(sk), nil
}
// NewKeyFromSeed derives a public/private key pair using the given seed.
//
// Note: does not include the domain separation of ML-KEM (line 1, algorithm 13
// of FIPS 203). For that use NewKeyFromSeedMLKEM().
//
// Panics if seed is not of length KeySeedSize.
func NewKeyFromSeed(seed []byte) (*PublicKey, *PrivateKey) {
if len(seed) != KeySeedSize {
panic("seed must be of length KeySeedSize")
}
pk, sk := internal.NewKeyFromSeed(seed)
return (*PublicKey)(pk), (*PrivateKey)(sk)
}
// NewKeyFromSeedMLKEM derives a public/private key pair using the given seed
// using the domain separation of ML-KEM.
//
// Panics if seed is not of length KeySeedSize.
func NewKeyFromSeedMLKEM(seed []byte) (*PublicKey, *PrivateKey) {
if len(seed) != KeySeedSize {
panic("seed must be of length KeySeedSize")
}
var seed2 [33]byte
copy(seed2[:32], seed)
seed2[32] = byte(internal.K)
pk, sk := internal.NewKeyFromSeed(seed2[:])
return (*PublicKey)(pk), (*PrivateKey)(sk)
}
// EncryptTo encrypts message pt for the public key and writes the ciphertext
// to ct using randomness from seed.
//
// This function panics if the lengths of pt, seed, and ct are not
// PlaintextSize, EncryptionSeedSize, and CiphertextSize respectively.
func (pk *PublicKey) EncryptTo(ct []byte, pt []byte, seed []byte) {
if len(pt) != PlaintextSize {
panic("pt must be of length PlaintextSize")
}
if len(ct) != CiphertextSize {
panic("ct must be of length CiphertextSize")
}
if len(seed) != EncryptionSeedSize {
panic("seed must be of length EncryptionSeedSize")
}
(*internal.PublicKey)(pk).EncryptTo(ct, pt, seed)
}
// DecryptTo decrypts message ct for the private key and writes the
// plaintext to pt.
//
// This function panics if the lengths of ct and pt are not
// CiphertextSize and PlaintextSize respectively.
func (sk *PrivateKey) DecryptTo(pt []byte, ct []byte) {
if len(pt) != PlaintextSize {
panic("pt must be of length PlaintextSize")
}
if len(ct) != CiphertextSize {
panic("ct must be of length CiphertextSize")
}
(*internal.PrivateKey)(sk).DecryptTo(pt, ct)
}
// Packs pk into the given buffer.
//
// Panics if buf is not of length PublicKeySize.
func (pk *PublicKey) Pack(buf []byte) {
if len(buf) != PublicKeySize {
panic("buf must be of size PublicKeySize")
}
(*internal.PublicKey)(pk).Pack(buf)
}
// Packs sk into the given buffer.
//
// Panics if buf is not of length PrivateKeySize.
func (sk *PrivateKey) Pack(buf []byte) {
if len(buf) != PrivateKeySize {
panic("buf must be of size PrivateKeySize")
}
(*internal.PrivateKey)(sk).Pack(buf)
}
// Unpacks pk from the given buffer.
//
// Panics if buf is not of length PublicKeySize.
func (pk *PublicKey) Unpack(buf []byte) {
if len(buf) != PublicKeySize {
panic("buf must be of size PublicKeySize")
}
(*internal.PublicKey)(pk).Unpack(buf)
}
// Unpacks pk from the given buffer.
//
// Returns an error if the buffer is not of the right size, or the public
// key is not normalized.
func (pk *PublicKey) UnpackMLKEM(buf []byte) error {
if len(buf) != PublicKeySize {
return kem.ErrPubKeySize
}
return (*internal.PublicKey)(pk).UnpackMLKEM(buf)
}
// Unpacks sk from the given buffer.
//
// Panics if buf is not of length PrivateKeySize.
func (sk *PrivateKey) Unpack(buf []byte) {
if len(buf) != PrivateKeySize {
panic("buf must be of size PrivateKeySize")
}
(*internal.PrivateKey)(sk).Unpack(buf)
}
// Returns whether the two private keys are equal.
func (sk *PrivateKey) Equal(other *PrivateKey) bool {
return (*internal.PrivateKey)(sk).Equal((*internal.PrivateKey)(other))
}
@@ -0,0 +1,192 @@
// Code generated from kyber512/internal/cpapke.go by gen.go
package internal
import (
"bytes"
"github.com/cloudflare/circl/internal/sha3"
"github.com/cloudflare/circl/kem"
"github.com/cloudflare/circl/pke/kyber/internal/common"
)
// A Kyber.CPAPKE private key.
type PrivateKey struct {
sh Vec // NTT(s), normalized
}
// A Kyber.CPAPKE public key.
type PublicKey struct {
rho [32]byte // ρ, the seed for the matrix A
th Vec // NTT(t), normalized
// cached values
aT Mat // the matrix Aᵀ
}
// Packs the private key to buf.
func (sk *PrivateKey) Pack(buf []byte) {
sk.sh.Pack(buf)
}
// Unpacks the private key from buf.
func (sk *PrivateKey) Unpack(buf []byte) {
sk.sh.Unpack(buf)
sk.sh.Normalize()
}
// Packs the public key to buf.
func (pk *PublicKey) Pack(buf []byte) {
pk.th.Pack(buf)
copy(buf[K*common.PolySize:], pk.rho[:])
}
// Unpacks the public key from buf. Checks if the public key is normalized.
func (pk *PublicKey) UnpackMLKEM(buf []byte) error {
pk.Unpack(buf)
// FIPS 203 §7.2 "encapsulation key check" (2).
var buf2 [K * common.PolySize]byte
pk.th.Pack(buf2[:])
if !bytes.Equal(buf[:len(buf2)], buf2[:]) {
return kem.ErrPubKey
}
return nil
}
// Unpacks the public key from buf.
func (pk *PublicKey) Unpack(buf []byte) {
pk.th.Unpack(buf)
pk.th.Normalize()
copy(pk.rho[:], buf[K*common.PolySize:])
pk.aT.Derive(&pk.rho, true)
}
// Derives a new Kyber.CPAPKE keypair from the given seed.
func NewKeyFromSeed(seed []byte) (*PublicKey, *PrivateKey) {
var pk PublicKey
var sk PrivateKey
var expandedSeed [64]byte
h := sha3.New512()
_, _ = h.Write(seed)
// This writes hash into expandedSeed. Yes, this is idiomatic Go.
_, _ = h.Read(expandedSeed[:])
copy(pk.rho[:], expandedSeed[:32])
sigma := expandedSeed[32:] // σ, the noise seed
pk.aT.Derive(&pk.rho, false) // Expand ρ to matrix A; we'll transpose later
var eh Vec
sk.sh.DeriveNoise(sigma, 0, Eta1) // Sample secret vector s
sk.sh.NTT()
sk.sh.Normalize()
eh.DeriveNoise(sigma, K, Eta1) // Sample blind e
eh.NTT()
// Next, we compute t = A s + e.
for i := 0; i < K; i++ {
// Note that coefficients of s are bounded by q and those of A
// are bounded by 4.5q and so their product is bounded by 2¹⁵q
// as required for multiplication.
PolyDotHat(&pk.th[i], &pk.aT[i], &sk.sh)
// A and s were not in Montgomery form, so the Montgomery
// multiplications in the inner product added a factor R⁻¹ which
// we'll cancel out now. This will also ensure the coefficients of
// t are bounded in absolute value by q.
pk.th[i].ToMont()
}
pk.th.Add(&pk.th, &eh) // bounded by 8q.
pk.th.Normalize()
pk.aT.Transpose()
return &pk, &sk
}
// Decrypts ciphertext ct meant for private key sk to plaintext pt.
func (sk *PrivateKey) DecryptTo(pt, ct []byte) {
var u Vec
var v, m common.Poly
u.Decompress(ct, DU)
v.Decompress(ct[K*compressedPolySize(DU):], DV)
// Compute m = v - <s, u>
u.NTT()
PolyDotHat(&m, &sk.sh, &u)
m.BarrettReduce()
m.InvNTT()
m.Sub(&v, &m)
m.Normalize()
// Compress polynomial m to original message
m.CompressMessageTo(pt)
}
// Encrypts message pt for the public key to ciphertext ct using randomness
// from seed.
//
// seed has to be of length SeedSize, pt of PlaintextSize and ct of
// CiphertextSize.
func (pk *PublicKey) EncryptTo(ct, pt, seed []byte) {
var rh, e1, u Vec
var e2, v, m common.Poly
// Sample r, e₁ and e₂ from B_η
rh.DeriveNoise(seed, 0, Eta1)
rh.NTT()
rh.BarrettReduce()
e1.DeriveNoise(seed, K, common.Eta2)
e2.DeriveNoise(seed, 2*K, common.Eta2)
// Next we compute u = Aᵀ r + e₁. First Aᵀ.
for i := 0; i < K; i++ {
// Note that coefficients of r are bounded by q and those of Aᵀ
// are bounded by 4.5q and so their product is bounded by 2¹⁵q
// as required for multiplication.
PolyDotHat(&u[i], &pk.aT[i], &rh)
}
u.BarrettReduce()
// Aᵀ and r were not in Montgomery form, so the Montgomery
// multiplications in the inner product added a factor R⁻¹ which
// the InvNTT cancels out.
u.InvNTT()
u.Add(&u, &e1) // u = Aᵀ r + e₁
// Next compute v = <t, r> + e₂ + Decompress_q(m, 1).
PolyDotHat(&v, &pk.th, &rh)
v.BarrettReduce()
v.InvNTT()
m.DecompressMessage(pt)
v.Add(&v, &m)
v.Add(&v, &e2) // v = <t, r> + e₂ + Decompress_q(m, 1)
// Pack ciphertext
u.Normalize()
v.Normalize()
u.CompressTo(ct, DU)
v.CompressTo(ct[K*compressedPolySize(DU):], DV)
}
// Returns whether sk equals other.
func (sk *PrivateKey) Equal(other *PrivateKey) bool {
ret := int16(0)
for i := 0; i < K; i++ {
for j := 0; j < common.N; j++ {
ret |= sk.sh[i][j] ^ other.sh[i][j]
}
}
return ret == 0
}
+85
View File
@@ -0,0 +1,85 @@
// Code generated from kyber512/internal/mat.go by gen.go
package internal
import (
"github.com/cloudflare/circl/pke/kyber/internal/common"
)
// A k by k matrix of polynomials.
type Mat [K]Vec
// Expands the given seed to the corresponding matrix A or its transpose Aᵀ.
func (m *Mat) Derive(seed *[32]byte, transpose bool) {
if !common.DeriveX4Available {
if transpose {
for i := 0; i < K; i++ {
for j := 0; j < K; j++ {
m[i][j].DeriveUniform(seed, uint8(i), uint8(j))
}
}
} else {
for i := 0; i < K; i++ {
for j := 0; j < K; j++ {
m[i][j].DeriveUniform(seed, uint8(j), uint8(i))
}
}
}
return
}
var ps [4]*common.Poly
var xs [4]uint8
var ys [4]uint8
x := uint8(0)
y := uint8(0)
for x != K {
idx := 0
for ; idx < 4; idx++ {
ps[idx] = &m[x][y]
if transpose {
xs[idx] = x
ys[idx] = y
} else {
xs[idx] = y
ys[idx] = x
}
y++
if y == K {
x++
y = 0
if x == K {
if idx == 0 {
// If there is just one left, then a plain DeriveUniform
// is quicker than the X4 variant.
ps[0].DeriveUniform(seed, xs[0], ys[0])
return
}
for idx++; idx < 4; idx++ {
ps[idx] = nil
}
break
}
}
}
common.PolyDeriveUniformX4(ps, seed, xs, ys)
}
}
// Transposes A in place.
func (m *Mat) Transpose() {
for i := 0; i < K-1; i++ {
for j := i + 1; j < K; j++ {
t := m[i][j]
m[i][j] = m[j][i]
m[j][i] = t
}
}
}
@@ -0,0 +1,21 @@
// Code generated from params.templ.go. DO NOT EDIT.
package internal
import (
"github.com/cloudflare/circl/pke/kyber/internal/common"
)
const (
K = 3
Eta1 = 2
DU = 10
DV = 4
PublicKeySize = 32 + K*common.PolySize
PrivateKeySize = K * common.PolySize
PlaintextSize = common.PlaintextSize
SeedSize = 32
CiphertextSize = 1088
)
+125
View File
@@ -0,0 +1,125 @@
// Code generated from kyber512/internal/vec.go by gen.go
package internal
import (
"github.com/cloudflare/circl/pke/kyber/internal/common"
)
// A vector of K polynomials
type Vec [K]common.Poly
// Samples v[i] from a centered binomial distribution with given η,
// seed and nonce+i.
//
// Essentially CBD_η(PRF(seed, nonce+i)) from the specification.
func (v *Vec) DeriveNoise(seed []byte, nonce uint8, eta int) {
for i := 0; i < K; i++ {
v[i].DeriveNoise(seed, nonce+uint8(i), eta)
}
}
// Sets p to the inner product of a and b using "pointwise" multiplication.
//
// See MulHat() and NTT() for a description of the multiplication.
// Assumes a and b are in Montgomery form. p will be in Montgomery form,
// and its coefficients will be bounded in absolute value by 2kq.
// If a and b are not in Montgomery form, then the action is the same
// as "pointwise" multiplication followed by multiplying by R⁻¹, the inverse
// of the Montgomery factor.
func PolyDotHat(p *common.Poly, a, b *Vec) {
var t common.Poly
*p = common.Poly{} // set p to zero
for i := 0; i < K; i++ {
t.MulHat(&a[i], &b[i])
p.Add(&t, p)
}
}
// Almost normalizes coefficients in-place.
//
// Ensures each coefficient is in {0, …, q}.
func (v *Vec) BarrettReduce() {
for i := 0; i < K; i++ {
v[i].BarrettReduce()
}
}
// Normalizes coefficients in-place.
//
// Ensures each coefficient is in {0, …, q-1}.
func (v *Vec) Normalize() {
for i := 0; i < K; i++ {
v[i].Normalize()
}
}
// Applies in-place inverse NTT(). See Poly.InvNTT() for assumptions.
func (v *Vec) InvNTT() {
for i := 0; i < K; i++ {
v[i].InvNTT()
}
}
// Applies in-place forward NTT(). See Poly.NTT() for assumptions.
func (v *Vec) NTT() {
for i := 0; i < K; i++ {
v[i].NTT()
}
}
// Sets v to a + b.
func (v *Vec) Add(a, b *Vec) {
for i := 0; i < K; i++ {
v[i].Add(&a[i], &b[i])
}
}
// Packs v into buf, which must be of length K*PolySize.
func (v *Vec) Pack(buf []byte) {
for i := 0; i < K; i++ {
v[i].Pack(buf[common.PolySize*i:])
}
}
// Unpacks v from buf which must be of length K*PolySize.
func (v *Vec) Unpack(buf []byte) {
for i := 0; i < K; i++ {
v[i].Unpack(buf[common.PolySize*i:])
}
}
// Writes Compress_q(v, d) to m.
//
// Assumes v is normalized and d is in {3, 4, 5, 10, 11}.
func (v *Vec) CompressTo(m []byte, d int) {
size := compressedPolySize(d)
for i := 0; i < K; i++ {
v[i].CompressTo(m[size*i:], d)
}
}
// Set v to Decompress_q(m, 1).
//
// Assumes d is in {3, 4, 5, 10, 11}. v will be normalized.
func (v *Vec) Decompress(m []byte, d int) {
size := compressedPolySize(d)
for i := 0; i < K; i++ {
v[i].Decompress(m[size*i:], d)
}
}
// ⌈(256 d)/8⌉
func compressedPolySize(d int) int {
switch d {
case 4:
return 128
case 5:
return 160
case 10:
return 320
case 11:
return 352
}
panic("unsupported d")
}
+175
View File
@@ -0,0 +1,175 @@
// Code generated from pkg.templ.go. DO NOT EDIT.
// kyber768 implements the IND-CPA-secure Public Key Encryption
// scheme Kyber768.CPAPKE as submitted to round 3 of the NIST PQC competition
// and described in
//
// https://pq-crystals.org/kyber/data/kyber-specification-round3.pdf
package kyber768
import (
cryptoRand "crypto/rand"
"io"
"github.com/cloudflare/circl/kem"
"github.com/cloudflare/circl/pke/kyber/kyber768/internal"
)
const (
// Size of seed for NewKeyFromSeed
KeySeedSize = internal.SeedSize
// Size of seed for EncryptTo
EncryptionSeedSize = internal.SeedSize
// Size of a packed PublicKey
PublicKeySize = internal.PublicKeySize
// Size of a packed PrivateKey
PrivateKeySize = internal.PrivateKeySize
// Size of a ciphertext
CiphertextSize = internal.CiphertextSize
// Size of a plaintext
PlaintextSize = internal.PlaintextSize
)
// PublicKey is the type of Kyber768.CPAPKE public key
type PublicKey internal.PublicKey
// PrivateKey is the type of Kyber768.CPAPKE private key
type PrivateKey internal.PrivateKey
// GenerateKey generates a public/private key pair using entropy from rand.
// If rand is nil, crypto/rand.Reader will be used.
func GenerateKey(rand io.Reader) (*PublicKey, *PrivateKey, error) {
var seed [KeySeedSize]byte
if rand == nil {
rand = cryptoRand.Reader
}
_, err := io.ReadFull(rand, seed[:])
if err != nil {
return nil, nil, err
}
pk, sk := internal.NewKeyFromSeed(seed[:])
return (*PublicKey)(pk), (*PrivateKey)(sk), nil
}
// NewKeyFromSeed derives a public/private key pair using the given seed.
//
// Note: does not include the domain separation of ML-KEM (line 1, algorithm 13
// of FIPS 203). For that use NewKeyFromSeedMLKEM().
//
// Panics if seed is not of length KeySeedSize.
func NewKeyFromSeed(seed []byte) (*PublicKey, *PrivateKey) {
if len(seed) != KeySeedSize {
panic("seed must be of length KeySeedSize")
}
pk, sk := internal.NewKeyFromSeed(seed)
return (*PublicKey)(pk), (*PrivateKey)(sk)
}
// NewKeyFromSeedMLKEM derives a public/private key pair using the given seed
// using the domain separation of ML-KEM.
//
// Panics if seed is not of length KeySeedSize.
func NewKeyFromSeedMLKEM(seed []byte) (*PublicKey, *PrivateKey) {
if len(seed) != KeySeedSize {
panic("seed must be of length KeySeedSize")
}
var seed2 [33]byte
copy(seed2[:32], seed)
seed2[32] = byte(internal.K)
pk, sk := internal.NewKeyFromSeed(seed2[:])
return (*PublicKey)(pk), (*PrivateKey)(sk)
}
// EncryptTo encrypts message pt for the public key and writes the ciphertext
// to ct using randomness from seed.
//
// This function panics if the lengths of pt, seed, and ct are not
// PlaintextSize, EncryptionSeedSize, and CiphertextSize respectively.
func (pk *PublicKey) EncryptTo(ct []byte, pt []byte, seed []byte) {
if len(pt) != PlaintextSize {
panic("pt must be of length PlaintextSize")
}
if len(ct) != CiphertextSize {
panic("ct must be of length CiphertextSize")
}
if len(seed) != EncryptionSeedSize {
panic("seed must be of length EncryptionSeedSize")
}
(*internal.PublicKey)(pk).EncryptTo(ct, pt, seed)
}
// DecryptTo decrypts message ct for the private key and writes the
// plaintext to pt.
//
// This function panics if the lengths of ct and pt are not
// CiphertextSize and PlaintextSize respectively.
func (sk *PrivateKey) DecryptTo(pt []byte, ct []byte) {
if len(pt) != PlaintextSize {
panic("pt must be of length PlaintextSize")
}
if len(ct) != CiphertextSize {
panic("ct must be of length CiphertextSize")
}
(*internal.PrivateKey)(sk).DecryptTo(pt, ct)
}
// Packs pk into the given buffer.
//
// Panics if buf is not of length PublicKeySize.
func (pk *PublicKey) Pack(buf []byte) {
if len(buf) != PublicKeySize {
panic("buf must be of size PublicKeySize")
}
(*internal.PublicKey)(pk).Pack(buf)
}
// Packs sk into the given buffer.
//
// Panics if buf is not of length PrivateKeySize.
func (sk *PrivateKey) Pack(buf []byte) {
if len(buf) != PrivateKeySize {
panic("buf must be of size PrivateKeySize")
}
(*internal.PrivateKey)(sk).Pack(buf)
}
// Unpacks pk from the given buffer.
//
// Panics if buf is not of length PublicKeySize.
func (pk *PublicKey) Unpack(buf []byte) {
if len(buf) != PublicKeySize {
panic("buf must be of size PublicKeySize")
}
(*internal.PublicKey)(pk).Unpack(buf)
}
// Unpacks pk from the given buffer.
//
// Returns an error if the buffer is not of the right size, or the public
// key is not normalized.
func (pk *PublicKey) UnpackMLKEM(buf []byte) error {
if len(buf) != PublicKeySize {
return kem.ErrPubKeySize
}
return (*internal.PublicKey)(pk).UnpackMLKEM(buf)
}
// Unpacks sk from the given buffer.
//
// Panics if buf is not of length PrivateKeySize.
func (sk *PrivateKey) Unpack(buf []byte) {
if len(buf) != PrivateKeySize {
panic("buf must be of size PrivateKeySize")
}
(*internal.PrivateKey)(sk).Unpack(buf)
}
// Returns whether the two private keys are equal.
func (sk *PrivateKey) Equal(other *PrivateKey) bool {
return (*internal.PrivateKey)(sk).Equal((*internal.PrivateKey)(other))
}
+8 -1
View File
@@ -1,7 +1,7 @@
// Package ed25519 implements Ed25519 signature scheme as described in RFC-8032.
//
// This package provides optimized implementations of the three signature
// variants and maintaining closer compatibility with crypto/ed25519.
// variants and maintaining almost full compatibility with crypto/ed25519.
//
// | Scheme Name | Sign Function | Verification | Context |
// |-------------|-------------------|---------------|-------------------|
@@ -28,6 +28,13 @@
// operations with the same key more efficient. This package refers to the
// RFC-8032 private key as the “seed”.
//
// Public keys are decoded following RFC 8032 strictly: an encoding whose
// y-coordinate is not reduced modulo 2^255-19, or whose x-coordinate is zero
// with the sign bit set, is rejected. crypto/ed25519 accepts those
// non-canonical encodings instead, so for a small, fixed set of malformed
// public keys the verification functions here return false where
// crypto/ed25519 returns true.
//
// References
//
// - RFC-8032: https://rfc-editor.org/rfc/rfc8032.txt
+4 -2
View File
@@ -69,7 +69,8 @@ func (*scheme) DeriveKey(seed []byte) (sign.PublicKey, sign.PrivateKey) {
}
func (*scheme) UnmarshalBinaryPublicKey(buf []byte) (sign.PublicKey, error) {
if len(buf) < PublicKeySize {
// NOTE Old version of CIRCL accepted trailing data.
if len(buf) != PublicKeySize {
return nil, sign.ErrPubKeySize
}
pub := make(PublicKey, PublicKeySize)
@@ -78,7 +79,8 @@ func (*scheme) UnmarshalBinaryPublicKey(buf []byte) (sign.PublicKey, error) {
}
func (*scheme) UnmarshalBinaryPrivateKey(buf []byte) (sign.PrivateKey, error) {
if len(buf) < PrivateKeySize {
// NOTE Old version of CIRCL accepted trailing data.
if len(buf) != PrivateKeySize {
return nil, sign.ErrPrivKeySize
}
priv := make(PrivateKey, PrivateKeySize)
+6
View File
@@ -345,6 +345,8 @@ func verify(public PublicKey, message, signature, ctx []byte, preHash bool) bool
// The opts.HashFunc() must return zero, this can be achieved by passing
// crypto.Hash(0) as the value for opts.
// Use a SignerOptions struct to pass a context string for signing.
//
// Uses the default RFC 8032 cofactor-clearing verification.
func VerifyAny(public PublicKey, message, signature []byte, opts crypto.SignerOpts) bool {
var ctx string
var scheme SchemeID
@@ -367,6 +369,8 @@ func VerifyAny(public PublicKey, message, signature []byte, opts crypto.SignerOp
// signature, or when the public key cannot be decoded.
// This function supports the signature variant defined in RFC-8032: Ed448,
// also known as the pure version of EdDSA.
//
// Uses the default RFC 8032 cofactor-clearing verification.
func Verify(public PublicKey, message, signature []byte, ctx string) bool {
return verify(public, message, signature, []byte(ctx), false)
}
@@ -377,6 +381,8 @@ func Verify(public PublicKey, message, signature []byte, ctx string) bool {
// meaning it internally hashes the message using SHAKE-256.
// Context could be passed to this function, which length should be no more than
// 255. It can be empty.
//
// Uses the default RFC 8032 cofactor-clearing verification.
func VerifyPh(public PublicKey, message, signature []byte, ctx string) bool {
return verify(public, message, signature, []byte(ctx), true)
}
+4 -2
View File
@@ -69,7 +69,8 @@ func (*scheme) DeriveKey(seed []byte) (sign.PublicKey, sign.PrivateKey) {
}
func (*scheme) UnmarshalBinaryPublicKey(buf []byte) (sign.PublicKey, error) {
if len(buf) < PublicKeySize {
// NOTE Old version of CIRCL accepted trailing data.
if len(buf) != PublicKeySize {
return nil, sign.ErrPubKeySize
}
pub := make(PublicKey, PublicKeySize)
@@ -78,7 +79,8 @@ func (*scheme) UnmarshalBinaryPublicKey(buf []byte) (sign.PublicKey, error) {
}
func (*scheme) UnmarshalBinaryPrivateKey(buf []byte) (sign.PrivateKey, error) {
if len(buf) < PrivateKeySize {
// NOTE Old version of CIRCL accepted trailing data.
if len(buf) != PrivateKeySize {
return nil, sign.ErrPrivKeySize
}
priv := make(PrivateKey, PrivateKeySize)
+162
View File
@@ -0,0 +1,162 @@
//go:build amd64 && !purego
// +build amd64,!purego
package dilithium
import (
"golang.org/x/sys/cpu"
)
// Execute an in-place forward NTT on as.
//
// Assumes the coefficients are in Montgomery representation and bounded
// by 2*Q. The resulting coefficients are again in Montgomery representation,
// but are only bounded bt 18*Q.
func (p *Poly) NTT() {
if cpu.X86.HasAVX2 {
nttAVX2(
(*[N]uint32)(p),
)
} else {
p.nttGeneric()
}
}
// Execute an in-place inverse NTT and multiply by Montgomery factor R
//
// Assumes the coefficients are in Montgomery representation and bounded
// by 2*Q. The resulting coefficients are again in Montgomery representation
// and bounded by 2*Q.
func (p *Poly) InvNTT() {
if cpu.X86.HasAVX2 {
invNttAVX2(
(*[N]uint32)(p),
)
} else {
p.invNttGeneric()
}
}
// Sets p to the polynomial whose coefficients are the pointwise multiplication
// of those of a and b. The coefficients of p are bounded by 2q.
//
// Assumes a and b are in Montgomery form and that the pointwise product
// of each coefficient is below 2³² q.
func (p *Poly) MulHat(a, b *Poly) {
if cpu.X86.HasAVX2 {
mulHatAVX2(
(*[N]uint32)(p),
(*[N]uint32)(a),
(*[N]uint32)(b),
)
} else {
p.mulHatGeneric(a, b)
}
}
// Sets p to a + b. Does not normalize polynomials.
func (p *Poly) Add(a, b *Poly) {
if cpu.X86.HasAVX2 {
addAVX2(
(*[N]uint32)(p),
(*[N]uint32)(a),
(*[N]uint32)(b),
)
} else {
p.addGeneric(a, b)
}
}
// Sets p to a - b.
//
// Warning: assumes coefficients of b are less than 2q.
// Sets p to a + b. Does not normalize polynomials.
func (p *Poly) Sub(a, b *Poly) {
if cpu.X86.HasAVX2 {
subAVX2(
(*[N]uint32)(p),
(*[N]uint32)(a),
(*[N]uint32)(b),
)
} else {
p.subGeneric(a, b)
}
}
// Writes p whose coefficients are in [0, 16) to buf, which must be of
// length N/2.
func (p *Poly) PackLe16(buf []byte) {
if cpu.X86.HasAVX2 {
if len(buf) < PolyLe16Size {
panic("buf too small")
}
packLe16AVX2(
(*[N]uint32)(p),
&buf[0],
)
} else {
p.packLe16Generic(buf)
}
}
// Reduces each of the coefficients to <2q.
func (p *Poly) ReduceLe2Q() {
if cpu.X86.HasAVX2 {
reduceLe2QAVX2((*[N]uint32)(p))
} else {
p.reduceLe2QGeneric()
}
}
// Reduce each of the coefficients to <q.
func (p *Poly) Normalize() {
if cpu.X86.HasAVX2 {
p.ReduceLe2Q()
p.NormalizeAssumingLe2Q()
} else {
p.normalizeGeneric()
}
}
// Normalize the coefficients in this polynomial assuming they are already
// bounded by 2q.
func (p *Poly) NormalizeAssumingLe2Q() {
if cpu.X86.HasAVX2 {
le2qModQAVX2((*[N]uint32)(p))
} else {
p.normalizeAssumingLe2QGeneric()
}
}
// Checks whether the "supnorm" (see sec 2.1 of the spec) of p is equal
// or greater than the given bound.
//
// Requires the coefficients of p to be normalized.
func (p *Poly) Exceeds(bound uint32) bool {
if cpu.X86.HasAVX2 {
return exceedsAVX2((*[N]uint32)(p), bound) == 1
}
return p.exceedsGeneric(bound)
}
// Sets p to 2ᵈ q without reducing.
//
// So it requires the coefficients of p to be less than 2³²⁻ᴰ.
func (p *Poly) MulBy2toD(q *Poly) {
if cpu.X86.HasAVX2 {
mulBy2toDAVX2(
(*[N]uint32)(p),
(*[N]uint32)(q),
)
} else {
p.mulBy2toDGeneric(q)
}
}
// Splits p into p1 and p0 such that [i]p1 * 2ᴰ + [i]p0 = [i]p
// with -2ᴰ⁻¹ < [i]p0 ≤ 2ᴰ⁻¹. Returns p0 + Q and p1.
//
// Requires the coefficients of p to be normalized.
func (p *Poly) Power2Round(p0PlusQ, p1 *Poly) {
p.power2RoundGeneric(p0PlusQ, p1)
}
File diff suppressed because it is too large. Load diff
+106
View File
@@ -0,0 +1,106 @@
//go:build arm64 && !purego
// +build arm64,!purego
package dilithium
// Execute an in-place forward NTT on as.
//
// Assumes the coefficients are in Montgomery representation and bounded
// by 2*Q. The resulting coefficients are again in Montgomery representation,
// but are only bounded bt 18*Q.
func (p *Poly) NTT() {
p.nttGeneric()
}
// Execute an in-place inverse NTT and multiply by Montgomery factor R
//
// Assumes the coefficients are in Montgomery representation and bounded
// by 2*Q. The resulting coefficients are again in Montgomery representation
// and bounded by 2*Q.
func (p *Poly) InvNTT() {
p.invNttGeneric()
}
// Sets p to the polynomial whose coefficients are the pointwise multiplication
// of those of a and b. The coefficients of p are bounded by 2q.
//
// Assumes a and b are in Montgomery form and that the pointwise product
// of each coefficient is below 2³² q.
func (p *Poly) MulHat(a, b *Poly) {
p.mulHatGeneric(a, b)
}
// Sets p to a + b. Does not normalize polynomials.
func (p *Poly) Add(a, b *Poly) {
polyAddARM64(p, a, b)
}
// Sets p to a - b.
//
// Warning: assumes coefficients of b are less than 2q.
// Sets p to a + b. Does not normalize polynomials.
func (p *Poly) Sub(a, b *Poly) {
polySubARM64(p, a, b)
}
// Writes p whose coefficients are in [0, 16) to buf, which must be of
// length N/2.
func (p *Poly) PackLe16(buf []byte) {
// early bounds so we don't have to in assembly code
// compiler may inline this func, so it may remove the bounds check
_ = buf[PolyLe16Size-1]
polyPackLe16ARM64(p, &buf[0])
}
// Reduces each of the coefficients to <2q.
func (p *Poly) ReduceLe2Q() {
p.reduceLe2QGeneric()
}
// Reduce each of the coefficients to <q.
func (p *Poly) Normalize() {
p.normalizeGeneric()
}
// Normalize the coefficients in this polynomial assuming they are already
// bounded by 2q.
func (p *Poly) NormalizeAssumingLe2Q() {
p.normalizeAssumingLe2QGeneric()
}
// Checks whether the "supnorm" (see sec 2.1 of the spec) of p is equal
// or greater than the given bound.
//
// Requires the coefficients of p to be normalized.
func (p *Poly) Exceeds(bound uint32) bool {
return p.exceedsGeneric(bound)
}
// Sets p to 2ᵈ q without reducing.
//
// So it requires the coefficients of p to be less than 2³²⁻ᴰ.
func (p *Poly) MulBy2toD(q *Poly) {
polyMulBy2toDARM64(p, q)
}
// Splits p into p1 and p0 such that [i]p1 * 2ᴰ + [i]p0 = [i]p
// with -2ᴰ⁻¹ < [i]p0 ≤ 2ᴰ⁻¹. Returns p0 + Q and p1.
//
// Requires the coefficients of p to be normalized.
func (p *Poly) Power2Round(p0PlusQ, p1 *Poly) {
// implementation in assembly follows
p.power2RoundGeneric(p0PlusQ, p1)
}
//go:noescape
func polyAddARM64(p, a, b *Poly)
//go:noescape
func polyPackLe16ARM64(p *Poly, buf *byte)
//go:noescape
func polyMulBy2toDARM64(p, q *Poly)
//go:noescape
func polySubARM64(p, a, b *Poly)
+126
View File
@@ -0,0 +1,126 @@
//go:build arm64 && !purego
#include "go_asm.h"
#include "textflag.h"
// func polyAddARM64(p, a, b *Poly)
TEXT ·polyAddARM64(SB), NOSPLIT|NOFRAME, $0-24
MOVD p+0(FP), R0
MOVD a+8(FP), R1
MOVD b+16(FP), R2
MOVW $(const_N / 16), R3 // loop iterations (for each iteration we emit 16 elements)
loop:
VLD1.P (64)(R1), [V0.S4, V1.S4, V2.S4, V3.S4]
VLD1.P (64)(R2), [V4.S4, V5.S4, V6.S4, V7.S4]
VADD V4.S4, V0.S4, V0.S4
VADD V5.S4, V1.S4, V1.S4
VADD V6.S4, V2.S4, V2.S4
VADD V7.S4, V3.S4, V3.S4
VST1.P [V0.S4, V1.S4, V2.S4, V3.S4], (64)(R0)
SUBS $1, R3, R3
BGT loop
RET
// func polyPackLe16ARM64(p *Poly, buf *byte)
TEXT ·polyPackLe16ARM64(SB), NOSPLIT|NOFRAME, $0-16
MOVD p+0(FP), R0
MOVD buf+8(FP), R1
MOVW $(const_PolyLe16Size / 16), R3 // loop iterations (for each iteration we emit 16 elements)
VMOVQ $0x1c0c180814041000, $0x3c2c382834243020, V15 // value explained at VTBL call
// on the first iteration we have:
// V0 = (p[0], p[4], p[8], p[12]) V1 = (p[1], p[5], p[9], p[13])
// V2 = (p[2], p[6], p[10], p[14]) V3 = (p[3], p[7], p[11], p[15])
// V4 = (p[16], p[20], p[24], p[28]) V5 = (p[17], p[21], p[25], p[29])
// V6 = (p[18], p[22], p[26], p[30]) V7 = (p[19], p[23], p[27], p[31])
loop:
VLD4.P (64)(R0), [V0.S4, V1.S4, V2.S4, V3.S4]
VLD4.P (64)(R0), [V4.S4, V5.S4, V6.S4, V7.S4]
VSHL $4, V1.S4, V1.S4
VSHL $4, V3.S4, V3.S4
VSHL $4, V5.S4, V5.S4
VSHL $4, V7.S4, V7.S4
// tmp = p[even] | (p[odd] << 4)
VORR V1.B16, V0.B16, V10.B16
VORR V3.B16, V2.B16, V11.B16
VORR V5.B16, V4.B16, V12.B16
VORR V7.B16, V6.B16, V13.B16
// so now we need to pick elements based on order:
// first from V10; first from V11; second from V10; second from V11;
// ...
// first from V12; first from V13; second from V12; second from V13;
// V15 contains the indices which correspond to the pick order above
VTBL V15.B16, [V10.B16, V11.B16, V12.B16, V13.B16], V16.B16
VST1.P [V16.B16], (16)(R1)
SUBS $1, R3, R3
BGT loop
RET
// func polyMulBy2toDARM64(p, q *Poly)
TEXT ·polyMulBy2toDARM64(SB), NOSPLIT|NOFRAME, $0-16
MOVD p+0(FP), R0
MOVD q+8(FP), R1
MOVW $(const_N / 16), R2
loop:
VLD1.P (64)(R1), [V0.S4, V1.S4, V2.S4, V3.S4]
VSHL $(const_D), V0.S4, V0.S4
VSHL $(const_D), V1.S4, V1.S4
VSHL $(const_D), V2.S4, V2.S4
VSHL $(const_D), V3.S4, V3.S4
VST1.P [V0.S4, V1.S4, V2.S4, V3.S4], (64)(R0)
SUBS $1, R2, R2
BGT loop
RET
// func polySubARM64(p, a, b *Poly)
TEXT ·polySubARM64(SB), NOSPLIT|NOFRAME, $0-24
MOVD p+0(FP), R0
MOVD a+8(FP), R1
MOVD b+16(FP), R2
MOVW $(const_N / 16), R3
MOVW $(const_Q << 1), R4
VDUP R4, V8.S4
// p = a + (2q - b)
loop:
VLD1.P (64)(R1), [V0.S4, V1.S4, V2.S4, V3.S4]
VLD1.P (64)(R2), [V4.S4, V5.S4, V6.S4, V7.S4]
VSUB V4.S4, V8.S4, V4.S4
VSUB V5.S4, V8.S4, V5.S4
VSUB V6.S4, V8.S4, V6.S4
VSUB V7.S4, V8.S4, V7.S4
VADD V4.S4, V0.S4, V0.S4
VADD V5.S4, V1.S4, V1.S4
VADD V6.S4, V2.S4, V2.S4
VADD V7.S4, V3.S4, V3.S4
VST1.P [V0.S4, V1.S4, V2.S4, V3.S4], (64)(R0)
SUBS $1, R3, R3
BGT loop
RET
+52
View File
@@ -0,0 +1,52 @@
package dilithium
// Returns a y with y < 2q and y = x mod q.
// Note that in general *not*: ReduceLe2Q(ReduceLe2Q(x)) == x.
func ReduceLe2Q(x uint32) uint32 {
// Note 2²³ = 2¹³ - 1 mod q. So, writing x = x₁ 2²³ + x₂ with x₂ < 2²³
// and x₁ < 2⁹, we have x = y (mod q) where
// y = x₂ + x₁ 2¹³ - x₁ ≤ 2²³ + 2¹³ < 2q.
x1 := x >> 23
x2 := x & 0x7FFFFF // 2²³-1
return x2 + (x1 << 13) - x1
}
// Returns x mod q.
func modQ(x uint32) uint32 {
return le2qModQ(ReduceLe2Q(x))
}
// For x R ≤ q 2³², find y ≤ 2q with y = x mod q.
func montReduceLe2Q(x uint64) uint32 {
// Qinv = 4236238847 = -(q⁻¹) mod 2³²
m := (x * Qinv) & 0xffffffff
return uint32((x + m*uint64(Q)) >> 32)
}
// Returns x mod q for 0 ≤ x < 2q.
func le2qModQ(x uint32) uint32 {
x -= Q
mask := uint32(int32(x) >> 31) // mask is 2³²-1 if x was neg.; 0 otherwise
return x + (mask & Q)
}
// Splits 0 ≤ a < Q into a0 and a1 with a = a1*2ᴰ + a0
// and -2ᴰ⁻¹ < a0 < 2ᴰ⁻¹. Returns a0 + Q and a1.
func power2round(a uint32) (a0plusQ, a1 uint32) {
// We effectively compute a0 = a mod± 2ᵈ
// and a1 = (a - a0) / 2ᵈ.
a0 := a & ((1 << D) - 1) // a mod 2ᵈ
// a0 is one of 0, 1, ..., 2ᵈ⁻¹-1, 2ᵈ⁻¹, 2ᵈ⁻¹+1, ..., 2ᵈ-1
a0 -= (1 << (D - 1)) + 1
// now a0 is -2ᵈ⁻¹-1, -2ᵈ⁻¹, ..., -2, -1, 0, ..., 2ᵈ⁻¹-2
// Next, we add 2ᴰ to those a0 that are negative (seen as int32).
a0 += uint32(int32(a0)>>31) & (1 << D)
// now a0 is 2ᵈ⁻¹-1, 2ᵈ⁻¹, ..., 2ᵈ-2, 2ᵈ-1, 0, ..., 2ᵈ⁻¹-2
a0 -= (1 << (D - 1)) - 1
// now a0 id 0, 1, 2, ..., 2ᵈ⁻¹-1, 2ᵈ⁻¹-1, -2ᵈ⁻¹-1, ...
// which is what we want.
a0plusQ = Q + a0
a1 = (a - a0) >> D
return
}
+89
View File
@@ -0,0 +1,89 @@
//go:build (!amd64 && !arm64) || purego
// +build !amd64,!arm64 purego
package dilithium
// Execute an in-place forward NTT on as.
//
// Assumes the coefficients are in Montgomery representation and bounded
// by 2*Q. The resulting coefficients are again in Montgomery representation,
// but are only bounded bt 18*Q.
func (p *Poly) NTT() {
p.nttGeneric()
}
// Execute an in-place inverse NTT and multiply by Montgomery factor R
//
// Assumes the coefficients are in Montgomery representation and bounded
// by 2*Q. The resulting coefficients are again in Montgomery representation
// and bounded by 2*Q.
func (p *Poly) InvNTT() {
p.invNttGeneric()
}
// Sets p to the polynomial whose coefficients are the pointwise multiplication
// of those of a and b. The coefficients of p are bounded by 2q.
//
// Assumes a and b are in Montgomery form and that the pointwise product
// of each coefficient is below 2³² q.
func (p *Poly) MulHat(a, b *Poly) {
p.mulHatGeneric(a, b)
}
// Sets p to a + b. Does not normalize polynomials.
func (p *Poly) Add(a, b *Poly) {
p.addGeneric(a, b)
}
// Sets p to a - b.
//
// Warning: assumes coefficients of b are less than 2q.
// Sets p to a + b. Does not normalize polynomials.
func (p *Poly) Sub(a, b *Poly) {
p.subGeneric(a, b)
}
// Writes p whose coefficients are in [0, 16) to buf, which must be of
// length N/2.
func (p *Poly) PackLe16(buf []byte) {
p.packLe16Generic(buf)
}
// Reduces each of the coefficients to <2q.
func (p *Poly) ReduceLe2Q() {
p.reduceLe2QGeneric()
}
// Reduce each of the coefficients to <q.
func (p *Poly) Normalize() {
p.normalizeGeneric()
}
// Normalize the coefficients in this polynomial assuming they are already
// bounded by 2q.
func (p *Poly) NormalizeAssumingLe2Q() {
p.normalizeAssumingLe2QGeneric()
}
// Checks whether the "supnorm" (see sec 2.1 of the spec) of p is equal
// or greater than the given bound.
//
// Requires the coefficients of p to be normalized.
func (p *Poly) Exceeds(bound uint32) bool {
return p.exceedsGeneric(bound)
}
// Sets p to 2ᵈ q without reducing.
//
// So it requires the coefficients of p to be less than 2³²⁻ᴰ.
func (p *Poly) MulBy2toD(q *Poly) {
p.mulBy2toDGeneric(q)
}
// Splits p into p1 and p0 such that [i]p1 * 2ᴰ + [i]p0 = [i]p
// with -2ᴰ⁻¹ < [i]p0 ≤ 2ᴰ⁻¹. Returns p0 + Q and p1.
//
// Requires the coefficients of p to be normalized.
func (p *Poly) Power2Round(p0PlusQ, p1 *Poly) {
p.power2RoundGeneric(p0PlusQ, p1)
}
+217
View File
@@ -0,0 +1,217 @@
package dilithium
// Zetas lists precomputed powers of the root of unity in Montgomery
// representation used for the NTT:
//
// Zetas[i] = zetaᵇʳᵛ⁽ⁱ⁾ R mod q,
//
// where zeta = 1753, brv(i) is the bitreversal of a 8-bit number
// and R=2³² mod q.
//
// The following Python code generates the Zetas (and InvZetas) lists:
//
// q = 2**23 - 2**13 + 1; zeta = 1753
// R = 2**32 % q # Montgomery const.
// def brv(x): return int(''.join(reversed(bin(x)[2:].zfill(8))),2)
// def inv(x): return pow(x, q-2, q) # inverse in F(q)
// print([(pow(zeta, brv(i), q)*R)%q for i in range(256)])
// print([(pow(inv(zeta), -(brv(255-i)-256), q)*R)%q for i in range(256)])
var Zetas = [N]uint32{
4193792, 25847, 5771523, 7861508, 237124, 7602457, 7504169,
466468, 1826347, 2353451, 8021166, 6288512, 3119733, 5495562,
3111497, 2680103, 2725464, 1024112, 7300517, 3585928, 7830929,
7260833, 2619752, 6271868, 6262231, 4520680, 6980856, 5102745,
1757237, 8360995, 4010497, 280005, 2706023, 95776, 3077325,
3530437, 6718724, 4788269, 5842901, 3915439, 4519302, 5336701,
3574422, 5512770, 3539968, 8079950, 2348700, 7841118, 6681150,
6736599, 3505694, 4558682, 3507263, 6239768, 6779997, 3699596,
811944, 531354, 954230, 3881043, 3900724, 5823537, 2071892,
5582638, 4450022, 6851714, 4702672, 5339162, 6927966, 3475950,
2176455, 6795196, 7122806, 1939314, 4296819, 7380215, 5190273,
5223087, 4747489, 126922, 3412210, 7396998, 2147896, 2715295,
5412772, 4686924, 7969390, 5903370, 7709315, 7151892, 8357436,
7072248, 7998430, 1349076, 1852771, 6949987, 5037034, 264944,
508951, 3097992, 44288, 7280319, 904516, 3958618, 4656075,
8371839, 1653064, 5130689, 2389356, 8169440, 759969, 7063561,
189548, 4827145, 3159746, 6529015, 5971092, 8202977, 1315589,
1341330, 1285669, 6795489, 7567685, 6940675, 5361315, 4499357,
4751448, 3839961, 2091667, 3407706, 2316500, 3817976, 5037939,
2244091, 5933984, 4817955, 266997, 2434439, 7144689, 3513181,
4860065, 4621053, 7183191, 5187039, 900702, 1859098, 909542,
819034, 495491, 6767243, 8337157, 7857917, 7725090, 5257975,
2031748, 3207046, 4823422, 7855319, 7611795, 4784579, 342297,
286988, 5942594, 4108315, 3437287, 5038140, 1735879, 203044,
2842341, 2691481, 5790267, 1265009, 4055324, 1247620, 2486353,
1595974, 4613401, 1250494, 2635921, 4832145, 5386378, 1869119,
1903435, 7329447, 7047359, 1237275, 5062207, 6950192, 7929317,
1312455, 3306115, 6417775, 7100756, 1917081, 5834105, 7005614,
1500165, 777191, 2235880, 3406031, 7838005, 5548557, 6709241,
6533464, 5796124, 4656147, 594136, 4603424, 6366809, 2432395,
2454455, 8215696, 1957272, 3369112, 185531, 7173032, 5196991,
162844, 1616392, 3014001, 810149, 1652634, 4686184, 6581310,
5341501, 3523897, 3866901, 269760, 2213111, 7404533, 1717735,
472078, 7953734, 1723600, 6577327, 1910376, 6712985, 7276084,
8119771, 4546524, 5441381, 6144432, 7959518, 6094090, 183443,
7403526, 1612842, 4834730, 7826001, 3919660, 8332111, 7018208,
3937738, 1400424, 7534263, 1976782,
}
// InvZetas lists precomputed powers of the inverse root of unity in Montgomery
// representation used for the inverse NTT:
//
// InvZetas[i] = zetaᵇʳᵛ⁽²⁵⁵⁻ⁱ⁾⁻²⁵⁶ R mod q,
//
// where zeta = 1753, brv(i) is the bitreversal of a 8-bit number
// and R=2³² mod q.
var InvZetas = [N]uint32{
6403635, 846154, 6979993, 4442679, 1362209, 48306, 4460757,
554416, 3545687, 6767575, 976891, 8196974, 2286327, 420899,
2235985, 2939036, 3833893, 260646, 1104333, 1667432, 6470041,
1803090, 6656817, 426683, 7908339, 6662682, 975884, 6167306,
8110657, 4513516, 4856520, 3038916, 1799107, 3694233, 6727783,
7570268, 5366416, 6764025, 8217573, 3183426, 1207385, 8194886,
5011305, 6423145, 164721, 5925962, 5948022, 2013608, 3776993,
7786281, 3724270, 2584293, 1846953, 1671176, 2831860, 542412,
4974386, 6144537, 7603226, 6880252, 1374803, 2546312, 6463336,
1279661, 1962642, 5074302, 7067962, 451100, 1430225, 3318210,
7143142, 1333058, 1050970, 6476982, 6511298, 2994039, 3548272,
5744496, 7129923, 3767016, 6784443, 5894064, 7132797, 4325093,
7115408, 2590150, 5688936, 5538076, 8177373, 6644538, 3342277,
4943130, 4272102, 2437823, 8093429, 8038120, 3595838, 768622,
525098, 3556995, 5173371, 6348669, 3122442, 655327, 522500,
43260, 1613174, 7884926, 7561383, 7470875, 6521319, 7479715,
3193378, 1197226, 3759364, 3520352, 4867236, 1235728, 5945978,
8113420, 3562462, 2446433, 6136326, 3342478, 4562441, 6063917,
4972711, 6288750, 4540456, 3628969, 3881060, 3019102, 1439742,
812732, 1584928, 7094748, 7039087, 7064828, 177440, 2409325,
1851402, 5220671, 3553272, 8190869, 1316856, 7620448, 210977,
5991061, 3249728, 6727353, 8578, 3724342, 4421799, 7475901,
1100098, 8336129, 5282425, 7871466, 8115473, 3343383, 1430430,
6527646, 7031341, 381987, 1308169, 22981, 1228525, 671102,
2477047, 411027, 3693493, 2967645, 5665122, 6232521, 983419,
4968207, 8253495, 3632928, 3157330, 3190144, 1000202, 4083598,
6441103, 1257611, 1585221, 6203962, 4904467, 1452451, 3041255,
3677745, 1528703, 3930395, 2797779, 6308525, 2556880, 4479693,
4499374, 7426187, 7849063, 7568473, 4680821, 1600420, 2140649,
4873154, 3821735, 4874723, 1643818, 1699267, 539299, 6031717,
300467, 4840449, 2867647, 4805995, 3043716, 3861115, 4464978,
2537516, 3592148, 1661693, 4849980, 5303092, 8284641, 5674394,
8100412, 4369920, 19422, 6623180, 3277672, 1399561, 3859737,
2118186, 2108549, 5760665, 1119584, 549488, 4794489, 1079900,
7356305, 5654953, 5700314, 5268920, 2884855, 5260684, 2091905,
359251, 6026966, 6554070, 7913949, 876248, 777960, 8143293,
518909, 2608894, 8354570, 4186625,
}
// Execute an in-place forward NTT on as.
//
// Assumes the coefficients are in Montgomery representation and bounded
// by 2*Q. The resulting coefficients are again in Montgomery representation,
// but are only bounded bt 18*Q.
func (p *Poly) nttGeneric() {
// Writing z := zeta for our root of unity zeta := 1753, note z²⁵⁶=-1
// (otherwise the order of z wouldn't be 512) and so
//
// x²⁵⁶ + 1 = x²⁵⁶ - z²⁵⁶
// = (x¹²⁸ - z¹²⁸)(x¹²⁸ + z¹²⁸)
// = (x⁶⁴ - z⁶⁴)(x⁶⁴ + z⁶⁴)(x⁶⁴ + z¹⁹²)(x⁶⁴ - z¹⁹²)
// ...
// = (x-z)(x+z)(x - z¹²⁹)(x + z¹²⁹) ... (x - z²⁵⁵)(x + z²⁵⁵)
//
// Note that the powers of z that appear (from the second line) are
// in binary
//
// 01000000 11000000
// 00100000 10100000 01100000 11100000
// 00010000 10010000 01010000 11010000 00110000 10110000 01110000 11110000
// ...
//
// i.e. brv(2), brv(3), brv(4), ... and these powers of z are given by
// the Zetas array.
//
// The polynomials x ± zⁱ are irreducible and coprime, hence by the
// Chinese Remainder Theorem we know
//
// R[x]/(x²⁵⁶+1) → R[x] / (x-z) x ... x R[x] / (x+z²⁵⁵)
// ~= ∏_i R
//
// given by
//
// a ↦ ( a mod x-z, ..., a mod x+z²⁵⁵ )
// ~ ( a(z), a(-z), a(z¹²⁹), a(-z¹²⁹), ..., a(z²⁵⁵), a(-z²⁵⁵) )
//
// is an isomorphism, which is the forward NTT. It can be computed
// efficiently by computing
//
// a ↦ ( a mod x¹²⁸ - z¹²⁸, a mod x¹²⁸ + z¹²⁸ )
// ↦ ( a mod x⁶⁴ - z⁶⁴, a mod x⁶⁴ + z⁶⁴,
// a mod x⁶⁴ - z¹⁹², a mod x⁶⁴ + z¹⁹² )
// et cetera
//
// If N was 8 then this can be pictured in the following diagram:
//
// https://cnx.org/resources/17ee4dfe517a6adda05377b25a00bf6e6c93c334/File0026.png
//
// Each cross is a Cooley--Tukey butterfly: it's the map
//
// (a, b) ↦ (a + ζ, a - ζ)
//
// for the appropriate ζ for that column and row group.
k := 0 // Index into Zetas
// l runs effectively over the columns in the diagram above; it is
// half the height of a row group, i.e. the number of butterflies in
// each row group. In the diagram above it would be 4, 2, 1.
for l := uint(N / 2); l > 0; l >>= 1 {
// On the n-th iteration of the l-loop, the coefficients start off
// bounded by n*2*Q.
//
// offset effectively loops over the row groups in this column; it
// is the first row in the row group.
for offset := uint(0); offset < N-l; offset += 2 * l {
k++
zeta := uint64(Zetas[k])
// j loops over each butterfly in the row group.
for j := offset; j < offset+l; j++ {
t := montReduceLe2Q(zeta * uint64(p[j+l]))
p[j+l] = p[j] + (2*Q - t) // Cooley--Tukey butterfly
p[j] += t
}
}
}
}
// Execute an in-place inverse NTT and multiply by Montgomery factor R
//
// Assumes the coefficients are in Montgomery representation and bounded
// by 2*Q. The resulting coefficients are again in Montgomery representation
// and bounded by 2*Q.
func (p *Poly) invNttGeneric() {
k := 0 // Index into InvZetas
// We basically do the opposite of NTT, but postpone dividing by 2 in the
// inverse of the Cooley--Tukey butterfly and accumulate that to a big
// division by 2⁸ at the end. See comments in the NTT() function.
for l := uint(1); l < N; l <<= 1 {
// On the n-th iteration of the l-loop, the coefficients start off
// bounded by 2ⁿ⁻¹*2*Q, so by 256*Q on the last.
for offset := uint(0); offset < N-l; offset += 2 * l {
zeta := uint64(InvZetas[k])
k++
for j := offset; j < offset+l; j++ {
t := p[j] // Gentleman--Sande butterfly
p[j] = t + p[j+l]
t += 256*Q - p[j+l]
p[j+l] = montReduceLe2Q(zeta * uint64(t))
}
}
}
for j := uint(0); j < N; j++ {
// ROver256 = 41978 = (256)⁻¹ R²
p[j] = montReduceLe2Q(ROver256 * uint64(p[j]))
}
}
+160
View File
@@ -0,0 +1,160 @@
package dilithium
// Sets p to the polynomial whose coefficients are less than 1024 encoded
// into buf (which must be of size PolyT1Size).
//
// p will be normalized.
func (p *Poly) UnpackT1(buf []byte) {
j := 0
for i := 0; i < PolyT1Size; i += 5 {
p[j] = (uint32(buf[i]) | (uint32(buf[i+1]) << 8)) & 0x3ff
p[j+1] = (uint32(buf[i+1]>>2) | (uint32(buf[i+2]) << 6)) & 0x3ff
p[j+2] = (uint32(buf[i+2]>>4) | (uint32(buf[i+3]) << 4)) & 0x3ff
p[j+3] = (uint32(buf[i+3]>>6) | (uint32(buf[i+4]) << 2)) & 0x3ff
j += 4
}
}
// Writes p whose coefficients are in (-2ᵈ⁻¹, 2ᵈ⁻¹] into buf which
// has to be of length at least PolyT0Size.
//
// Assumes that the coefficients are not normalized, but lie in the
// range (q-2ᵈ⁻¹, q+2ᵈ⁻¹].
func (p *Poly) PackT0(buf []byte) {
j := 0
for i := 0; i < PolyT0Size; i += 13 {
p0 := Q + (1 << (D - 1)) - p[j]
p1 := Q + (1 << (D - 1)) - p[j+1]
p2 := Q + (1 << (D - 1)) - p[j+2]
p3 := Q + (1 << (D - 1)) - p[j+3]
p4 := Q + (1 << (D - 1)) - p[j+4]
p5 := Q + (1 << (D - 1)) - p[j+5]
p6 := Q + (1 << (D - 1)) - p[j+6]
p7 := Q + (1 << (D - 1)) - p[j+7]
buf[i] = byte(p0 >> 0)
buf[i+1] = byte(p0>>8) | byte(p1<<5)
buf[i+2] = byte(p1 >> 3)
buf[i+3] = byte(p1>>11) | byte(p2<<2)
buf[i+4] = byte(p2>>6) | byte(p3<<7)
buf[i+5] = byte(p3 >> 1)
buf[i+6] = byte(p3>>9) | byte(p4<<4)
buf[i+7] = byte(p4 >> 4)
buf[i+8] = byte(p4>>12) | byte(p5<<1)
buf[i+9] = byte(p5>>7) | byte(p6<<6)
buf[i+10] = byte(p6 >> 2)
buf[i+11] = byte(p6>>10) | byte(p7<<3)
buf[i+12] = byte(p7 >> 5)
j += 8
}
}
// Sets p to the polynomial packed into buf by PackT0.
//
// The coefficients of p will not be normalized, but will lie
// in (-2ᵈ⁻¹, 2ᵈ⁻¹].
func (p *Poly) UnpackT0(buf []byte) {
j := 0
for i := 0; i < PolyT0Size; i += 13 {
p[j] = Q + (1 << (D - 1)) - ((uint32(buf[i]) |
(uint32(buf[i+1]) << 8)) & 0x1fff)
p[j+1] = Q + (1 << (D - 1)) - (((uint32(buf[i+1]) >> 5) |
(uint32(buf[i+2]) << 3) |
(uint32(buf[i+3]) << 11)) & 0x1fff)
p[j+2] = Q + (1 << (D - 1)) - (((uint32(buf[i+3]) >> 2) |
(uint32(buf[i+4]) << 6)) & 0x1fff)
p[j+3] = Q + (1 << (D - 1)) - (((uint32(buf[i+4]) >> 7) |
(uint32(buf[i+5]) << 1) |
(uint32(buf[i+6]) << 9)) & 0x1fff)
p[j+4] = Q + (1 << (D - 1)) - (((uint32(buf[i+6]) >> 4) |
(uint32(buf[i+7]) << 4) |
(uint32(buf[i+8]) << 12)) & 0x1fff)
p[j+5] = Q + (1 << (D - 1)) - (((uint32(buf[i+8]) >> 1) |
(uint32(buf[i+9]) << 7)) & 0x1fff)
p[j+6] = Q + (1 << (D - 1)) - (((uint32(buf[i+9]) >> 6) |
(uint32(buf[i+10]) << 2) |
(uint32(buf[i+11]) << 10)) & 0x1fff)
p[j+7] = Q + (1 << (D - 1)) - ((uint32(buf[i+11]) >> 3) |
(uint32(buf[i+12]) << 5))
j += 8
}
}
// Writes p whose coefficients are less than 1024 into buf, which must be
// of size at least PolyT1Size .
//
// Assumes coefficients of p are normalized.
func (p *Poly) PackT1(buf []byte) {
j := 0
for i := 0; i < PolyT1Size; i += 5 {
buf[i] = byte(p[j])
buf[i+1] = byte(p[j]>>8) | byte(p[j+1]<<2)
buf[i+2] = byte(p[j+1]>>6) | byte(p[j+2]<<4)
buf[i+3] = byte(p[j+2]>>4) | byte(p[j+3]<<6)
buf[i+4] = byte(p[j+3] >> 2)
j += 4
}
}
// Writes p whose coefficients are in [0, 16) to buf, which must be of
// length N/2.
func (p *Poly) packLe16Generic(buf []byte) {
j := 0
for i := 0; i < PolyLe16Size; i++ {
buf[i] = byte(p[j]) | byte(p[j+1]<<4)
j += 2
}
}
// Writes p with 60 non-zero coefficients {-1,1} to buf, which must have
// length 40.
func (p *Poly) PackB60(buf []byte) {
// We start with a mask of the non-zero positions of p (which is 32 bytes)
// and then append 60 packed bits, where a one indicates a negative
// coefficients.
var signs uint64
mask := uint64(1)
for i := 0; i < 32; i++ {
buf[i] = 0
for j := 0; j < 8; j++ {
if p[8*i+j] != 0 {
buf[i] |= 1 << uint(j)
if p[8*i+j] == Q-1 {
signs |= mask
}
mask <<= 1
}
}
}
for i := uint64(0); i < 8; i++ {
buf[i+32] = uint8(signs >> (8 * i))
}
}
// UnpackB60 sets p to the polynomial packed into buf with Poly.PackB60().
//
// Returns whether unpacking was successful.
func (p *Poly) UnpackB60(buf []byte) bool {
*p = Poly{} // zero p
signs := (uint64(buf[32]) | (uint64(buf[33]) << 8) |
(uint64(buf[34]) << 16) | (uint64(buf[35]) << 24) |
(uint64(buf[36]) << 32) | (uint64(buf[37]) << 40) |
(uint64(buf[38]) << 48) | (uint64(buf[39]) << 56))
if signs>>60 != 0 {
return false // ensure unused bits are zero for strong unforgeability
}
for i := 0; i < 32; i++ {
for j := 0; j < 8; j++ {
if (buf[i]>>uint(j))&1 == 1 {
p[8*i+j] = 1
// Note 1 ^ (1 | (Q-1)) = Q-1 and (-1)&x = x
p[8*i+j] ^= uint32(-(signs & 1)) & (1 | (Q - 1))
signs >>= 1
}
}
}
return true
}
+18
View File
@@ -0,0 +1,18 @@
package dilithium
import (
"github.com/cloudflare/circl/sign/internal/dilithium/params"
)
const (
SeedSize = params.SeedSize
N = params.N
Q = params.Q
QBits = params.QBits
Qinv = params.Qinv
ROver256 = params.ROver256
D = params.D
PolyT1Size = params.PolyT1Size
PolyT0Size = params.PolyT0Size
PolyLe16Size = params.PolyLe16Size
)
@@ -0,0 +1,25 @@
package params
// We put these parameters in a separate package so that the Go code,
// such as ntt_amd64_src.go, that generates assembler can import it.
const (
SeedSize = 32
N = 256
Q = 8380417 // 2²³ - 2¹³ + 1
QBits = 23
Qinv = 4236238847 // = -(q^-1) mod 2³²
ROver256 = 41978 // = (256)⁻¹ R² mod q, where R=2³²
D = 13
// Size of T1 packed. (Note that the formula is not valid in general,
// but it is for the parameters used in the modes of Dilithium.)
PolyT1Size = (N * (QBits - D)) / 8
// Size of T0 packed. (Note that the formula is not valid in general,
// but it is for the parameters used in the modes of Dilithium.)
PolyT0Size = (N * D) / 8
// Size of a packed polynomial whose coefficients are in [0,16).
PolyLe16Size = N / 2
)
+101
View File
@@ -0,0 +1,101 @@
package dilithium
// An element of our base ring R which are polynomials over Z_q modulo
// the equation Xᴺ = -1, where q=2²³ - 2¹³ + 1 and N=256.
//
// Coefficients aren't always reduced. See Normalize().
type Poly [N]uint32
// Reduces each of the coefficients to <2q.
func (p *Poly) reduceLe2QGeneric() {
for i := uint(0); i < N; i++ {
p[i] = ReduceLe2Q(p[i])
}
}
// Reduce each of the coefficients to <q.
func (p *Poly) normalizeGeneric() {
for i := uint(0); i < N; i++ {
p[i] = modQ(p[i])
}
}
// Normalize the coefficients in this polynomial assuming they are already
// bounded by 2q.
func (p *Poly) normalizeAssumingLe2QGeneric() {
for i := 0; i < N; i++ {
p[i] = le2qModQ(p[i])
}
}
// Sets p to a + b. Does not normalize polynomials.
func (p *Poly) addGeneric(a, b *Poly) {
for i := uint(0); i < N; i++ {
p[i] = a[i] + b[i]
}
}
// Sets p to a - b.
//
// Warning: assumes coefficients of b are less than 2q.
func (p *Poly) subGeneric(a, b *Poly) {
for i := uint(0); i < N; i++ {
p[i] = a[i] + (2*Q - b[i])
}
}
// Checks whether the "supnorm" (see sec 2.1 of the spec) of p is equal
// or greater than the given bound.
//
// Requires the coefficients of p to be normalized.
func (p *Poly) exceedsGeneric(bound uint32) bool {
// Note that we are allowed to leak which coefficients break the bound,
// but not their sign.
for i := 0; i < N; i++ {
// The central. reps. of {0, 1, ..., (Q-1)/2, (Q+1)/2, ..., Q-1}
// are given by {0, 1, ..., (Q-1)/2, -(Q-1)/2, ..., -1}
// so their norms are {0, 1, ..., (Q-1)/2, (Q-1)/2, ..., 1}.
// We'll compute them in a different way though.
// Sets x to {(Q-1)/2, (Q-3)/2, ..., 0, -1, ..., -(Q-1)/2}
x := int32((Q-1)/2) - int32(p[i])
// Sets x to {(Q-1)/2, (Q-3)/2, ..., 0, 0, ..., (Q-3)/2}
x ^= (x >> 31)
// Sets x to {0, 1, ..., (Q-1)/2, (Q-1)/2, ..., 1}
x = int32((Q-1)/2) - x
if uint32(x) >= bound {
return true
}
}
return false
}
// Splits p into p1 and p0 such that [i]p1 * 2ᴰ + [i]p0 = [i]p
// with -2ᴰ⁻¹ < [i]p0 ≤ 2ᴰ⁻¹. Returns p0 + Q and p1.
//
// Requires the coefficients of p to be normalized.
func (p *Poly) power2RoundGeneric(p0PlusQ, p1 *Poly) {
for i := 0; i < N; i++ {
p0PlusQ[i], p1[i] = power2round(p[i])
}
}
// Sets p to the polynomial whose coefficients are the pointwise multiplication
// of those of a and b. The coefficients of p are bounded by 2q.
//
// Assumes a and b are in Montgomery form and that the pointwise product
// of each coefficient is below 2³² q.
func (p *Poly) mulHatGeneric(a, b *Poly) {
for i := 0; i < N; i++ {
p[i] = montReduceLe2Q(uint64(a[i]) * uint64(b[i]))
}
}
// Sets p to 2ᵈ q without reducing.
//
// So it requires the coefficients of p to be less than 2³²⁻ᴰ.
func (p *Poly) mulBy2toDGeneric(q *Poly) {
for i := 0; i < N; i++ {
p[i] = q[i] << D
}
}
@@ -0,0 +1,35 @@
// Code generated by command: go run src.go -out ../amd64.s -stubs ../stubs_amd64.go -pkg dilithium. DO NOT EDIT.
//go:build amd64 && !purego
package dilithium
//go:noescape
func nttAVX2(p *[256]uint32)
//go:noescape
func invNttAVX2(p *[256]uint32)
//go:noescape
func mulHatAVX2(p *[256]uint32, a *[256]uint32, b *[256]uint32)
//go:noescape
func addAVX2(p *[256]uint32, a *[256]uint32, b *[256]uint32)
//go:noescape
func subAVX2(p *[256]uint32, a *[256]uint32, b *[256]uint32)
//go:noescape
func packLe16AVX2(p *[256]uint32, buf *byte)
//go:noescape
func reduceLe2QAVX2(p *[256]uint32)
//go:noescape
func le2qModQAVX2(p *[256]uint32)
//go:noescape
func exceedsAVX2(p *[256]uint32, bound uint32) uint8
//go:noescape
func mulBy2toDAVX2(p *[256]uint32, q *[256]uint32)
+366
View File
@@ -0,0 +1,366 @@
// Code generated from pkg.templ.go. DO NOT EDIT.
// mldsa65 implements NIST signature scheme ML-DSA-65 as defined in FIPS204.
package mldsa65
import (
"crypto"
cryptoRand "crypto/rand"
"encoding/asn1"
"errors"
"io"
"github.com/cloudflare/circl/sign"
common "github.com/cloudflare/circl/sign/internal/dilithium"
"github.com/cloudflare/circl/sign/mldsa/mldsa65/internal"
)
const (
// Size of seed for NewKeyFromSeed
SeedSize = common.SeedSize
// Size of a packed PublicKey
PublicKeySize = internal.PublicKeySize
// Size of a packed PrivateKey
PrivateKeySize = internal.PrivateKeySize
// Size of a signature
SignatureSize = internal.SignatureSize
)
// PublicKey is the type of ML-DSA-65 public key
type PublicKey internal.PublicKey
// PrivateKey is the type of ML-DSA-65 private key
type PrivateKey internal.PrivateKey
// GenerateKey generates a public/private key pair using entropy from rand.
// If rand is nil, crypto/rand.Reader will be used.
func GenerateKey(rand io.Reader) (*PublicKey, *PrivateKey, error) {
pk, sk, err := internal.GenerateKey(rand)
return (*PublicKey)(pk), (*PrivateKey)(sk), err
}
// NewKeyFromSeed derives a public/private key pair using the given seed.
func NewKeyFromSeed(seed *[SeedSize]byte) (*PublicKey, *PrivateKey) {
pk, sk := internal.NewKeyFromSeed(seed)
return (*PublicKey)(pk), (*PrivateKey)(sk)
}
// SignTo signs the given message and writes the signature into signature.
// It will panic if signature is not of length at least SignatureSize.
//
// ctx is the optional context string. Errors if ctx is larger than 255 bytes.
// A nil context string is equivalent to an empty context string.
func SignTo(sk *PrivateKey, msg, ctx []byte, randomized bool, sig []byte) error {
var rnd [32]byte
if randomized {
_, err := cryptoRand.Read(rnd[:])
if err != nil {
return err
}
}
if len(ctx) > 255 {
return sign.ErrContextTooLong
}
internal.SignTo(
(*internal.PrivateKey)(sk),
func(w io.Writer) {
_, _ = w.Write([]byte{0})
_, _ = w.Write([]byte{byte(len(ctx))})
if ctx != nil {
_, _ = w.Write(ctx)
}
w.Write(msg)
},
rnd,
sig,
)
return nil
}
// Do not use. Implements ML-DSA.Sign_internal used for compatibility tests.
func (sk *PrivateKey) unsafeSignInternal(msg []byte, rnd [32]byte) []byte {
var ret [SignatureSize]byte
internal.SignTo(
(*internal.PrivateKey)(sk),
func(w io.Writer) {
_, _ = w.Write(msg)
},
rnd,
ret[:],
)
return ret[:]
}
// Do not use. Implements ML-DSA.Verify_internal used for compatibility tests.
func unsafeVerifyInternal(pk *PublicKey, msg, sig []byte) bool {
return internal.Verify(
(*internal.PublicKey)(pk),
func(w io.Writer) {
_, _ = w.Write(msg)
},
sig,
)
}
// Verify checks whether the given signature by pk on msg is valid.
//
// ctx is the optional context string. Fails if ctx is larger than 255 bytes.
// A nil context string is equivalent to an empty context string.
func Verify(pk *PublicKey, msg, ctx, sig []byte) bool {
if len(ctx) > 255 {
return false
}
return internal.Verify(
(*internal.PublicKey)(pk),
func(w io.Writer) {
_, _ = w.Write([]byte{0})
_, _ = w.Write([]byte{byte(len(ctx))})
if ctx != nil {
_, _ = w.Write(ctx)
}
_, _ = w.Write(msg)
},
sig,
)
}
// Sets pk to the public key encoded in buf.
func (pk *PublicKey) Unpack(buf *[PublicKeySize]byte) {
(*internal.PublicKey)(pk).Unpack(buf)
}
// Sets sk to the private key encoded in buf.
func (sk *PrivateKey) Unpack(buf *[PrivateKeySize]byte) {
(*internal.PrivateKey)(sk).Unpack(buf)
}
// Packs the public key into buf.
func (pk *PublicKey) Pack(buf *[PublicKeySize]byte) {
(*internal.PublicKey)(pk).Pack(buf)
}
// Packs the private key into buf.
func (sk *PrivateKey) Pack(buf *[PrivateKeySize]byte) {
(*internal.PrivateKey)(sk).Pack(buf)
}
// Packs the public key.
func (pk *PublicKey) Bytes() []byte {
var buf [PublicKeySize]byte
pk.Pack(&buf)
return buf[:]
}
// Packs the private key.
func (sk *PrivateKey) Bytes() []byte {
var buf [PrivateKeySize]byte
sk.Pack(&buf)
return buf[:]
}
// Packs the public key.
func (pk *PublicKey) MarshalBinary() ([]byte, error) {
return pk.Bytes(), nil
}
// Packs the private key.
func (sk *PrivateKey) MarshalBinary() ([]byte, error) {
return sk.Bytes(), nil
}
// Unpacks the public key from data.
func (pk *PublicKey) UnmarshalBinary(data []byte) error {
if len(data) != PublicKeySize {
return errors.New("packed public key must be of mldsa65.PublicKeySize bytes")
}
var buf [PublicKeySize]byte
copy(buf[:], data)
pk.Unpack(&buf)
return nil
}
// Unpacks the private key from data.
func (sk *PrivateKey) UnmarshalBinary(data []byte) error {
if len(data) != PrivateKeySize {
return errors.New("packed private key must be of mldsa65.PrivateKeySize bytes")
}
var buf [PrivateKeySize]byte
copy(buf[:], data)
sk.Unpack(&buf)
return nil
}
// Returns seed used to generate PrivateKey, and nil if not retained.
func (sk *PrivateKey) Seed() []byte {
return (*internal.PrivateKey)(sk).Seed()
}
// Sign signs the given message.
//
// opts.HashFunc() must return zero, which can be achieved by passing
// crypto.Hash(0) or nil for opts. rand is ignored. Will only return an error
// if opts.HashFunc() is non-zero.
//
// This function is used to make PrivateKey implement the crypto.Signer
// interface. The package-level SignTo function might be more convenient
// to use.
func (sk *PrivateKey) Sign(rand io.Reader, msg []byte, opts crypto.SignerOpts) (
sig []byte, err error) {
var ret [SignatureSize]byte
if opts != nil && opts.HashFunc() != crypto.Hash(0) {
return nil, errors.New("dilithium: cannot sign hashed message")
}
if err = SignTo(sk, msg, nil, false, ret[:]); err != nil {
return nil, err
}
return ret[:], nil
}
// Computes the public key corresponding to this private key.
//
// Returns a *PublicKey. The type crypto.PublicKey is used to make
// PrivateKey implement the crypto.Signer interface.
func (sk *PrivateKey) Public() crypto.PublicKey {
return (*PublicKey)((*internal.PrivateKey)(sk).Public())
}
// Equal returns whether the two private keys equal.
func (sk *PrivateKey) Equal(other crypto.PrivateKey) bool {
castOther, ok := other.(*PrivateKey)
if !ok {
return false
}
return (*internal.PrivateKey)(sk).Equal((*internal.PrivateKey)(castOther))
}
// Equal returns whether the two public keys equal.
func (pk *PublicKey) Equal(other crypto.PublicKey) bool {
castOther, ok := other.(*PublicKey)
if !ok {
return false
}
return (*internal.PublicKey)(pk).Equal((*internal.PublicKey)(castOther))
}
// Boilerplate for generic signatures API
type scheme struct{}
var sch sign.Scheme = &scheme{}
// Scheme returns a generic signature interface for ML-DSA-65.
func Scheme() sign.Scheme { return sch }
func (*scheme) Name() string { return "ML-DSA-65" }
func (*scheme) PublicKeySize() int { return PublicKeySize }
func (*scheme) PrivateKeySize() int { return PrivateKeySize }
func (*scheme) SignatureSize() int { return SignatureSize }
func (*scheme) SeedSize() int { return SeedSize }
// TODO TLSIdentifier()
func (*scheme) Oid() asn1.ObjectIdentifier {
return asn1.ObjectIdentifier{2, 16, 840, 1, 101, 3, 4, 3, 18}
}
func (*scheme) SupportsContext() bool {
return true
}
func (*scheme) GenerateKey() (sign.PublicKey, sign.PrivateKey, error) {
return GenerateKey(nil)
}
func (*scheme) Sign(
sk sign.PrivateKey,
msg []byte,
opts *sign.SignatureOpts,
) []byte {
var ctx []byte
sig := make([]byte, SignatureSize)
priv, ok := sk.(*PrivateKey)
if !ok {
panic(sign.ErrTypeMismatch)
}
if opts != nil && opts.Context != "" {
ctx = []byte(opts.Context)
}
err := SignTo(priv, msg, ctx, false, sig)
if err != nil {
panic(err)
}
return sig
}
func (*scheme) Verify(
pk sign.PublicKey,
msg, sig []byte,
opts *sign.SignatureOpts,
) bool {
var ctx []byte
pub, ok := pk.(*PublicKey)
if !ok {
panic(sign.ErrTypeMismatch)
}
if opts != nil && opts.Context != "" {
ctx = []byte(opts.Context)
}
return Verify(pub, msg, ctx, sig)
}
func (*scheme) DeriveKey(seed []byte) (sign.PublicKey, sign.PrivateKey) {
if len(seed) != SeedSize {
panic(sign.ErrSeedSize)
}
var seed2 [SeedSize]byte
copy(seed2[:], seed)
return NewKeyFromSeed(&seed2)
}
func (*scheme) UnmarshalBinaryPublicKey(buf []byte) (sign.PublicKey, error) {
if len(buf) != PublicKeySize {
return nil, sign.ErrPubKeySize
}
var (
buf2 [PublicKeySize]byte
ret PublicKey
)
copy(buf2[:], buf)
ret.Unpack(&buf2)
return &ret, nil
}
func (*scheme) UnmarshalBinaryPrivateKey(buf []byte) (sign.PrivateKey, error) {
if len(buf) != PrivateKeySize {
return nil, sign.ErrPrivKeySize
}
var (
buf2 [PrivateKeySize]byte
ret PrivateKey
)
copy(buf2[:], buf)
ret.Unpack(&buf2)
return &ret, nil
}
func (sk *PrivateKey) Scheme() sign.Scheme {
return sch
}
func (sk *PublicKey) Scheme() sign.Scheme {
return sch
}
@@ -0,0 +1,509 @@
// Code generated from mode3/internal/dilithium.go by gen.go
package internal
import (
cryptoRand "crypto/rand"
"crypto/subtle"
"io"
"github.com/cloudflare/circl/internal/sha3"
common "github.com/cloudflare/circl/sign/internal/dilithium"
)
const (
// Size of a packed polynomial of norm ≤η.
// (Note that the formula is not valid in general.)
PolyLeqEtaSize = (common.N * DoubleEtaBits) / 8
// β = τη, the maximum size of c s₂.
Beta = Tau * Eta
// γ₁ range of y
Gamma1 = 1 << Gamma1Bits
// Size of packed polynomial of norm <γ₁ such as z
PolyLeGamma1Size = (Gamma1Bits + 1) * common.N / 8
// α = 2γ₂ parameter for decompose
Alpha = 2 * Gamma2
// Size of a packed private key
PrivateKeySize = 32 + 32 + TRSize + PolyLeqEtaSize*(L+K) + common.PolyT0Size*K
// Size of a packed public key
PublicKeySize = 32 + common.PolyT1Size*K
// Size of a packed signature
SignatureSize = L*PolyLeGamma1Size + Omega + K + CTildeSize
// Size of packed w₁
PolyW1Size = (common.N * (common.QBits - Gamma1Bits)) / 8
)
// PublicKey is the type of Dilithium public keys.
type PublicKey struct {
rho [32]byte
t1 VecK
// Cached values
t1p [common.PolyT1Size * K]byte
A *Mat
tr [TRSize]byte
}
// PrivateKey is the type of Dilithium private keys.
type PrivateKey struct {
rho [32]byte
key [32]byte
s1 VecL
s2 VecK
t0 VecK
tr [TRSize]byte
// Cached values
A Mat // ExpandA(ρ)
s1h VecL // NTT(s₁)
s2h VecK // NTT(s₂)
t0h VecK // NTT(t₀)
seed [common.SeedSize]byte
seedSet bool
}
type unpackedSignature struct {
z VecL
hint VecK
c [CTildeSize]byte
}
// Packs the signature into buf.
func (sig *unpackedSignature) Pack(buf []byte) {
copy(buf[:], sig.c[:])
sig.z.PackLeGamma1(buf[CTildeSize:])
sig.hint.PackHint(buf[CTildeSize+L*PolyLeGamma1Size:])
}
// Sets sig to the signature encoded in the buffer.
//
// Returns whether buf contains a properly packed signature.
func (sig *unpackedSignature) Unpack(buf []byte) bool {
// NOTE: Previously the Dilithium (but not ML-DSA) implementation accepted
// signatures with trailing data.
if len(buf) != SignatureSize {
return false
}
copy(sig.c[:], buf[:])
sig.z.UnpackLeGamma1(buf[CTildeSize:])
if sig.z.Exceeds(Gamma1 - Beta) {
return false
}
if !sig.hint.UnpackHint(buf[CTildeSize+L*PolyLeGamma1Size:]) {
return false
}
return true
}
// Packs the public key into buf.
func (pk *PublicKey) Pack(buf *[PublicKeySize]byte) {
copy(buf[:32], pk.rho[:])
copy(buf[32:], pk.t1p[:])
}
// Sets pk to the public key encoded in buf.
func (pk *PublicKey) Unpack(buf *[PublicKeySize]byte) {
copy(pk.rho[:], buf[:32])
copy(pk.t1p[:], buf[32:])
pk.t1.UnpackT1(pk.t1p[:])
pk.A = new(Mat)
pk.A.Derive(&pk.rho)
// tr = CRH(ρ ‖ t1) = CRH(pk)
h := sha3.NewShake256()
_, _ = h.Write(buf[:])
_, _ = h.Read(pk.tr[:])
}
// Packs the private key into buf.
func (sk *PrivateKey) Pack(buf *[PrivateKeySize]byte) {
copy(buf[:32], sk.rho[:])
copy(buf[32:64], sk.key[:])
copy(buf[64:64+TRSize], sk.tr[:])
offset := 64 + TRSize
sk.s1.PackLeqEta(buf[offset:])
offset += PolyLeqEtaSize * L
sk.s2.PackLeqEta(buf[offset:])
offset += PolyLeqEtaSize * K
sk.t0.PackT0(buf[offset:])
}
// Sets sk to the private key encoded in buf.
func (sk *PrivateKey) Unpack(buf *[PrivateKeySize]byte) {
sk.seedSet = false
copy(sk.rho[:], buf[:32])
copy(sk.key[:], buf[32:64])
copy(sk.tr[:], buf[64:64+TRSize])
offset := 64 + TRSize
sk.s1.UnpackLeqEta(buf[offset:])
offset += PolyLeqEtaSize * L
sk.s2.UnpackLeqEta(buf[offset:])
offset += PolyLeqEtaSize * K
sk.t0.UnpackT0(buf[offset:])
// Cached values
sk.A.Derive(&sk.rho)
sk.t0h = sk.t0
sk.t0h.NTT()
sk.s1h = sk.s1
sk.s1h.NTT()
sk.s2h = sk.s2
sk.s2h.NTT()
}
// GenerateKey generates a public/private key pair using entropy from rand.
// If rand is nil, crypto/rand.Reader will be used.
func GenerateKey(rand io.Reader) (*PublicKey, *PrivateKey, error) {
var seed [32]byte
if rand == nil {
rand = cryptoRand.Reader
}
_, err := io.ReadFull(rand, seed[:])
if err != nil {
return nil, nil, err
}
pk, sk := NewKeyFromSeed(&seed)
return pk, sk, nil
}
// NewKeyFromSeed derives a public/private key pair using the given seed.
func NewKeyFromSeed(seed *[common.SeedSize]byte) (*PublicKey, *PrivateKey) {
var eSeed [128]byte // expanded seed
var pk PublicKey
var sk PrivateKey
var sSeed [64]byte
sk.seedSet = true
copy(sk.seed[:], seed[:])
h := sha3.NewShake256()
_, _ = h.Write(seed[:])
if NIST {
_, _ = h.Write([]byte{byte(K), byte(L)})
}
_, _ = h.Read(eSeed[:])
copy(pk.rho[:], eSeed[:32])
copy(sSeed[:], eSeed[32:96])
copy(sk.key[:], eSeed[96:])
copy(sk.rho[:], pk.rho[:])
sk.A.Derive(&pk.rho)
for i := uint16(0); i < L; i++ {
PolyDeriveUniformLeqEta(&sk.s1[i], &sSeed, i)
}
for i := uint16(0); i < K; i++ {
PolyDeriveUniformLeqEta(&sk.s2[i], &sSeed, i+L)
}
sk.s1h = sk.s1
sk.s1h.NTT()
sk.s2h = sk.s2
sk.s2h.NTT()
sk.computeT0andT1(&sk.t0, &pk.t1)
sk.t0h = sk.t0
sk.t0h.NTT()
// Complete public key far enough to be packed
pk.t1.PackT1(pk.t1p[:])
pk.A = &sk.A
// Finish private key
var packedPk [PublicKeySize]byte
pk.Pack(&packedPk)
// tr = CRH(ρ ‖ t1) = CRH(pk)
h.Reset()
_, _ = h.Write(packedPk[:])
_, _ = h.Read(sk.tr[:])
// Finish cache of public key
pk.tr = sk.tr
return &pk, &sk
}
func (sk *PrivateKey) Seed() []byte {
if !sk.seedSet {
return nil
}
var ret [common.SeedSize]byte
copy(ret[:], sk.seed[:])
return ret[:]
}
// Computes t0 and t1 from sk.s1h, sk.s2 and sk.A.
func (sk *PrivateKey) computeT0andT1(t0, t1 *VecK) {
var t VecK
// Set t to A s₁ + s₂
for i := 0; i < K; i++ {
PolyDotHat(&t[i], &sk.A[i], &sk.s1h)
t[i].ReduceLe2Q()
t[i].InvNTT()
}
t.Add(&t, &sk.s2)
t.Normalize()
// Compute t₀, t₁ = Power2Round(t)
t.Power2Round(t0, t1)
}
// Verify checks whether the given signature by pk on msg is valid.
//
// For Dilithium this is the top-level verification function.
// In ML-DSA, this is ML-DSA.Verify_internal.
func Verify(pk *PublicKey, msg func(io.Writer), signature []byte) bool {
var sig unpackedSignature
var mu [64]byte
var zh VecL
var Az, Az2dct1, w1 VecK
var ch common.Poly
var cp [CTildeSize]byte
var w1Packed [PolyW1Size * K]byte
// Note that Unpack() checked whether ‖z‖_∞ < γ₁ - β
// and ensured that there at most ω ones in pk.hint.
if !sig.Unpack(signature) {
return false
}
// μ = CRH(tr ‖ msg)
h := sha3.NewShake256()
_, _ = h.Write(pk.tr[:])
msg(&h)
_, _ = h.Read(mu[:])
// Compute Az
zh = sig.z
zh.NTT()
for i := 0; i < K; i++ {
PolyDotHat(&Az[i], &pk.A[i], &zh)
}
// Next, we compute Az - 2ᵈ·c·t₁.
// Note that the coefficients of t₁ are bounded by 256 = 2⁹,
// so the coefficients of Az2dct1 will bounded by 2⁹⁺ᵈ = 2²³ < 2q,
// which is small enough for NTT().
Az2dct1.MulBy2toD(&pk.t1)
Az2dct1.NTT()
PolyDeriveUniformBall(&ch, sig.c[:])
ch.NTT()
for i := 0; i < K; i++ {
Az2dct1[i].MulHat(&Az2dct1[i], &ch)
}
Az2dct1.Sub(&Az, &Az2dct1)
Az2dct1.ReduceLe2Q()
Az2dct1.InvNTT()
Az2dct1.NormalizeAssumingLe2Q()
// UseHint(pk.hint, Az - 2ᵈ·c·t₁)
// = UseHint(pk.hint, w - c·s₂ + c·t₀)
// = UseHint(pk.hint, r + c·t₀)
// = r₁ = w₁.
w1.UseHint(&Az2dct1, &sig.hint)
w1.PackW1(w1Packed[:])
// c' = H(μ, w₁)
h.Reset()
_, _ = h.Write(mu[:])
_, _ = h.Write(w1Packed[:])
_, _ = h.Read(cp[:])
return sig.c == cp
}
// SignTo signs the given message and writes the signature into signature.
//
// For Dilithium this is the top-level signing function. For ML-DSA
// this is ML-DSA.Sign_internal.
//
//nolint:funlen
func SignTo(sk *PrivateKey, msg func(io.Writer), rnd [32]byte, signature []byte) {
var mu, rhop [64]byte
var w1Packed [PolyW1Size * K]byte
var y, yh VecL
var w, w0, w1, w0mcs2, ct0, w0mcs2pct0 VecK
var ch common.Poly
var yNonce uint16
var sig unpackedSignature
if len(signature) < SignatureSize {
panic("Signature does not fit in that byteslice")
}
// μ = CRH(tr ‖ msg)
h := sha3.NewShake256()
_, _ = h.Write(sk.tr[:])
msg(&h)
_, _ = h.Read(mu[:])
// ρ' = CRH(key ‖ μ)
h.Reset()
_, _ = h.Write(sk.key[:])
if NIST {
_, _ = h.Write(rnd[:])
}
_, _ = h.Write(mu[:])
_, _ = h.Read(rhop[:])
// Main rejection loop
attempt := 0
for {
attempt++
if attempt >= 576 {
// Depending on the mode, one try has a chance between 1/7 and 1/4
// of succeeding. Thus it is safe to say that 576 iterations
// are enough as (6/7)⁵⁷⁶ < 2⁻¹²⁸.
panic("This should only happen 1 in 2^{128}: something is wrong.")
}
// y = ExpandMask(ρ', key)
VecLDeriveUniformLeGamma1(&y, &rhop, yNonce)
yNonce += uint16(L)
// Set w to A y
yh = y
yh.NTT()
for i := 0; i < K; i++ {
PolyDotHat(&w[i], &sk.A[i], &yh)
w[i].ReduceLe2Q()
w[i].InvNTT()
}
// Decompose w into w₀ and w₁
w.NormalizeAssumingLe2Q()
w.Decompose(&w0, &w1)
// c~ = H(μ ‖ w₁)
w1.PackW1(w1Packed[:])
h.Reset()
_, _ = h.Write(mu[:])
_, _ = h.Write(w1Packed[:])
_, _ = h.Read(sig.c[:])
PolyDeriveUniformBall(&ch, sig.c[:])
ch.NTT()
// Ensure ‖ w₀ - c·s2 ‖_∞ < γ₂ - β.
//
// By Lemma 3 of the specification this is equivalent to checking that
// both ‖ r₀ ‖_∞ < γ₂ - β and r₁ = w₁, for the decomposition
// w - c·s₂ = r₁ α + r₀ as computed by decompose().
// See also §4.1 of the specification.
for i := 0; i < K; i++ {
w0mcs2[i].MulHat(&ch, &sk.s2h[i])
w0mcs2[i].InvNTT()
}
w0mcs2.Sub(&w0, &w0mcs2)
w0mcs2.Normalize()
if w0mcs2.Exceeds(Gamma2 - Beta) {
continue
}
// z = y + c·s₁
for i := 0; i < L; i++ {
sig.z[i].MulHat(&ch, &sk.s1h[i])
sig.z[i].InvNTT()
}
sig.z.Add(&sig.z, &y)
sig.z.Normalize()
// Ensure ‖z‖_∞ < γ₁ - β
if sig.z.Exceeds(Gamma1 - Beta) {
continue
}
// Compute c·t₀
for i := 0; i < K; i++ {
ct0[i].MulHat(&ch, &sk.t0h[i])
ct0[i].InvNTT()
}
ct0.NormalizeAssumingLe2Q()
// Ensure ‖c·t₀‖_∞ < γ₂.
if ct0.Exceeds(Gamma2) {
continue
}
// Create the hint to be able to reconstruct w₁ from w - c·s₂ + c·t0.
// Note that we're not using makeHint() in the obvious way as we
// do not know whether ‖ sc·s₂ - c·t₀ ‖_∞ < γ₂. Instead we note
// that our makeHint() is actually the same as a makeHint for a
// different decomposition:
//
// Earlier we ensured indirectly with a check that r₁ = w₁ where
// r = w - c·s₂. Hence r₀ = r - r₁ α = w - c·s₂ - w₁ α = w₀ - c·s₂.
// Thus MakeHint(w₀ - c·s₂ + c·t₀, w₁) = MakeHint(r0 + c·t₀, r₁)
// and UseHint(w - c·s₂ + c·t₀, w₁) = UseHint(r + c·t₀, r₁).
// As we just ensured that ‖ c·t₀ ‖_∞ < γ₂ our usage is correct.
w0mcs2pct0.Add(&w0mcs2, &ct0)
w0mcs2pct0.NormalizeAssumingLe2Q()
hintPop := sig.hint.MakeHint(&w0mcs2pct0, &w1)
if hintPop > Omega {
continue
}
break
}
sig.Pack(signature[:])
}
// Computes the public key corresponding to this private key.
func (sk *PrivateKey) Public() *PublicKey {
var t0 VecK
pk := &PublicKey{
rho: sk.rho,
A: &sk.A,
tr: sk.tr,
}
sk.computeT0andT1(&t0, &pk.t1)
pk.t1.PackT1(pk.t1p[:])
return pk
}
// Equal returns whether the two public keys are equal
func (pk *PublicKey) Equal(other *PublicKey) bool {
return pk.rho == other.rho && pk.t1 == other.t1
}
// Equal returns whether the two private keys are equal
func (sk *PrivateKey) Equal(other *PrivateKey) bool {
ret := (subtle.ConstantTimeCompare(sk.rho[:], other.rho[:]) &
subtle.ConstantTimeCompare(sk.key[:], other.key[:]) &
subtle.ConstantTimeCompare(sk.tr[:], other.tr[:]))
acc := uint32(0)
for i := 0; i < L; i++ {
for j := 0; j < common.N; j++ {
acc |= sk.s1[i][j] ^ other.s1[i][j]
}
}
for i := 0; i < K; i++ {
for j := 0; j < common.N; j++ {
acc |= sk.s2[i][j] ^ other.s2[i][j]
acc |= sk.t0[i][j] ^ other.t0[i][j]
}
}
return (ret & subtle.ConstantTimeEq(int32(acc), 0)) == 1
}
+59
View File
@@ -0,0 +1,59 @@
// Code generated from mode3/internal/mat.go by gen.go
package internal
import (
common "github.com/cloudflare/circl/sign/internal/dilithium"
)
// A k by l matrix of polynomials.
type Mat [K]VecL
// Expands the given seed to a complete matrix.
//
// This function is called ExpandA in the specification.
func (m *Mat) Derive(seed *[32]byte) {
if !DeriveX4Available {
for i := uint16(0); i < K; i++ {
for j := uint16(0); j < L; j++ {
PolyDeriveUniform(&m[i][j], seed, (i<<8)+j)
}
}
return
}
idx := 0
var nonces [4]uint16
var ps [4]*common.Poly
for i := uint16(0); i < K; i++ {
for j := uint16(0); j < L; j++ {
nonces[idx] = (i << 8) + j
ps[idx] = &m[i][j]
idx++
if idx == 4 {
idx = 0
PolyDeriveUniformX4(ps, seed, nonces)
}
}
}
if idx != 0 {
for i := idx; i < 4; i++ {
ps[i] = nil
}
PolyDeriveUniformX4(ps, seed, nonces)
}
}
// Set p to the inner product of a and b using pointwise multiplication.
//
// Assumes a and b are in Montgomery form and their coefficients are
// pairwise sufficiently small to multiply, see Poly.MulHat(). Resulting
// coefficients are bounded by 2Lq.
func PolyDotHat(p *common.Poly, a, b *VecL) {
var t common.Poly
*p = common.Poly{} // zero p
for i := 0; i < L; i++ {
t.MulHat(&a[i], &b[i])
p.Add(&t, p)
}
}
+270
View File
@@ -0,0 +1,270 @@
// Code generated from mode3/internal/pack.go by gen.go
package internal
import (
common "github.com/cloudflare/circl/sign/internal/dilithium"
)
// Writes p with norm less than or equal η into buf, which must be of
// size PolyLeqEtaSize.
//
// Assumes coefficients of p are not normalized, but in [q-η,q+η].
func PolyPackLeqEta(p *common.Poly, buf []byte) { //#nosec G602 -- buf length is fixed (PolyLeqEtaSize)
if DoubleEtaBits == 4 { // compiler eliminates branch
j := 0
for i := 0; i < PolyLeqEtaSize; i++ {
buf[i] = (byte(common.Q+Eta-p[j]) |
byte(common.Q+Eta-p[j+1])<<4)
j += 2
}
} else if DoubleEtaBits == 3 {
j := 0
for i := 0; i < PolyLeqEtaSize; i += 3 {
buf[i] = (byte(common.Q+Eta-p[j]) |
(byte(common.Q+Eta-p[j+1]) << 3) |
(byte(common.Q+Eta-p[j+2]) << 6))
buf[i+1] = ((byte(common.Q+Eta-p[j+2]) >> 2) |
(byte(common.Q+Eta-p[j+3]) << 1) |
(byte(common.Q+Eta-p[j+4]) << 4) |
(byte(common.Q+Eta-p[j+5]) << 7))
buf[i+2] = ((byte(common.Q+Eta-p[j+5]) >> 1) |
(byte(common.Q+Eta-p[j+6]) << 2) |
(byte(common.Q+Eta-p[j+7]) << 5))
j += 8
}
} else {
panic("eta not supported")
}
}
// Sets p to the polynomial of norm less than or equal η encoded in the
// given buffer of size PolyLeqEtaSize.
//
// Output coefficients of p are not normalized, but in [q-η,q+η] provided
// buf was created using PackLeqEta.
//
// Beware, for arbitrary buf the coefficients of p might end up in
// the interval [q-2^b,q+2^b] where b is the least b with η≤2^b.
func PolyUnpackLeqEta(p *common.Poly, buf []byte) { //#nosec G602 -- buf length is fixed (PolyLeqEtaSize)
if DoubleEtaBits == 4 { // compiler eliminates branch
j := 0
for i := 0; i < PolyLeqEtaSize; i++ {
p[j] = common.Q + Eta - uint32(buf[i]&15)
p[j+1] = common.Q + Eta - uint32(buf[i]>>4)
j += 2
}
} else if DoubleEtaBits == 3 {
j := 0
for i := 0; i < PolyLeqEtaSize; i += 3 {
p[j] = common.Q + Eta - uint32(buf[i]&7)
p[j+1] = common.Q + Eta - uint32((buf[i]>>3)&7)
p[j+2] = common.Q + Eta - uint32((buf[i]>>6)|((buf[i+1]<<2)&7))
p[j+3] = common.Q + Eta - uint32((buf[i+1]>>1)&7)
p[j+4] = common.Q + Eta - uint32((buf[i+1]>>4)&7)
p[j+5] = common.Q + Eta - uint32((buf[i+1]>>7)|((buf[i+2]<<1)&7))
p[j+6] = common.Q + Eta - uint32((buf[i+2]>>2)&7)
p[j+7] = common.Q + Eta - uint32((buf[i+2]>>5)&7)
j += 8
}
} else {
panic("eta not supported")
}
}
// Writes v with coefficients in {0, 1} of which at most ω non-zero
// to buf, which must have length ω+k.
func (v *VecK) PackHint(buf []byte) {
// The packed hint starts with the indices of the non-zero coefficients
// For instance:
//
// (x⁵⁶ + x¹⁰⁰, x²⁵⁵, 0, x² + x²³, x¹)
//
// Yields
//
// 56, 100, 255, 2, 23, 1
//
// Then we pad with zeroes until we have a list of ω items:
// // 56, 100, 255, 2, 23, 1, 0, 0, ..., 0
//
// Then we finish with a list of the switch-over-indices in this
// list between polynomials, so:
//
// 56, 100, 255, 2, 23, 1, 0, 0, ..., 0, 2, 3, 3, 5, 6
off := uint8(0)
for i := 0; i < K; i++ {
for j := uint16(0); j < common.N; j++ {
if v[i][j] != 0 {
buf[off] = uint8(j)
off++
}
}
buf[Omega+i] = off
}
for ; off < Omega; off++ {
buf[off] = 0
}
}
// Sets v to the vector encoded using VecK.PackHint()
//
// Returns whether unpacking was successful.
func (v *VecK) UnpackHint(buf []byte) bool {
// A priori, there would be several reasonable ways to encode the same
// hint vector. We take care to only allow only one encoding, to ensure
// "strong unforgeability".
//
// See PackHint() source for description of the encoding.
*v = VecK{} // zero v
prevSOP := uint8(0) // previous switch-over-point
for i := 0; i < K; i++ {
SOP := buf[Omega+i]
if SOP < prevSOP || SOP > Omega {
return false // ensures switch-over-points are increasing
}
for j := prevSOP; j < SOP; j++ {
if j > prevSOP && buf[j] <= buf[j-1] {
return false // ensures indices are increasing (within a poly)
}
v[i][buf[j]] = 1
}
prevSOP = SOP
}
for j := prevSOP; j < Omega; j++ {
if buf[j] != 0 {
return false // ensures padding indices are zero
}
}
return true
}
// Sets p to the polynomial packed into buf by PolyPackLeGamma1.
//
// p will be normalized.
func PolyUnpackLeGamma1(p *common.Poly, buf []byte) { //#nosec G602 -- buf length is fixed (PolyLeGamma1Size)
if Gamma1Bits == 17 {
j := 0
for i := 0; i < PolyLeGamma1Size; i += 9 {
p0 := uint32(buf[i]) | (uint32(buf[i+1]) << 8) |
(uint32(buf[i+2]&0x3) << 16)
p1 := uint32(buf[i+2]>>2) | (uint32(buf[i+3]) << 6) |
(uint32(buf[i+4]&0xf) << 14)
p2 := uint32(buf[i+4]>>4) | (uint32(buf[i+5]) << 4) |
(uint32(buf[i+6]&0x3f) << 12)
p3 := uint32(buf[i+6]>>6) | (uint32(buf[i+7]) << 2) |
(uint32(buf[i+8]) << 10)
// coefficients in [0,…,2γ₁)
p0 = Gamma1 - p0 // (-γ₁,…,γ₁]
p1 = Gamma1 - p1
p2 = Gamma1 - p2
p3 = Gamma1 - p3
p0 += uint32(int32(p0)>>31) & common.Q // normalize
p1 += uint32(int32(p1)>>31) & common.Q
p2 += uint32(int32(p2)>>31) & common.Q
p3 += uint32(int32(p3)>>31) & common.Q
p[j] = p0
p[j+1] = p1
p[j+2] = p2
p[j+3] = p3
j += 4
}
} else if Gamma1Bits == 19 {
j := 0
for i := 0; i < PolyLeGamma1Size; i += 5 {
p0 := uint32(buf[i]) | (uint32(buf[i+1]) << 8) |
(uint32(buf[i+2]&0xf) << 16)
p1 := uint32(buf[i+2]>>4) | (uint32(buf[i+3]) << 4) |
(uint32(buf[i+4]) << 12)
p0 = Gamma1 - p0
p1 = Gamma1 - p1
p0 += uint32(int32(p0)>>31) & common.Q
p1 += uint32(int32(p1)>>31) & common.Q
p[j] = p0
p[j+1] = p1
j += 2
}
} else {
panic("γ₁ not supported")
}
}
// Writes p whose coefficients are in (-γ₁,γ₁] into buf
// which has to be of length PolyLeGamma1Size.
//
// Assumes p is normalized.
func PolyPackLeGamma1(p *common.Poly, buf []byte) { //#nosec G602 -- buf length is fixed (PolyLeGamma1Size)
if Gamma1Bits == 17 {
j := 0
// coefficients in [0,…,γ₁] ∪ (q-γ₁,…,q)
for i := 0; i < PolyLeGamma1Size; i += 9 {
p0 := Gamma1 - p[j] // [0,…,γ₁] ∪ (γ₁-q,…,2γ₁-q)
p0 += uint32(int32(p0)>>31) & common.Q // [0,…,2γ₁)
p1 := Gamma1 - p[j+1]
p1 += uint32(int32(p1)>>31) & common.Q
p2 := Gamma1 - p[j+2]
p2 += uint32(int32(p2)>>31) & common.Q
p3 := Gamma1 - p[j+3]
p3 += uint32(int32(p3)>>31) & common.Q
buf[i+0] = byte(p0)
buf[i+1] = byte(p0 >> 8)
buf[i+2] = byte(p0>>16) | byte(p1<<2)
buf[i+3] = byte(p1 >> 6)
buf[i+4] = byte(p1>>14) | byte(p2<<4)
buf[i+5] = byte(p2 >> 4)
buf[i+6] = byte(p2>>12) | byte(p3<<6)
buf[i+7] = byte(p3 >> 2)
buf[i+8] = byte(p3 >> 10)
j += 4
}
} else if Gamma1Bits == 19 {
j := 0
for i := 0; i < PolyLeGamma1Size; i += 5 {
// Coefficients are in [0, γ₁] ∪ (Q-γ₁, Q)
p0 := Gamma1 - p[j]
p0 += uint32(int32(p0)>>31) & common.Q
p1 := Gamma1 - p[j+1]
p1 += uint32(int32(p1)>>31) & common.Q
buf[i+0] = byte(p0)
buf[i+1] = byte(p0 >> 8)
buf[i+2] = byte(p0>>16) | byte(p1<<4)
buf[i+3] = byte(p1 >> 4)
buf[i+4] = byte(p1 >> 12)
j += 2
}
} else {
panic("γ₁ not supported")
}
}
// Pack w₁ into buf, which must be of length PolyW1Size.
//
// Assumes w₁ is normalized.
func PolyPackW1(p *common.Poly, buf []byte) {
if Gamma1Bits == 19 {
p.PackLe16(buf)
} else if Gamma1Bits == 17 {
j := 0
for i := 0; i < PolyW1Size; i += 3 {
buf[i] = byte(p[j]) | byte(p[j+1]<<6)
buf[i+1] = byte(p[j+1]>>2) | byte(p[j+2]<<4)
buf[i+2] = byte(p[j+2]>>4) | byte(p[j+3]<<2)
j += 4
}
} else {
panic("unsupported γ₁")
}
}
@@ -0,0 +1,18 @@
// Code generated from params.templ.go. DO NOT EDIT.
package internal
const (
Name = "ML-DSA-65"
K = 6
L = 5
Eta = 4
DoubleEtaBits = 4
Omega = 55
Tau = 49
Gamma1Bits = 19
Gamma2 = 261888
NIST = true
TRSize = 64
CTildeSize = 48
)
@@ -0,0 +1,142 @@
// Code generated from mode3/internal/rounding.go by gen.go
package internal
import (
common "github.com/cloudflare/circl/sign/internal/dilithium"
)
// Splits 0 ≤ a < q into a₀ and a₁ with a = a₁*α + a₀ with -α/2 < a₀ ≤ α/2,
// except for when we would have a₁ = (q-1)/α in which case a₁=0 is taken
// and -α/2 ≤ a₀ < 0. Returns a₀ + q. Note 0 ≤ a₁ < (q-1)/α.
// Recall α = 2γ₂.
func decompose(a uint32) (a0plusQ, a1 uint32) {
// a₁ = ⌈a / 128⌉
a1 = (a + 127) >> 7
if Alpha == 523776 {
// 1025/2²² is close enough to 1/4092 so that a₁
// becomes a/α rounded down.
a1 = ((a1*1025 + (1 << 21)) >> 22)
// For the corner-case a₁ = (q-1)/α = 16, we have to set a₁=0.
a1 &= 15
} else if Alpha == 190464 {
// 1488/2²⁴ is close enough to 1/1488 so that a₁
// becomes a/α rounded down.
a1 = ((a1 * 11275) + (1 << 23)) >> 24
// For the corner-case a₁ = (q-1)/α = 44, we have to set a₁=0.
a1 ^= uint32(int32(43-a1)>>31) & a1
} else {
panic("unsupported α")
}
a0plusQ = a - a1*Alpha
// In the corner-case, when we set a₁=0, we will incorrectly
// have a₀ > (q-1)/2 and we'll need to subtract q. As we
// return a₀ + q, that comes down to adding q if a₀ < (q-1)/2.
a0plusQ += uint32(int32(a0plusQ-(common.Q-1)/2)>>31) & common.Q
return
}
// Assume 0 ≤ r, f < Q with ‖f‖_∞ ≤ α/2. Decompose r as r = r1*α + r0 as
// computed by decompose(). Write r' := r - f (mod Q). Now, decompose
// r'=r-f again as r' = r'1*α + r'0 using decompose(). As f is small, we
// have r'1 = r1 + h, where h ∈ {-1, 0, 1}. makeHint() computes |h|
// given z0 := r0 - f (mod Q) and r1. With |h|, which is called the hint,
// we can reconstruct r1 using only r' = r - f, which is done by useHint().
// To wit:
//
// useHint( r - f, makeHint( r0 - f, r1 ) ) = r1.
//
// Assumes 0 ≤ z0 < Q.
func makeHint(z0, r1 uint32) uint32 {
// If -α/2 < r0 - f ≤ α/2, then r1*α + r0 - f is a valid decomposition of r'
// with the restrictions of decompose() and so r'1 = r1. So the hint
// should be 0. This is covered by the first two inequalities.
// There is one other case: if r0 - f = -α/2, then r1*α + r0 - f is also
// a valid decomposition if r1 = 0. In the other cases a one is carried
// and the hint should be 1.
if z0 <= Gamma2 || z0 > common.Q-Gamma2 || (z0 == common.Q-Gamma2 && r1 == 0) {
return 0
}
return 1
}
// Uses the hint created by makeHint() to reconstruct r1 from r'=r-f; see
// documentation of makeHint() for context.
// Assumes 0 ≤ r' < Q.
func useHint(rp uint32, hint uint32) uint32 {
rp0plusQ, rp1 := decompose(rp)
if hint == 0 {
return rp1
}
if rp0plusQ > common.Q {
return (rp1 + 1) & 15
}
return (rp1 - 1) & 15
}
// Sets p to the hint polynomial for p0 the modified low bits and p1
// the unmodified high bits --- see makeHint().
//
// Returns the number of ones in the hint polynomial.
func PolyMakeHint(p, p0, p1 *common.Poly) (pop uint32) {
for i := 0; i < common.N; i++ {
h := makeHint(p0[i], p1[i])
pop += h
p[i] = h
}
return
}
// Computes corrections to the high bits of the polynomial q according
// to the hints in h and sets p to the corrected high bits. Returns p.
func PolyUseHint(p, q, hint *common.Poly) {
var q0PlusQ common.Poly
// See useHint() and makeHint() for an explanation. We reimplement it
// here so that we can call Poly.Decompose(), which might be way faster
// than calling decompose() in a loop (for instance when having AVX2.)
PolyDecompose(q, &q0PlusQ, p)
for i := 0; i < common.N; i++ {
if hint[i] == 0 {
continue
}
if Gamma2 == 261888 {
if q0PlusQ[i] > common.Q {
p[i] = (p[i] + 1) & 15
} else {
p[i] = (p[i] - 1) & 15
}
} else if Gamma2 == 95232 {
if q0PlusQ[i] > common.Q {
if p[i] == 43 {
p[i] = 0
} else {
p[i]++
}
} else {
if p[i] == 0 {
p[i] = 43
} else {
p[i]--
}
}
} else {
panic("unsupported γ₂")
}
}
}
// Splits each of the coefficients of p using decompose.
func PolyDecompose(p, p0PlusQ, p1 *common.Poly) {
for i := 0; i < common.N; i++ {
p0PlusQ[i], p1[i] = decompose(p[i])
}
}
@@ -0,0 +1,339 @@
// Code generated from mode3/internal/sample.go by gen.go
package internal
import (
"encoding/binary"
"github.com/cloudflare/circl/internal/sha3"
common "github.com/cloudflare/circl/sign/internal/dilithium"
"github.com/cloudflare/circl/simd/keccakf1600"
)
// DeriveX4Available indicates whether the system supports the quick fourway
// sampling variants like PolyDeriveUniformX4.
var DeriveX4Available = keccakf1600.IsEnabledX4()
// For each i, sample ps[i] uniformly from the given seed and nonces[i].
// ps[i] may be nil and is ignored in that case.
//
// Can only be called when DeriveX4Available is true.
func PolyDeriveUniformX4(ps [4]*common.Poly, seed *[32]byte, nonces [4]uint16) {
var perm keccakf1600.StateX4
state := perm.Initialize(false)
// Absorb the seed in the four states
for i := 0; i < 4; i++ {
v := binary.LittleEndian.Uint64(seed[8*i : 8*(i+1)])
for j := 0; j < 4; j++ {
state[i*4+j] = v
}
}
// Absorb the nonces, the SHAKE128 domain separator (0b1111), the
// start of the padding (0b...001) and the end of the padding 0b100...
// Recall that the rate of SHAKE128 is 168 --- i.e. 21 uint64s.
for j := 0; j < 4; j++ {
state[4*4+j] = uint64(nonces[j]) | (0x1f << 16)
state[20*4+j] = 0x80 << 56
}
var idx [4]int // indices into ps
for j := 0; j < 4; j++ {
if ps[j] == nil {
idx[j] = common.N // mark nil polynomial as completed
}
}
done := false
for !done {
// Applies KeccaK-f[1600] to state to get the next 21 uint64s of each
// of the four SHAKE128 streams.
perm.Permute()
done = true
PolyLoop:
for j := 0; j < 4; j++ {
if idx[j] == common.N {
continue
}
for i := 0; i < 7; i++ {
var t [8]uint32
t[0] = uint32(state[i*3*4+j] & 0x7fffff)
t[1] = uint32((state[i*3*4+j] >> 24) & 0x7fffff)
t[2] = uint32((state[i*3*4+j] >> 48) |
((state[(i*3+1)*4+j] & 0x7f) << 16))
t[3] = uint32((state[(i*3+1)*4+j] >> 8) & 0x7fffff)
t[4] = uint32((state[(i*3+1)*4+j] >> 32) & 0x7fffff)
t[5] = uint32((state[(i*3+1)*4+j] >> 56) |
((state[(i*3+2)*4+j] & 0x7fff) << 8))
t[6] = uint32((state[(i*3+2)*4+j] >> 16) & 0x7fffff)
t[7] = uint32((state[(i*3+2)*4+j] >> 40) & 0x7fffff)
for k := 0; k < 8; k++ {
if t[k] < common.Q {
ps[j][idx[j]] = t[k]
idx[j]++
if idx[j] == common.N {
continue PolyLoop
}
}
}
}
done = false
}
}
}
// Sample p uniformly from the given seed and nonce.
//
// p will be normalized.
func PolyDeriveUniform(p *common.Poly, seed *[32]byte, nonce uint16) {
var i, length int
var buf [12 * 16]byte // fits 168B SHAKE-128 rate
length = 168
sample := func() {
// Note that 3 divides into 168 and 12*16, so we use up buf completely.
for j := 0; j < length && i < common.N; j += 3 {
t := (uint32(buf[j]) | (uint32(buf[j+1]) << 8) |
(uint32(buf[j+2]) << 16)) & 0x7fffff
// We use rejection sampling
if t < common.Q {
p[i] = t
i++
}
}
}
var iv [32 + 2]byte // 32 byte seed + uint16 nonce
h := sha3.NewShake128()
copy(iv[:32], seed[:])
iv[32] = uint8(nonce)
iv[33] = uint8(nonce >> 8)
_, _ = h.Write(iv[:])
for i < common.N {
_, _ = h.Read(buf[:168])
sample()
}
}
// Sample p uniformly with coefficients of norm less than or equal η,
// using the given seed and nonce.
//
// p will not be normalized, but will have coefficients in [q-η,q+η].
func PolyDeriveUniformLeqEta(p *common.Poly, seed *[64]byte, nonce uint16) {
// Assumes 2 < η < 8.
var i, length int
var buf [9 * 16]byte // fits 136B SHAKE-256 rate
length = 136
sample := func() {
// We use rejection sampling
for j := 0; j < length && i < common.N; j++ {
t1 := uint32(buf[j]) & 15
t2 := uint32(buf[j]) >> 4
if Eta == 2 { // branch is eliminated by compiler
if t1 <= 14 {
t1 -= ((205 * t1) >> 10) * 5 // reduce mod 5
p[i] = common.Q + Eta - t1
i++
}
if t2 <= 14 && i < common.N {
t2 -= ((205 * t2) >> 10) * 5 // reduce mod 5
p[i] = common.Q + Eta - t2
i++
}
} else if Eta == 4 {
if t1 <= 2*Eta {
p[i] = common.Q + Eta - t1
i++
}
if t2 <= 2*Eta && i < common.N {
p[i] = common.Q + Eta - t2
i++
}
} else {
panic("unsupported η")
}
}
}
var iv [64 + 2]byte // 64 byte seed + uint16 nonce
h := sha3.NewShake256()
copy(iv[:64], seed[:])
iv[64] = uint8(nonce)
iv[65] = uint8(nonce >> 8)
// 136 is SHAKE-256 rate
_, _ = h.Write(iv[:])
for i < common.N {
_, _ = h.Read(buf[:136])
sample()
}
}
// Sample v[i] uniformly with coefficients in (-γ₁,…,γ₁] using the
// given seed and nonce+i
//
// p will be normalized.
func VecLDeriveUniformLeGamma1(v *VecL, seed *[64]byte, nonce uint16) {
for i := 0; i < L; i++ {
PolyDeriveUniformLeGamma1(&v[i], seed, nonce+uint16(i))
}
}
// Sample p uniformly with coefficients in (-γ₁,…,γK1s] using the
// given seed and nonce.
//
// p will be normalized.
func PolyDeriveUniformLeGamma1(p *common.Poly, seed *[64]byte, nonce uint16) {
var buf [PolyLeGamma1Size]byte
var iv [66]byte
h := sha3.NewShake256()
copy(iv[:64], seed[:])
iv[64] = uint8(nonce)
iv[65] = uint8(nonce >> 8)
_, _ = h.Write(iv[:])
_, _ = h.Read(buf[:])
PolyUnpackLeGamma1(p, buf[:])
}
// For each i, sample ps[i] uniformly with τ non-zero coefficients in {q-1,1}
// using the given seed and w1[i]. ps[i] may be nil and is ignored
// in that case. ps[i] will be normalized.
//
// Can only be called when DeriveX4Available is true.
//
// This function is currently not used (yet).
func PolyDeriveUniformBallX4(ps [4]*common.Poly, seed []byte) {
var perm keccakf1600.StateX4
state := perm.Initialize(false)
// Absorb the seed in the four states
for i := 0; i < CTildeSize/8; i++ {
v := binary.LittleEndian.Uint64(seed[8*i : 8*(i+1)])
for j := 0; j < 4; j++ {
state[i*4+j] = v
}
}
// SHAKE256 domain separator and padding
for j := 0; j < 4; j++ {
state[(CTildeSize/8)*4+j] ^= 0x1f
state[16*4+j] ^= 0x80 << 56
}
perm.Permute()
var signs [4]uint64
var idx [4]uint16 // indices into ps
for j := 0; j < 4; j++ {
if ps[j] != nil {
signs[j] = state[j]
*ps[j] = common.Poly{} // zero ps[j]
idx[j] = common.N - Tau
} else {
idx[j] = common.N // mark as completed
}
}
stateOffset := 1
for {
done := true
PolyLoop:
for j := 0; j < 4; j++ {
if idx[j] == common.N {
continue
}
for i := stateOffset; i < 17; i++ {
var bs [8]byte
binary.LittleEndian.PutUint64(bs[:], state[4*i+j])
for k := 0; k < 8; k++ {
b := uint16(bs[k])
if b > idx[j] {
continue
}
ps[j][idx[j]] = ps[j][b]
ps[j][b] = 1
// Takes least significant bit of signs and uses it for the sign.
// Note 1 ^ (1 | (Q-1)) = Q-1.
ps[j][b] ^= uint32((-(signs[j] & 1)) & (1 | (common.Q - 1)))
signs[j] >>= 1
idx[j]++
if idx[j] == common.N {
continue PolyLoop
}
}
}
done = false
}
if done {
break
}
perm.Permute()
stateOffset = 0
}
}
// Samples p uniformly with τ non-zero coefficients in {q-1,1}.
//
// The polynomial p will be normalized.
func PolyDeriveUniformBall(p *common.Poly, seed []byte) {
var buf [136]byte // SHAKE-256 rate is 136
h := sha3.NewShake256()
_, _ = h.Write(seed[:])
_, _ = h.Read(buf[:])
// Essentially we generate a sequence of τ ones or minus ones,
// prepend 196 zeroes and shuffle the concatenation using the
// usual algorithm (Fisher--Yates.)
signs := binary.LittleEndian.Uint64(buf[:])
bufOff := 8 // offset into buf
*p = common.Poly{} // zero p
for i := uint16(common.N - Tau); i < common.N; i++ {
var b uint16
// Find location of where to move the new coefficient to using
// rejection sampling.
for {
if bufOff >= 136 {
_, _ = h.Read(buf[:])
bufOff = 0
}
b = uint16(buf[bufOff])
bufOff++
if b <= i {
break
}
}
p[i] = p[b]
p[b] = 1
// Takes least significant bit of signs and uses it for the sign.
// Note 1 ^ (1 | (Q-1)) = Q-1.
p[b] ^= uint32((-(signs & 1)) & (1 | (common.Q - 1)))
signs >>= 1
}
}
+281
View File
@@ -0,0 +1,281 @@
// Code generated from mode3/internal/vec.go by gen.go
package internal
import (
common "github.com/cloudflare/circl/sign/internal/dilithium"
)
// A vector of L polynomials.
type VecL [L]common.Poly
// A vector of K polynomials.
type VecK [K]common.Poly
// Normalize the polynomials in this vector.
func (v *VecL) Normalize() {
for i := 0; i < L; i++ {
v[i].Normalize()
}
}
// Normalize the polynomials in this vector assuming their coefficients
// are already bounded by 2q.
func (v *VecL) NormalizeAssumingLe2Q() {
for i := 0; i < L; i++ {
v[i].NormalizeAssumingLe2Q()
}
}
// Sets v to w + u. Does not normalize.
func (v *VecL) Add(w, u *VecL) {
for i := 0; i < L; i++ {
v[i].Add(&w[i], &u[i])
}
}
// Applies NTT componentwise. See Poly.NTT() for details.
func (v *VecL) NTT() {
for i := 0; i < L; i++ {
v[i].NTT()
}
}
// Checks whether any of the coefficients exceeds the given bound in supnorm
//
// Requires the vector to be normalized.
func (v *VecL) Exceeds(bound uint32) bool {
for i := 0; i < L; i++ {
if v[i].Exceeds(bound) {
return true
}
}
return false
}
// Applies Poly.Power2Round componentwise.
//
// Requires the vector to be normalized.
func (v *VecL) Power2Round(v0PlusQ, v1 *VecL) {
for i := 0; i < L; i++ {
v[i].Power2Round(&v0PlusQ[i], &v1[i])
}
}
// Applies Poly.Decompose componentwise.
//
// Requires the vector to be normalized.
func (v *VecL) Decompose(v0PlusQ, v1 *VecL) {
for i := 0; i < L; i++ {
PolyDecompose(&v[i], &v0PlusQ[i], &v1[i])
}
}
// Sequentially packs each polynomial using Poly.PackLeqEta().
func (v *VecL) PackLeqEta(buf []byte) {
offset := 0
for i := 0; i < L; i++ {
PolyPackLeqEta(&v[i], buf[offset:])
offset += PolyLeqEtaSize
}
}
// Sets v to the polynomials packed in buf using VecL.PackLeqEta().
func (v *VecL) UnpackLeqEta(buf []byte) {
offset := 0
for i := 0; i < L; i++ {
PolyUnpackLeqEta(&v[i], buf[offset:])
offset += PolyLeqEtaSize
}
}
// Sequentially packs each polynomial using PolyPackLeGamma1().
func (v *VecL) PackLeGamma1(buf []byte) {
offset := 0
for i := 0; i < L; i++ {
PolyPackLeGamma1(&v[i], buf[offset:])
offset += PolyLeGamma1Size
}
}
// Sets v to the polynomials packed in buf using VecL.PackLeGamma1().
func (v *VecL) UnpackLeGamma1(buf []byte) {
offset := 0
for i := 0; i < L; i++ {
PolyUnpackLeGamma1(&v[i], buf[offset:])
offset += PolyLeGamma1Size
}
}
// Normalize the polynomials in this vector.
func (v *VecK) Normalize() {
for i := 0; i < K; i++ {
v[i].Normalize()
}
}
// Normalize the polynomials in this vector assuming their coefficients
// are already bounded by 2q.
func (v *VecK) NormalizeAssumingLe2Q() {
for i := 0; i < K; i++ {
v[i].NormalizeAssumingLe2Q()
}
}
// Sets v to w + u. Does not normalize.
func (v *VecK) Add(w, u *VecK) {
for i := 0; i < K; i++ {
v[i].Add(&w[i], &u[i])
}
}
// Checks whether any of the coefficients exceeds the given bound in supnorm
//
// Requires the vector to be normalized.
func (v *VecK) Exceeds(bound uint32) bool {
for i := 0; i < K; i++ {
if v[i].Exceeds(bound) {
return true
}
}
return false
}
// Applies Poly.Power2Round componentwise.
//
// Requires the vector to be normalized.
func (v *VecK) Power2Round(v0PlusQ, v1 *VecK) {
for i := 0; i < K; i++ {
v[i].Power2Round(&v0PlusQ[i], &v1[i])
}
}
// Applies Poly.Decompose componentwise.
//
// Requires the vector to be normalized.
func (v *VecK) Decompose(v0PlusQ, v1 *VecK) {
for i := 0; i < K; i++ {
PolyDecompose(&v[i], &v0PlusQ[i], &v1[i])
}
}
// Sets v to the hint vector for v0 the modified low bits and v1
// the unmodified high bits --- see makeHint().
//
// Returns the number of ones in the hint vector.
func (v *VecK) MakeHint(v0, v1 *VecK) (pop uint32) {
for i := 0; i < K; i++ {
pop += PolyMakeHint(&v[i], &v0[i], &v1[i])
}
return
}
// Computes corrections to the high bits of the polynomials in the vector
// w using the hints in h and sets v to the corrected high bits. Returns v.
// See useHint().
func (v *VecK) UseHint(q, hint *VecK) *VecK {
for i := 0; i < K; i++ {
PolyUseHint(&v[i], &q[i], &hint[i])
}
return v
}
// Sequentially packs each polynomial using Poly.PackT1().
func (v *VecK) PackT1(buf []byte) {
offset := 0
for i := 0; i < K; i++ {
v[i].PackT1(buf[offset:])
offset += common.PolyT1Size
}
}
// Sets v to the vector packed into buf by PackT1().
func (v *VecK) UnpackT1(buf []byte) {
offset := 0
for i := 0; i < K; i++ {
v[i].UnpackT1(buf[offset:])
offset += common.PolyT1Size
}
}
// Sequentially packs each polynomial using Poly.PackT0().
func (v *VecK) PackT0(buf []byte) {
offset := 0
for i := 0; i < K; i++ {
v[i].PackT0(buf[offset:])
offset += common.PolyT0Size
}
}
// Sets v to the vector packed into buf by PackT0().
func (v *VecK) UnpackT0(buf []byte) {
offset := 0
for i := 0; i < K; i++ {
v[i].UnpackT0(buf[offset:])
offset += common.PolyT0Size
}
}
// Sequentially packs each polynomial using Poly.PackLeqEta().
func (v *VecK) PackLeqEta(buf []byte) {
offset := 0
for i := 0; i < K; i++ {
PolyPackLeqEta(&v[i], buf[offset:])
offset += PolyLeqEtaSize
}
}
// Sets v to the polynomials packed in buf using VecK.PackLeqEta().
func (v *VecK) UnpackLeqEta(buf []byte) {
offset := 0
for i := 0; i < K; i++ {
PolyUnpackLeqEta(&v[i], buf[offset:])
offset += PolyLeqEtaSize
}
}
// Applies NTT componentwise. See Poly.NTT() for details.
func (v *VecK) NTT() {
for i := 0; i < K; i++ {
v[i].NTT()
}
}
// Sequentially packs each polynomial using PolyPackW1().
func (v *VecK) PackW1(buf []byte) {
offset := 0
for i := 0; i < K; i++ {
PolyPackW1(&v[i], buf[offset:])
offset += PolyW1Size
}
}
// Sets v to a - b.
//
// Warning: assumes coefficients of the polynomials of b are less than 2q.
func (v *VecK) Sub(a, b *VecK) {
for i := 0; i < K; i++ {
v[i].Sub(&a[i], &b[i])
}
}
// Sets v to 2ᵈ w without reducing.
func (v *VecK) MulBy2toD(w *VecK) {
for i := 0; i < K; i++ {
v[i].MulBy2toD(&w[i])
}
}
// Applies InvNTT componentwise. See Poly.InvNTT() for details.
func (v *VecK) InvNTT() {
for i := 0; i < K; i++ {
v[i].InvNTT()
}
}
// Applies Poly.ReduceLe2Q() componentwise.
func (v *VecK) ReduceLe2Q() {
for i := 0; i < K; i++ {
v[i].ReduceLe2Q()
}
}
+366
View File
@@ -0,0 +1,366 @@
// Code generated from pkg.templ.go. DO NOT EDIT.
// mldsa87 implements NIST signature scheme ML-DSA-87 as defined in FIPS204.
package mldsa87
import (
"crypto"
cryptoRand "crypto/rand"
"encoding/asn1"
"errors"
"io"
"github.com/cloudflare/circl/sign"
common "github.com/cloudflare/circl/sign/internal/dilithium"
"github.com/cloudflare/circl/sign/mldsa/mldsa87/internal"
)
const (
// Size of seed for NewKeyFromSeed
SeedSize = common.SeedSize
// Size of a packed PublicKey
PublicKeySize = internal.PublicKeySize
// Size of a packed PrivateKey
PrivateKeySize = internal.PrivateKeySize
// Size of a signature
SignatureSize = internal.SignatureSize
)
// PublicKey is the type of ML-DSA-87 public key
type PublicKey internal.PublicKey
// PrivateKey is the type of ML-DSA-87 private key
type PrivateKey internal.PrivateKey
// GenerateKey generates a public/private key pair using entropy from rand.
// If rand is nil, crypto/rand.Reader will be used.
func GenerateKey(rand io.Reader) (*PublicKey, *PrivateKey, error) {
pk, sk, err := internal.GenerateKey(rand)
return (*PublicKey)(pk), (*PrivateKey)(sk), err
}
// NewKeyFromSeed derives a public/private key pair using the given seed.
func NewKeyFromSeed(seed *[SeedSize]byte) (*PublicKey, *PrivateKey) {
pk, sk := internal.NewKeyFromSeed(seed)
return (*PublicKey)(pk), (*PrivateKey)(sk)
}
// SignTo signs the given message and writes the signature into signature.
// It will panic if signature is not of length at least SignatureSize.
//
// ctx is the optional context string. Errors if ctx is larger than 255 bytes.
// A nil context string is equivalent to an empty context string.
func SignTo(sk *PrivateKey, msg, ctx []byte, randomized bool, sig []byte) error {
var rnd [32]byte
if randomized {
_, err := cryptoRand.Read(rnd[:])
if err != nil {
return err
}
}
if len(ctx) > 255 {
return sign.ErrContextTooLong
}
internal.SignTo(
(*internal.PrivateKey)(sk),
func(w io.Writer) {
_, _ = w.Write([]byte{0})
_, _ = w.Write([]byte{byte(len(ctx))})
if ctx != nil {
_, _ = w.Write(ctx)
}
w.Write(msg)
},
rnd,
sig,
)
return nil
}
// Do not use. Implements ML-DSA.Sign_internal used for compatibility tests.
func (sk *PrivateKey) unsafeSignInternal(msg []byte, rnd [32]byte) []byte {
var ret [SignatureSize]byte
internal.SignTo(
(*internal.PrivateKey)(sk),
func(w io.Writer) {
_, _ = w.Write(msg)
},
rnd,
ret[:],
)
return ret[:]
}
// Do not use. Implements ML-DSA.Verify_internal used for compatibility tests.
func unsafeVerifyInternal(pk *PublicKey, msg, sig []byte) bool {
return internal.Verify(
(*internal.PublicKey)(pk),
func(w io.Writer) {
_, _ = w.Write(msg)
},
sig,
)
}
// Verify checks whether the given signature by pk on msg is valid.
//
// ctx is the optional context string. Fails if ctx is larger than 255 bytes.
// A nil context string is equivalent to an empty context string.
func Verify(pk *PublicKey, msg, ctx, sig []byte) bool {
if len(ctx) > 255 {
return false
}
return internal.Verify(
(*internal.PublicKey)(pk),
func(w io.Writer) {
_, _ = w.Write([]byte{0})
_, _ = w.Write([]byte{byte(len(ctx))})
if ctx != nil {
_, _ = w.Write(ctx)
}
_, _ = w.Write(msg)
},
sig,
)
}
// Sets pk to the public key encoded in buf.
func (pk *PublicKey) Unpack(buf *[PublicKeySize]byte) {
(*internal.PublicKey)(pk).Unpack(buf)
}
// Sets sk to the private key encoded in buf.
func (sk *PrivateKey) Unpack(buf *[PrivateKeySize]byte) {
(*internal.PrivateKey)(sk).Unpack(buf)
}
// Packs the public key into buf.
func (pk *PublicKey) Pack(buf *[PublicKeySize]byte) {
(*internal.PublicKey)(pk).Pack(buf)
}
// Packs the private key into buf.
func (sk *PrivateKey) Pack(buf *[PrivateKeySize]byte) {
(*internal.PrivateKey)(sk).Pack(buf)
}
// Packs the public key.
func (pk *PublicKey) Bytes() []byte {
var buf [PublicKeySize]byte
pk.Pack(&buf)
return buf[:]
}
// Packs the private key.
func (sk *PrivateKey) Bytes() []byte {
var buf [PrivateKeySize]byte
sk.Pack(&buf)
return buf[:]
}
// Packs the public key.
func (pk *PublicKey) MarshalBinary() ([]byte, error) {
return pk.Bytes(), nil
}
// Packs the private key.
func (sk *PrivateKey) MarshalBinary() ([]byte, error) {
return sk.Bytes(), nil
}
// Unpacks the public key from data.
func (pk *PublicKey) UnmarshalBinary(data []byte) error {
if len(data) != PublicKeySize {
return errors.New("packed public key must be of mldsa87.PublicKeySize bytes")
}
var buf [PublicKeySize]byte
copy(buf[:], data)
pk.Unpack(&buf)
return nil
}
// Unpacks the private key from data.
func (sk *PrivateKey) UnmarshalBinary(data []byte) error {
if len(data) != PrivateKeySize {
return errors.New("packed private key must be of mldsa87.PrivateKeySize bytes")
}
var buf [PrivateKeySize]byte
copy(buf[:], data)
sk.Unpack(&buf)
return nil
}
// Returns seed used to generate PrivateKey, and nil if not retained.
func (sk *PrivateKey) Seed() []byte {
return (*internal.PrivateKey)(sk).Seed()
}
// Sign signs the given message.
//
// opts.HashFunc() must return zero, which can be achieved by passing
// crypto.Hash(0) or nil for opts. rand is ignored. Will only return an error
// if opts.HashFunc() is non-zero.
//
// This function is used to make PrivateKey implement the crypto.Signer
// interface. The package-level SignTo function might be more convenient
// to use.
func (sk *PrivateKey) Sign(rand io.Reader, msg []byte, opts crypto.SignerOpts) (
sig []byte, err error) {
var ret [SignatureSize]byte
if opts != nil && opts.HashFunc() != crypto.Hash(0) {
return nil, errors.New("dilithium: cannot sign hashed message")
}
if err = SignTo(sk, msg, nil, false, ret[:]); err != nil {
return nil, err
}
return ret[:], nil
}
// Computes the public key corresponding to this private key.
//
// Returns a *PublicKey. The type crypto.PublicKey is used to make
// PrivateKey implement the crypto.Signer interface.
func (sk *PrivateKey) Public() crypto.PublicKey {
return (*PublicKey)((*internal.PrivateKey)(sk).Public())
}
// Equal returns whether the two private keys equal.
func (sk *PrivateKey) Equal(other crypto.PrivateKey) bool {
castOther, ok := other.(*PrivateKey)
if !ok {
return false
}
return (*internal.PrivateKey)(sk).Equal((*internal.PrivateKey)(castOther))
}
// Equal returns whether the two public keys equal.
func (pk *PublicKey) Equal(other crypto.PublicKey) bool {
castOther, ok := other.(*PublicKey)
if !ok {
return false
}
return (*internal.PublicKey)(pk).Equal((*internal.PublicKey)(castOther))
}
// Boilerplate for generic signatures API
type scheme struct{}
var sch sign.Scheme = &scheme{}
// Scheme returns a generic signature interface for ML-DSA-87.
func Scheme() sign.Scheme { return sch }
func (*scheme) Name() string { return "ML-DSA-87" }
func (*scheme) PublicKeySize() int { return PublicKeySize }
func (*scheme) PrivateKeySize() int { return PrivateKeySize }
func (*scheme) SignatureSize() int { return SignatureSize }
func (*scheme) SeedSize() int { return SeedSize }
// TODO TLSIdentifier()
func (*scheme) Oid() asn1.ObjectIdentifier {
return asn1.ObjectIdentifier{2, 16, 840, 1, 101, 3, 4, 3, 19}
}
func (*scheme) SupportsContext() bool {
return true
}
func (*scheme) GenerateKey() (sign.PublicKey, sign.PrivateKey, error) {
return GenerateKey(nil)
}
func (*scheme) Sign(
sk sign.PrivateKey,
msg []byte,
opts *sign.SignatureOpts,
) []byte {
var ctx []byte
sig := make([]byte, SignatureSize)
priv, ok := sk.(*PrivateKey)
if !ok {
panic(sign.ErrTypeMismatch)
}
if opts != nil && opts.Context != "" {
ctx = []byte(opts.Context)
}
err := SignTo(priv, msg, ctx, false, sig)
if err != nil {
panic(err)
}
return sig
}
func (*scheme) Verify(
pk sign.PublicKey,
msg, sig []byte,
opts *sign.SignatureOpts,
) bool {
var ctx []byte
pub, ok := pk.(*PublicKey)
if !ok {
panic(sign.ErrTypeMismatch)
}
if opts != nil && opts.Context != "" {
ctx = []byte(opts.Context)
}
return Verify(pub, msg, ctx, sig)
}
func (*scheme) DeriveKey(seed []byte) (sign.PublicKey, sign.PrivateKey) {
if len(seed) != SeedSize {
panic(sign.ErrSeedSize)
}
var seed2 [SeedSize]byte
copy(seed2[:], seed)
return NewKeyFromSeed(&seed2)
}
func (*scheme) UnmarshalBinaryPublicKey(buf []byte) (sign.PublicKey, error) {
if len(buf) != PublicKeySize {
return nil, sign.ErrPubKeySize
}
var (
buf2 [PublicKeySize]byte
ret PublicKey
)
copy(buf2[:], buf)
ret.Unpack(&buf2)
return &ret, nil
}
func (*scheme) UnmarshalBinaryPrivateKey(buf []byte) (sign.PrivateKey, error) {
if len(buf) != PrivateKeySize {
return nil, sign.ErrPrivKeySize
}
var (
buf2 [PrivateKeySize]byte
ret PrivateKey
)
copy(buf2[:], buf)
ret.Unpack(&buf2)
return &ret, nil
}
func (sk *PrivateKey) Scheme() sign.Scheme {
return sch
}
func (sk *PublicKey) Scheme() sign.Scheme {
return sch
}
Loaded 100 of 396 files, more files were not shown because too many files have changed in this diff. Show more