use alloc::vec::Vec;
use miden_air::{
BaseAir, MidenAir,
logup::{BusId, MIDEN_MAX_MESSAGE_WIDTH},
lookup::{Challenges, LookupFractions, LookupMessage, build_lookup_fractions},
};
use miden_core::{field::QuadFelt, utils::RowMajorMatrix};
use miden_utils_testing::rand::rand_array;
use super::{Felt, VmTrace};
pub(super) struct InteractionLog {
pub challenges: Challenges<QuadFelt>,
rows: Vec<Vec<(Felt, QuadFelt)>>,
}
impl InteractionLog {
pub fn new(trace: &VmTrace) -> Self {
let (core_matrix, chip_matrix, poseidon2_matrix) = trace.main_trace().to_air_matrices();
Self::from_air_matrices(&core_matrix, &chip_matrix, &poseidon2_matrix)
}
pub(super) fn from_air_matrices(
core_matrix: &RowMajorMatrix<Felt>,
chip_matrix: &RowMajorMatrix<Felt>,
poseidon2_matrix: &RowMajorMatrix<Felt>,
) -> Self {
let chip_periodic = MidenAir::Chiplets.periodic_columns();
let poseidon2_periodic = MidenAir::Poseidon2Permutation.periodic_columns();
let raw = rand_array::<Felt, 4>();
let alpha = QuadFelt::new([raw[0], raw[1]]);
let beta = QuadFelt::new([raw[2], raw[3]]);
let challenges =
Challenges::<QuadFelt>::new(alpha, beta, MIDEN_MAX_MESSAGE_WIDTH, BusId::COUNT);
let core_fractions = build_lookup_fractions(&MidenAir::Core, core_matrix, &[], &challenges);
let chip_fractions =
build_lookup_fractions(&MidenAir::Chiplets, chip_matrix, &chip_periodic, &challenges);
let poseidon2_fractions = build_lookup_fractions(
&MidenAir::Poseidon2Permutation,
poseidon2_matrix,
&poseidon2_periodic,
&challenges,
);
let rows = merge_rows(vec![
split_rows(&core_fractions),
split_rows(&chip_fractions),
split_rows(&poseidon2_fractions),
]);
Self { challenges, rows }
}
pub fn assert_contains(&self, expected: &Expectations) {
for &entry in &expected.entries {
let (row, mult, denom) = entry;
let want = expected.entries.iter().filter(|&&e| e == entry).count();
let have = self.rows[row].iter().filter(|&&(m, d)| m == mult && d == denom).count();
assert!(
have >= want,
"row {row}: expected at least {want}× (mult={mult:?}, denom={denom:?}), saw {have}.\n\
actual row bag: {:?}",
self.rows[row],
);
}
}
pub fn net_multiplicity<M>(&self, message: &M) -> Felt
where
M: LookupMessage<Felt, QuadFelt>,
{
let denominator = message.encode(&self.challenges);
self.rows
.iter()
.flatten()
.filter_map(|&(multiplicity, encoded)| (encoded == denominator).then_some(multiplicity))
.sum()
}
}
pub(super) struct Expectations<'a> {
challenges: &'a Challenges<QuadFelt>,
entries: Vec<(usize, Felt, QuadFelt)>,
}
impl<'a> Expectations<'a> {
pub fn new(log: &'a InteractionLog) -> Self {
Self {
challenges: &log.challenges,
entries: Vec::new(),
}
}
pub fn add<M>(&mut self, row: usize, msg: &M) -> &mut Self
where
M: LookupMessage<Felt, QuadFelt>,
{
self.push(row, Felt::ONE, msg)
}
pub fn remove<M>(&mut self, row: usize, msg: &M) -> &mut Self
where
M: LookupMessage<Felt, QuadFelt>,
{
self.push(row, -Felt::ONE, msg)
}
pub fn push<M>(&mut self, row: usize, mult: Felt, msg: &M) -> &mut Self
where
M: LookupMessage<Felt, QuadFelt>,
{
let denom = msg.encode(self.challenges);
self.entries.push((row, mult, denom));
self
}
pub fn count_adds(&self) -> usize {
self.entries.iter().filter(|(_, m, _)| *m == Felt::ONE).count()
}
pub fn count_removes(&self) -> usize {
self.entries.iter().filter(|(_, m, _)| *m == -Felt::ONE).count()
}
}
fn split_rows(fractions: &LookupFractions<Felt, QuadFelt>) -> Vec<Vec<(Felt, QuadFelt)>> {
let num_cols = fractions.num_columns();
let counts = fractions.counts();
let flat = fractions.fractions();
let num_rows = counts.len() / num_cols;
let mut rows = Vec::with_capacity(num_rows);
let mut cursor = 0usize;
for per_row in counts.chunks(num_cols) {
let total: usize = per_row.iter().sum();
rows.push(flat[cursor..cursor + total].to_vec());
cursor += total;
}
rows
}
fn merge_rows(airs: Vec<Vec<Vec<(Felt, QuadFelt)>>>) -> Vec<Vec<(Felt, QuadFelt)>> {
let num_rows = airs.iter().map(Vec::len).max().unwrap_or(0);
(0..num_rows)
.map(|row| {
let mut bag = Vec::new();
for air_rows in &airs {
if let Some(row_bag) = air_rows.get(row) {
bag.extend(row_bag.iter().copied());
}
}
bag
})
.collect()
}