mod public_inputs;
use public_inputs::{add_public_inputs_struct, public_input_type_to_string};
mod periodic_columns;
use periodic_columns::add_fn_get_periodic_column_values;
mod graph;
use graph::Codegen;
mod boundary_constraints;
use boundary_constraints::{add_fn_get_assertions, add_fn_get_aux_assertions};
mod transition_constraints;
use transition_constraints::{add_fn_evaluate_aux_transition, add_fn_evaluate_transition};
use air_ir::{Air, BusBoundary, BusType, ConstraintDomain, Identifier, TraceSegmentId};
use super::{Impl, Scope};
#[derive(Debug, Clone, Copy)]
pub enum ElemType {
Base,
Ext,
}
pub(super) fn add_air(scope: &mut Scope, ir: &Air) {
add_public_inputs_struct(scope, ir);
let name = ir.name();
add_air_struct(scope, ir, name);
add_air_trait(scope, ir, name);
}
fn add_air_struct(scope: &mut Scope, ir: &Air, name: &str) {
let air_struct = scope
.new_struct(name)
.vis("pub")
.field("context", "AirContext<Felt>");
for public_input in ir.public_inputs() {
air_struct.field(
public_input.name().as_str(),
public_input_type_to_string(public_input),
);
}
let base_impl = scope.new_impl(name);
base_impl
.new_fn("last_step")
.arg_ref_self()
.vis("pub")
.ret("usize")
.line("self.trace_length() - self.context().num_transition_exemptions()");
let (mut add_bus_multiset_boundary_varlen, mut add_bus_logup_boundary_varlen) = (false, false);
for bus in ir.buses.values() {
let bus_constraints = [&bus.first, &bus.last];
for fl in bus_constraints {
if let BusBoundary::PublicInputTable(_) = fl {
match bus.bus_type {
BusType::Multiset => {
add_bus_multiset_boundary_varlen = true;
}
BusType::Logup => {
add_bus_logup_boundary_varlen = true;
}
}
}
}
}
if add_bus_multiset_boundary_varlen {
impl_bus_multiset_boundary_varlen(base_impl);
}
if add_bus_logup_boundary_varlen {
impl_bus_logup_boundary_varlen(base_impl);
}
}
fn impl_bus_multiset_boundary_varlen(base_impl: &mut Impl) {
base_impl
.new_fn("bus_multiset_boundary_varlen")
.generic("'a")
.generic("const N: usize")
.generic("I: IntoIterator<Item = &'a [Felt; N]> + Clone")
.generic("E: FieldElement<BaseField = Felt>")
.arg("aux_rand_elements", "&AuxRandElements<E>")
.arg("public_inputs", "&I")
.ret("E")
.vis("pub")
.line("let mut bus_p_last: E = E::ONE;")
.line("let rand = aux_rand_elements.rand_elements();")
.line("for row in public_inputs.clone().into_iter() {")
.line(" let mut p_last = rand[0];")
.line(" for (c, p_i) in row.iter().enumerate() {")
.line(" p_last += E::from(*p_i) * rand[c + 1];")
.line(" }")
.line(" bus_p_last *= p_last;")
.line("}")
.line("bus_p_last");
}
fn impl_bus_logup_boundary_varlen(base_impl: &mut Impl) {
base_impl
.new_fn("bus_logup_boundary_varlen")
.generic("'a")
.generic("const N: usize")
.generic("I: IntoIterator<Item = &'a [Felt; N]> + Clone")
.generic("E: FieldElement<BaseField = Felt>")
.arg("aux_rand_elements", "&AuxRandElements<E>")
.arg("public_inputs", "&I")
.ret("E")
.vis("pub")
.line("let mut bus_q_last = E::ZERO;")
.line("let rand = aux_rand_elements.rand_elements();")
.line("for row in public_inputs.clone().into_iter() {")
.line(" let mut q_last = rand[0];")
.line(" for (c, p_i) in row.iter().enumerate() {")
.line(" let p_i = *p_i;")
.line(" q_last += E::from(p_i) * rand[c + 1];")
.line(" }")
.line(" bus_q_last += q_last.inv();")
.line("}")
.line("bus_q_last");
}
fn add_air_trait(scope: &mut Scope, ir: &Air, name: &str) {
let air_impl = scope
.new_impl(name)
.impl_trait("Air")
.associate_type("BaseField", "Felt")
.associate_type("PublicInputs", "PublicInputs");
let fn_context = air_impl
.new_fn("context")
.arg_ref_self()
.ret("&AirContext<Felt>");
fn_context.line("&self.context");
add_fn_new(air_impl, ir);
add_fn_get_periodic_column_values(air_impl, ir);
add_fn_get_assertions(air_impl, ir);
add_fn_get_aux_assertions(air_impl, ir);
add_fn_evaluate_transition(air_impl, ir);
add_fn_evaluate_aux_transition(air_impl, ir);
}
fn add_fn_new(impl_ref: &mut Impl, ir: &Air) {
let new = impl_ref
.new_fn("new")
.arg("trace_info", "TraceInfo")
.arg("public_inputs", "PublicInputs")
.arg("options", "WinterProofOptions")
.ret("Self");
add_constraint_degrees(new, ir, 0, "main_degrees");
add_constraint_degrees(new, ir, 1, "aux_degrees");
new.line(format!(
"let num_main_assertions = {};",
ir.num_boundary_constraints(0)
));
new.line(format!(
"let num_aux_assertions = {};",
num_bus_boundary_constraints(ir)
));
let context = "
let context = AirContext::new_multi_segment(
trace_info,
main_degrees,
aux_degrees,
num_main_assertions,
num_aux_assertions,
options,
)
.set_num_transition_exemptions(2);";
new.line(context);
let mut pub_inputs = Vec::new();
for public_input in ir.public_inputs() {
pub_inputs.push(format!("{0}: public_inputs.{0}", public_input.name()));
}
new.line(format!("Self {{ context, {} }}", pub_inputs.join(", ")));
}
fn add_constraint_degrees(
func_body: &mut codegen::Function,
ir: &Air,
trace_segment: TraceSegmentId,
decl_name: &str,
) {
let degrees = ir
.integrity_constraint_degrees(trace_segment)
.iter()
.map(|degree| degree.to_string(ir, ElemType::Ext, trace_segment))
.collect::<Vec<_>>();
func_body.line(format!("let {decl_name} = vec![{}];", degrees.join(", ")));
}
fn call_bus_boundary_varlen_pubinput(
ir: &Air,
bus_name: Identifier,
table_name: Identifier,
) -> String {
let bus = ir.buses.get(&bus_name).expect("bus not found");
match bus.bus_type {
BusType::Multiset => format!(
"Self::bus_multiset_boundary_varlen(aux_rand_elements, &self.{table_name}.iter())",
),
BusType::Logup => {
format!("Self::bus_logup_boundary_varlen(aux_rand_elements, &self.{table_name}.iter())",)
}
}
}
fn num_bus_boundary_constraints(ir: &Air) -> usize {
let mut num_bus_boundary_constraints = 0;
let domains = [ConstraintDomain::FirstRow, ConstraintDomain::LastRow];
for domain in &domains {
for bus in ir.buses.values() {
let bus_boundary = match domain {
ConstraintDomain::FirstRow => &bus.first,
ConstraintDomain::LastRow => &bus.last,
_ => unreachable!("Invalid domain for bus boundary constraint"),
};
match bus_boundary {
air_ir::BusBoundary::PublicInputTable(_) | air_ir::BusBoundary::Null => {
num_bus_boundary_constraints += 1;
}
air_ir::BusBoundary::Unconstrained => {}
}
}
}
num_bus_boundary_constraints
}