diff --git a/boring/src/mldsa.rs b/boring/src/mldsa.rs index db3ace0a2..8c33d1b31 100644 --- a/boring/src/mldsa.rs +++ b/boring/src/mldsa.rs @@ -221,6 +221,48 @@ impl MlDsaPrivateKey { &self.seed } + /// Returns the public key corresponding to this private key. + pub fn public_key(&self) -> Result { + unsafe { + ffi::init(); + match &self.inner { + PrivateKeyInner::MlDsa44(key) => { + let mut pub_key: MaybeUninit = MaybeUninit::uninit(); + cvt(ffi::MLDSA44_public_from_private( + pub_key.as_mut_ptr(), + key.as_ref(), + ))?; + Ok(MlDsaPublicKey { + algorithm: self.algorithm, + inner: PublicKeyInner::MlDsa44(Box::new(pub_key.assume_init())), + }) + } + PrivateKeyInner::MlDsa65(key) => { + let mut pub_key: MaybeUninit = MaybeUninit::uninit(); + cvt(ffi::MLDSA65_public_from_private( + pub_key.as_mut_ptr(), + key.as_ref(), + ))?; + Ok(MlDsaPublicKey { + algorithm: self.algorithm, + inner: PublicKeyInner::MlDsa65(Box::new(pub_key.assume_init())), + }) + } + PrivateKeyInner::MlDsa87(key) => { + let mut pub_key: MaybeUninit = MaybeUninit::uninit(); + cvt(ffi::MLDSA87_public_from_private( + pub_key.as_mut_ptr(), + key.as_ref(), + ))?; + Ok(MlDsaPublicKey { + algorithm: self.algorithm, + inner: PublicKeyInner::MlDsa87(Box::new(pub_key.assume_init())), + }) + } + } + } + } + /// Signs `msg` and returns the signature bytes. pub fn sign(&self, msg: &[u8]) -> Result, ErrorStack> { unsafe { @@ -330,6 +372,41 @@ impl MlDsaPublicKey { self.algorithm } + /// Returns the serialized form of this public key. + pub fn to_bytes(&self) -> Result, ErrorStack> { + unsafe { + ffi::init(); + let mut bytes = vec![0u8; self.algorithm.public_key_bytes()]; + let mut cbb: MaybeUninit = MaybeUninit::uninit(); + cvt(ffi::CBB_init_fixed( + cbb.as_mut_ptr(), + bytes.as_mut_ptr(), + bytes.len(), + ))?; + match &self.inner { + PublicKeyInner::MlDsa44(key) => { + cvt(ffi::MLDSA44_marshal_public_key( + cbb.as_mut_ptr(), + key.as_ref(), + ))?; + } + PublicKeyInner::MlDsa65(key) => { + cvt(ffi::MLDSA65_marshal_public_key( + cbb.as_mut_ptr(), + key.as_ref(), + ))?; + } + PublicKeyInner::MlDsa87(key) => { + cvt(ffi::MLDSA87_marshal_public_key( + cbb.as_mut_ptr(), + key.as_ref(), + ))?; + } + } + Ok(bytes) + } + } + /// Verifies `signature` over `msg` using this public key. pub fn verify(&self, msg: &[u8], signature: &[u8]) -> Result<(), ErrorStack> { unsafe { @@ -447,6 +524,26 @@ mod tests { assert!(pk.verify(msg, &sig2).is_ok()); } + #[test] + fn public_key_roundtrip() { + let (pk, sk) = MlDsaPrivateKey::generate($alg).unwrap(); + let bytes = pk.to_bytes().unwrap(); + assert_eq!(bytes.len(), $alg.public_key_bytes()); + let pk2 = MlDsaPublicKey::from_slice($alg, &bytes).unwrap(); + assert_eq!(pk2.to_bytes().unwrap(), bytes); + let msg = b"public key roundtrip"; + let sig = sk.sign(msg).unwrap(); + assert!(pk2.verify(msg, &sig).is_ok()); + } + + #[test] + fn public_from_private() { + let (pk, sk) = MlDsaPrivateKey::generate($alg).unwrap(); + let derived = sk.public_key().unwrap(); + assert_eq!(derived.algorithm(), $alg); + assert_eq!(derived.to_bytes().unwrap(), pk.to_bytes().unwrap()); + } + #[test] fn debug_redacts_seed() { let (_, sk) = MlDsaPrivateKey::generate($alg).unwrap();