use p3_field::PrimeField32;
use p3_poseidon2::DiffusionPermutation;
use p3_symmetric::Permutation;
use serde::{Deserialize, Serialize};
use crate::{monty_reduce, to_babybear_array, BabyBear};
pub const MONTY_INVERSE: BabyBear = BabyBear { value: 1 };
pub const POSEIDON2_INTERNAL_MATRIX_DIAG_16_BABYBEAR_MONTY: [BabyBear; 16] = to_babybear_array([
BabyBear::ORDER_U32 - 2,
1,
1 << 1,
1 << 2,
1 << 3,
1 << 4,
1 << 5,
1 << 6,
1 << 7,
1 << 8,
1 << 9,
1 << 10,
1 << 11,
1 << 12,
1 << 13,
1 << 15,
]);
const POSEIDON2_INTERNAL_MATRIX_DIAG_16_MONTY_SHIFTS: [u8; 15] =
[0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 15];
pub const POSEIDON2_INTERNAL_MATRIX_DIAG_24_BABYBEAR_MONTY: [BabyBear; 24] = to_babybear_array([
BabyBear::ORDER_U32 - 2,
1,
1 << 1,
1 << 2,
1 << 3,
1 << 4,
1 << 5,
1 << 6,
1 << 7,
1 << 8,
1 << 9,
1 << 10,
1 << 11,
1 << 12,
1 << 13,
1 << 14,
1 << 15,
1 << 16,
1 << 18,
1 << 19,
1 << 20,
1 << 21,
1 << 22,
1 << 23,
]);
const POSEIDON2_INTERNAL_MATRIX_DIAG_24_MONTY_SHIFTS: [u8; 23] = [
0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 18, 19, 20, 21, 22, 23,
];
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct DiffusionMatrixBabyBear;
impl Permutation<[BabyBear; 16]> for DiffusionMatrixBabyBear {
#[inline]
fn permute_mut(&self, state: &mut [BabyBear; 16]) {
let part_sum: u64 = state.iter().skip(1).map(|x| x.value as u64).sum();
let full_sum = part_sum + (state[0].value as u64);
let s0 = part_sum + (-state[0]).value as u64;
state[0] = BabyBear {
value: monty_reduce(s0),
};
for i in 1..16 {
let si = full_sum
+ ((state[i].value as u64)
<< POSEIDON2_INTERNAL_MATRIX_DIAG_16_MONTY_SHIFTS[i - 1]);
state[i] = BabyBear {
value: monty_reduce(si),
};
}
}
}
impl DiffusionPermutation<BabyBear, 16> for DiffusionMatrixBabyBear {}
impl Permutation<[BabyBear; 24]> for DiffusionMatrixBabyBear {
#[inline]
fn permute_mut(&self, state: &mut [BabyBear; 24]) {
let part_sum: u64 = state.iter().skip(1).map(|x| x.value as u64).sum();
let full_sum = part_sum + (state[0].value as u64);
let s0 = part_sum + (-state[0]).value as u64;
state[0] = BabyBear {
value: monty_reduce(s0),
};
for i in 1..24 {
let si = full_sum
+ ((state[i].value as u64)
<< POSEIDON2_INTERNAL_MATRIX_DIAG_24_MONTY_SHIFTS[i - 1]);
state[i] = BabyBear {
value: monty_reduce(si),
};
}
}
}
impl DiffusionPermutation<BabyBear, 24> for DiffusionMatrixBabyBear {}
#[cfg(test)]
mod tests {
use p3_field::AbstractField;
use p3_poseidon2::{Poseidon2, Poseidon2ExternalMatrixGeneral};
use rand::SeedableRng;
use rand_xoshiro::Xoroshiro128Plus;
use super::*;
type F = BabyBear;
fn poseidon2_babybear<const WIDTH: usize, const D: u64, DiffusionMatrix>(
input: &mut [F; WIDTH],
diffusion_matrix: DiffusionMatrix,
) where
DiffusionMatrix: DiffusionPermutation<F, WIDTH>,
{
let mut rng = Xoroshiro128Plus::seed_from_u64(1);
let poseidon2: Poseidon2<F, Poseidon2ExternalMatrixGeneral, DiffusionMatrix, WIDTH, D> =
Poseidon2::new_from_rng_128(Poseidon2ExternalMatrixGeneral, diffusion_matrix, &mut rng);
poseidon2.permute_mut(input);
}
#[test]
fn test_poseidon2_width_16_random() {
let mut input: [F; 16] = [
894848333, 1437655012, 1200606629, 1690012884, 71131202, 1749206695, 1717947831,
120589055, 19776022, 42382981, 1831865506, 724844064, 171220207, 1299207443, 227047920,
1783754913,
]
.map(F::from_canonical_u32);
let expected: [F; 16] = [
512585766, 975869435, 1921378527, 1238606951, 899635794, 132650430, 1426417547,
1734425242, 57415409, 67173027, 1535042492, 1318033394, 1070659233, 17258943,
856719028, 1500534995,
]
.map(F::from_canonical_u32);
poseidon2_babybear::<16, 7, _>(&mut input, DiffusionMatrixBabyBear);
assert_eq!(input, expected);
}
#[test]
fn test_poseidon2_width_24_random() {
let mut input: [F; 24] = [
886409618, 1327899896, 1902407911, 591953491, 648428576, 1844789031, 1198336108,
355597330, 1799586834, 59617783, 790334801, 1968791836, 559272107, 31054313,
1042221543, 474748436, 135686258, 263665994, 1962340735, 1741539604, 449439011,
1131357108, 50869465, 1589724894,
]
.map(F::from_canonical_u32);
let expected: [F; 24] = [
162275163, 462059149, 1096991565, 924509284, 300323988, 608502870, 427093935,
733126108, 1676785000, 669115065, 441326760, 60861458, 124006210, 687842154, 270552480,
1279931581, 1030167257, 126690434, 1291783486, 669126431, 1320670824, 1121967237,
458234203, 142219603,
]
.map(F::from_canonical_u32);
poseidon2_babybear::<24, 7, _>(&mut input, DiffusionMatrixBabyBear);
assert_eq!(input, expected);
}
}