use core::marker::PhantomData;
use p3_mds::MdsPermutation;
use p3_mds::karatsuba_convolution::Convolve;
use p3_mds::util::dot_product;
use p3_symmetric::Permutation;
use crate::{BarrettParameters, MontyField31, MontyParameters};
pub trait MDSUtils: Clone + Sync {
const MATRIX_CIRC_MDS_8_COL: [i64; 8];
const MATRIX_CIRC_MDS_12_COL: [i64; 12];
const MATRIX_CIRC_MDS_16_COL: [i64; 16];
const MATRIX_CIRC_MDS_24_COL: [i64; 24];
const MATRIX_CIRC_MDS_32_COL: [i64; 32];
const MATRIX_CIRC_MDS_64_COL: [i64; 64];
}
#[derive(Clone, Debug, Default)]
pub struct MdsMatrixMontyField31<MU: MDSUtils> {
_phantom: PhantomData<MU>,
}
struct SmallConvolveMontyField31;
impl<FP: MontyParameters> Convolve<MontyField31<FP>, i64, i64, i64> for SmallConvolveMontyField31 {
#[inline(always)]
fn read(input: MontyField31<FP>) -> i64 {
input.value as i64
}
#[inline(always)]
fn parity_dot<const N: usize>(u: [i64; N], v: [i64; N]) -> i64 {
dot_product(u, v)
}
#[inline(always)]
fn reduce(z: i64) -> MontyField31<FP> {
debug_assert!(z >= 0);
MontyField31::new_monty((z as u64 % FP::PRIME as u64) as u32)
}
}
#[inline(always)]
const fn barrett_red_monty31<BP: BarrettParameters>(input: i128) -> i64 {
let input_high = (input >> BP::N) as i64;
let quot = (((input_high as i128) * (BP::PSEUDO_INV as i128)) >> BP::N) as i64;
let quot_2adic = quot & BP::MASK;
let sub = (quot_2adic as i128) * BP::PRIME_I128;
(input - sub) as i64
}
#[derive(Debug, Clone, Default)]
struct LargeConvolveMontyField31;
impl<FP> Convolve<MontyField31<FP>, i64, i64, i64> for LargeConvolveMontyField31
where
FP: BarrettParameters,
{
#[inline(always)]
fn read(input: MontyField31<FP>) -> i64 {
input.value as i64
}
#[inline(always)]
fn parity_dot<const N: usize>(u: [i64; N], v: [i64; N]) -> i64 {
let mut dp = 0i128;
for i in 0..N {
dp += u[i] as i128 * v[i] as i128;
}
barrett_red_monty31::<FP>(dp)
}
#[inline(always)]
fn reduce(z: i64) -> MontyField31<FP> {
debug_assert!(z > -(1i64 << 55));
debug_assert!(z < (1i64 << 55));
let red = (z % (FP::PRIME as i64)) as u32;
let (corr, over) = red.overflowing_add(FP::PRIME);
let value = if over { corr } else { red };
MontyField31::new_monty(value)
}
}
impl<FP: MontyParameters, MU: MDSUtils> Permutation<[MontyField31<FP>; 8]>
for MdsMatrixMontyField31<MU>
{
fn permute(&self, input: [MontyField31<FP>; 8]) -> [MontyField31<FP>; 8] {
SmallConvolveMontyField31::apply(
input,
MU::MATRIX_CIRC_MDS_8_COL,
<SmallConvolveMontyField31 as Convolve<MontyField31<FP>, i64, i64, i64>>::conv8,
)
}
}
impl<FP: MontyParameters, MU: MDSUtils> MdsPermutation<MontyField31<FP>, 8>
for MdsMatrixMontyField31<MU>
{
}
impl<FP: MontyParameters, MU: MDSUtils> Permutation<[MontyField31<FP>; 12]>
for MdsMatrixMontyField31<MU>
{
fn permute(&self, input: [MontyField31<FP>; 12]) -> [MontyField31<FP>; 12] {
SmallConvolveMontyField31::apply(
input,
MU::MATRIX_CIRC_MDS_12_COL,
<SmallConvolveMontyField31 as Convolve<MontyField31<FP>, i64, i64, i64>>::conv12,
)
}
}
impl<FP: MontyParameters, MU: MDSUtils> MdsPermutation<MontyField31<FP>, 12>
for MdsMatrixMontyField31<MU>
{
}
impl<FP: MontyParameters, MU: MDSUtils> Permutation<[MontyField31<FP>; 16]>
for MdsMatrixMontyField31<MU>
{
fn permute(&self, input: [MontyField31<FP>; 16]) -> [MontyField31<FP>; 16] {
SmallConvolveMontyField31::apply(
input,
MU::MATRIX_CIRC_MDS_16_COL,
<SmallConvolveMontyField31 as Convolve<MontyField31<FP>, i64, i64, i64>>::conv16,
)
}
}
impl<FP: MontyParameters, MU: MDSUtils> MdsPermutation<MontyField31<FP>, 16>
for MdsMatrixMontyField31<MU>
{
}
impl<FP, MU: MDSUtils> Permutation<[MontyField31<FP>; 24]> for MdsMatrixMontyField31<MU>
where
FP: BarrettParameters,
{
fn permute(&self, input: [MontyField31<FP>; 24]) -> [MontyField31<FP>; 24] {
LargeConvolveMontyField31::apply(
input,
MU::MATRIX_CIRC_MDS_24_COL,
<LargeConvolveMontyField31 as Convolve<MontyField31<FP>, i64, i64, i64>>::conv24,
)
}
}
impl<FP: BarrettParameters, MU: MDSUtils> MdsPermutation<MontyField31<FP>, 24>
for MdsMatrixMontyField31<MU>
{
}
impl<FP: BarrettParameters, MU: MDSUtils> Permutation<[MontyField31<FP>; 32]>
for MdsMatrixMontyField31<MU>
{
fn permute(&self, input: [MontyField31<FP>; 32]) -> [MontyField31<FP>; 32] {
LargeConvolveMontyField31::apply(
input,
MU::MATRIX_CIRC_MDS_32_COL,
<LargeConvolveMontyField31 as Convolve<MontyField31<FP>, i64, i64, i64>>::conv32,
)
}
}
impl<FP: BarrettParameters, MU: MDSUtils> MdsPermutation<MontyField31<FP>, 32>
for MdsMatrixMontyField31<MU>
{
}
impl<FP: BarrettParameters, MU: MDSUtils> Permutation<[MontyField31<FP>; 64]>
for MdsMatrixMontyField31<MU>
{
fn permute(&self, input: [MontyField31<FP>; 64]) -> [MontyField31<FP>; 64] {
LargeConvolveMontyField31::apply(
input,
MU::MATRIX_CIRC_MDS_64_COL,
<LargeConvolveMontyField31 as Convolve<MontyField31<FP>, i64, i64, i64>>::conv64,
)
}
}
impl<FP: BarrettParameters, MU: MDSUtils> MdsPermutation<MontyField31<FP>, 64>
for MdsMatrixMontyField31<MU>
{
}