use crate::{
crypto_hash::PoseidonGrainLFSR,
traits::{AlgebraicSponge, DefaultCapacityAlgebraicSponge, DuplexSpongeMode, SpongeParameters},
CryptoHash,
};
use snarkvm_fields::{
Fp256,
Fp256Parameters,
Fp384,
Fp384Parameters,
Fp768,
Fp768Parameters,
PoseidonDefaultParameters,
PrimeField,
};
use snarkvm_utilities::{FromBytes, ToBytes};
use smallvec::SmallVec;
use std::{
io::{Read, Result as IoResult, Write},
ops::{Index, IndexMut, Range},
sync::Arc,
};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PoseidonParameters<F: PrimeField, const RATE: usize, const CAPACITY: usize> {
pub full_rounds: usize,
pub partial_rounds: usize,
pub alpha: u64,
pub ark: Vec<Vec<F>>,
pub mds: Vec<Vec<F>>,
}
impl<F: PrimeField, const RATE: usize, const CAPACITY: usize> PoseidonParameters<F, RATE, CAPACITY> {
pub fn new(full_rounds: usize, partial_rounds: usize, alpha: u64, mds: Vec<Vec<F>>, ark: Vec<Vec<F>>) -> Self {
assert_eq!(ark.len(), full_rounds + partial_rounds);
for item in &ark {
assert_eq!(item.len(), RATE + CAPACITY);
}
assert_eq!(mds.len(), RATE + CAPACITY);
for item in &mds {
assert_eq!(item.len(), RATE + CAPACITY);
}
Self {
full_rounds,
partial_rounds,
alpha,
mds,
ark,
}
}
}
impl<F: PrimeField, const RATE: usize, const CAPACITY: usize> ToBytes for PoseidonParameters<F, RATE, CAPACITY> {
#[inline]
fn write_le<W: Write>(&self, mut writer: W) -> IoResult<()> {
(self.full_rounds as u32).write_le(&mut writer)?;
(self.partial_rounds as u32).write_le(&mut writer)?;
self.alpha.write_le(&mut writer)?;
(self.ark.len() as u32).write_le(&mut writer)?;
for fields in &self.ark {
(fields.len() as u32).write_le(&mut writer)?;
for field in fields {
field.write_le(&mut writer)?;
}
}
(self.mds.len() as u32).write_le(&mut writer)?;
for fields in &self.mds {
(fields.len() as u32).write_le(&mut writer)?;
for field in fields {
field.write_le(&mut writer)?;
}
}
(RATE as u32).write_le(&mut writer)?;
(CAPACITY as u32).write_le(&mut writer)
}
}
impl<F: PrimeField, const RATE: usize, const CAPACITY: usize> FromBytes for PoseidonParameters<F, RATE, CAPACITY> {
#[inline]
fn read_le<R: Read>(mut reader: R) -> IoResult<Self> {
let full_rounds: u32 = FromBytes::read_le(&mut reader)?;
let partial_rounds: u32 = FromBytes::read_le(&mut reader)?;
let alpha: u64 = FromBytes::read_le(&mut reader)?;
let ark_length: u32 = FromBytes::read_le(&mut reader)?;
let mut ark = Vec::with_capacity(ark_length as usize);
for _ in 0..ark_length {
let num_fields: u32 = FromBytes::read_le(&mut reader)?;
let mut fields = Vec::with_capacity(num_fields as usize);
for _ in 0..num_fields {
let field: F = FromBytes::read_le(&mut reader)?;
fields.push(field);
}
ark.push(fields);
}
let mds_length: u32 = FromBytes::read_le(&mut reader)?;
let mut mds = Vec::with_capacity(mds_length as usize);
for _ in 0..mds_length {
let num_fields: u32 = FromBytes::read_le(&mut reader)?;
let mut fields = Vec::with_capacity(num_fields as usize);
for _ in 0..num_fields {
let field: F = FromBytes::read_le(&mut reader)?;
fields.push(field);
}
mds.push(fields);
}
let rate: u32 = FromBytes::read_le(&mut reader)?;
let capacity: u32 = FromBytes::read_le(&mut reader)?;
if rate != RATE as u32 || capacity != CAPACITY as u32 {
return Err(std::io::ErrorKind::Other.into());
}
Ok(Self::new(
full_rounds as usize,
partial_rounds as usize,
alpha,
mds,
ark,
))
}
}
#[derive(Clone, Debug)]
pub struct PoseidonSponge<F: PrimeField, const RATE: usize, const CAPACITY: usize> {
pub parameters: Arc<PoseidonParameters<F, RATE, CAPACITY>>,
pub state: State<F, RATE, CAPACITY>,
pub mode: DuplexSpongeMode,
}
#[derive(Copy, Clone, Debug)]
pub struct State<F: PrimeField, const RATE: usize, const CAPACITY: usize> {
capacity_state: [F; CAPACITY],
rate_state: [F; RATE],
}
impl<F: PrimeField, const RATE: usize, const CAPACITY: usize> Default for State<F, RATE, CAPACITY> {
fn default() -> Self {
Self {
capacity_state: [F::zero(); CAPACITY],
rate_state: [F::zero(); RATE],
}
}
}
impl<F: PrimeField, const RATE: usize, const CAPACITY: usize> State<F, RATE, CAPACITY> {
pub fn iter(&self) -> impl Iterator<Item = &F> {
self.capacity_state.iter().chain(self.rate_state.iter())
}
pub fn iter_mut(&mut self) -> impl Iterator<Item = &mut F> {
self.capacity_state.iter_mut().chain(self.rate_state.iter_mut())
}
pub fn range(&self, range: Range<usize>) -> impl Iterator<Item = &F> {
let start = range.start;
let end = range.end;
assert!(
start < end,
"start < end in range: start is {} but end is {}",
start,
end
);
assert!(
end <= RATE + CAPACITY,
"Range out of bounds: range is {:?} but length is {}",
range,
RATE + CAPACITY
);
if start >= CAPACITY {
self.rate_state[(start - CAPACITY)..(end - CAPACITY)].iter().chain(&[]) } else if end > CAPACITY {
self.capacity_state[start..]
.iter()
.chain(self.rate_state[..(end - CAPACITY)].iter())
} else {
debug_assert!(end <= CAPACITY);
debug_assert!(start < CAPACITY);
self.capacity_state[start..end].iter().chain(&[])
}
}
}
impl<F: PrimeField, const RATE: usize, const CAPACITY: usize> Index<usize> for State<F, RATE, CAPACITY> {
type Output = F;
fn index(&self, index: usize) -> &Self::Output {
assert!(
index < RATE + CAPACITY,
"Index out of bounds: index is {} but length is {}",
index,
RATE + CAPACITY
);
if index < CAPACITY {
&self.capacity_state[index]
} else {
&self.rate_state[index - CAPACITY]
}
}
}
impl<F: PrimeField, const RATE: usize, const CAPACITY: usize> IndexMut<usize> for State<F, RATE, CAPACITY> {
fn index_mut(&mut self, index: usize) -> &mut Self::Output {
assert!(
index < RATE + CAPACITY,
"Index out of bounds: index is {} but length is {}",
index,
RATE + CAPACITY
);
if index < CAPACITY {
&mut self.capacity_state[index]
} else {
&mut self.rate_state[index - CAPACITY]
}
}
}
impl<F: PrimeField, const RATE: usize, const CAPACITY: usize> PoseidonSponge<F, RATE, CAPACITY> {
#[inline]
fn apply_s_box(&self, state: &mut State<F, RATE, CAPACITY>, is_full_round: bool) {
if is_full_round {
for elem in state.iter_mut() {
*elem = elem.pow(&[self.parameters.alpha]);
}
}
else {
state[0] = state[0].pow(&[self.parameters.alpha]);
}
}
#[inline]
fn apply_ark(&self, state: &mut State<F, RATE, CAPACITY>, round_number: usize) {
for (state_elem, ark_elem) in state.iter_mut().zip(&self.parameters.ark[round_number]) {
*state_elem += ark_elem;
}
}
#[inline]
fn apply_mds(&self, state: &mut State<F, RATE, CAPACITY>) {
let mut new_state = State::default();
new_state
.iter_mut()
.zip(&self.parameters.mds)
.for_each(|(new_elem, mds_row)| {
*new_elem = state
.iter()
.zip(mds_row)
.map(|(state_elem, &mds_elem)| mds_elem * state_elem)
.sum::<F>();
});
*state = new_state;
}
fn permute(&mut self) {
let full_rounds_over_2 = self.parameters.full_rounds / 2;
let partial_round_range = full_rounds_over_2..(full_rounds_over_2 + self.parameters.partial_rounds);
let mut state = self.state;
for i in 0..(self.parameters.partial_rounds + self.parameters.full_rounds) {
let is_full_round = !partial_round_range.contains(&i);
self.apply_ark(&mut state, i);
self.apply_s_box(&mut state, is_full_round);
self.apply_mds(&mut state);
}
self.state = state;
}
fn absorb_internal(&mut self, mut rate_start: usize, elements: &[F]) {
if elements.is_empty() {
return;
}
let first_chunk_size = std::cmp::min(RATE - rate_start, elements.len());
let num_elements_remaining = elements.len() - first_chunk_size;
let (first_chunk, rest_chunk) = elements.split_at(first_chunk_size);
let rest_chunks = rest_chunk.chunks(RATE);
let total_num_chunks = 1 + (num_elements_remaining / RATE) +
usize::from((num_elements_remaining % RATE) != 0);
for (i, chunk) in std::iter::once(first_chunk).chain(rest_chunks).enumerate() {
for (element, state_elem) in chunk.iter().zip(&mut self.state.rate_state[rate_start..]) {
*state_elem += element;
}
if i == total_num_chunks - 1 {
self.mode = DuplexSpongeMode::Absorbing {
next_absorb_index: rate_start + chunk.len(),
};
return;
} else {
self.permute();
}
rate_start = 0;
}
}
fn squeeze_helper(&mut self, output: &mut [F]) {
match self.mode {
DuplexSpongeMode::Absorbing { next_absorb_index: _ } => {
self.permute();
self.squeeze_internal(0, output);
}
DuplexSpongeMode::Squeezing { mut next_squeeze_index } => {
if next_squeeze_index == RATE {
self.permute();
next_squeeze_index = 0;
}
self.squeeze_internal(next_squeeze_index, output);
}
};
}
fn squeeze_internal(&mut self, mut rate_start: usize, output: &mut [F]) {
let output_length = output.len();
if output_length == 0 {
return;
}
let first_chunk_size = std::cmp::min(RATE - rate_start, output.len());
let num_output_remaining = output.len() - first_chunk_size;
let (first_chunk, rest_chunk) = output.split_at_mut(first_chunk_size);
assert_eq!(rest_chunk.len(), num_output_remaining);
let rest_chunks = rest_chunk.chunks_mut(RATE);
let total_num_chunks = 1 + (num_output_remaining / RATE) +
usize::from((num_output_remaining % RATE) != 0);
for (i, chunk) in std::iter::once(first_chunk).chain(rest_chunks).enumerate() {
let range = rate_start..(rate_start + chunk.len());
debug_assert_eq!(
chunk.len(),
self.state.rate_state[range.clone()].len(),
"failed with squeeze {} at rate {} and rate_start {}",
output_length,
RATE,
rate_start
);
chunk.copy_from_slice(&self.state.rate_state[range]);
if i == total_num_chunks - 1 {
self.mode = DuplexSpongeMode::Squeezing {
next_squeeze_index: (rate_start + chunk.len()),
};
return;
} else {
self.permute();
}
rate_start = 0;
}
}
}
impl<F: PoseidonDefaultParametersField, const RATE: usize> PoseidonSponge<F, RATE, 1> {
pub fn sample_default_parameters() -> Arc<PoseidonParameters<F, RATE, 1>> {
Arc::new(F::get_default_poseidon_parameters::<RATE>(false).unwrap())
}
pub fn with_default_parameters() -> Self {
let parameters = Arc::new(F::get_default_poseidon_parameters::<RATE>(false).unwrap());
let state = State::default();
let mode = DuplexSpongeMode::Absorbing { next_absorb_index: 0 };
Self {
parameters,
state,
mode,
}
}
}
impl<F: PoseidonDefaultParametersField, const RATE: usize, const CAPACITY: usize> SpongeParameters<RATE, CAPACITY>
for Arc<PoseidonParameters<F, RATE, CAPACITY>>
{
}
impl<F: PoseidonDefaultParametersField, const RATE: usize, const CAPACITY: usize> AlgebraicSponge<F, RATE, CAPACITY>
for PoseidonSponge<F, RATE, CAPACITY>
{
type Parameters = Arc<PoseidonParameters<F, RATE, CAPACITY>>;
fn with_parameters(parameters: &Self::Parameters) -> Self {
let state = State::default();
let mode = DuplexSpongeMode::Absorbing { next_absorb_index: 0 };
Self {
parameters: parameters.clone(),
state,
mode,
}
}
fn absorb(&mut self, input: &[F]) {
if input.is_empty() {
return;
}
match self.mode {
DuplexSpongeMode::Absorbing { mut next_absorb_index } => {
if next_absorb_index == RATE {
self.permute();
next_absorb_index = 0;
}
self.absorb_internal(next_absorb_index, input);
}
DuplexSpongeMode::Squeezing { next_squeeze_index: _ } => {
self.permute();
self.absorb_internal(0, input);
}
};
}
fn squeeze_field_elements(&mut self, num_elements: usize) -> SmallVec<[F; 10]> {
if num_elements == 0 {
return SmallVec::new();
}
let mut buf = if num_elements <= 10 {
smallvec::smallvec_inline![F::zero(); 10]
} else {
smallvec::smallvec![F::zero(); num_elements]
};
self.squeeze_helper(&mut buf[..num_elements]);
buf.truncate(num_elements);
buf
}
}
impl<F: PoseidonDefaultParametersField, const RATE: usize> DefaultCapacityAlgebraicSponge<F, RATE>
for PoseidonSponge<F, RATE, 1>
{
fn sample_parameters() -> Arc<PoseidonParameters<F, RATE, 1>> {
Arc::new(F::get_default_poseidon_parameters::<RATE>(false).unwrap())
}
fn with_default_parameters() -> Self {
let parameters = Arc::new(F::get_default_poseidon_parameters::<RATE>(false).unwrap());
let state = State::default();
let mode = DuplexSpongeMode::Absorbing { next_absorb_index: 0 };
Self {
parameters,
state,
mode,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PoseidonCryptoHash<
F: PrimeField + PoseidonDefaultParametersField,
const RATE: usize,
const OPTIMIZED_FOR_WEIGHTS: bool,
> {
parameters: Arc<PoseidonParameters<F, RATE, 1>>,
}
impl<F: PrimeField + PoseidonDefaultParametersField, const RATE: usize, const OPTIMIZED_FOR_WEIGHTS: bool> CryptoHash
for PoseidonCryptoHash<F, RATE, OPTIMIZED_FOR_WEIGHTS>
{
type Input = F;
type Output = F;
type Parameters = Arc<PoseidonParameters<F, RATE, 1>>;
fn setup() -> Self {
Self {
parameters: Arc::new(F::get_default_poseidon_parameters::<RATE>(OPTIMIZED_FOR_WEIGHTS).unwrap()),
}
}
fn evaluate(&self, input: &[Self::Input]) -> Self::Output {
let mut sponge = PoseidonSponge::<F, RATE, 1>::with_parameters(&self.parameters);
sponge.absorb(input);
sponge.squeeze_field_elements(1)[0]
}
fn parameters(&self) -> &Self::Parameters {
&self.parameters
}
}
impl<F: PrimeField + PoseidonDefaultParametersField, const RATE: usize, const OPTIMIZED_FOR_WEIGHTS: bool>
From<PoseidonParameters<F, RATE, 1>> for PoseidonCryptoHash<F, RATE, OPTIMIZED_FOR_WEIGHTS>
{
fn from(parameters: PoseidonParameters<F, RATE, 1>) -> Self {
Self {
parameters: Arc::new(parameters),
}
}
}
impl<F: PrimeField + PoseidonDefaultParametersField, const RATE: usize, const OPTIMIZED_FOR_WEIGHTS: bool>
From<Arc<PoseidonParameters<F, RATE, 1>>> for PoseidonCryptoHash<F, RATE, OPTIMIZED_FOR_WEIGHTS>
{
fn from(parameters: Arc<PoseidonParameters<F, RATE, 1>>) -> Self {
Self { parameters }
}
}
impl<F: PrimeField + PoseidonDefaultParametersField, const RATE: usize, const OPTIMIZED_FOR_WEIGHTS: bool> ToBytes
for PoseidonCryptoHash<F, RATE, OPTIMIZED_FOR_WEIGHTS>
{
#[inline]
fn write_le<W: Write>(&self, mut writer: W) -> IoResult<()> {
self.parameters.write_le(&mut writer)
}
}
impl<F: PrimeField + PoseidonDefaultParametersField, const RATE: usize, const OPTIMIZED_FOR_WEIGHTS: bool> FromBytes
for PoseidonCryptoHash<F, RATE, OPTIMIZED_FOR_WEIGHTS>
{
#[inline]
fn read_le<R: Read>(mut reader: R) -> IoResult<Self> {
let parameters: PoseidonParameters<F, RATE, 1> = FromBytes::read_le(&mut reader)?;
Ok(Self::from(parameters))
}
}
pub trait PoseidonDefaultParametersField: PrimeField {
fn get_default_poseidon_parameters<const RATE: usize>(
optimized_for_weights: bool,
) -> Option<PoseidonParameters<Self, RATE, 1>>;
}
pub fn get_default_poseidon_parameters_internal<F: PrimeField, P: PoseidonDefaultParameters, const RATE: usize>(
optimized_for_weights: bool,
) -> Option<PoseidonParameters<F, RATE, 1>> {
let params_set = if !optimized_for_weights {
P::PARAMS_OPT_FOR_CONSTRAINTS
} else {
P::PARAMS_OPT_FOR_WEIGHTS
};
params_set.iter().find(|p| p.rate == RATE).map(|p| {
let (ark, mds) = find_poseidon_ark_and_mds::<F, RATE>(
P::MODULUS_BITS as u64,
p.full_rounds as u64,
p.partial_rounds as u64,
p.skip_matrices as u64,
);
PoseidonParameters {
full_rounds: p.full_rounds,
partial_rounds: p.partial_rounds,
alpha: p.alpha as u64,
ark,
mds,
}
})
}
pub fn find_poseidon_ark_and_mds<F: PrimeField, const RATE: usize>(
prime_bits: u64,
full_rounds: u64,
partial_rounds: u64,
skip_matrices: u64,
) -> (Vec<Vec<F>>, Vec<Vec<F>>) {
let mut lfsr = PoseidonGrainLFSR::new(false, prime_bits, (RATE + 1) as u64, full_rounds, partial_rounds);
let mut ark = Vec::<Vec<F>>::new();
for _ in 0..(full_rounds + partial_rounds) {
ark.push(lfsr.get_field_elements_rejection_sampling(RATE + 1));
}
let mut mds = vec![vec![F::zero(); RATE + 1]; RATE + 1];
for _ in 0..skip_matrices {
let _ = lfsr.get_field_elements_mod_p::<F>(2 * (RATE + 1));
}
let xs = lfsr.get_field_elements_mod_p::<F>(RATE + 1);
let ys = lfsr.get_field_elements_mod_p::<F>(RATE + 1);
for (i, x) in xs.iter().enumerate().take(RATE + 1) {
for (j, y) in ys.iter().enumerate().take(RATE + 1) {
mds[i][j] = (*x + y).inverse().unwrap();
}
}
(ark, mds)
}
macro_rules! impl_poseidon_default_parameters_field {
($field: ident, $params: ident) => {
impl<P: $params + PoseidonDefaultParameters> PoseidonDefaultParametersField for $field<P> {
fn get_default_poseidon_parameters<const RATE: usize>(
optimized_for_weights: bool,
) -> Option<PoseidonParameters<Self, RATE, 1>> {
get_default_poseidon_parameters_internal::<Self, P, RATE>(optimized_for_weights)
}
}
};
}
impl_poseidon_default_parameters_field!(Fp256, Fp256Parameters);
impl_poseidon_default_parameters_field!(Fp384, Fp384Parameters);
impl_poseidon_default_parameters_field!(Fp768, Fp768Parameters);