diff --git a/key_schedule.go b/key_schedule.go index b07a86f..b798bdd 100644 --- a/key_schedule.go +++ b/key_schedule.go @@ -75,16 +75,16 @@ type keyExchange interface { } func keyExchangeForCurveID(id CurveID) (keyExchange, error) { - newMLKEMPrivateKey768 := func(b []byte) (crypto.Decapsulator, error) { - return mlkem.NewDecapsulationKey768(b) + mlkemGenerateKey768 := func() (crypto.Decapsulator, error) { + return mlkem.GenerateKey768() } - newMLKEMPrivateKey1024 := func(b []byte) (crypto.Decapsulator, error) { - return mlkem.NewDecapsulationKey1024(b) + mlkemGenerateKey1024 := func() (crypto.Decapsulator, error) { + return mlkem.GenerateKey1024() } - newMLKEMPublicKey768 := func(b []byte) (crypto.Encapsulator, error) { + mlkemNewPublicKey768 := func(b []byte) (crypto.Encapsulator, error) { return mlkem.NewEncapsulationKey768(b) } - newMLKEMPublicKey1024 := func(b []byte) (crypto.Encapsulator, error) { + mlkemNewPublicKey1024 := func(b []byte) (crypto.Encapsulator, error) { return mlkem.NewEncapsulationKey1024(b) } switch id { @@ -99,15 +99,15 @@ func keyExchangeForCurveID(id CurveID) (keyExchange, error) { case X25519MLKEM768: return &hybridKeyExchange{id, ecdhKeyExchange{X25519, ecdh.X25519()}, 32, mlkem.EncapsulationKeySize768, mlkem.CiphertextSize768, - newMLKEMPrivateKey768, newMLKEMPublicKey768}, nil + mlkemGenerateKey768, mlkemNewPublicKey768}, nil case SecP256r1MLKEM768: return &hybridKeyExchange{id, ecdhKeyExchange{CurveP256, ecdh.P256()}, 65, mlkem.EncapsulationKeySize768, mlkem.CiphertextSize768, - newMLKEMPrivateKey768, newMLKEMPublicKey768}, nil + mlkemGenerateKey768, mlkemNewPublicKey768}, nil case SecP384r1MLKEM1024: return &hybridKeyExchange{id, ecdhKeyExchange{CurveP384, ecdh.P384()}, 97, mlkem.EncapsulationKeySize1024, mlkem.CiphertextSize1024, - newMLKEMPrivateKey1024, newMLKEMPublicKey1024}, nil + mlkemGenerateKey1024, mlkemNewPublicKey1024}, nil default: return nil, errors.New("tls: unsupported key exchange") } @@ -162,8 +162,8 @@ type hybridKeyExchange struct { mlkemPublicKeySize int mlkemCiphertextSize int - newMLKEMPrivateKey func([]byte) (crypto.Decapsulator, error) - newMLKEMPublicKey func([]byte) (crypto.Encapsulator, error) + mlkemGenerateKey func() (crypto.Decapsulator, error) + mlkemNewPublicKey func([]byte) (crypto.Encapsulator, error) } func (ke *hybridKeyExchange) keyShares(rand io.Reader) (*keySharePrivateKeys, []keyShare, error) { @@ -178,11 +178,7 @@ func (ke *hybridKeyExchange) keyShares(rand io.Reader) (*keySharePrivateKeys, [] if err != nil { return nil, nil, err } - seed := make([]byte, mlkem.SeedSize) - if _, err := io.ReadFull(rand, seed); err != nil { - return nil, nil, err - } - priv.mlkem, err = ke.newMLKEMPrivateKey(seed) + priv.mlkem, err = ke.mlkemGenerateKey() if err != nil { return nil, nil, err } @@ -221,7 +217,7 @@ func (ke *hybridKeyExchange) serverSharedSecret(rand io.Reader, clientKeyShare [ if err != nil { return nil, keyShare{}, err } - mlkemPeerKey, err := ke.newMLKEMPublicKey(mlkemShareData) + mlkemPeerKey, err := ke.mlkemNewPublicKey(mlkemShareData) if err != nil { return nil, keyShare{}, err }