use super::arithmetic::w1_bits_needed;
use super::polyvec::{PolyVecK, PolyVecL};
use crate::error::Error as SignError;
#[cfg(not(feature = "std"))]
use alloc::{format, vec, vec::Vec};
use dcrypt_algorithms::poly::serialize::{
CoefficientPacker, CoefficientUnpacker, DefaultCoefficientSerde,
};
use dcrypt_api::SecretVec;
use dcrypt_internal::{boxed_bytes_zeroed, Zeroizing, ZeroizingBytes};
use dcrypt_params::pqc::ml_dsa::{MlDsaSchemeParams, ML_DSA_N, ML_DSA_Q};
#[inline]
fn centered_coefficient(coefficient: u32) -> i32 {
let reduced = coefficient % ML_DSA_Q;
if reduced > ML_DSA_Q / 2 {
reduced as i32 - ML_DSA_Q as i32
} else {
reduced as i32
}
}
#[inline]
fn signed_to_mod_q(value: i32) -> u32 {
(value as i64).rem_euclid(ML_DSA_Q as i64) as u32
}
fn pack_hints_bitpacked<P: MlDsaSchemeParams>(
h_hint_poly: &PolyVecK<P>,
) -> Result<ZeroizingBytes, SignError> {
let omega = P::OMEGA_PARAM as usize;
let mut packed = Zeroizing::new(boxed_bytes_zeroed(omega + P::K_DIM));
let mut index = 0usize;
for (row, poly) in h_hint_poly.polys.iter().enumerate() {
for (col, &bit) in poly.coeffs.iter().enumerate() {
match bit {
0 => {}
1 => {
if index >= omega {
return Err(SignError::Serialization(
"too many ML-DSA hint coefficients".into(),
));
}
packed[index] = col as u8;
index += 1;
}
_ => {
return Err(SignError::Serialization(
"ML-DSA hint coefficients must be zero or one".into(),
));
}
}
}
packed[omega + row] = index as u8;
}
Ok(packed)
}
fn unpack_hints_bitpacked<P: MlDsaSchemeParams>(
bytes: &[u8],
) -> Result<(PolyVecK<P>, usize), SignError> {
let omega = P::OMEGA_PARAM as usize;
if bytes.len() != omega + P::K_DIM {
return Err(SignError::Deserialization(
"invalid ML-DSA hint length".into(),
));
}
let (idx_bytes, boundaries) = bytes.split_at(omega);
let mut h_poly = PolyVecK::<P>::zero();
let mut start = 0usize;
for (row, &boundary) in boundaries.iter().enumerate() {
let end = usize::from(boundary);
if end < start || end > omega {
return Err(SignError::Deserialization(
"non-monotonic ML-DSA hint boundaries".into(),
));
}
if !idx_bytes[start..end]
.windows(2)
.all(|pair| pair[0] < pair[1])
{
return Err(SignError::Deserialization(
"duplicate or unsorted ML-DSA hint indices".into(),
));
}
for &idx in &idx_bytes[start..end] {
h_poly.polys[row].coeffs[usize::from(idx)] = 1;
}
start = end;
}
if idx_bytes[start..].iter().any(|&byte| byte != 0) {
return Err(SignError::Deserialization(
"nonzero unused ML-DSA hint bytes".into(),
));
}
Ok((h_poly, start))
}
pub fn pack_public_key<P: MlDsaSchemeParams>(
rho_seed: &[u8; 32], t1_vec: &PolyVecK<P>,
) -> Result<Vec<u8>, SignError> {
let mut pk_bytes = Vec::with_capacity(P::PUBLIC_KEY_BYTES);
pk_bytes.extend_from_slice(rho_seed);
for i in 0..P::K_DIM {
let packed_poly = DefaultCoefficientSerde::pack_coeffs(&t1_vec.polys[i], 10)
.map_err(SignError::from_algo)?;
pk_bytes.extend_from_slice(&packed_poly);
}
if pk_bytes.len() != P::PUBLIC_KEY_BYTES {
return Err(SignError::Serialization(format!(
"Public key size mismatch: expected {}, got {}",
P::PUBLIC_KEY_BYTES,
pk_bytes.len()
)));
}
Ok(pk_bytes)
}
pub fn unpack_public_key<P: MlDsaSchemeParams>(
pk_bytes: &[u8],
) -> Result<([u8; 32], PolyVecK<P>), SignError> {
if pk_bytes.len() != P::PUBLIC_KEY_BYTES {
return Err(SignError::Deserialization(format!(
"Public key size mismatch: expected {}, got {}",
P::PUBLIC_KEY_BYTES,
pk_bytes.len()
)));
}
let mut rho_seed = [0u8; 32];
rho_seed.copy_from_slice(&pk_bytes[0..32]);
let mut t1_vec = PolyVecK::<P>::zero();
let mut offset = P::SEED_RHO_BYTES;
let bytes_per_poly = ML_DSA_N * 10 / 8;
for i in 0..P::K_DIM {
let poly_bytes = &pk_bytes[offset..offset + bytes_per_poly];
t1_vec.polys[i] =
DefaultCoefficientSerde::unpack_coeffs(poly_bytes, 10).map_err(SignError::from_algo)?;
offset += bytes_per_poly;
}
Ok((rho_seed, t1_vec))
}
pub fn pack_secret_key<P: MlDsaSchemeParams>(
rho_seed: &[u8; 32], k_seed: &[u8; 32],
tr_hash: &[u8; 64],
s1_vec: &PolyVecL<P>,
s2_vec: &PolyVecK<P>,
t0_vec: &PolyVecK<P>,
) -> Result<SecretVec, SignError> {
let mut sk_bytes = SecretVec::empty();
sk_bytes.extend_from_slice(rho_seed);
sk_bytes.extend_from_slice(k_seed);
sk_bytes.extend_from_slice(tr_hash);
let eta_bits = if P::ETA_S1S2 == 2 { 3 } else { 4 }; let bytes_per_s_poly = ML_DSA_N * eta_bits / 8;
let bytes_per_t0_poly = ML_DSA_N * P::D_PARAM as usize / 8;
for i in 0..P::L_DIM {
let mut temp_poly = s1_vec.polys[i].clone();
for c in temp_poly.coeffs.iter_mut() {
let centered = centered_coefficient(*c);
if !(-(P::ETA_S1S2 as i32)..=P::ETA_S1S2 as i32).contains(¢ered) {
return Err(SignError::Serialization(
"s1 coefficient out of range".into(),
));
}
*c = (P::ETA_S1S2 as i32 - centered) as u32;
}
let mut packed = Zeroizing::new(boxed_bytes_zeroed(bytes_per_s_poly));
DefaultCoefficientSerde::pack_coeffs_into(&temp_poly, eta_bits, &mut packed)
.map_err(SignError::from_algo)?;
sk_bytes.extend_from_slice(&packed);
}
for i in 0..P::K_DIM {
let mut temp_poly = s2_vec.polys[i].clone();
for c in temp_poly.coeffs.iter_mut() {
let centered = centered_coefficient(*c);
if !(-(P::ETA_S1S2 as i32)..=P::ETA_S1S2 as i32).contains(¢ered) {
return Err(SignError::Serialization(
"s2 coefficient out of range".into(),
));
}
*c = (P::ETA_S1S2 as i32 - centered) as u32;
}
let mut packed = Zeroizing::new(boxed_bytes_zeroed(bytes_per_s_poly));
DefaultCoefficientSerde::pack_coeffs_into(&temp_poly, eta_bits, &mut packed)
.map_err(SignError::from_algo)?;
sk_bytes.extend_from_slice(&packed);
}
let t0_offset = 1 << (P::D_PARAM - 1);
for i in 0..P::K_DIM {
let mut temp_poly = t0_vec.polys[i].clone();
for c in temp_poly.coeffs.iter_mut() {
let centered = centered_coefficient(*c);
if !(-(t0_offset - 1)..=t0_offset).contains(¢ered) {
return Err(SignError::Serialization(
"t0 coefficient out of range".into(),
));
}
*c = (t0_offset - centered) as u32;
}
let mut packed = Zeroizing::new(boxed_bytes_zeroed(bytes_per_t0_poly));
DefaultCoefficientSerde::pack_coeffs_into(&temp_poly, P::D_PARAM as usize, &mut packed)
.map_err(SignError::from_algo)?;
sk_bytes.extend_from_slice(&packed);
}
if sk_bytes.len() != P::SECRET_KEY_BYTES {
return Err(SignError::Serialization(format!(
"secret key size mismatch: expected {}, got {}",
P::SECRET_KEY_BYTES,
sk_bytes.len()
)));
}
debug_assert_eq!(sk_bytes.len(), P::SECRET_KEY_BYTES);
Ok(sk_bytes)
}
pub type UnpackedSecretKey<P> = (
[u8; 32], Zeroizing<[u8; 32]>, Zeroizing<[u8; 64]>, PolyVecL<P>,
PolyVecK<P>,
PolyVecK<P>,
);
pub fn unpack_secret_key<P: MlDsaSchemeParams>(
sk_bytes: &[u8],
) -> Result<UnpackedSecretKey<P>, SignError> {
if sk_bytes.len() != P::SECRET_KEY_BYTES {
return Err(SignError::Deserialization(format!(
"Secret key size mismatch: expected {}, got {}",
P::SECRET_KEY_BYTES,
sk_bytes.len()
)));
}
let mut offset = 0;
let mut rho_seed = [0u8; 32];
rho_seed.copy_from_slice(&sk_bytes[offset..offset + 32]);
offset += 32;
let mut k_seed = Zeroizing::new([0u8; 32]);
k_seed.copy_from_slice(&sk_bytes[offset..offset + 32]);
offset += 32;
let mut tr_hash = Zeroizing::new([0u8; 64]);
tr_hash.copy_from_slice(&sk_bytes[offset..offset + 64]);
offset += 64;
let eta_bits = if P::ETA_S1S2 == 2 { 3 } else { 4 };
let bytes_per_s_poly = ML_DSA_N * eta_bits / 8;
let bytes_per_t0_poly = ML_DSA_N * P::D_PARAM as usize / 8;
let mut s1_vec = PolyVecL::<P>::zero();
for i in 0..P::L_DIM {
let poly_bytes = &sk_bytes[offset..offset + bytes_per_s_poly];
let mut temp_poly = DefaultCoefficientSerde::unpack_coeffs(poly_bytes, eta_bits)
.map_err(SignError::from_algo)?;
for c in temp_poly.coeffs.iter_mut() {
if *c > 2 * P::ETA_S1S2 {
return Err(SignError::InvalidKey(
"ML-DSA s1 coefficient out of range".into(),
));
}
let signed = P::ETA_S1S2 as i32 - *c as i32;
*c = signed_to_mod_q(signed);
}
s1_vec.polys[i] = temp_poly;
offset += bytes_per_s_poly;
}
let mut s2_vec = PolyVecK::<P>::zero();
for i in 0..P::K_DIM {
let poly_bytes = &sk_bytes[offset..offset + bytes_per_s_poly];
let mut temp_poly = DefaultCoefficientSerde::unpack_coeffs(poly_bytes, eta_bits)
.map_err(SignError::from_algo)?;
for c in temp_poly.coeffs.iter_mut() {
if *c > 2 * P::ETA_S1S2 {
return Err(SignError::InvalidKey(
"ML-DSA s2 coefficient out of range".into(),
));
}
let signed = P::ETA_S1S2 as i32 - *c as i32;
*c = signed_to_mod_q(signed);
}
s2_vec.polys[i] = temp_poly;
offset += bytes_per_s_poly;
}
let mut t0_vec = PolyVecK::<P>::zero();
let t0_offset = 1 << (P::D_PARAM - 1);
for i in 0..P::K_DIM {
let poly_bytes = &sk_bytes[offset..offset + bytes_per_t0_poly];
let mut temp_poly = DefaultCoefficientSerde::unpack_coeffs(poly_bytes, P::D_PARAM as usize)
.map_err(SignError::from_algo)?;
for c in temp_poly.coeffs.iter_mut() {
let signed = t0_offset - *c as i32;
*c = signed_to_mod_q(signed);
}
t0_vec.polys[i] = temp_poly;
offset += bytes_per_t0_poly;
}
if offset != sk_bytes.len() {
return Err(SignError::Deserialization(format!(
"secret key decoding consumed {offset} of {} bytes",
sk_bytes.len(),
)));
}
Ok((rho_seed, k_seed, tr_hash, s1_vec, s2_vec, t0_vec))
}
pub fn pack_signature<P: MlDsaSchemeParams>(
c_tilde_seed: &[u8], z_vec: &PolyVecL<P>,
h_hint_poly: &PolyVecK<P>,
) -> Result<Vec<u8>, SignError> {
if c_tilde_seed.len() != P::CHALLENGE_BYTES {
return Err(SignError::Serialization(format!(
"Challenge seed size mismatch: expected {}, got {}",
P::CHALLENGE_BYTES,
c_tilde_seed.len()
)));
}
let mut sig_bytes = Zeroizing::new(boxed_bytes_zeroed(P::SIGNATURE_SIZE));
let mut offset = 0usize;
sig_bytes[..c_tilde_seed.len()].copy_from_slice(c_tilde_seed);
offset += c_tilde_seed.len();
let bytes_per_z_poly = ML_DSA_N * P::Z_BITS / 8;
for i in 0..P::L_DIM {
let mut temp_poly = z_vec.polys[i].clone();
for c in temp_poly.coeffs.iter_mut() {
let centered = centered_coefficient(*c);
let lower = -(P::GAMMA1_PARAM as i32) + 1;
let upper = P::GAMMA1_PARAM as i32;
if !(lower..=upper).contains(¢ered) {
return Err(SignError::Serialization(
"z coefficient out of range".into(),
));
}
*c = (P::GAMMA1_PARAM as i32 - centered) as u32;
}
let end = offset
.checked_add(bytes_per_z_poly)
.ok_or_else(|| SignError::Serialization("signature length overflow".into()))?;
if end > sig_bytes.len() {
return Err(SignError::Serialization(
"signature components exceed the standardized size".into(),
));
}
DefaultCoefficientSerde::pack_coeffs_into(
&temp_poly,
P::Z_BITS,
&mut sig_bytes[offset..end],
)
.map_err(SignError::from_algo)?;
offset = end;
}
let hint_bytes = pack_hints_bitpacked::<P>(h_hint_poly)?;
let end = offset
.checked_add(hint_bytes.len())
.ok_or_else(|| SignError::Serialization("signature length overflow".into()))?;
if end > sig_bytes.len() {
return Err(SignError::Serialization(
"signature components exceed the standardized size".into(),
));
}
sig_bytes[offset..end].copy_from_slice(&hint_bytes);
offset = end;
if offset != P::SIGNATURE_SIZE {
return Err(SignError::Serialization(format!(
"Signature size mismatch: expected {}, got {}",
P::SIGNATURE_SIZE,
offset,
)));
}
Ok(sig_bytes.into_inner().into_vec())
}
pub fn pack_polyveck_w1<P: MlDsaSchemeParams>(
w1_vec: &PolyVecK<P>,
) -> Result<ZeroizingBytes, SignError> {
let bits_per_coeff = w1_bits_needed::<P>();
let maximum = (ML_DSA_Q - 1) / (2 * P::GAMMA2_PARAM) - 1;
let bytes_per_poly = ML_DSA_N * bits_per_coeff as usize / 8;
let total_len = P::K_DIM
.checked_mul(bytes_per_poly)
.ok_or_else(|| SignError::Serialization("w1 encoding length overflow".into()))?;
let mut packed = Zeroizing::new(boxed_bytes_zeroed(total_len));
for (index, poly) in w1_vec.polys.iter().enumerate() {
for &coeff in &poly.coeffs {
if coeff > maximum {
return Err(SignError::Serialization(
"w1 coefficient out of range".into(),
));
}
}
let start = index * bytes_per_poly;
DefaultCoefficientSerde::pack_coeffs_into(
poly,
bits_per_coeff as usize,
&mut packed[start..start + bytes_per_poly],
)
.map_err(SignError::from_algo)?;
}
Ok(packed)
}
pub type UnpackedSignature<P> = (Vec<u8>, PolyVecL<P>, PolyVecK<P>);
pub fn unpack_signature<P: MlDsaSchemeParams>(
sig_bytes: &[u8],
) -> Result<UnpackedSignature<P>, SignError> {
if sig_bytes.len() != P::SIGNATURE_SIZE {
return Err(SignError::Deserialization(format!(
"Signature size mismatch: expected {}, got {}",
P::SIGNATURE_SIZE,
sig_bytes.len()
)));
}
let mut offset = 0;
let mut c_tilde_seed = vec![0u8; P::CHALLENGE_BYTES];
c_tilde_seed.copy_from_slice(&sig_bytes[offset..offset + P::CHALLENGE_BYTES]);
offset += P::CHALLENGE_BYTES;
let mut z_vec = PolyVecL::<P>::zero();
let bytes_per_z_poly = ML_DSA_N * P::Z_BITS / 8;
for i in 0..P::L_DIM {
let poly_bytes = &sig_bytes[offset..offset + bytes_per_z_poly];
let mut temp_poly = DefaultCoefficientSerde::unpack_coeffs(poly_bytes, P::Z_BITS)
.map_err(SignError::from_algo)?;
for c in temp_poly.coeffs.iter_mut() {
let value = P::GAMMA1_PARAM as i32 - *c as i32;
*c = signed_to_mod_q(value);
}
z_vec.polys[i] = temp_poly;
offset += bytes_per_z_poly;
}
let hint_bytes = &sig_bytes[offset..];
let (h_hint_poly, _hint_cnt) = unpack_hints_bitpacked::<P>(hint_bytes)?;
Ok((c_tilde_seed, z_vec, h_hint_poly))
}
#[cfg(test)]
mod tests {
use super::*;
use dcrypt_params::pqc::ml_dsa::MlDsa44Params;
#[test]
fn test_roundtrip_hints_basic() {
let mut h = PolyVecK::<MlDsa44Params>::zero();
h.polys[1].coeffs[5] = 1;
h.polys[2].coeffs[20] = 1;
let packed = pack_hints_bitpacked::<MlDsa44Params>(&h).unwrap();
let (unpacked, cnt) = unpack_hints_bitpacked::<MlDsa44Params>(&packed).unwrap();
assert_eq!(cnt, 2, "Hint count mismatch");
assert_eq!(
unpacked.polys[1].coeffs[5], 1,
"Lost hint at poly[1].coeff[5]"
);
assert_eq!(
unpacked.polys[2].coeffs[20], 1,
"Lost hint at poly[2].coeff[20]"
);
for i in 0..MlDsa44Params::K_DIM {
for j in 0..256 {
if !((i == 1 && j == 5) || (i == 2 && j == 20)) {
assert_eq!(
unpacked.polys[i].coeffs[j], 0,
"Spurious hint at poly[{}].coeff[{}]",
i, j
);
}
}
}
}
#[test]
fn signing_encodings_use_exact_sized_storage() {
let w1 = PolyVecK::<MlDsa44Params>::zero();
let packed_w1 = pack_polyveck_w1::<MlDsa44Params>(&w1).unwrap();
assert_eq!(packed_w1.capacity(), packed_w1.len());
let hints = PolyVecK::<MlDsa44Params>::zero();
let packed_hints = pack_hints_bitpacked::<MlDsa44Params>(&hints).unwrap();
assert_eq!(packed_hints.capacity(), packed_hints.len());
let z = PolyVecL::<MlDsa44Params>::zero();
let challenge = [0u8; MlDsa44Params::CHALLENGE_BYTES];
let signature = pack_signature::<MlDsa44Params>(&challenge, &z, &hints).unwrap();
assert_eq!(signature.len(), MlDsa44Params::SIGNATURE_SIZE);
assert_eq!(signature.capacity(), signature.len());
}
}