use alloc::{vec, vec::Vec};
use core::fmt;
use zeroize::Zeroize;
use zeroize::Zeroizing;
use crate::packing;
pub use crate::params::DilithiumMode;
use crate::params::*;
use crate::polyvec::*;
use crate::sign;
use crate::symmetric::shake256;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum DilithiumError {
RandomError,
FormatError,
BadSignature,
BadArgument,
InvalidKey,
}
impl fmt::Display for DilithiumError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::RandomError => write!(f, "random number generation failed"),
Self::FormatError => write!(f, "invalid format"),
Self::BadSignature => write!(f, "invalid signature"),
Self::BadArgument => write!(f, "invalid argument"),
Self::InvalidKey => write!(f, "key validation failed"),
}
}
}
#[cfg(feature = "std")]
impl std::error::Error for DilithiumError {}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct DilithiumKeyPair {
#[cfg_attr(
feature = "serde",
serde(
serialize_with = "serde_zeroizing::serialize",
deserialize_with = "serde_zeroizing::deserialize"
)
)]
privkey: Zeroizing<Vec<u8>>,
pubkey: Vec<u8>,
mode: DilithiumMode,
}
#[cfg(feature = "serde")]
mod serde_zeroizing {
use super::*;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
pub fn serialize<S: Serializer>(val: &Zeroizing<Vec<u8>>, s: S) -> Result<S::Ok, S::Error> {
let inner: &Vec<u8> = val;
inner.serialize(s)
}
pub fn deserialize<'de, D: Deserializer<'de>>(d: D) -> Result<Zeroizing<Vec<u8>>, D::Error> {
let v = Vec::<u8>::deserialize(d)?;
Ok(Zeroizing::new(v))
}
}
pub type MlDsaKeyPair = DilithiumKeyPair;
#[derive(Debug, Clone, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct DilithiumSignature {
data: Vec<u8>,
}
pub type MlDsaSignature = DilithiumSignature;
impl DilithiumKeyPair {
#[cfg(feature = "getrandom")]
pub fn generate(mode: DilithiumMode) -> Result<Self, DilithiumError> {
let mut seed = [0u8; SEEDBYTES];
getrandom(&mut seed).map_err(|()| DilithiumError::RandomError)?;
let result = Self::generate_deterministic(mode, &seed);
seed.zeroize();
Ok(result)
}
#[must_use]
pub fn generate_deterministic(mode: DilithiumMode, seed: &[u8; SEEDBYTES]) -> Self {
let (pk, sk) = sign::keypair(mode, seed);
DilithiumKeyPair {
privkey: Zeroizing::new(sk),
pubkey: pk,
mode,
}
}
#[cfg(feature = "getrandom")]
pub fn sign(&self, msg: &[u8], ctx: &[u8]) -> Result<DilithiumSignature, DilithiumError> {
if ctx.len() > 255 {
return Err(DilithiumError::BadArgument);
}
let mut rnd = [0u8; RNDBYTES];
getrandom(&mut rnd).map_err(|()| DilithiumError::RandomError)?;
let mut sig = vec![0u8; self.mode.signature_bytes()];
let ret = sign::sign_signature(self.mode, &mut sig, msg, ctx, &rnd, &self.privkey);
rnd.zeroize();
if ret != 0 {
return Err(DilithiumError::BadArgument);
}
Ok(DilithiumSignature { data: sig })
}
#[cfg(feature = "getrandom")]
pub fn sign_prehash(
&self,
msg: &[u8],
ctx: &[u8],
) -> Result<DilithiumSignature, DilithiumError> {
if ctx.len() > 255 {
return Err(DilithiumError::BadArgument);
}
let mut rnd = [0u8; RNDBYTES];
getrandom(&mut rnd).map_err(|()| DilithiumError::RandomError)?;
let mut sig = vec![0u8; self.mode.signature_bytes()];
let ret = sign::sign_hash(self.mode, &mut sig, msg, ctx, &rnd, &self.privkey);
rnd.zeroize();
if ret != 0 {
return Err(DilithiumError::BadArgument);
}
Ok(DilithiumSignature { data: sig })
}
pub fn sign_deterministic(
&self,
msg: &[u8],
ctx: &[u8],
rnd: &[u8; RNDBYTES],
) -> Result<DilithiumSignature, DilithiumError> {
if ctx.len() > 255 {
return Err(DilithiumError::BadArgument);
}
let mut sig = vec![0u8; self.mode.signature_bytes()];
let ret = sign::sign_signature(self.mode, &mut sig, msg, ctx, rnd, &self.privkey);
if ret != 0 {
return Err(DilithiumError::BadArgument);
}
Ok(DilithiumSignature { data: sig })
}
#[must_use]
pub fn verify(
pk: &[u8],
sig: &DilithiumSignature,
msg: &[u8],
ctx: &[u8],
mode: DilithiumMode,
) -> bool {
if pk.len() != mode.public_key_bytes() {
return false;
}
if sig.data.len() != mode.signature_bytes() {
return false;
}
sign::verify(mode, &sig.data, msg, ctx, pk)
}
#[must_use]
pub fn verify_prehash(
pk: &[u8],
sig: &DilithiumSignature,
msg: &[u8],
ctx: &[u8],
mode: DilithiumMode,
) -> bool {
if pk.len() != mode.public_key_bytes() {
return false;
}
if sig.data.len() != mode.signature_bytes() {
return false;
}
sign::verify_hash(mode, &sig.data, msg, ctx, pk)
}
#[must_use]
pub fn public_key(&self) -> &[u8] {
&self.pubkey
}
#[must_use]
pub fn private_key(&self) -> &[u8] {
&self.privkey
}
#[must_use]
pub fn mode(&self) -> DilithiumMode {
self.mode
}
pub fn from_keys(
privkey: &[u8],
pubkey: &[u8],
mode: DilithiumMode,
) -> Result<Self, DilithiumError> {
if privkey.len() != mode.secret_key_bytes() {
return Err(DilithiumError::FormatError);
}
if pubkey.len() != mode.public_key_bytes() {
return Err(DilithiumError::FormatError);
}
let sk_rho = &privkey[..SEEDBYTES];
let pk_rho = &pubkey[..SEEDBYTES];
if sk_rho != pk_rho {
return Err(DilithiumError::InvalidKey);
}
let tr_offset = 2 * SEEDBYTES;
let sk_tr = &privkey[tr_offset..tr_offset + TRBYTES];
let mut expected_tr = [0u8; TRBYTES];
shake256(&mut expected_tr, pubkey);
if sk_tr != &expected_tr[..] {
return Err(DilithiumError::InvalidKey);
}
let mut rho = [0u8; SEEDBYTES];
let mut tr = [0u8; TRBYTES];
let mut key = [0u8; SEEDBYTES];
let mut t0 = PolyVecK::default();
let mut s1 = PolyVecL::default();
let mut s2 = PolyVecK::default();
packing::unpack_sk(
mode, &mut rho, &mut tr, &mut key, &mut t0, &mut s1, &mut s2, privkey,
);
let mut mat = vec![PolyVecL::default(); K_MAX];
matrix_expand(mode, &mut mat, &rho);
polyvecl_ntt(mode, &mut s1);
let mut t = PolyVecK::default();
matrix_pointwise_montgomery(mode, &mut t, &mat, &s1);
polyveck_reduce(mode, &mut t);
polyveck_invntt_tomont(mode, &mut t);
polyveck_add_assign(mode, &mut t, &s2);
polyveck_caddq(mode, &mut t);
let mut t1 = PolyVecK::default();
let mut t0_expected = PolyVecK::default();
polyveck_power2round(mode, &mut t1, &mut t0_expected, &t);
let mut pk_expected = vec![0u8; mode.public_key_bytes()];
packing::pack_pk(mode, &mut pk_expected, &rho, &t1);
let pk_ok = pk_expected == pubkey;
let mut t0_ok = true;
for i in 0..mode.k() {
if t0.vec[i].coeffs != t0_expected.vec[i].coeffs {
t0_ok = false;
}
}
key.zeroize();
s1.zeroize();
s2.zeroize();
t0.zeroize();
t0_expected.zeroize();
t.zeroize();
if !pk_ok || !t0_ok {
return Err(DilithiumError::InvalidKey);
}
Ok(DilithiumKeyPair {
privkey: Zeroizing::new(privkey.to_vec()),
pubkey: pubkey.to_vec(),
mode,
})
}
#[must_use]
pub fn to_bytes(&self) -> Vec<u8> {
let mut buf = Vec::with_capacity(1 + self.pubkey.len() + self.privkey.len());
buf.push(self.mode.mode_tag());
buf.extend_from_slice(&self.pubkey);
buf.extend_from_slice(&self.privkey);
buf
}
pub fn from_bytes(data: &[u8]) -> Result<Self, DilithiumError> {
if data.is_empty() {
return Err(DilithiumError::FormatError);
}
let mode = DilithiumMode::from_tag(data[0]).ok_or(DilithiumError::FormatError)?;
let pk_len = mode.public_key_bytes();
let sk_len = mode.secret_key_bytes();
if data.len() != 1 + pk_len + sk_len {
return Err(DilithiumError::FormatError);
}
let pk = &data[1..=pk_len];
let sk = &data[1 + pk_len..];
Self::from_keys(sk, pk, mode)
}
#[must_use]
pub fn public_key_bytes(&self) -> Vec<u8> {
let mut buf = Vec::with_capacity(1 + self.pubkey.len());
buf.push(self.mode.mode_tag());
buf.extend_from_slice(&self.pubkey);
buf
}
pub fn from_public_key(data: &[u8]) -> Result<(DilithiumMode, Vec<u8>), DilithiumError> {
if data.is_empty() {
return Err(DilithiumError::FormatError);
}
let mode = DilithiumMode::from_tag(data[0]).ok_or(DilithiumError::FormatError)?;
if data.len() != 1 + mode.public_key_bytes() {
return Err(DilithiumError::FormatError);
}
Ok((mode, data[1..].to_vec()))
}
}
impl DilithiumSignature {
#[must_use]
pub fn as_bytes(&self) -> &[u8] {
&self.data
}
#[must_use]
pub fn from_bytes(data: Vec<u8>) -> Self {
Self { data }
}
#[must_use]
pub fn from_slice(data: &[u8]) -> Self {
Self {
data: data.to_vec(),
}
}
#[must_use]
pub fn len(&self) -> usize {
self.data.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.data.is_empty()
}
}
#[cfg(feature = "getrandom")]
fn getrandom(buf: &mut [u8]) -> Result<(), ()> {
::getrandom::getrandom(buf).map_err(|_| ())
}