Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
97 changes: 97 additions & 0 deletions boring/src/mldsa.rs
Original file line number Diff line number Diff line change
Expand Up @@ -221,6 +221,48 @@ impl MlDsaPrivateKey {
&self.seed
}

/// Returns the public key corresponding to this private key.
pub fn public_key(&self) -> Result<MlDsaPublicKey, ErrorStack> {
unsafe {
ffi::init();
match &self.inner {
PrivateKeyInner::MlDsa44(key) => {
let mut pub_key: MaybeUninit<ffi::MLDSA44_public_key> = 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<ffi::MLDSA65_public_key> = 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<ffi::MLDSA87_public_key> = 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<Vec<u8>, ErrorStack> {
unsafe {
Expand Down Expand Up @@ -330,6 +372,41 @@ impl MlDsaPublicKey {
self.algorithm
}

/// Returns the serialized form of this public key.
pub fn to_bytes(&self) -> Result<Vec<u8>, ErrorStack> {
unsafe {
ffi::init();
let mut bytes = vec![0u8; self.algorithm.public_key_bytes()];
let mut cbb: MaybeUninit<ffi::CBB> = 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 {
Expand Down Expand Up @@ -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();
Expand Down
Loading