use alloc::vec::Vec;
use core::ops::Range;
use miden_core::{Felt, deferred::Node, field::QuadFelt, utils::RowMajorMatrix};
use crate::{
hash::chunk::{ChunkAir, NUM_F, NUM_MAIN_COLS},
logup::build_logup_aux_trace,
transcript::poseidon2::{
digest::{P2Cap, P2Digest},
trace::{PermSpan, Poseidon2Requires},
},
};
#[derive(Debug, Clone)]
pub struct Invocation {
pub input: Vec<u8>,
}
#[derive(Debug, Clone)]
pub struct ChunkOutput {
pub digest: P2Digest,
pub chunk_head: ChunkSeqId,
pub perm_span: PermSpan,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ChunkSeqId(u32);
impl ChunkSeqId {
pub fn seq(self) -> u32 {
self.0
}
pub fn ptr(self) -> u32 {
self.0 * 4
}
#[cfg(test)]
pub(crate) fn forged(seq: u32) -> Self {
Self(seq)
}
}
#[derive(Debug, Clone)]
struct ChunkRecord {
f_per_chunk: Vec<[Felt; NUM_F]>,
chunk_seq_id_range: Range<u32>,
perm_span: PermSpan,
}
#[derive(Debug, Clone, Default)]
pub struct ChunkRequires {
records: Vec<ChunkRecord>,
next_chunk_seq: u32,
}
impl ChunkRequires {
pub fn new() -> Self {
Self::default()
}
pub fn require(&mut self, inv: &Invocation, p2: &mut Poseidon2Requires) -> ChunkOutput {
let f_per_chunk = Node::chunks_from_bytes(&inv.input)
.payload()
.as_data()
.expect("chunks_from_bytes creates data payload")
.to_vec();
debug_assert!(!f_per_chunk.is_empty(), "chunks_from_bytes guarantees ≥1 chunk",);
let rate_pairs: Vec<([Felt; 4], [Felt; 4])> = f_per_chunk
.iter()
.map(|f| {
let rate0: [Felt; 4] = f[0..4].try_into().expect("rate0 slice");
let rate1: [Felt; 4] = f[4..8].try_into().expect("rate1 slice");
(rate0, rate1)
})
.collect();
let p2_out = p2.require_absorption(P2Cap::chunk(), rate_pairs.iter().copied());
let n = f_per_chunk.len() as u32;
let chunk_head = ChunkSeqId(self.next_chunk_seq);
let chunk_seq_id_range = self.next_chunk_seq..self.next_chunk_seq + n;
self.next_chunk_seq += n;
self.records.push(ChunkRecord {
f_per_chunk,
chunk_seq_id_range,
perm_span: p2_out.span,
});
ChunkOutput {
digest: p2_out.digest,
chunk_head,
perm_span: p2_out.span,
}
}
pub fn total_chunks(&self) -> u32 {
self.next_chunk_seq
}
}
pub fn generate_trace(requires: ChunkRequires) -> RowMajorMatrix<Felt> {
generate_trace_padded_to(requires, 0)
}
pub(crate) fn generate_trace_padded_to(
requires: ChunkRequires,
min_height: usize,
) -> RowMajorMatrix<Felt> {
let total_chunks = requires.total_chunks() as usize;
let height = total_chunks.next_power_of_two().max(2).max(min_height);
let mut trace = Vec::with_capacity(height * NUM_MAIN_COLS);
let mut chunk_seq_id = 0u32;
let mut next_perm_seq_id = 0u32;
for rec in &requires.records {
debug_assert_eq!(chunk_seq_id, rec.chunk_seq_id_range.start, "contiguous chunk_seq_id");
let perm_start = rec.perm_span.head().seq();
for (c, f) in rec.f_per_chunk.iter().enumerate() {
trace.extend([
Felt::new(chunk_seq_id as u64).expect("chunk_seq_id fits"),
Felt::new((perm_start + c as u32) as u64).expect("perm_seq_id fits"),
Felt::ONE,
Felt::from(u8::from(c == 0)),
]);
trace.extend(*f);
chunk_seq_id += 1;
}
next_perm_seq_id = rec.perm_span.tail().seq() + 1;
}
for _ in total_chunks..height {
trace.extend([
Felt::new(chunk_seq_id as u64).expect("chunk_seq_id fits"),
Felt::new(next_perm_seq_id as u64).expect("perm_seq_id fits"),
Felt::ZERO,
Felt::ZERO,
]);
trace.extend([Felt::ZERO; NUM_F]);
chunk_seq_id += 1;
next_perm_seq_id += 1;
}
debug_assert_eq!(trace.len(), height * NUM_MAIN_COLS);
RowMajorMatrix::new(trace, NUM_MAIN_COLS)
}
pub(crate) fn build_aux(
main: &RowMajorMatrix<Felt>,
challenges: &[QuadFelt],
) -> (RowMajorMatrix<QuadFelt>, Vec<QuadFelt>) {
build_logup_aux_trace(&ChunkAir, main, challenges)
}