use alloc::vec::Vec;
use miden_core::{Felt, field::QuadFelt, utils::RowMajorMatrix};
use crate::{
hash::{
chunk::trace::{ChunkRequires, ChunkSeqId, Invocation as ChunkInvocation},
keccak::{
digest::KeccakDigest,
reference::{KECCAK_RC, keccak_f1600, keccak_round},
round::{NUM_ROUNDS, RoundRequires},
sponge::{
CHUNK_BYTES_RANGE, CLEARED_BYTES_RANGE, COL_ACT, COL_B_BEGIN, COL_BYTES_LEFT,
COL_CHUNK_LO, COL_CHUNK_PTR, COL_CLEARED_LO, COL_IS_CHUNK_AVAIL,
COL_IS_FIRST_BLOCK_OF_INVOCATION, COL_IS_ZERO, COL_PADDED_LO, COL_SPONGE_SEQ_ID,
COL_STATE_NEW_LO, COL_STATE_OUT_LO, COL_STATE_PREV_LO, KeccakSpongeAir,
NUM_MAIN_COLS, PADDED_BYTES_RANGE, SPONGE_PERIOD, STATE_NEW_BYTES_RANGE,
STATE_PREV_BYTES_RANGE,
program::{EXTRA_BLOCK_BEGIN, NOP_SLACK_BEGIN},
},
},
},
logup::build_logup_aux_trace,
primitives::byte_pair_lut::{BytePairLutRequires, BytePairOp, require_logic64},
transcript::poseidon2::{
digest::P2Digest,
trace::{PermSpan, Poseidon2Requires},
},
utils::split_u64,
};
const RATE_BYTES: usize = 136;
const RATE_LANES: usize = 17;
const CHUNK_BYTES: usize = 32;
const CHUNK_LANES: usize = CHUNK_BYTES / 8;
const LANE_16: usize = 16;
const PAD_CONST: u64 = 0x8000_0000_0000_0000;
#[derive(Debug, Clone)]
pub struct Invocation {
pub input: Vec<u8>,
}
impl Invocation {
pub fn num_blocks(&self) -> usize {
(self.input.len() + RATE_BYTES) / RATE_BYTES
}
pub fn chunk_lanes(&self) -> usize {
self.input.len().div_ceil(CHUNK_BYTES).max(1) * CHUNK_LANES
}
}
#[derive(Debug, Clone)]
struct InvocationLayout {
num_blocks: usize,
pad_lane_idx: usize,
byte_offset: usize,
chunk_lanes: usize,
}
impl InvocationLayout {
fn of(inv: &Invocation) -> Self {
let num_blocks = inv.num_blocks();
let bytes_in_last_block = inv.input.len() - RATE_BYTES * (num_blocks - 1);
Self {
num_blocks,
pad_lane_idx: bytes_in_last_block / 8,
byte_offset: bytes_in_last_block % 8,
chunk_lanes: inv.chunk_lanes(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct SpongeSeqId(u32);
impl SpongeSeqId {
pub fn seq(self) -> u32 {
self.0
}
#[cfg(test)]
pub(crate) fn forged(seq: u32) -> Self {
Self(seq)
}
}
#[derive(Debug, Clone)]
pub struct SpongeOutput {
pub keccak_digest: KeccakDigest,
pub chunk_content_digest: P2Digest,
pub chunk_content_perm_span: PermSpan,
pub sponge_head: SpongeSeqId,
pub chunk_head: ChunkSeqId,
}
#[derive(Debug, Clone)]
struct BlockSnapshot {
state_at_block_start: [u64; 25],
post_xorin: [u64; 25],
perm_out: [u64; 25],
}
#[derive(Debug, Clone)]
struct SpongeRecord {
input: Vec<u8>,
layout: InvocationLayout,
chunk_head: ChunkSeqId,
sponge_head: SpongeSeqId,
blocks: Vec<BlockSnapshot>,
}
#[derive(Debug, Clone, Default)]
pub struct SpongeRequires {
invocations: Vec<SpongeRecord>,
next_sponge_seq: u32,
}
impl SpongeRequires {
pub fn new() -> Self {
Self::default()
}
pub fn require(
&mut self,
inv: &Invocation,
chunk_req: &mut ChunkRequires,
round_req: &mut RoundRequires,
bpl_req: &mut BytePairLutRequires,
p2: &mut Poseidon2Requires,
) -> SpongeOutput {
let chunk_out = chunk_req.require(&ChunkInvocation { input: inv.input.clone() }, p2);
let (chunk_head, chunk_content_digest, chunk_content_perm_span) =
(chunk_out.chunk_head, chunk_out.digest, chunk_out.perm_span);
let layout = InvocationLayout::of(inv);
let blocks = compute_block_snapshots_driving(inv, &layout, round_req, bpl_req);
let keccak_digest =
KeccakDigest::from_state(&blocks.last().expect("≥1 block per invocation").perm_out);
let sponge_head = SpongeSeqId(self.next_sponge_seq);
self.next_sponge_seq += (layout.num_blocks * SPONGE_PERIOD) as u32;
self.invocations.push(SpongeRecord {
input: inv.input.clone(),
layout,
chunk_head,
sponge_head,
blocks,
});
SpongeOutput {
keccak_digest,
chunk_content_digest,
chunk_content_perm_span,
sponge_head,
chunk_head,
}
}
pub fn total_active_rows(&self) -> u32 {
self.next_sponge_seq
}
}
pub fn keccak_oracle(input: &[u8]) -> KeccakDigest {
let inv = Invocation { input: input.to_vec() };
let layout = InvocationLayout::of(&inv);
let blocks = compute_block_snapshots(&inv, &layout);
KeccakDigest::from_state(
&blocks.last().expect("compute_block_snapshots yields ≥1 block").perm_out,
)
}
fn compute_block_snapshots_driving(
inv: &Invocation,
layout: &InvocationLayout,
round_req: &mut RoundRequires,
bpl_req: &mut BytePairLutRequires,
) -> Vec<BlockSnapshot> {
let mut state = [0u64; 25];
let mut tape = pack_chunk_tape(inv);
(0..layout.num_blocks)
.map(|block_n| {
let is_last_block = block_n + 1 == layout.num_blocks;
let state_at_block_start = state;
for (k, lane) in state.iter_mut().enumerate().take(RATE_LANES) {
let chunk_lane = tape.next().unwrap_or(0);
let is_verbatim = !is_last_block || k < layout.pad_lane_idx;
let is_pad_row = is_last_block && k == layout.pad_lane_idx;
if is_verbatim {
*lane = require_logic64(bpl_req, BytePairOp::Xor, *lane, chunk_lane);
} else if is_pad_row {
let andnot_mask_val = andnot_mask(layout.byte_offset);
let padding_mask_val = padding_mask(layout.byte_offset);
let cleared =
require_logic64(bpl_req, BytePairOp::AndNot, andnot_mask_val, chunk_lane);
let padded =
require_logic64(bpl_req, BytePairOp::Xor, cleared, padding_mask_val);
*lane = require_logic64(bpl_req, BytePairOp::Xor, *lane, padded);
}
}
let post_xorin = state;
if is_last_block {
state[LANE_16] =
require_logic64(bpl_req, BytePairOp::Xor, state[LANE_16], PAD_CONST);
}
for &rc in &KECCAK_RC[..NUM_ROUNDS] {
round_req.require_round(state);
keccak_round(&mut state, rc);
}
let perm_out = state;
BlockSnapshot {
state_at_block_start,
post_xorin,
perm_out,
}
})
.collect()
}
fn compute_block_snapshots(inv: &Invocation, layout: &InvocationLayout) -> Vec<BlockSnapshot> {
let mut state = [0u64; 25];
let mut tape = pack_chunk_tape(inv);
(0..layout.num_blocks)
.map(|block_n| {
let is_last_block = block_n + 1 == layout.num_blocks;
let state_at_block_start = state;
for (k, lane) in state.iter_mut().enumerate().take(RATE_LANES) {
let chunk_lane = tape.next().unwrap_or(0);
let is_verbatim = !is_last_block || k < layout.pad_lane_idx;
let is_pad_row = is_last_block && k == layout.pad_lane_idx;
if is_verbatim {
*lane ^= chunk_lane;
} else if is_pad_row {
let cleared = !andnot_mask(layout.byte_offset) & chunk_lane;
let padded = cleared ^ padding_mask(layout.byte_offset);
*lane ^= padded;
}
}
let post_xorin = state;
if is_last_block {
state[LANE_16] ^= PAD_CONST;
}
let perm_out = keccak_f1600(state);
state = perm_out;
BlockSnapshot {
state_at_block_start,
post_xorin,
perm_out,
}
})
.collect()
}
pub fn generate_trace(requires: SpongeRequires) -> RowMajorMatrix<Felt> {
generate_trace_padded_to(requires, 0)
}
pub(crate) fn generate_trace_padded_to(
requires: SpongeRequires,
min_height: usize,
) -> RowMajorMatrix<Felt> {
let active_rows = requires.total_active_rows() as usize;
let min_height = min_height
.checked_next_power_of_two()
.expect("minimum sponge trace height exceeds the host power-of-two range");
let height = active_rows.next_power_of_two().max(SPONGE_PERIOD).max(min_height);
let mut trace = Vec::with_capacity(height * NUM_MAIN_COLS);
let mut row = 0usize;
let mut chunk_ptr: u64 = 0;
let mut bytes_left = Felt::ZERO;
let eight = Felt::from(8u8);
for record in &requires.invocations {
bytes_left = Felt::new(record.input.len() as u64).expect("input.len() < p");
chunk_ptr = record.chunk_head.ptr() as u64;
let layout = &record.layout;
let mut tape = pack_chunk_tape_from_bytes(&record.input, layout.chunk_lanes);
let mut chunks_consumed_in_inv = 0usize;
for (block_n, block) in record.blocks.iter().enumerate() {
let is_last_block = block_n + 1 == layout.num_blocks;
let chunks_in_block = if is_last_block {
layout.chunk_lanes - chunks_consumed_in_inv
} else {
RATE_LANES
};
let rate_avail = chunks_in_block.min(RATE_LANES);
let overshoot = chunks_in_block - rate_avail;
for slot in 0..SPONGE_PERIOD {
let mut r = [Felt::ZERO; NUM_MAIN_COLS];
r[COL_SPONGE_SEQ_ID] = Felt::new(row as u64).expect("row index fits");
r[COL_ACT] = Felt::ONE;
r[COL_BYTES_LEFT] = bytes_left;
r[COL_IS_FIRST_BLOCK_OF_INVOCATION] =
if block_n == 0 { Felt::ONE } else { Felt::ZERO };
r[COL_CHUNK_PTR] = Felt::new(chunk_ptr).expect("chunk_ptr fits");
let is_rate_slot = slot < RATE_LANES;
let is_zero = is_last_block && slot > layout.pad_lane_idx;
r[COL_IS_ZERO] = Felt::from(is_zero as u8);
let is_extra_slot = (EXTRA_BLOCK_BEGIN..NOP_SLACK_BEGIN).contains(&slot);
let consume = (is_rate_slot && slot < rate_avail)
|| (is_extra_slot && slot - EXTRA_BLOCK_BEGIN < overshoot);
let avail_end = if overshoot > 0 {
EXTRA_BLOCK_BEGIN + overshoot
} else {
rate_avail
};
r[COL_IS_CHUNK_AVAIL] = Felt::from((slot < avail_end) as u8);
if is_last_block {
r[COL_B_BEGIN + layout.byte_offset] = Felt::ONE;
}
let chunk_lane = if consume { tape.next().unwrap_or(0) } else { 0 };
write_u64_with_bytes(&mut r, COL_CHUNK_LO, CHUNK_BYTES_RANGE.start, chunk_lane);
fill_state_lane_row(
&mut r,
slot,
is_last_block,
layout,
chunk_lane,
&block.state_at_block_start,
&block.post_xorin,
&block.perm_out,
);
trace.extend(r);
if consume {
chunk_ptr += 1;
chunks_consumed_in_inv += 1;
}
if is_rate_slot {
bytes_left -= eight;
}
row += 1;
}
}
debug_assert_eq!(
row,
(record.sponge_head.seq() + (layout.num_blocks * SPONGE_PERIOD) as u32) as usize
);
}
while row < height {
let mut r = [Felt::ZERO; NUM_MAIN_COLS];
r[COL_SPONGE_SEQ_ID] = Felt::new(row as u64).expect("row index fits");
r[COL_BYTES_LEFT] = bytes_left;
r[COL_CHUNK_PTR] = Felt::new(chunk_ptr).expect("chunk_ptr fits");
trace.extend(r);
if row % SPONGE_PERIOD < RATE_LANES {
bytes_left -= eight;
}
row += 1;
}
debug_assert_eq!(trace.len(), height * NUM_MAIN_COLS);
RowMajorMatrix::new(trace, NUM_MAIN_COLS)
}
fn pack_chunk_tape(inv: &Invocation) -> impl Iterator<Item = u64> + '_ {
pack_chunk_tape_from_bytes(&inv.input, inv.chunk_lanes())
}
fn pack_chunk_tape_from_bytes(input: &[u8], chunk_lanes: usize) -> impl Iterator<Item = u64> + '_ {
input
.chunks(8)
.map(|c| {
let mut buf = [0u8; 8];
buf[..c.len()].copy_from_slice(c);
u64::from_le_bytes(buf)
})
.chain(core::iter::repeat(0u64))
.take(chunk_lanes)
}
#[allow(clippy::too_many_arguments)]
fn fill_state_lane_row(
r: &mut [Felt],
slot: usize,
is_last_block: bool,
layout: &InvocationLayout,
chunk_lane: u64,
state_at_block_start: &[u64; 25],
post_xorin_this_block: &[u64; 25],
perm_out_last_block: &[u64; 25],
) {
let is_rate_slot = slot < RATE_LANES;
let is_capacity_slot = (RATE_LANES..RATE_LANES + 8).contains(&slot);
let is_lane16_0x80 = slot == 25;
if is_rate_slot {
let state_prev = state_at_block_start[slot];
write_u64_with_bytes(r, COL_STATE_PREV_LO, STATE_PREV_BYTES_RANGE.start, state_prev);
let (state_new, cleared, padded) = if is_last_block && slot == layout.pad_lane_idx {
let cleared = !andnot_mask(layout.byte_offset) & chunk_lane;
let padded = cleared ^ padding_mask(layout.byte_offset);
(state_prev ^ padded, cleared, padded)
} else if is_last_block && slot > layout.pad_lane_idx {
(state_prev, 0, 0)
} else {
(state_prev ^ chunk_lane, 0, 0)
};
write_u64_with_bytes(r, COL_STATE_NEW_LO, STATE_NEW_BYTES_RANGE.start, state_new);
write_u64_with_bytes(r, COL_CLEARED_LO, CLEARED_BYTES_RANGE.start, cleared);
write_u64_with_bytes(r, COL_PADDED_LO, PADDED_BYTES_RANGE.start, padded);
if is_last_block {
write_u64(r, COL_STATE_OUT_LO, perm_out_last_block[slot]);
}
} else if is_capacity_slot {
let state_prev = state_at_block_start[slot];
write_u64_with_bytes(r, COL_STATE_PREV_LO, STATE_PREV_BYTES_RANGE.start, state_prev);
write_u64_with_bytes(r, COL_STATE_NEW_LO, STATE_NEW_BYTES_RANGE.start, state_prev);
if is_last_block {
write_u64(r, COL_STATE_OUT_LO, perm_out_last_block[slot]);
}
} else if is_lane16_0x80 {
let state_prev = post_xorin_this_block[LANE_16];
write_u64_with_bytes(r, COL_STATE_PREV_LO, STATE_PREV_BYTES_RANGE.start, state_prev);
let state_new = if is_last_block {
state_prev ^ PAD_CONST
} else {
state_prev
};
write_u64_with_bytes(r, COL_STATE_NEW_LO, STATE_NEW_BYTES_RANGE.start, state_new);
}
}
fn write_u64(row: &mut [Felt], col_lo: usize, value: u64) {
let [lo, hi] = split_u64(value);
row[col_lo] = lo;
row[col_lo + 1] = hi;
}
fn write_u64_with_bytes(row: &mut [Felt], col_lo: usize, bytes_start: usize, value: u64) {
write_u64(row, col_lo, value);
for (i, b) in value.to_le_bytes().into_iter().enumerate() {
row[bytes_start + i] = Felt::from(b);
}
}
fn andnot_mask(byte_offset: usize) -> u64 {
u64::MAX << (8 * byte_offset)
}
fn padding_mask(byte_offset: usize) -> u64 {
1u64 << (8 * byte_offset)
}
pub(crate) fn build_aux(
main: &RowMajorMatrix<Felt>,
challenges: &[QuadFelt],
) -> (RowMajorMatrix<QuadFelt>, Vec<QuadFelt>) {
build_logup_aux_trace(&KeccakSpongeAir, main, challenges)
}
#[cfg(test)]
mod tests {
use std::vec;
use miden_core::utils::Matrix;
use super::*;
#[test]
fn num_blocks_matches_fips_202_rule() {
assert_eq!(Invocation { input: vec![] }.num_blocks(), 1);
assert_eq!(Invocation { input: vec![0; 7] }.num_blocks(), 1);
assert_eq!(Invocation { input: vec![0; 135] }.num_blocks(), 1);
assert_eq!(Invocation { input: vec![0; 136] }.num_blocks(), 2);
assert_eq!(Invocation { input: vec![0; 200] }.num_blocks(), 2);
assert_eq!(Invocation { input: vec![0; 272] }.num_blocks(), 3);
}
#[test]
fn chunk_lanes_round_up_to_32_byte_granularity() {
assert_eq!(Invocation { input: vec![] }.chunk_lanes(), 4);
assert_eq!(Invocation { input: vec![0] }.chunk_lanes(), 4);
assert_eq!(Invocation { input: vec![0; 32] }.chunk_lanes(), 4);
assert_eq!(Invocation { input: vec![0; 33] }.chunk_lanes(), 8);
assert_eq!(Invocation { input: vec![0; 200] }.chunk_lanes(), 28);
}
#[test]
fn padded_height_rounds_the_floor_to_a_power_of_two() {
let trace = generate_trace_padded_to(SpongeRequires::new(), 33);
assert_eq!(trace.height(), 64);
}
}