use alloc::{collections::BTreeMap, vec, vec::Vec};
use miden_core::{
Felt,
deferred::{Digest, Node},
field::QuadFelt,
utils::RowMajorMatrix,
};
use miden_precompiles::Keccak256Precompile;
use crate::{
hash::{
chunk::trace::{ChunkRequires, ChunkSeqId},
keccak::{
digest::KeccakDigest,
node::{KeccakNodeAir, NUM_HASH, NUM_MAIN_COLS},
round::RoundRequires,
sponge::trace::{
Invocation as SpongeInvocation, SpongeRequires, SpongeSeqId, keccak_oracle,
},
},
},
logup::build_logup_aux_trace,
primitives::byte_pair_lut::BytePairLutRequires,
relations::ProvideMult,
transcript::poseidon2::{
digest::{P2Cap, P2Digest},
trace::{PermSeqId, Poseidon2Requires},
},
};
#[derive(Debug, Clone)]
pub struct KeccakNodeInvocation {
pub len_bytes: u32,
pub d: [u32; 8],
pub h_input_chunks: [Felt; 4],
pub chunk_seq_id_head: ChunkSeqId,
pub perm_seq_id_chunks: PermSeqId,
pub perm_seq_id_digest_chunks: PermSeqId,
pub perm_seq_id_keccak: PermSeqId,
pub sponge_seq_id_head: SpongeSeqId,
pub out_mult: ProvideMult,
}
impl KeccakNodeInvocation {
pub fn n_sponge_perms(&self) -> u64 {
u64::from(self.len_bytes) / 136 + 1
}
pub fn n_chunks(&self) -> u64 {
u64::from(self.len_bytes).div_ceil(Node::PACKED_BYTES_PER_CHUNK as u64).max(1)
}
}
pub fn generate_trace(requires: KeccakNodeRequires) -> RowMajorMatrix<Felt> {
let active_rows = requires.total_rows() as usize;
let height = active_rows.next_power_of_two().max(2);
let mut trace = Vec::with_capacity(height * NUM_MAIN_COLS);
for rec in &requires.records {
push_row(&mut trace, &rec.invocation);
}
trace.resize(height * NUM_MAIN_COLS, Felt::ZERO);
RowMajorMatrix::new(trace, NUM_MAIN_COLS)
}
pub fn generate_trace_from_invocations(
invocations: &[KeccakNodeInvocation],
) -> RowMajorMatrix<Felt> {
let height = invocations.len().next_power_of_two().max(2);
let mut trace = Vec::with_capacity(height * NUM_MAIN_COLS);
for inv in invocations {
push_row(&mut trace, inv);
}
trace.resize(height * NUM_MAIN_COLS, Felt::ZERO);
RowMajorMatrix::new(trace, NUM_MAIN_COLS)
}
fn push_row(trace: &mut Vec<Felt>, inv: &KeccakNodeInvocation) {
let len_bytes = Felt::from(inv.len_bytes);
let d_felts: [Felt; 8] = inv.d.map(Felt::from);
let h_digest_chunks = Node::chunks(vec![d_felts])
.expect("Keccak digest chunks are non-empty")
.digest()
.into_elements();
let h_keccak = Keccak256Precompile::assert_node(
inv.len_bytes,
Digest::new(inv.h_input_chunks),
Digest::new(h_digest_chunks),
)
.digest()
.into_elements();
trace.extend([
Felt::ONE,
Felt::from(inv.sponge_seq_id_head.seq()),
Felt::new(inv.n_sponge_perms()).expect("n_sponge_perms fits"),
Felt::from(inv.chunk_seq_id_head.seq()),
Felt::new(inv.n_chunks()).expect("n_chunks fits"),
Felt::from(inv.perm_seq_id_chunks.seq()),
len_bytes,
Felt::from(inv.perm_seq_id_digest_chunks.seq()),
Felt::from(inv.perm_seq_id_keccak.seq()),
]);
trace.extend(d_felts);
trace.extend(inv.h_input_chunks);
trace.extend(h_digest_chunks);
trace.extend(h_keccak);
trace.extend([Felt::from(inv.out_mult)]);
}
#[derive(Debug, Clone)]
pub struct KeccakNodeOutput {
pub keccak_digest: KeccakDigest,
pub h_keccak: P2Digest,
pub node_row: u32,
}
#[derive(Debug, Clone)]
struct NodeRecord {
invocation: KeccakNodeInvocation,
h_keccak: P2Digest,
#[allow(dead_code)]
keccak_digest: KeccakDigest,
}
#[derive(Debug, Clone, Default)]
pub struct KeccakNodeRequires {
records: Vec<NodeRecord>,
by_keccak: BTreeMap<KeccakDigest, usize>,
next_row: u32,
}
impl KeccakNodeRequires {
pub fn new() -> Self {
Self::default()
}
pub(crate) fn add_consumers(&mut self, row: u32, consumers: ProvideMult) {
let invocation = &mut self.records[row as usize].invocation;
invocation.out_mult = invocation
.out_mult
.checked_add(consumers)
.expect("too many Keccak claim consumers");
}
pub fn require(
&mut self,
input: &[u8],
sponge_req: &mut SpongeRequires,
chunk_req: &mut ChunkRequires,
round_req: &mut RoundRequires,
bpl_req: &mut BytePairLutRequires,
p2: &mut Poseidon2Requires,
) -> KeccakNodeOutput {
let keccak_digest = keccak_oracle(input);
if let Some(&idx) = self.by_keccak.get(&keccak_digest) {
let rec = &mut self.records[idx];
rec.invocation.out_mult =
rec.invocation.out_mult.checked_add(1).expect("too many Keccak claim consumers");
return KeccakNodeOutput {
keccak_digest,
h_keccak: rec.h_keccak,
node_row: idx as u32,
};
}
let sponge_inv = SpongeInvocation { input: input.to_vec() };
let sponge_out = sponge_req.require(&sponge_inv, chunk_req, round_req, bpl_req, p2);
debug_assert_eq!(sponge_out.keccak_digest, keccak_digest);
let h_input_chunks_digest = sponge_out.chunk_content_digest;
let h_input_chunks: [Felt; NUM_HASH] = h_input_chunks_digest.as_array();
let _ = p2.require_digest(h_input_chunks_digest);
let d_felts = sponge_out.keccak_digest.to_felts();
let d_rate0: [Felt; 4] = d_felts[0..4].try_into().expect("rate0 slice");
let d_rate1: [Felt; 4] = d_felts[4..8].try_into().expect("rate1 slice");
let digest_chunks_out = p2.require_one_shot(P2Cap::chunk(), d_rate0, d_rate1);
let h_digest_chunks = digest_chunks_out.digest;
let _ = p2.require_digest(h_digest_chunks);
let len_bytes = u32::try_from(input.len()).expect("len_bytes fits in u32");
let keccak_out = p2.require_one_shot(
P2Cap::keccak256_assertion(len_bytes),
h_input_chunks,
h_digest_chunks.as_array(),
);
let h_keccak = keccak_out.digest;
let _ = p2.require_digest(h_keccak);
let invocation = KeccakNodeInvocation {
len_bytes,
d: sponge_out.keccak_digest.to_u32s(),
h_input_chunks,
chunk_seq_id_head: sponge_out.chunk_head,
perm_seq_id_chunks: sponge_out.chunk_content_perm_span.head(),
perm_seq_id_digest_chunks: digest_chunks_out.head(),
perm_seq_id_keccak: keccak_out.head(),
sponge_seq_id_head: sponge_out.sponge_head,
out_mult: 1,
};
let node_row = self.next_row;
self.next_row += 1;
let idx = self.records.len();
self.records.push(NodeRecord { invocation, h_keccak, keccak_digest });
self.by_keccak.insert(keccak_digest, idx);
KeccakNodeOutput { keccak_digest, h_keccak, node_row }
}
pub fn total_rows(&self) -> u32 {
self.next_row
}
}
pub(crate) fn build_aux(
main: &RowMajorMatrix<Felt>,
challenges: &[QuadFelt],
) -> (RowMajorMatrix<QuadFelt>, Vec<QuadFelt>) {
build_logup_aux_trace(&KeccakNodeAir, main, challenges)
}