use crate::error::Error as SignError;
#[cfg(not(feature = "std"))]
use alloc::{boxed::Box, vec::Vec};
use core::marker::PhantomData;
use dcrypt_algorithms::error::Result as AlgoResult;
use dcrypt_algorithms::poly::params::{MlDsaParams, Modulus};
use dcrypt_algorithms::poly::polynomial::Polynomial;
use dcrypt_algorithms::xof::shake::ShakeXof128;
use dcrypt_algorithms::xof::ExtendableOutputFunction;
use dcrypt_internal::{Choice, ConditionallySelectable, Zeroize, ZeroizeOnDrop};
use dcrypt_params::pqc::ml_dsa::MlDsaSchemeParams;
#[derive(Debug)]
pub struct PolyVecL<P: MlDsaSchemeParams> {
pub(crate) polys: Box<[Polynomial<MlDsaParams>]>,
_params: PhantomData<P>,
}
#[derive(Debug)]
pub struct PolyVecK<P: MlDsaSchemeParams> {
pub(crate) polys: Box<[Polynomial<MlDsaParams>]>,
_params: PhantomData<P>,
}
impl<P: MlDsaSchemeParams> Clone for PolyVecL<P> {
fn clone(&self) -> Self {
Self {
polys: self.polys.clone(),
_params: PhantomData,
}
}
}
impl<P: MlDsaSchemeParams> Clone for PolyVecK<P> {
fn clone(&self) -> Self {
Self {
polys: self.polys.clone(),
_params: PhantomData,
}
}
}
impl<P: MlDsaSchemeParams> Zeroize for PolyVecL<P> {
fn zeroize(&mut self) {
for poly in self.polys.iter_mut() {
poly.coeffs.as_mut().zeroize(); }
}
}
impl<P: MlDsaSchemeParams> Zeroize for PolyVecK<P> {
fn zeroize(&mut self) {
for poly in self.polys.iter_mut() {
poly.coeffs.as_mut().zeroize(); }
}
}
impl<P: MlDsaSchemeParams> Drop for PolyVecL<P> {
fn drop(&mut self) {
self.zeroize();
}
}
impl<P: MlDsaSchemeParams> ZeroizeOnDrop for PolyVecL<P> {}
impl<P: MlDsaSchemeParams> Drop for PolyVecK<P> {
fn drop(&mut self) {
self.zeroize();
}
}
impl<P: MlDsaSchemeParams> ZeroizeOnDrop for PolyVecK<P> {}
impl<P: MlDsaSchemeParams> PolyVecL<P> {
pub fn zero() -> Self {
let mut polys = Vec::with_capacity(P::L_DIM);
for _ in 0..P::L_DIM {
polys.push(Polynomial::<MlDsaParams>::zero());
}
Self {
polys: polys.into_boxed_slice(),
_params: PhantomData,
}
}
pub fn ntt_inplace(&mut self) -> AlgoResult<()> {
for p in self.polys.iter_mut() {
p.ntt_inplace()?;
}
Ok(())
}
pub fn inv_ntt_inplace(&mut self) -> AlgoResult<()> {
for p in self.polys.iter_mut() {
p.from_ntt_inplace()?;
}
Ok(())
}
pub fn pointwise_dot_product(&self, other: &PolyVecL<P>) -> Polynomial<MlDsaParams> {
let mut acc = Polynomial::<MlDsaParams>::zero();
for i in 0..P::L_DIM {
let prod = self.polys[i].ntt_mul(&other.polys[i]);
acc = acc.add(&prod);
}
acc
}
pub fn conditional_select(a: &Self, b: &Self, choice: Choice) -> Self {
let mut out = Self::zero();
for i in 0..P::L_DIM {
for j in 0..MlDsaParams::N {
out.polys[i].coeffs[j] =
u32::conditional_select(&a.polys[i].coeffs[j], &b.polys[i].coeffs[j], choice);
}
}
out
}
}
impl<P: MlDsaSchemeParams> PolyVecK<P> {
pub fn zero() -> Self {
let mut polys = Vec::with_capacity(P::K_DIM);
for _ in 0..P::K_DIM {
polys.push(Polynomial::<MlDsaParams>::zero());
}
Self {
polys: polys.into_boxed_slice(),
_params: PhantomData,
}
}
pub fn ntt_inplace(&mut self) -> AlgoResult<()> {
for p in self.polys.iter_mut() {
p.ntt_inplace()?;
}
Ok(())
}
pub fn inv_ntt_inplace(&mut self) -> AlgoResult<()> {
for p in self.polys.iter_mut() {
p.from_ntt_inplace()?;
}
Ok(())
}
pub fn add(&self, other: &Self) -> Self {
let mut res = Self::zero();
for i in 0..P::K_DIM {
res.polys[i] = self.polys[i].add(&other.polys[i]);
}
res
}
pub fn neg_mod_q(&self) -> Self {
let mut res = Self::zero();
for i in 0..P::K_DIM {
for j in 0..MlDsaParams::N {
let coeff = self.polys[i].coeffs[j];
res.polys[i].coeffs[j] = if coeff == 0 {
0
} else {
MlDsaParams::Q - coeff
};
}
}
res
}
pub fn sub(&self, other: &Self) -> Self {
let mut res = Self::zero();
for i in 0..P::K_DIM {
res.polys[i] = self.polys[i].sub(&other.polys[i]);
}
res
}
pub fn conditional_select(a: &Self, b: &Self, choice: Choice) -> Self {
let mut out = Self::zero();
for i in 0..P::K_DIM {
for j in 0..MlDsaParams::N {
out.polys[i].coeffs[j] =
u32::conditional_select(&a.polys[i].coeffs[j], &b.polys[i].coeffs[j], choice);
}
}
out
}
}
pub fn matrix_polyvecl_mul<P: MlDsaSchemeParams>(
matrix_a_hat: &[PolyVecL<P>], vector_l_hat: &PolyVecL<P>, ) -> PolyVecK<P> {
let mut result_veck = PolyVecK::<P>::zero();
for (i, row) in matrix_a_hat.iter().enumerate() {
result_veck.polys[i] = row.pointwise_dot_product(vector_l_hat);
}
result_veck
}
pub fn expand_matrix_a<P: MlDsaSchemeParams>(
rho_seed: &[u8; 32], ) -> Result<Vec<PolyVecL<P>>, SignError> {
let mut matrix_a = Vec::with_capacity(P::K_DIM);
for i in 0..P::K_DIM {
let mut row = PolyVecL::<P>::zero();
for j in 0..P::L_DIM {
let mut xof = ShakeXof128::new();
xof.update(rho_seed).map_err(SignError::from_algo)?;
xof.update(&[j as u8]).map_err(SignError::from_algo)?;
xof.update(&[i as u8]).map_err(SignError::from_algo)?;
let mut poly = Polynomial::<MlDsaParams>::zero();
let mut ctr = 0;
let mut temp_buf = [0u8; 3];
while ctr < MlDsaParams::N {
xof.squeeze(&mut temp_buf).map_err(SignError::from_algo)?;
let candidate = u32::from(temp_buf[0])
| (u32::from(temp_buf[1]) << 8)
| (u32::from(temp_buf[2] & 0x7f) << 16);
if candidate < MlDsaParams::Q {
poly.coeffs[ctr] = candidate;
ctr += 1;
}
}
row.polys[j] = poly;
}
matrix_a.push(row);
}
Ok(matrix_a)
}