use alloc::{vec, vec::Vec};
use crate::packing;
use crate::params::*;
use crate::poly::Poly;
use crate::polyvec::*;
use crate::symmetric::{shake256, shake256_multi};
use subtle::ConstantTimeEq;
use zeroize::Zeroize;
#[must_use]
pub fn keypair(mode: DilithiumMode, random_seed: &[u8; SEEDBYTES]) -> (Vec<u8>, Vec<u8>) {
let k = mode.k();
let l = mode.l();
let mut seedbuf = [0u8; 2 * SEEDBYTES + CRHBYTES];
let mut expanded = [0u8; 2 * SEEDBYTES + CRHBYTES];
seedbuf[..SEEDBYTES].copy_from_slice(random_seed);
seedbuf[SEEDBYTES] = k as u8;
seedbuf[SEEDBYTES + 1] = l as u8;
shake256(&mut expanded, &seedbuf[..SEEDBYTES + 2]);
seedbuf.zeroize();
let rho: [u8; SEEDBYTES] = expanded[..SEEDBYTES].try_into().unwrap();
let mut rhoprime: [u8; CRHBYTES] = expanded[SEEDBYTES..SEEDBYTES + CRHBYTES]
.try_into()
.unwrap();
let mut key: [u8; SEEDBYTES] = expanded[SEEDBYTES + CRHBYTES..].try_into().unwrap();
expanded.zeroize();
let mut mat = vec![PolyVecL::default(); K_MAX];
matrix_expand(mode, &mut mat, &rho);
let mut s1 = PolyVecL::default();
let mut s2 = PolyVecK::default();
polyvecl_uniform_eta(mode, &mut s1, &rhoprime, 0);
polyveck_uniform_eta(mode, &mut s2, &rhoprime, l as u16);
rhoprime.zeroize();
let mut s1hat = s1.clone();
polyvecl_ntt(mode, &mut s1hat);
let mut t1 = PolyVecK::default();
matrix_pointwise_montgomery(mode, &mut t1, &mat, &s1hat);
polyveck_reduce(mode, &mut t1);
polyveck_invntt_tomont(mode, &mut t1);
polyveck_add_assign(mode, &mut t1, &s2);
polyveck_caddq(mode, &mut t1);
let mut t1_high = PolyVecK::default();
let mut t0 = PolyVecK::default();
polyveck_power2round(mode, &mut t1_high, &mut t0, &t1);
let mut pk = vec![0u8; mode.public_key_bytes()];
packing::pack_pk(mode, &mut pk, &rho, &t1_high);
let mut tr = [0u8; TRBYTES];
shake256(&mut tr, &pk);
let mut sk = vec![0u8; mode.secret_key_bytes()];
packing::pack_sk(mode, &mut sk, &rho, &tr, &key, &t0, &s1, &s2);
key.zeroize();
s1.zeroize();
s1hat.zeroize();
s2.zeroize();
t0.zeroize();
t1.zeroize();
(pk, sk)
}
pub fn sign_signature_internal(
mode: DilithiumMode,
sig: &mut [u8],
m: &[u8],
pre: &[u8],
rnd: &[u8; RNDBYTES],
sk: &[u8],
) -> usize {
if sk.len() != mode.secret_key_bytes() || sig.len() < mode.signature_bytes() {
return 0;
}
let k = mode.k();
let l = mode.l();
let beta = mode.beta();
let gamma1 = mode.gamma1();
let gamma2 = mode.gamma2();
let omega = mode.omega();
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, sk,
);
let mut mu = [0u8; CRHBYTES];
shake256_multi(&mut mu, &[&tr, pre, m]);
let mut rhoprime = [0u8; CRHBYTES];
shake256_multi(&mut rhoprime, &[&key, rnd, &mu]);
key.zeroize();
let mut mat = vec![PolyVecL::default(); K_MAX];
matrix_expand(mode, &mut mat, &rho);
polyvecl_ntt(mode, &mut s1);
polyveck_ntt(mode, &mut s2);
polyveck_ntt(mode, &mut t0);
let mut nonce: u16 = 0;
let mut h = PolyVecK::default();
let mut y = PolyVecL::default();
let mut y_ntt = PolyVecL::default();
let mut z = PolyVecL::default();
let mut w = PolyVecK::default();
let mut w0 = PolyVecK::default();
let mut cp = Poly::zero();
let siglen = loop {
polyvecl_uniform_gamma1(mode, &mut y, &rhoprime, nonce);
nonce = nonce.wrapping_add(l as u16);
y_ntt.clone_from(&y);
polyvecl_ntt(mode, &mut y_ntt);
matrix_pointwise_montgomery(mode, &mut w, &mat, &y_ntt);
polyveck_reduce(mode, &mut w);
polyveck_invntt_tomont(mode, &mut w);
polyveck_caddq(mode, &mut w);
let mut w1_high = PolyVecK::default();
polyveck_decompose(mode, &mut w1_high, &mut w0, &w);
let mut w1_packed = vec![0u8; k * mode.polyw1_packedbytes()];
polyveck_pack_w1(mode, &mut w1_packed, &w1_high);
let ctilde = mode.ctildebytes();
let mut ctilde_buf = vec![0u8; ctilde];
shake256_multi(&mut ctilde_buf, &[&mu, &w1_packed]);
Poly::challenge(mode, &mut cp, &ctilde_buf);
cp.ntt();
polyvecl_pointwise_poly_montgomery(mode, &mut z, &cp, &s1);
polyvecl_invntt_tomont(mode, &mut z);
polyvecl_add_assign(mode, &mut z, &y);
polyvecl_reduce(mode, &mut z);
if polyvecl_chknorm(mode, &z, gamma1 - beta) {
continue;
}
polyveck_pointwise_poly_montgomery(mode, &mut h, &cp, &s2);
polyveck_invntt_tomont(mode, &mut h);
polyveck_sub_assign(mode, &mut w0, &h);
polyveck_reduce(mode, &mut w0);
if polyveck_chknorm(mode, &w0, gamma2 - beta) {
continue;
}
polyveck_pointwise_poly_montgomery(mode, &mut h, &cp, &t0);
polyveck_invntt_tomont(mode, &mut h);
polyveck_reduce(mode, &mut h);
if polyveck_chknorm(mode, &h, gamma2) {
continue;
}
polyveck_add_assign(mode, &mut w0, &h);
let n = polyveck_make_hint(mode, &mut h, &w0, &w1_high);
if n > omega {
continue;
}
packing::pack_sig(mode, sig, &ctilde_buf, &z, &h);
break mode.signature_bytes();
};
s1.zeroize();
s2.zeroize();
t0.zeroize();
rhoprime.zeroize();
y.zeroize();
y_ntt.zeroize();
w.zeroize();
w0.zeroize();
siglen
}
pub fn sign_signature(
mode: DilithiumMode,
sig: &mut [u8],
m: &[u8],
ctx: &[u8],
rnd: &[u8; RNDBYTES],
sk: &[u8],
) -> i32 {
if ctx.len() > 255 {
return -1;
}
let mut pre = vec![0u8; 2 + ctx.len()];
pre[0] = 0;
pre[1] = ctx.len() as u8;
pre[2..].copy_from_slice(ctx);
if sign_signature_internal(mode, sig, m, &pre, rnd, sk) == 0 {
return -1;
}
0
}
#[must_use]
pub fn verify_internal(mode: DilithiumMode, sig: &[u8], m: &[u8], pre: &[u8], pk: &[u8]) -> bool {
let k = mode.k();
let beta = mode.beta();
let gamma1 = mode.gamma1();
let ctilde_len = mode.ctildebytes();
if sig.len() != mode.signature_bytes() {
return false;
}
if pk.len() != mode.public_key_bytes() {
return false;
}
let mut rho = [0u8; SEEDBYTES];
let mut t1 = PolyVecK::default();
packing::unpack_pk(mode, &mut rho, &mut t1, pk);
let mut c = vec![0u8; ctilde_len];
let mut z = PolyVecL::default();
let mut h = PolyVecK::default();
if packing::unpack_sig(mode, &mut c, &mut z, &mut h, sig) {
return false;
}
if polyvecl_chknorm(mode, &z, gamma1 - beta) {
return false;
}
let mut mu = [0u8; CRHBYTES];
let mut tr = [0u8; TRBYTES];
shake256(&mut tr, pk);
shake256_multi(&mut mu, &[&tr, pre, m]);
let mut cp = Poly::zero();
Poly::challenge(mode, &mut cp, &c);
let mut mat = vec![PolyVecL::default(); K_MAX];
matrix_expand(mode, &mut mat, &rho);
polyvecl_ntt(mode, &mut z);
let mut w1 = PolyVecK::default();
matrix_pointwise_montgomery(mode, &mut w1, &mat, &z);
cp.ntt();
polyveck_shiftl(mode, &mut t1);
polyveck_ntt(mode, &mut t1);
let t1_clone = t1.clone();
polyveck_pointwise_poly_montgomery(mode, &mut t1, &cp, &t1_clone);
let w1_copy = w1.clone();
polyveck_sub(mode, &mut w1, &w1_copy, &t1);
polyveck_reduce(mode, &mut w1);
polyveck_invntt_tomont(mode, &mut w1);
polyveck_caddq(mode, &mut w1);
let w1_copy2 = w1.clone();
polyveck_use_hint(mode, &mut w1, &w1_copy2, &h);
let mut buf = vec![0u8; k * mode.polyw1_packedbytes()];
polyveck_pack_w1(mode, &mut buf, &w1);
let mut c2 = vec![0u8; ctilde_len];
shake256_multi(&mut c2, &[&mu, &buf]);
c.ct_eq(&c2).into()
}
#[must_use]
pub fn verify(mode: DilithiumMode, sig: &[u8], m: &[u8], ctx: &[u8], pk: &[u8]) -> bool {
if ctx.len() > 255 {
return false;
}
let mut pre = vec![0u8; 2 + ctx.len()];
pre[0] = 0;
pre[1] = ctx.len() as u8;
pre[2..].copy_from_slice(ctx);
verify_internal(mode, sig, m, &pre, pk)
}
pub fn sign_hash(
mode: DilithiumMode,
sig: &mut [u8],
msg: &[u8],
ctx: &[u8],
rnd: &[u8; RNDBYTES],
sk: &[u8],
) -> i32 {
if ctx.len() > 255 {
return -1;
}
use sha2::Digest;
let ph_m = sha2::Sha512::digest(msg);
let oid = mode.hash_oid();
let mut pre = vec![0u8; 2 + ctx.len() + oid.len() + ph_m.len()];
pre[0] = 1; pre[1] = ctx.len() as u8;
let mut off = 2;
pre[off..off + ctx.len()].copy_from_slice(ctx);
off += ctx.len();
pre[off..off + oid.len()].copy_from_slice(oid);
off += oid.len();
pre[off..off + ph_m.len()].copy_from_slice(&ph_m);
if sign_signature_internal(mode, sig, &[], &pre, rnd, sk) == 0 {
return -1;
}
0
}
#[must_use]
pub fn verify_hash(mode: DilithiumMode, sig: &[u8], msg: &[u8], ctx: &[u8], pk: &[u8]) -> bool {
if ctx.len() > 255 {
return false;
}
use sha2::Digest;
let ph_m = sha2::Sha512::digest(msg);
let oid = mode.hash_oid();
let mut pre = vec![0u8; 2 + ctx.len() + oid.len() + ph_m.len()];
pre[0] = 1;
pre[1] = ctx.len() as u8;
let mut off = 2;
pre[off..off + ctx.len()].copy_from_slice(ctx);
off += ctx.len();
pre[off..off + oid.len()].copy_from_slice(oid);
off += oid.len();
pre[off..off + ph_m.len()].copy_from_slice(&ph_m);
verify_internal(mode, sig, &[], &pre, pk)
}