use alloc::{collections::BTreeMap, vec::Vec};
use core::ops::Range;
use miden_core::{
Felt,
chiplets::hasher::Hasher,
field::{PrimeCharacteristicRing, QuadFelt},
utils::RowMajorMatrix,
};
use crate::{
logup::build_logup_aux_trace,
relations::ProvideMult,
transcript::poseidon2::{
NUM_MAIN_COLS, NUM_WITNESSES, Poseidon2Air,
digest::{P2Cap, P2Digest},
math::{NUM_CUBE_REGS, STATE_WIDTH},
program::PERIOD,
},
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct PermSeqId(u32);
impl PermSeqId {
pub fn seq(self) -> u32 {
self.0
}
#[cfg(test)]
pub(crate) fn forged(seq: u32) -> Self {
Self(seq)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PermSpan {
start: u32,
len: u32,
}
impl PermSpan {
fn new(range: Range<u32>) -> Self {
Self {
start: range.start,
len: range.end - range.start,
}
}
pub fn head(self) -> PermSeqId {
PermSeqId(self.start)
}
pub fn tail(self) -> PermSeqId {
PermSeqId(self.start + self.len - 1)
}
pub fn n_cycles(self) -> u32 {
self.len
}
}
#[derive(Debug, Clone)]
pub struct AbsorptionOutput {
pub digest: P2Digest,
pub span: PermSpan,
}
impl AbsorptionOutput {
pub fn head(&self) -> PermSeqId {
self.span.head()
}
pub fn tail(&self) -> PermSeqId {
self.span.tail()
}
}
pub fn apply_permutation(state_in: [Felt; STATE_WIDTH]) -> [Felt; STATE_WIDTH] {
let mut state = state_in;
Hasher::apply_permutation(&mut state);
state
}
#[derive(Debug)]
pub(crate) struct PreparedAbsorption {
cap: P2Cap,
blocks: Vec<([Felt; 4], [Felt; 4])>,
digests: Vec<P2Digest>,
}
impl PreparedAbsorption {
fn new(cap: P2Cap, blocks: Vec<([Felt; 4], [Felt; 4])>) -> Self {
assert!(!blocks.is_empty(), "absorption needs at least one block");
let mut current_cap = cap.as_array();
let mut digests = Vec::with_capacity(blocks.len());
for &(rate0, rate1) in &blocks {
let state_out = apply_permutation(state_from_chunks(rate0, rate1, current_cap));
digests.push(P2Digest(chunk_from_state(&state_out, 0)));
current_cap = chunk_from_state(&state_out, 8);
}
Self { cap, blocks, digests }
}
pub(crate) fn digest(&self) -> P2Digest {
*self.digests.last().expect("prepared absorption is non-empty")
}
}
fn absorb_oracle(cap: P2Cap, blocks: &[([Felt; 4], [Felt; 4])]) -> P2Digest {
let mut cap = cap.as_array();
let mut digest = [Felt::ZERO; 4];
for &(rate0, rate1) in blocks {
let state_out = apply_permutation(state_from_chunks(rate0, rate1, cap));
digest = chunk_from_state(&state_out, 0);
cap = chunk_from_state(&state_out, 8);
}
P2Digest(digest)
}
#[derive(Debug, Clone)]
struct RecordedAbsorption {
cap: P2Cap,
blocks: Vec<([Felt; 4], [Felt; 4])>,
#[allow(dead_code)]
digest: P2Digest,
range: Range<u32>,
in_mult: ProvideMult,
out_mult: ProvideMult,
}
#[derive(Debug, Clone, Default)]
pub struct Poseidon2Requires {
absorptions: Vec<RecordedAbsorption>,
by_digest: BTreeMap<P2Digest, usize>,
next_seq: u32,
}
impl Poseidon2Requires {
pub fn new() -> Self {
Self::default()
}
pub fn digest_of(cap: P2Cap, blocks: &[([Felt; 4], [Felt; 4])]) -> P2Digest {
absorb_oracle(cap, blocks)
}
pub(crate) fn prepare_absorption(
cap: P2Cap,
blocks: Vec<([Felt; 4], [Felt; 4])>,
) -> PreparedAbsorption {
PreparedAbsorption::new(cap, blocks)
}
pub(crate) fn require_prepared_absorption(
&mut self,
prepared: PreparedAbsorption,
) -> (AbsorptionOutput, Vec<P2Digest>) {
let PreparedAbsorption { cap, blocks, digests } = prepared;
let digest = *digests.last().expect("prepared absorption is non-empty");
let output = self.require_absorption_with_digest(cap, blocks, digest);
(output, digests)
}
pub fn require_absorption(
&mut self,
cap: P2Cap,
blocks: impl IntoIterator<Item = ([Felt; 4], [Felt; 4])>,
) -> AbsorptionOutput {
let blocks: Vec<_> = blocks.into_iter().collect();
assert!(!blocks.is_empty(), "absorption needs at least one block");
let digest = absorb_oracle(cap, &blocks);
self.require_absorption_with_digest(cap, blocks, digest)
}
fn require_absorption_with_digest(
&mut self,
cap: P2Cap,
blocks: Vec<([Felt; 4], [Felt; 4])>,
digest: P2Digest,
) -> AbsorptionOutput {
if let Some(&idx) = self.by_digest.get(&digest) {
let rec = &mut self.absorptions[idx];
rec.in_mult += 1;
return AbsorptionOutput {
digest,
span: PermSpan::new(rec.range.clone()),
};
}
self.lay_span(cap, blocks, digest, 1, 0)
}
pub fn require_one_shot(
&mut self,
cap: P2Cap,
rate0: [Felt; 4],
rate1: [Felt; 4],
) -> AbsorptionOutput {
self.require_absorption(cap, core::iter::once((rate0, rate1)))
}
pub fn require_digest(&mut self, digest: P2Digest) -> Option<PermSpan> {
let &idx = self.by_digest.get(&digest)?;
let rec = &mut self.absorptions[idx];
rec.out_mult += 1;
Some(PermSpan::new(rec.range.clone()))
}
pub fn lookup(&self, digest: P2Digest) -> Option<PermSpan> {
self.by_digest
.get(&digest)
.map(|&idx| PermSpan::new(self.absorptions[idx].range.clone()))
}
pub fn total_cycles(&self) -> u32 {
self.next_seq
}
fn lay_span(
&mut self,
cap: P2Cap,
blocks: Vec<([Felt; 4], [Felt; 4])>,
digest: P2Digest,
in_mult: ProvideMult,
out_mult: ProvideMult,
) -> AbsorptionOutput {
let n = blocks.len() as u32;
let range = self.next_seq..self.next_seq + n;
self.next_seq += n;
let idx = self.absorptions.len();
self.absorptions.push(RecordedAbsorption {
cap,
blocks,
digest,
range: range.clone(),
in_mult,
out_mult,
});
self.by_digest.insert(digest, idx);
AbsorptionOutput { digest, span: PermSpan::new(range) }
}
}
pub fn generate_trace(requires: Poseidon2Requires) -> RowMajorMatrix<Felt> {
let total_cycles = requires.next_seq as usize;
let height = (total_cycles * PERIOD).next_power_of_two().max(PERIOD);
let num_cycles = height / PERIOD;
let mut trace = Vec::with_capacity(height * NUM_MAIN_COLS);
for rec in &requires.absorptions {
let mut cap = rec.cap.as_array();
for (block_idx, &(rate0, rate1)) in rec.blocks.iter().enumerate() {
let cycle_idx = rec.range.start as usize + block_idx;
debug_assert_eq!(
trace.len(),
cycle_idx * PERIOD * NUM_MAIN_COLS,
"cycles laid in contiguous global order",
);
let is_absorb = block_idx > 0;
let state_in = state_from_chunks(rate0, rate1, cap);
let state_out =
write_cycle(&mut trace, cycle_idx, state_in, rec.in_mult, rec.out_mult, is_absorb);
cap = chunk_from_state(&state_out, 8);
}
}
for cycle in total_cycles..num_cycles {
let perm_seq_id =
Felt::new(cycle as u64).expect("perm_seq_id fits in canonical Goldilocks");
for _ in 0..PERIOD {
trace.push(perm_seq_id); trace.extend([Felt::ZERO; NUM_MAIN_COLS - 1]); }
}
debug_assert_eq!(trace.len(), height * NUM_MAIN_COLS);
RowMajorMatrix::new(trace, NUM_MAIN_COLS)
}
fn state_from_chunks(rate0: [Felt; 4], rate1: [Felt; 4], cap: [Felt; 4]) -> [Felt; STATE_WIDTH] {
let mut state = [Felt::ZERO; STATE_WIDTH];
state[0..4].copy_from_slice(&rate0);
state[4..8].copy_from_slice(&rate1);
state[8..12].copy_from_slice(&cap);
state
}
fn chunk_from_state(state: &[Felt; STATE_WIDTH], offset: usize) -> [Felt; 4] {
state[offset..offset + 4].try_into().expect("4-felt slice fits")
}
fn ext_cube_regs(sbox_in: &[Felt; STATE_WIDTH]) -> [Felt; NUM_CUBE_REGS] {
let mut regs = [Felt::ZERO; NUM_CUBE_REGS];
for (reg, x) in regs.iter_mut().zip(sbox_in.iter()) {
*reg = x.cube();
}
regs
}
fn write_cycle(
trace: &mut Vec<Felt>,
cycle_idx: usize,
initial_state: [Felt; STATE_WIDTH],
in_multiplicity: ProvideMult,
out_multiplicity: ProvideMult,
is_absorb: bool,
) -> [Felt; STATE_WIDTH] {
let perm_seq_id =
Felt::new(cycle_idx as u64).expect("perm_seq_id fits in canonical Goldilocks");
let in_mult = Felt::from(in_multiplicity);
let out_mult = Felt::from(out_multiplicity);
let absorb = Felt::from(is_absorb as u8);
let mut state = initial_state;
let mut row0_sbox_in = state;
Hasher::apply_matmul_external(&mut row0_sbox_in);
Hasher::add_rc(&mut row0_sbox_in, &Hasher::ARK_EXT_INITIAL[0]);
push_row(
trace,
&state,
&[Felt::ZERO; 3],
&ext_cube_regs(&row0_sbox_in),
perm_seq_id,
in_mult,
out_mult,
absorb,
);
Hasher::apply_matmul_external(&mut state);
Hasher::add_rc(&mut state, &Hasher::ARK_EXT_INITIAL[0]);
Hasher::apply_sbox(&mut state);
Hasher::apply_matmul_external(&mut state);
for r in 1..=3 {
let mut sbox_in = state;
Hasher::add_rc(&mut sbox_in, &Hasher::ARK_EXT_INITIAL[r]);
push_row(
trace,
&state,
&[Felt::ZERO; 3],
&ext_cube_regs(&sbox_in),
perm_seq_id,
in_mult,
out_mult,
absorb,
);
Hasher::add_rc(&mut state, &Hasher::ARK_EXT_INITIAL[r]);
Hasher::apply_sbox(&mut state);
Hasher::apply_matmul_external(&mut state);
}
for triple in 0..7_usize {
let base = triple * 3;
let pre_state = state;
let mut witnesses = [Felt::ZERO; 3];
let mut cube_regs = [Felt::ZERO; NUM_CUBE_REGS];
for (k, witness) in witnesses.iter_mut().enumerate() {
let sbox_in = state[0] + Hasher::ARK_INT[base + k];
cube_regs[k] = sbox_in.cube();
let sbox_out = sbox_in.exp_const_u64::<7>();
*witness = sbox_out;
state[0] = sbox_out;
Hasher::matmul_internal(&mut state, Hasher::MAT_DIAG);
}
push_row(
trace,
&pre_state,
&witnesses,
&cube_regs,
perm_seq_id,
in_mult,
out_mult,
absorb,
);
}
let pre_state = state;
let w0_in = state[0] + Hasher::ARK_INT[21];
let w0 = w0_in.exp_const_u64::<7>();
state[0] = w0;
Hasher::matmul_internal(&mut state, Hasher::MAT_DIAG);
let mut ext_sbox_in = state;
Hasher::add_rc(&mut ext_sbox_in, &Hasher::ARK_EXT_TERMINAL[0]);
let mut row11_cube_regs = ext_cube_regs(&ext_sbox_in);
row11_cube_regs[STATE_WIDTH] = w0_in.cube();
Hasher::add_rc(&mut state, &Hasher::ARK_EXT_TERMINAL[0]);
Hasher::apply_sbox(&mut state);
Hasher::apply_matmul_external(&mut state);
push_row(
trace,
&pre_state,
&[w0, Felt::ZERO, Felt::ZERO],
&row11_cube_regs,
perm_seq_id,
in_mult,
out_mult,
absorb,
);
for r in 1..=3 {
let mut sbox_in = state;
Hasher::add_rc(&mut sbox_in, &Hasher::ARK_EXT_TERMINAL[r]);
push_row(
trace,
&state,
&[Felt::ZERO; 3],
&ext_cube_regs(&sbox_in),
perm_seq_id,
in_mult,
out_mult,
absorb,
);
Hasher::add_rc(&mut state, &Hasher::ARK_EXT_TERMINAL[r]);
Hasher::apply_sbox(&mut state);
Hasher::apply_matmul_external(&mut state);
}
push_row(
trace,
&state,
&[Felt::ZERO; 3],
&[Felt::ZERO; NUM_CUBE_REGS],
perm_seq_id,
in_mult,
out_mult,
absorb,
);
state
}
fn push_row(
trace: &mut Vec<Felt>,
state: &[Felt; STATE_WIDTH],
witnesses: &[Felt; NUM_WITNESSES],
cube_regs: &[Felt; NUM_CUBE_REGS],
perm_seq_id: Felt,
in_multiplicity: Felt,
out_multiplicity: Felt,
is_absorb: Felt,
) {
trace.extend([perm_seq_id, in_multiplicity, out_multiplicity, is_absorb]);
trace.extend(*state);
trace.extend(*witnesses);
trace.extend(*cube_regs);
}
pub(crate) fn build_aux(
main: &RowMajorMatrix<Felt>,
challenges: &[QuadFelt],
) -> (RowMajorMatrix<QuadFelt>, Vec<QuadFelt>) {
build_logup_aux_trace(&Poseidon2Air, main, challenges)
}