use std::fmt;
use std::mem::MaybeUninit;
use crate::cvt;
use crate::error::ErrorStack;
use crate::ffi;
use crate::ffi::cbs_init;
pub const PRIVATE_KEY_SEED_BYTES: usize = ffi::MLDSA_SEED_BYTES as usize;
pub type MlDsaPrivateKeySeed = [u8; PRIVATE_KEY_SEED_BYTES];
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Algorithm {
MlDsa44,
MlDsa65,
MlDsa87,
}
impl Algorithm {
#[must_use]
pub const fn public_key_bytes(&self) -> usize {
match self {
Self::MlDsa44 => ffi::MLDSA44_PUBLIC_KEY_BYTES as usize,
Self::MlDsa65 => ffi::MLDSA65_PUBLIC_KEY_BYTES as usize,
Self::MlDsa87 => ffi::MLDSA87_PUBLIC_KEY_BYTES as usize,
}
}
#[must_use]
pub const fn signature_bytes(&self) -> usize {
match self {
Self::MlDsa44 => ffi::MLDSA44_SIGNATURE_BYTES as usize,
Self::MlDsa65 => ffi::MLDSA65_SIGNATURE_BYTES as usize,
Self::MlDsa87 => ffi::MLDSA87_SIGNATURE_BYTES as usize,
}
}
}
#[derive(Clone)]
pub struct MlDsaPublicKey {
algorithm: Algorithm,
inner: PublicKeyInner,
}
#[derive(Clone)]
enum PublicKeyInner {
MlDsa44(Box<ffi::MLDSA44_public_key>),
MlDsa65(Box<ffi::MLDSA65_public_key>),
MlDsa87(Box<ffi::MLDSA87_public_key>),
}
pub struct MlDsaPrivateKey {
algorithm: Algorithm,
seed: MlDsaPrivateKeySeed,
inner: PrivateKeyInner,
}
enum PrivateKeyInner {
MlDsa44(Box<ffi::MLDSA44_private_key>),
MlDsa65(Box<ffi::MLDSA65_private_key>),
MlDsa87(Box<ffi::MLDSA87_private_key>),
}
impl Clone for MlDsaPrivateKey {
fn clone(&self) -> Self {
Self::from_seed(self.algorithm, &self.seed).unwrap()
}
}
impl MlDsaPrivateKey {
pub fn generate(algorithm: Algorithm) -> Result<(MlDsaPublicKey, MlDsaPrivateKey), ErrorStack> {
unsafe {
ffi::init();
match algorithm {
Algorithm::MlDsa44 => {
let mut pub_bytes = [0u8; ffi::MLDSA44_PUBLIC_KEY_BYTES as usize];
let mut seed = [0u8; PRIVATE_KEY_SEED_BYTES];
let mut priv_key: MaybeUninit<ffi::MLDSA44_private_key> = MaybeUninit::uninit();
cvt(ffi::MLDSA44_generate_key(
pub_bytes.as_mut_ptr(),
seed.as_mut_ptr(),
priv_key.as_mut_ptr(),
))?;
let public_key = MlDsaPublicKey::from_slice(algorithm, &pub_bytes)?;
Ok((
public_key,
MlDsaPrivateKey {
algorithm,
seed,
inner: PrivateKeyInner::MlDsa44(Box::new(priv_key.assume_init())),
},
))
}
Algorithm::MlDsa65 => {
let mut pub_bytes = [0u8; ffi::MLDSA65_PUBLIC_KEY_BYTES as usize];
let mut seed = [0u8; PRIVATE_KEY_SEED_BYTES];
let mut priv_key: MaybeUninit<ffi::MLDSA65_private_key> = MaybeUninit::uninit();
cvt(ffi::MLDSA65_generate_key(
pub_bytes.as_mut_ptr(),
seed.as_mut_ptr(),
priv_key.as_mut_ptr(),
))?;
let public_key = MlDsaPublicKey::from_slice(algorithm, &pub_bytes)?;
Ok((
public_key,
MlDsaPrivateKey {
algorithm,
seed,
inner: PrivateKeyInner::MlDsa65(Box::new(priv_key.assume_init())),
},
))
}
Algorithm::MlDsa87 => {
let mut pub_bytes = [0u8; ffi::MLDSA87_PUBLIC_KEY_BYTES as usize];
let mut seed = [0u8; PRIVATE_KEY_SEED_BYTES];
let mut priv_key: MaybeUninit<ffi::MLDSA87_private_key> = MaybeUninit::uninit();
cvt(ffi::MLDSA87_generate_key(
pub_bytes.as_mut_ptr(),
seed.as_mut_ptr(),
priv_key.as_mut_ptr(),
))?;
let public_key = MlDsaPublicKey::from_slice(algorithm, &pub_bytes)?;
Ok((
public_key,
MlDsaPrivateKey {
algorithm,
seed,
inner: PrivateKeyInner::MlDsa87(Box::new(priv_key.assume_init())),
},
))
}
}
}
}
pub fn from_seed(algorithm: Algorithm, seed: &MlDsaPrivateKeySeed) -> Result<Self, ErrorStack> {
unsafe {
ffi::init();
match algorithm {
Algorithm::MlDsa44 => {
let mut priv_key: MaybeUninit<ffi::MLDSA44_private_key> = MaybeUninit::uninit();
cvt(ffi::MLDSA44_private_key_from_seed(
priv_key.as_mut_ptr(),
seed.as_ptr(),
seed.len(),
))?;
Ok(Self {
algorithm,
seed: *seed,
inner: PrivateKeyInner::MlDsa44(Box::new(priv_key.assume_init())),
})
}
Algorithm::MlDsa65 => {
let mut priv_key: MaybeUninit<ffi::MLDSA65_private_key> = MaybeUninit::uninit();
cvt(ffi::MLDSA65_private_key_from_seed(
priv_key.as_mut_ptr(),
seed.as_ptr(),
seed.len(),
))?;
Ok(Self {
algorithm,
seed: *seed,
inner: PrivateKeyInner::MlDsa65(Box::new(priv_key.assume_init())),
})
}
Algorithm::MlDsa87 => {
let mut priv_key: MaybeUninit<ffi::MLDSA87_private_key> = MaybeUninit::uninit();
cvt(ffi::MLDSA87_private_key_from_seed(
priv_key.as_mut_ptr(),
seed.as_ptr(),
seed.len(),
))?;
Ok(Self {
algorithm,
seed: *seed,
inner: PrivateKeyInner::MlDsa87(Box::new(priv_key.assume_init())),
})
}
}
}
}
pub fn algorithm(&self) -> Algorithm {
self.algorithm
}
pub fn seed_bytes(&self) -> &MlDsaPrivateKeySeed {
&self.seed
}
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())),
})
}
}
}
}
pub fn sign(&self, msg: &[u8]) -> Result<Vec<u8>, ErrorStack> {
unsafe {
ffi::init();
match &self.inner {
PrivateKeyInner::MlDsa44(key) => {
let mut sig = vec![0u8; ffi::MLDSA44_SIGNATURE_BYTES as usize];
cvt(ffi::MLDSA44_sign(
sig.as_mut_ptr(),
key.as_ref(),
msg.as_ptr(),
msg.len(),
core::ptr::null(),
0,
))?;
Ok(sig)
}
PrivateKeyInner::MlDsa65(key) => {
let mut sig = vec![0u8; ffi::MLDSA65_SIGNATURE_BYTES as usize];
cvt(ffi::MLDSA65_sign(
sig.as_mut_ptr(),
key.as_ref(),
msg.as_ptr(),
msg.len(),
core::ptr::null(),
0,
))?;
Ok(sig)
}
PrivateKeyInner::MlDsa87(key) => {
let mut sig = vec![0u8; ffi::MLDSA87_SIGNATURE_BYTES as usize];
cvt(ffi::MLDSA87_sign(
sig.as_mut_ptr(),
key.as_ref(),
msg.as_ptr(),
msg.len(),
core::ptr::null(),
0,
))?;
Ok(sig)
}
}
}
}
}
impl MlDsaPublicKey {
pub fn from_slice(
algorithm: Algorithm,
serialized_public_key: &[u8],
) -> Result<Self, ErrorStack> {
ffi::init();
if serialized_public_key.len() != algorithm.public_key_bytes() {
return Err(ErrorStack::internal_error_str("invalid public key length"));
}
let mut cbs = cbs_init(serialized_public_key);
unsafe {
match algorithm {
Algorithm::MlDsa44 => {
let mut key: MaybeUninit<ffi::MLDSA44_public_key> = MaybeUninit::uninit();
cvt(ffi::MLDSA44_parse_public_key(key.as_mut_ptr(), &mut cbs))?;
if cbs.len != 0 {
return Err(ErrorStack::internal_error_str(
"trailing bytes after ML-DSA-44 public key",
));
}
Ok(Self {
algorithm,
inner: PublicKeyInner::MlDsa44(Box::new(key.assume_init())),
})
}
Algorithm::MlDsa65 => {
let mut key: MaybeUninit<ffi::MLDSA65_public_key> = MaybeUninit::uninit();
cvt(ffi::MLDSA65_parse_public_key(key.as_mut_ptr(), &mut cbs))?;
if cbs.len != 0 {
return Err(ErrorStack::internal_error_str(
"trailing bytes after ML-DSA-65 public key",
));
}
Ok(Self {
algorithm,
inner: PublicKeyInner::MlDsa65(Box::new(key.assume_init())),
})
}
Algorithm::MlDsa87 => {
let mut key: MaybeUninit<ffi::MLDSA87_public_key> = MaybeUninit::uninit();
cvt(ffi::MLDSA87_parse_public_key(key.as_mut_ptr(), &mut cbs))?;
if cbs.len != 0 {
return Err(ErrorStack::internal_error_str(
"trailing bytes after ML-DSA-87 public key",
));
}
Ok(Self {
algorithm,
inner: PublicKeyInner::MlDsa87(Box::new(key.assume_init())),
})
}
}
}
}
pub fn algorithm(&self) -> Algorithm {
self.algorithm
}
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)
}
}
pub fn verify(&self, msg: &[u8], signature: &[u8]) -> Result<(), ErrorStack> {
unsafe {
ffi::init();
match &self.inner {
PublicKeyInner::MlDsa44(key) => {
cvt(ffi::MLDSA44_verify(
key.as_ref(),
signature.as_ptr(),
signature.len(),
msg.as_ptr(),
msg.len(),
core::ptr::null(),
0,
))?;
}
PublicKeyInner::MlDsa65(key) => {
cvt(ffi::MLDSA65_verify(
key.as_ref(),
signature.as_ptr(),
signature.len(),
msg.as_ptr(),
msg.len(),
core::ptr::null(),
0,
))?;
}
PublicKeyInner::MlDsa87(key) => {
cvt(ffi::MLDSA87_verify(
key.as_ref(),
signature.as_ptr(),
signature.len(),
msg.as_ptr(),
msg.len(),
core::ptr::null(),
0,
))?;
}
}
Ok(())
}
}
}
impl fmt::Debug for MlDsaPrivateKey {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("MlDsaPrivateKey")
.field("algorithm", &self.algorithm)
.field("seed", &"[redacted]")
.finish()
}
}
impl Drop for MlDsaPrivateKey {
fn drop(&mut self) {
unsafe {
ffi::OPENSSL_cleanse(self.seed.as_mut_ptr().cast(), self.seed.len());
}
}
}
impl fmt::Debug for MlDsaPublicKey {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("MlDsaPublicKey")
.field("algorithm", &self.algorithm)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
macro_rules! mldsa_tests {
($name:ident, $alg:expr) => {
mod $name {
use super::*;
#[test]
fn sign_and_verify() {
let (pk, sk) = MlDsaPrivateKey::generate($alg).unwrap();
let msg = b"test message";
let sig1 = sk.sign(msg).unwrap();
let sig2 = sk.clone().sign(msg).unwrap();
assert_eq!(sig1.len(), $alg.signature_bytes());
assert!(pk.verify(msg, &sig1).is_ok());
assert!(pk.verify(msg, &sig2).is_ok());
assert!(pk.clone().verify(msg, &sig1).is_ok());
}
#[test]
fn bad_signature_fails() {
let (pk, sk) = MlDsaPrivateKey::generate($alg).unwrap();
let msg = b"test message";
let mut sig = sk.sign(msg).unwrap();
sig[5] ^= 1;
assert!(pk.verify(msg, &sig).is_err());
}
#[test]
fn wrong_message_fails() {
let (pk, sk) = MlDsaPrivateKey::generate($alg).unwrap();
let sig = sk.sign(b"correct").unwrap();
assert!(pk.verify(b"wrong", &sig).is_err());
}
#[test]
fn seed_roundtrip() {
let (pk, sk) = MlDsaPrivateKey::generate($alg).unwrap();
let sk2 = MlDsaPrivateKey::from_seed($alg, sk.seed_bytes()).unwrap();
let msg = b"seed roundtrip";
let sig1 = sk2.sign(msg).unwrap();
let sig2 = sk2.clone().sign(msg).unwrap();
assert!(pk.verify(msg, &sig1).is_ok());
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();
let dbg = format!("{:?}", sk);
assert!(dbg.contains("redacted"));
assert!(!dbg.contains(&format!("{:?}", sk.seed_bytes())));
}
}
};
}
mldsa_tests!(mldsa44, Algorithm::MlDsa44);
mldsa_tests!(mldsa65, Algorithm::MlDsa65);
mldsa_tests!(mldsa87, Algorithm::MlDsa87);
}