diff --git a/common.go b/common.go index 9ea149a..c3ed291 100644 --- a/common.go +++ b/common.go @@ -155,11 +155,12 @@ const ( X25519MLKEM768 CurveID = 4588 SecP256r1MLKEM768 CurveID = 4587 SecP384r1MLKEM1024 CurveID = 4589 + MLKEM1024 CurveID = 514 ) func isTLS13OnlyKeyExchange(curve CurveID) bool { switch curve { - case X25519MLKEM768, SecP256r1MLKEM768, SecP384r1MLKEM1024: + case X25519MLKEM768, SecP256r1MLKEM768, SecP384r1MLKEM1024, MLKEM1024: return true default: return false @@ -168,7 +169,7 @@ func isTLS13OnlyKeyExchange(curve CurveID) bool { func isPQKeyExchange(curve CurveID) bool { switch curve { - case X25519MLKEM768, SecP256r1MLKEM768, SecP384r1MLKEM1024: + case X25519MLKEM768, SecP256r1MLKEM768, SecP384r1MLKEM1024, MLKEM1024: return true default: return false diff --git a/common_string.go b/common_string.go index bab3b84..09ce16c 100644 --- a/common_string.go +++ b/common_string.go @@ -82,17 +82,19 @@ func _() { _ = x[X25519MLKEM768-4588] _ = x[SecP256r1MLKEM768-4587] _ = x[SecP384r1MLKEM1024-4589] + _ = x[MLKEM1024-514] } const ( _CurveID_name_0 = "CurveP256CurveP384CurveP521" _CurveID_name_1 = "X25519" - _CurveID_name_2 = "SecP256r1MLKEM768X25519MLKEM768SecP384r1MLKEM1024" + _CurveID_name_2 = "MLKEM1024" + _CurveID_name_3 = "SecP256r1MLKEM768X25519MLKEM768SecP384r1MLKEM1024" ) var ( _CurveID_index_0 = [...]uint8{0, 9, 18, 27} - _CurveID_index_2 = [...]uint8{0, 17, 31, 49} + _CurveID_index_3 = [...]uint8{0, 17, 31, 49} ) func (i CurveID) String() string { @@ -102,9 +104,11 @@ func (i CurveID) String() string { return _CurveID_name_0[_CurveID_index_0[i]:_CurveID_index_0[i+1]] case i == 29: return _CurveID_name_1 + case i == 514: + return _CurveID_name_2 case 4587 <= i && i <= 4589: i -= 4587 - return _CurveID_name_2[_CurveID_index_2[i]:_CurveID_index_2[i+1]] + return _CurveID_name_3[_CurveID_index_3[i]:_CurveID_index_3[i+1]] default: return "CurveID(" + strconv.FormatInt(int64(i), 10) + ")" } diff --git a/defaults.go b/defaults.go index 17ab1d6..ee8f81b 100644 --- a/defaults.go +++ b/defaults.go @@ -37,7 +37,7 @@ func defaultCurveEnabled(c CurveID) bool { // include every supported key exchange. func curvePreferenceOrder() []CurveID { return []CurveID{ - X25519MLKEM768, SecP256r1MLKEM768, SecP384r1MLKEM1024, + X25519MLKEM768, SecP256r1MLKEM768, SecP384r1MLKEM1024, MLKEM1024, X25519, CurveP256, CurveP384, CurveP521, } } diff --git a/defaults_fips140.go b/defaults_fips140.go index 63bebd0..901e81d 100644 --- a/defaults_fips140.go +++ b/defaults_fips140.go @@ -35,6 +35,7 @@ var ( X25519MLKEM768, SecP256r1MLKEM768, SecP384r1MLKEM1024, + MLKEM1024, CurveP256, CurveP384, CurveP521, diff --git a/handshake_client.go b/handshake_client.go index 94cd56a..0f47b51 100644 --- a/handshake_client.go +++ b/handshake_client.go @@ -140,7 +140,7 @@ func (c *Conn) makeClientHello() (*clientHelloMsg, *keySharePrivateKeys, *echCli } if len(hello.supportedCurves) == 0 { - return nil, nil, nil, errors.New("tls: no supported elliptic curves for ECDHE") + return nil, nil, nil, errors.New("tls: no supported key exchange methods (CurveIDs)") } // Since the order is fixed, the first one is always the one to send a // key share for. All the PQ hybrids sort first, and produce a fallback diff --git a/handshake_client_tls13.go b/handshake_client_tls13.go index e6db8e9..c15490b 100644 --- a/handshake_client_tls13.go +++ b/handshake_client_tls13.go @@ -54,7 +54,8 @@ func (hs *clientHandshakeStateTLS13) handshake() error { } // Consistency check on the presence of a keyShare and its parameters. - if hs.keyShareKeys == nil || hs.keyShareKeys.ecdhe == nil || len(hs.hello.keyShares) == 0 { + if hs.keyShareKeys == nil || (hs.keyShareKeys.ecdhe == nil && hs.keyShareKeys.mlkem == nil) || + len(hs.hello.keyShares) == 0 { return c.sendAlert(alertInternalError) } diff --git a/key_schedule.go b/key_schedule.go index b798bdd..c2d09ca 100644 --- a/key_schedule.go +++ b/key_schedule.go @@ -108,11 +108,40 @@ func keyExchangeForCurveID(id CurveID) (keyExchange, error) { return &hybridKeyExchange{id, ecdhKeyExchange{CurveP384, ecdh.P384()}, 97, mlkem.EncapsulationKeySize1024, mlkem.CiphertextSize1024, mlkemGenerateKey1024, mlkemNewPublicKey1024}, nil + case MLKEM1024: + return &mlkem1024KeyExchange{}, nil default: return nil, errors.New("tls: unsupported key exchange") } } +type mlkem1024KeyExchange struct{} + +func (ke *mlkem1024KeyExchange) keyShares(_ io.Reader) (*keySharePrivateKeys, []keyShare, error) { + priv, err := mlkem.GenerateKey1024() + if err != nil { + return nil, nil, err + } + return &keySharePrivateKeys{mlkem: priv}, []keyShare{{MLKEM1024, priv.EncapsulationKey().Bytes()}}, nil +} + +func (ke *mlkem1024KeyExchange) serverSharedSecret(_ io.Reader, clientKeyShare []byte) ([]byte, keyShare, error) { + peerKey, err := mlkem.NewEncapsulationKey1024(clientKeyShare) + if err != nil { + return nil, keyShare{}, err + } + sharedKey, keyShareData := peerKey.Encapsulate() + return sharedKey, keyShare{MLKEM1024, keyShareData}, nil +} + +func (ke *mlkem1024KeyExchange) clientSharedSecret(priv *keySharePrivateKeys, serverKeyShare []byte) ([]byte, error) { + sharedKey, err := priv.mlkem.Decapsulate(serverKeyShare) + if err != nil { + return nil, err + } + return sharedKey, nil +} + type ecdhKeyExchange struct { id CurveID curve ecdh.Curve