use crate::error::{Error, Result};
use zeroize::{Zeroize, ZeroizeOnDrop, Zeroizing};
#[derive(Clone, PartialEq, Eq)]
pub struct PublicKey(pub(crate) Vec<u8>);
#[derive(Zeroize, ZeroizeOnDrop)]
pub struct SecretKey(pub(crate) Vec<u8>);
#[derive(Clone, PartialEq, Eq)]
pub struct Signature(pub(crate) Vec<u8>);
const PK_LEN: usize = 1952;
const SK_LEN: usize = 32;
const SIG_LEN: usize = 3309;
impl PublicKey {
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
pub fn from_bytes(bytes: Vec<u8>) -> Result<Self> {
if bytes.len() != PK_LEN {
return Err(Error::InvalidLength {
expected: PK_LEN,
got: bytes.len(),
});
}
Ok(Self(bytes))
}
pub(crate) fn from_bytes_unchecked(bytes: Vec<u8>) -> Self {
assert_eq!(
bytes.len(),
PK_LEN,
"mldsa::PublicKey::from_bytes_unchecked: wrong size"
);
Self(bytes)
}
}
impl SecretKey {
pub(crate) fn as_bytes(&self) -> &[u8] {
&self.0
}
pub fn from_bytes(bytes: Vec<u8>) -> Result<Self> {
let mut bytes = Zeroizing::new(bytes);
if bytes.len() != SK_LEN {
return Err(Error::InvalidLength {
expected: SK_LEN,
got: bytes.len(),
});
}
Ok(Self(std::mem::take(&mut *bytes)))
}
}
impl Signature {
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
pub fn from_bytes(bytes: Vec<u8>) -> Result<Self> {
if bytes.len() != SIG_LEN {
return Err(Error::InvalidLength {
expected: SIG_LEN,
got: bytes.len(),
});
}
Ok(Self(bytes))
}
pub(crate) fn from_bytes_unchecked(bytes: Vec<u8>) -> Self {
assert_eq!(
bytes.len(),
SIG_LEN,
"mldsa::Signature::from_bytes_unchecked: wrong size"
);
Self(bytes)
}
}
pub const fn pk_len() -> usize {
PK_LEN
}
pub const fn sk_len() -> usize {
SK_LEN
}
pub const fn sig_len() -> usize {
SIG_LEN
}
pub fn keygen() -> Result<(PublicKey, SecretKey)> {
use ml_dsa::{B32, KeyGen, MlDsa65};
let mut seed_bytes = [0u8; 32];
super::random::random_bytes(&mut seed_bytes);
let mut seed: B32 = seed_bytes.into();
seed_bytes.zeroize();
let kp = MlDsa65::from_seed(&seed);
seed.zeroize();
let pk_bytes = kp.verifying_key().encode().to_vec();
let mut seed_out = kp.to_seed();
let mut sk_bytes = zeroize::Zeroizing::new(seed_out.to_vec());
seed_out.zeroize();
assert_eq!(
pk_bytes.len(),
PK_LEN,
"ML-DSA-65 VK size mismatch: got {}, expected {PK_LEN}",
pk_bytes.len()
);
assert_eq!(
sk_bytes.len(),
SK_LEN,
"ML-DSA-65 seed size mismatch: got {}, expected {SK_LEN}",
sk_bytes.len()
);
Ok((
PublicKey(pk_bytes),
SecretKey(std::mem::take(&mut *sk_bytes)),
))
}
pub fn sign(sk: &SecretKey, message: &[u8]) -> Result<Signature> {
use ml_dsa::{B32, MlDsa65, Seed, SigningKey};
if sk.0.len() != SK_LEN {
return Err(Error::InvalidLength {
expected: SK_LEN,
got: sk.0.len(),
});
}
let seed: &Seed = sk.0.as_slice().try_into().map_err(|_| Error::Internal)?;
let signing_key = SigningKey::<MlDsa65>::from_seed(seed);
let mut rnd_bytes = [0u8; 32];
super::random::random_bytes(&mut rnd_bytes);
let mut rnd: B32 = rnd_bytes.into();
rnd_bytes.zeroize();
let sig = signing_key.sign_internal(&[message], &rnd);
rnd.zeroize();
let sig_bytes = sig.encode().to_vec();
assert_eq!(
sig_bytes.len(),
SIG_LEN,
"ML-DSA-65 sig size mismatch: got {}, expected {SIG_LEN}",
sig_bytes.len()
);
Ok(Signature(sig_bytes))
}
pub fn verify(pk: &PublicKey, message: &[u8], signature: &Signature) -> Result<()> {
use ml_dsa::{EncodedSignature, EncodedVerifyingKey, MlDsa65, VerifyingKey};
if pk.0.len() != PK_LEN {
return Err(Error::InvalidLength {
expected: PK_LEN,
got: pk.0.len(),
});
}
if signature.0.len() != SIG_LEN {
return Err(Error::InvalidLength {
expected: SIG_LEN,
got: signature.0.len(),
});
}
let vk_enc: &EncodedVerifyingKey<MlDsa65> =
pk.0.as_slice().try_into().map_err(|_| Error::Internal)?;
let vk = VerifyingKey::<MlDsa65>::decode(vk_enc);
let sig_enc: &EncodedSignature<MlDsa65> = signature
.0
.as_slice()
.try_into()
.map_err(|_| Error::Internal)?;
let sig = ml_dsa::Signature::<MlDsa65>::decode(sig_enc).ok_or(Error::VerificationFailed)?;
if vk.verify_internal(message, &sig) {
Ok(())
} else {
Err(Error::VerificationFailed)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::Error;
#[test]
fn keygen_sizes() {
let (pk, sk) = keygen().unwrap();
assert_eq!(pk.as_bytes().len(), 1952);
assert_eq!(sk.as_bytes().len(), 32);
}
#[test]
fn sign_verify_round_trip() {
let (pk, sk) = keygen().unwrap();
let sig = sign(&sk, b"hello mldsa").unwrap();
assert!(verify(&pk, b"hello mldsa", &sig).is_ok());
}
#[test]
fn sign_size() {
let (_, sk) = keygen().unwrap();
let sig = sign(&sk, b"test").unwrap();
assert_eq!(sig.as_bytes().len(), 3309);
}
#[test]
fn sign_verify_empty_message() {
let (pk, sk) = keygen().unwrap();
let sig = sign(&sk, b"").unwrap();
assert!(verify(&pk, b"", &sig).is_ok());
}
#[test]
fn verify_wrong_message() {
let (pk, sk) = keygen().unwrap();
let sig = sign(&sk, b"message one").unwrap();
assert!(matches!(
verify(&pk, b"message two", &sig),
Err(Error::VerificationFailed)
));
}
#[test]
fn verify_wrong_key() {
let (_, sk) = keygen().unwrap();
let (pk2, _) = keygen().unwrap();
let sig = sign(&sk, b"test").unwrap();
assert!(matches!(
verify(&pk2, b"test", &sig),
Err(Error::VerificationFailed)
));
}
#[test]
fn verify_tampered_signature() {
let (pk, sk) = keygen().unwrap();
let sig = sign(&sk, b"test").unwrap();
let mut bad = sig.as_bytes().to_vec();
bad[0] ^= 0xFF;
let bad_sig = Signature::from_bytes(bad).unwrap();
assert!(matches!(
verify(&pk, b"test", &bad_sig),
Err(Error::VerificationFailed)
));
}
#[test]
fn verify_all_zeros_signature() {
let (pk, _) = keygen().unwrap();
let zeros_sig = Signature::from_bytes(vec![0u8; 3309]).unwrap();
assert!(matches!(
verify(&pk, b"test", &zeros_sig),
Err(Error::VerificationFailed)
));
}
#[test]
fn pk_from_bytes_wrong_size() {
assert!(matches!(
PublicKey::from_bytes(vec![0u8; 100]),
Err(Error::InvalidLength {
expected: 1952,
got: 100
})
));
}
#[test]
fn sig_from_bytes_wrong_size() {
assert!(matches!(
Signature::from_bytes(vec![0u8; 100]),
Err(Error::InvalidLength {
expected: 3309,
got: 100
})
));
}
#[test]
fn sk_from_bytes_wrong_size() {
assert!(matches!(
SecretKey::from_bytes(vec![0u8; 100]),
Err(Error::InvalidLength {
expected: 32,
got: 100
})
));
}
#[test]
fn sign_hedged_both_verify() {
let (pk, sk) = keygen().unwrap();
let msg = b"hedged signing test";
let sig1 = sign(&sk, msg).unwrap();
let sig2 = sign(&sk, msg).unwrap();
assert_ne!(
sig1.as_bytes(),
sig2.as_bytes(),
"hedged signatures on the same message must differ"
);
assert!(verify(&pk, msg, &sig1).is_ok());
assert!(verify(&pk, msg, &sig2).is_ok());
}
proptest::proptest! {
#[test]
fn proptest_round_trip(
msg in proptest::collection::vec(proptest::prelude::any::<u8>(), 0..1024),
flip_byte in 0..3309usize,
) {
let kg = keygen();
proptest::prop_assert!(kg.is_ok());
let (pk, sk) = kg.unwrap();
let sign_result = sign(&sk, &msg);
proptest::prop_assert!(sign_result.is_ok());
let sig = sign_result.unwrap();
proptest::prop_assert!(verify(&pk, &msg, &sig).is_ok());
let mut bad_sig_bytes = sig.as_bytes().to_vec();
bad_sig_bytes[flip_byte] ^= 0x01;
let bad_sig = Signature::from_bytes(bad_sig_bytes).unwrap();
proptest::prop_assert!(verify(&pk, &msg, &bad_sig).is_err());
if !msg.is_empty() {
let mut bad_msg = msg.clone();
bad_msg[0] ^= 0x01;
proptest::prop_assert!(verify(&pk, &bad_msg, &sig).is_err());
}
}
}
}