use alloc::vec::Vec;
use core::borrow::BorrowMut;
use miden_air::{
AceCols, QuadFeltExpr,
trace::{RowIndex, chiplets::ace::ACE_CHIPLET_NUM_COLS},
};
use miden_core::{
Felt, Word,
field::{BasedVectorSpace, QuadFelt},
serde::{ByteReader, ByteWriter, Deserializable, DeserializationError, Serializable},
};
use super::{
MAX_NUM_ACE_WIRES,
instruction::{Op, decode_instruction},
};
use crate::{ContextId, errors::AceError};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct ReadNode {
ptr: Felt,
id_0: Felt,
v_0: QuadFelt,
id_1: Felt,
v_1: QuadFelt,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct EvalNode {
ptr: Felt,
eval_op: Felt,
id_0: Felt,
v_0: QuadFelt,
id_1: Felt,
v_1: QuadFelt,
id_2: Felt,
v_2: QuadFelt,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CircuitEvaluation {
ctx: ContextId,
clk: RowIndex,
wire_bus: WireBus,
read_nodes: Vec<ReadNode>,
eval_nodes: Vec<EvalNode>,
}
impl CircuitEvaluation {
pub fn new(ctx: ContextId, clk: RowIndex, num_read_rows: u32, num_eval_rows: u32) -> Self {
let num_wires = 2 * (num_read_rows as u64) + (num_eval_rows as u64);
assert!(num_wires <= MAX_NUM_ACE_WIRES as u64, "too many wires");
Self {
ctx,
clk,
wire_bus: WireBus::new(num_wires as u32),
read_nodes: Vec::with_capacity(num_read_rows as usize),
eval_nodes: Vec::with_capacity(num_eval_rows as usize),
}
}
pub fn num_rows(&self) -> usize {
self.read_nodes.len() + self.eval_nodes.len()
}
pub fn clk(&self) -> u32 {
self.clk.into()
}
pub fn ctx(&self) -> u32 {
self.ctx.into()
}
pub fn num_read_rows(&self) -> u32 {
self.read_nodes.len() as u32
}
pub fn num_eval_rows(&self) -> u32 {
self.eval_nodes.len() as u32
}
pub fn do_read(&mut self, ptr: Felt, word: Word) {
let v_0 = QuadFelt::from_basis_coefficients_fn(|i: usize| [word[0], word[1]][i]);
let id_0 = self.wire_bus.insert(v_0);
let v_1 = QuadFelt::from_basis_coefficients_fn(|i: usize| [word[2], word[3]][i]);
let id_1 = self.wire_bus.insert(v_1);
self.read_nodes.push(ReadNode { ptr, id_0, v_0, id_1, v_1 });
}
pub fn do_eval(&mut self, ptr: Felt, instruction: Felt) -> Result<(), AceError> {
let (id_l, id_r, op) = decode_instruction(instruction)
.ok_or(AceError("failed to decode instruction".into()))?;
let v_l = self
.wire_bus
.read_value(id_l)
.ok_or(AceError("failed to read from the wiring bus".into()))?;
let id_1 = Felt::from_u32(id_l);
let v_r = self
.wire_bus
.read_value(id_r)
.ok_or(AceError("failed to read from the wiring bus".into()))?;
let id_2 = Felt::from_u32(id_r);
let v_0 = match op {
Op::Sub => v_l - v_r,
Op::Mul => v_l * v_r,
Op::Add => v_l + v_r,
};
let id_0 = self.wire_bus.insert(v_0);
let eval_op = match op {
Op::Sub => -Felt::ONE,
Op::Mul => Felt::ZERO,
Op::Add => Felt::ONE,
};
self.eval_nodes.push(EvalNode {
ptr,
eval_op,
id_0,
v_0,
id_1,
v_1: v_l,
id_2,
v_2: v_r,
});
Ok(())
}
pub fn fill(&self, offset: usize, out: &mut [Felt]) {
const W: usize = ACE_CHIPLET_NUM_COLS;
let (out_rows, _) = out.as_chunks_mut::<W>();
let num_read_rows = self.read_nodes.len();
let num_eval_rows = self.eval_nodes.len();
let ctx_felt: Felt = self.ctx.into();
let clk_felt: Felt = self.clk.into();
let eval_section_first_idx = Felt::from_u32(num_eval_rows as u32 - 1);
let mut multiplicities_iter = self.wire_bus.wires.iter().map(|(_v, m)| Felt::from_u32(*m));
for (i, node) in self.read_nodes.iter().enumerate() {
let cols: &mut AceCols<Felt> = out_rows[offset + i].as_mut_slice().borrow_mut();
cols.s_start = if i == 0 { Felt::ONE } else { Felt::ZERO };
cols.s_block = Felt::ZERO;
cols.ctx = ctx_felt;
cols.clk = clk_felt;
cols.ptr = node.ptr;
cols.id_0 = node.id_0;
cols.v_0 = quad_to_expr(node.v_0);
cols.id_1 = node.id_1;
cols.v_1 = quad_to_expr(node.v_1);
let m_0 = multiplicities_iter
.next()
.expect("the m0 multiplicities were not constructed properly");
let m_1 = multiplicities_iter
.next()
.expect("the m1 multiplicities were not constructed properly");
let read = cols.read_mut();
read.num_eval = eval_section_first_idx;
read.m_0 = m_0;
read.m_1 = m_1;
}
for (i, node) in self.eval_nodes.iter().enumerate() {
let cols: &mut AceCols<Felt> =
out_rows[offset + num_read_rows + i].as_mut_slice().borrow_mut();
cols.s_start = Felt::ZERO;
cols.s_block = Felt::ONE;
cols.ctx = ctx_felt;
cols.clk = clk_felt;
cols.ptr = node.ptr;
cols.eval_op = node.eval_op;
cols.id_0 = node.id_0;
cols.v_0 = quad_to_expr(node.v_0);
cols.id_1 = node.id_1;
cols.v_1 = quad_to_expr(node.v_1);
let m_0 = multiplicities_iter
.next()
.expect("the m0 multiplicities were not constructed properly");
let eval = cols.eval_mut();
eval.id_2 = node.id_2;
eval.v_2 = quad_to_expr(node.v_2);
eval.m_0 = m_0;
}
let next = multiplicities_iter.next();
debug_assert!(next.is_none());
}
pub fn output_value(&self) -> Option<QuadFelt> {
if !self.wire_bus.is_finalized() {
return None;
}
self.wire_bus.wires.last().map(|(v, _m)| *v)
}
}
fn quad_to_expr(v: QuadFelt) -> QuadFeltExpr<Felt> {
let c = v.as_basis_coefficients_slice();
QuadFeltExpr(c[0], c[1])
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct WireBus {
id_next: Felt,
wires: Vec<(QuadFelt, u32)>,
num_wires: u32,
}
impl WireBus {
fn new(num_wires: u32) -> Self {
Self {
wires: Vec::with_capacity(num_wires as usize),
num_wires,
id_next: Felt::from_u32(num_wires - 1),
}
}
fn insert(&mut self, value: QuadFelt) -> Felt {
debug_assert!(!self.is_finalized());
self.wires.push((value, 0));
let id = self.id_next;
self.id_next -= Felt::ONE;
id
}
fn read_value(&mut self, id: u32) -> Option<QuadFelt> {
let (v, m) = self
.num_wires
.checked_sub(id + 1)
.and_then(|id| self.wires.get_mut(id as usize))?;
*m += 1;
Some(*v)
}
fn is_finalized(&self) -> bool {
self.wires.len() == self.num_wires as usize
}
}
impl Serializable for ReadNode {
fn write_into<W: ByteWriter>(&self, target: &mut W) {
self.ptr.write_into(target);
self.id_0.write_into(target);
self.v_0.write_into(target);
self.id_1.write_into(target);
self.v_1.write_into(target);
}
}
impl Deserializable for ReadNode {
fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
Ok(Self {
ptr: Felt::read_from(source)?,
id_0: Felt::read_from(source)?,
v_0: QuadFelt::read_from(source)?,
id_1: Felt::read_from(source)?,
v_1: QuadFelt::read_from(source)?,
})
}
fn min_serialized_size() -> usize {
Felt::min_serialized_size() * 3 + QuadFelt::min_serialized_size() * 2
}
}
impl Serializable for EvalNode {
fn write_into<W: ByteWriter>(&self, target: &mut W) {
self.ptr.write_into(target);
self.eval_op.write_into(target);
self.id_0.write_into(target);
self.v_0.write_into(target);
self.id_1.write_into(target);
self.v_1.write_into(target);
self.id_2.write_into(target);
self.v_2.write_into(target);
}
}
impl Deserializable for EvalNode {
fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
Ok(Self {
ptr: Felt::read_from(source)?,
eval_op: Felt::read_from(source)?,
id_0: Felt::read_from(source)?,
v_0: QuadFelt::read_from(source)?,
id_1: Felt::read_from(source)?,
v_1: QuadFelt::read_from(source)?,
id_2: Felt::read_from(source)?,
v_2: QuadFelt::read_from(source)?,
})
}
fn min_serialized_size() -> usize {
Felt::min_serialized_size() * 5 + QuadFelt::min_serialized_size() * 3
}
}
impl Serializable for WireBus {
fn write_into<W: ByteWriter>(&self, target: &mut W) {
self.id_next.write_into(target);
self.wires.write_into(target);
self.num_wires.write_into(target);
}
}
impl Deserializable for WireBus {
fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
let id_next = Felt::read_from(source)?;
let wires = Vec::<(QuadFelt, u32)>::read_from(source)?;
let wire_count = wires.len();
let num_wires = u32::read_from(source)?;
if num_wires == 0 {
return Err(DeserializationError::InvalidValue(
"ACE wire bus must contain at least one wire".into(),
));
}
if num_wires > MAX_NUM_ACE_WIRES {
return Err(DeserializationError::InvalidValue(format!(
"ACE declared wire count {num_wires} exceeds maximum {MAX_NUM_ACE_WIRES}"
)));
}
if wire_count != num_wires as usize {
return Err(DeserializationError::InvalidValue(format!(
"ACE wire count {wire_count} does not match declared wire count {num_wires}"
)));
}
Ok(Self { id_next, wires, num_wires })
}
fn min_serialized_size() -> usize {
Felt::min_serialized_size() + Vec::<u8>::min_serialized_size() + u32::min_serialized_size()
}
}
impl Serializable for CircuitEvaluation {
fn write_into<W: ByteWriter>(&self, target: &mut W) {
self.ctx.write_into(target);
self.clk.write_into(target);
self.wire_bus.write_into(target);
self.read_nodes.write_into(target);
self.eval_nodes.write_into(target);
}
}
impl Deserializable for CircuitEvaluation {
fn read_from<R: ByteReader>(source: &mut R) -> Result<Self, DeserializationError> {
let evaluation = Self {
ctx: ContextId::read_from(source)?,
clk: RowIndex::read_from(source)?,
wire_bus: WireBus::read_from(source)?,
read_nodes: Vec::<ReadNode>::read_from(source)?,
eval_nodes: Vec::<EvalNode>::read_from(source)?,
};
evaluation.validate_wire_count()?;
Ok(evaluation)
}
fn min_serialized_size() -> usize {
ContextId::min_serialized_size()
+ RowIndex::min_serialized_size()
+ WireBus::min_serialized_size()
+ Vec::<ReadNode>::min_serialized_size()
+ Vec::<EvalNode>::min_serialized_size()
}
}
impl CircuitEvaluation {
fn validate_wire_count(&self) -> Result<(), DeserializationError> {
if self.eval_nodes.is_empty() {
return Err(DeserializationError::InvalidValue(
"ACE circuit evaluation must contain at least one eval node".into(),
));
}
let read_wires = self.read_nodes.len().checked_mul(2).ok_or_else(|| {
DeserializationError::InvalidValue("ACE read-node wire count overflow".into())
})?;
let expected_wires = read_wires.checked_add(self.eval_nodes.len()).ok_or_else(|| {
DeserializationError::InvalidValue("ACE total wire count overflow".into())
})?;
if expected_wires == 0 {
return Err(DeserializationError::InvalidValue(
"ACE circuit evaluation must contain at least one wire".into(),
));
}
if expected_wires > MAX_NUM_ACE_WIRES as usize {
return Err(DeserializationError::InvalidValue(format!(
"ACE circuit evaluation wire count {expected_wires} exceeds maximum {MAX_NUM_ACE_WIRES}"
)));
}
if self.wire_bus.num_wires as usize != expected_wires {
return Err(DeserializationError::InvalidValue(format!(
"ACE wire bus count {} does not match read/eval node wire count {expected_wires}",
self.wire_bus.num_wires
)));
}
Ok(())
}
}
#[cfg(test)]
mod serialization_tests {
use alloc::vec;
use super::*;
fn sample_read_node(value: QuadFelt) -> ReadNode {
ReadNode {
ptr: Felt::ZERO,
id_0: Felt::ZERO,
v_0: value,
id_1: Felt::ONE,
v_1: value,
}
}
fn sample_eval_node(value: QuadFelt) -> EvalNode {
EvalNode {
ptr: Felt::ZERO,
eval_op: Felt::ZERO,
id_0: Felt::ZERO,
v_0: value,
id_1: Felt::ZERO,
v_1: value,
id_2: Felt::ONE,
v_2: value,
}
}
#[test]
fn circuit_evaluation_read_rejects_mismatched_wire_bus_count() {
let value = QuadFelt::new([Felt::ONE, Felt::ZERO]);
let evaluation = CircuitEvaluation {
ctx: ContextId::from(0),
clk: RowIndex::from(0_u32),
wire_bus: WireBus {
id_next: Felt::ZERO,
wires: vec![(value, 0), (value, 0)],
num_wires: 2,
},
read_nodes: vec![sample_read_node(value)],
eval_nodes: vec![sample_eval_node(value)],
};
let err = CircuitEvaluation::read_from_bytes(&evaluation.to_bytes()).unwrap_err();
let DeserializationError::InvalidValue(message) = err else {
panic!("expected invalid ACE wire count error");
};
assert!(message.contains("does not match read/eval node wire count"));
}
#[test]
fn circuit_evaluation_read_rejects_empty_eval_section() {
let value = QuadFelt::new([Felt::ONE, Felt::ZERO]);
let evaluation = CircuitEvaluation {
ctx: ContextId::from(0),
clk: RowIndex::from(0_u32),
wire_bus: WireBus {
id_next: Felt::ZERO,
wires: vec![(value, 0), (value, 0)],
num_wires: 2,
},
read_nodes: vec![sample_read_node(value)],
eval_nodes: Vec::new(),
};
let err = CircuitEvaluation::read_from_bytes(&evaluation.to_bytes()).unwrap_err();
let DeserializationError::InvalidValue(message) = err else {
panic!("expected invalid ACE eval section error");
};
assert!(message.contains("at least one eval node"));
}
}