pub mod require;
pub mod trace;
use alloc::vec::Vec;
use miden_core::{
Felt,
field::{PrimeCharacteristicRing, QuadFelt},
utils::RowMajorMatrix,
};
use miden_lifted_air::{AirBuilder, BaseAir, LiftedAir, LiftedAirBuilder};
use crate::{
ec::{
EcGroupMsg, EcPointMsg,
add::{EcGroupAddMsg, EcOnCurveCertMsg},
},
logup::{
Challenges, CyclicConstraintLookupBuilder, Deg, LookupAir, LookupBatch, LookupBuilder,
LookupColumn, LookupGroup, LookupMessage, NUM_PUBLIC_VALUES, NUM_RANDOMNESS,
NUM_SIGMA_VALUES, frac_col,
},
primitives::byte_pair_lut::Range16Msg,
relations::{BusId, MAX_MESSAGE_WIDTH, NUM_BUS_IDS},
uint::{UintValMsg, add::UintAddMsg},
utils::{current_main, next_main},
};
#[derive(Debug, Clone)]
pub struct MsmTermMsg<E> {
pub expr_ptr: E,
pub idx: E,
pub base_ptr: E,
pub scalar_ptr: E,
}
impl<E, EF> LookupMessage<E, EF> for MsmTermMsg<E>
where
E: miden_core::field::Algebra<E>,
EF: miden_core::field::Algebra<E>,
{
fn encode(&self, challenges: &Challenges<EF>) -> EF {
challenges.encode(
BusId::MsmTerm as usize,
[
self.expr_ptr.clone(),
self.idx.clone(),
self.base_ptr.clone(),
self.scalar_ptr.clone(),
],
)
}
}
#[derive(Debug, Clone)]
pub struct MsmExprMsg<E> {
pub expr_ptr: E,
pub group_ptr: E,
pub val_ptr: E,
pub k: E,
}
impl<E, EF> LookupMessage<E, EF> for MsmExprMsg<E>
where
E: miden_core::field::Algebra<E>,
EF: miden_core::field::Algebra<E>,
{
fn encode(&self, challenges: &Challenges<EF>) -> EF {
challenges.encode(
BusId::MsmExpr as usize,
[
self.expr_ptr.clone(),
self.group_ptr.clone(),
self.val_ptr.clone(),
self.k.clone(),
],
)
}
}
#[derive(Debug, Clone)]
pub struct MsmClaimTermMsg<E> {
pub expr_ptr: E,
pub base_ptr: E,
pub scalar_ptr: E,
}
impl<E, EF> LookupMessage<E, EF> for MsmClaimTermMsg<E>
where
E: miden_core::field::Algebra<E>,
EF: miden_core::field::Algebra<E>,
{
fn encode(&self, challenges: &Challenges<EF>) -> EF {
challenges.encode(
BusId::MsmClaimTerm as usize,
[self.expr_ptr.clone(), self.base_ptr.clone(), self.scalar_ptr.clone()],
)
}
}
pub const COL_ACT: usize = 0;
pub const COL_EXPR_PTR: usize = 1;
pub const COL_IS_BOUNDARY: usize = 2;
pub const COL_GROUP_PTR: usize = 3;
pub const COL_SBOUND_PTR: usize = 4;
pub const COL_IDX: usize = 5;
pub const COL_BASE: usize = 6;
pub const COL_SCALAR: usize = 7;
pub const COL_VAL: usize = 8;
pub const COL_MULT: usize = 9;
pub const COL_IS_INTRO: usize = 10;
pub const COL_IS_COMBINE: usize = 11;
pub const COL_A_EXPR: usize = 12;
pub const COL_B_EXPR: usize = 13;
pub const COL_I: usize = 14;
pub const COL_J: usize = 15;
pub const COL_TAKE_A: usize = 16;
pub const COL_TAKE_B: usize = 17;
pub const COL_TAKE_BOTH: usize = 18;
pub const COL_BASE_A: usize = 19;
pub const COL_S_A: usize = 20;
pub const COL_BASE_B: usize = 21;
pub const COL_S_B: usize = 22;
pub const COL_VAL_A: usize = 23;
pub const COL_VAL_B: usize = 24;
pub const COL_A_PTR: usize = 25;
pub const COL_B_PTR: usize = 26;
pub const COL_BOUND_PTR: usize = 27;
pub const COL_A_DIFF_LO: usize = 28;
pub const COL_A_DIFF_HI: usize = 29;
pub const COL_B_DIFF_LO: usize = 30;
pub const COL_B_DIFF_HI: usize = 31;
pub const COL_IS_NEG: usize = 32;
pub const COL_NEG_X: usize = 33;
pub const COL_CLAIM_MULT: usize = 34;
pub const COL_NEG_YA: usize = 35;
pub const COL_NEG_YR: usize = 36;
pub const COL_NEG_MINTED: usize = 37;
pub const NUM_MAIN_COLS: usize = 38;
const NUM_LOGUP_COLS: usize = 11;
const AUX_WIDTH: usize = 11;
const COLUMN_SHAPE: [usize; NUM_LOGUP_COLS] = [1, 2, 2, 1, 2, 2, 2, 2, 2, 2, 2];
const TWO16: u32 = 1 << 16;
#[derive(Debug, Default, Clone, Copy)]
pub struct EcMsmAir;
impl BaseAir<Felt> for EcMsmAir {
fn width(&self) -> usize {
NUM_MAIN_COLS
}
fn num_public_values(&self) -> usize {
NUM_PUBLIC_VALUES
}
}
impl LiftedAir<Felt, QuadFelt> for EcMsmAir {
fn num_randomness(&self) -> usize {
NUM_RANDOMNESS
}
fn aux_width(&self) -> usize {
AUX_WIDTH
}
fn num_aux_values(&self) -> usize {
NUM_SIGMA_VALUES
}
fn build_aux_trace(
&self,
main: &RowMajorMatrix<Felt>,
_air_inputs: &[Felt],
_aux_inputs: &[Felt],
challenges: &[QuadFelt],
) -> (RowMajorMatrix<QuadFelt>, Vec<QuadFelt>) {
trace::build_aux(main, challenges)
}
fn eval<AB: LiftedAirBuilder<F = Felt>>(&self, builder: &mut AB) {
let local: [AB::Var; NUM_MAIN_COLS] = current_main(builder.main(), 0);
let next: [AB::Var; NUM_MAIN_COLS] = next_main(builder.main(), 0);
let act: AB::Expr = local[COL_ACT].into();
let act_next: AB::Expr = next[COL_ACT].into();
let is_boundary: AB::Expr = local[COL_IS_BOUNDARY].into();
let is_intro: AB::Expr = local[COL_IS_INTRO].into();
let is_combine: AB::Expr = local[COL_IS_COMBINE].into();
let is_neg: AB::Expr = local[COL_IS_NEG].into();
let neg_minted: AB::Expr = local[COL_NEG_MINTED].into();
let idx: AB::Expr = local[COL_IDX].into();
builder.assert_bool(local[COL_ACT]);
builder.assert_bool(local[COL_IS_BOUNDARY]);
builder.assert_bool(local[COL_IS_INTRO]);
builder.assert_bool(local[COL_IS_COMBINE]);
builder.assert_bool(local[COL_IS_NEG]);
builder.assert_bool(local[COL_NEG_MINTED]);
builder.assert_zero((AB::Expr::ONE - is_neg.clone()) * neg_minted);
builder.when_transition().assert_zero((AB::Expr::ONE - act.clone()) * act_next);
builder.assert_zero(is_intro.clone() + is_combine.clone() + is_neg.clone() - act.clone());
builder.assert_zero((AB::Expr::ONE - act.clone()) * is_boundary.clone());
builder.assert_zero((AB::Expr::ONE - act.clone()) * local[COL_MULT].into());
builder.assert_zero((AB::Expr::ONE - act) * local[COL_CLAIM_MULT].into());
let expr_ptr: AB::Expr = local[COL_EXPR_PTR].into();
let expr_ptr_next: AB::Expr = next[COL_EXPR_PTR].into();
builder.when_first_row().assert_zero(expr_ptr - AB::Expr::ONE);
builder
.when_transition()
.assert_zero(expr_ptr_next - local[COL_EXPR_PTR].into() - is_boundary.clone());
let idx_next: AB::Expr = next[COL_IDX].into();
builder.when_first_row().assert_zero(idx.clone());
builder
.when_transition()
.assert_zero(idx_next - (AB::Expr::ONE - is_boundary.clone()) * (idx + AB::Expr::ONE));
let not_boundary = AB::Expr::ONE - is_boundary.clone();
for col in [
COL_GROUP_PTR,
COL_SBOUND_PTR,
COL_VAL,
COL_MULT,
COL_CLAIM_MULT,
COL_IS_INTRO,
COL_IS_COMBINE,
COL_IS_NEG,
COL_A_EXPR,
COL_B_EXPR,
COL_VAL_A,
COL_VAL_B,
COL_A_PTR,
COL_B_PTR,
COL_BOUND_PTR,
] {
let here: AB::Expr = local[col].into();
let there: AB::Expr = next[col].into();
builder.when_transition().assert_zero(not_boundary.clone() * (there - here));
}
builder.assert_zero(is_intro.clone() * (AB::Expr::ONE - is_boundary.clone()));
let base: AB::Expr = local[COL_BASE].into();
let val: AB::Expr = local[COL_VAL].into();
builder.assert_zero(is_intro * (val - base));
let take_a: AB::Expr = local[COL_TAKE_A].into();
let take_b: AB::Expr = local[COL_TAKE_B].into();
let take_both: AB::Expr = local[COL_TAKE_BOTH].into();
builder.assert_bool(local[COL_TAKE_A]);
builder.assert_bool(local[COL_TAKE_B]);
builder.assert_bool(local[COL_TAKE_BOTH]);
builder
.assert_zero(take_a.clone() + take_b.clone() + take_both.clone() - is_combine.clone());
let i_cur: AB::Expr = local[COL_I].into();
let j_cur: AB::Expr = local[COL_J].into();
let i_next: AB::Expr = next[COL_I].into();
let j_next: AB::Expr = next[COL_J].into();
builder.when_first_row().assert_zero(i_cur.clone());
builder.when_first_row().assert_zero(j_cur.clone());
let adv_i = take_a.clone() + take_both.clone() + is_neg.clone();
let adv_j = take_b.clone() + take_both.clone();
builder
.when_transition()
.assert_zero(i_next - (AB::Expr::ONE - is_boundary.clone()) * (i_cur + adv_i));
builder
.when_transition()
.assert_zero(j_next - (AB::Expr::ONE - is_boundary.clone()) * (j_cur + adv_j));
let base_a: AB::Expr = local[COL_BASE_A].into();
let base_b: AB::Expr = local[COL_BASE_B].into();
let s_a: AB::Expr = local[COL_S_A].into();
let s_b: AB::Expr = local[COL_S_B].into();
let out_base: AB::Expr = local[COL_BASE].into();
let out_scalar: AB::Expr = local[COL_SCALAR].into();
builder.assert_zero(
(take_a.clone() + take_both.clone() + is_neg.clone()) * (out_base - base_a.clone())
+ take_b.clone() * (local[COL_BASE].into() - base_b.clone()),
);
builder
.assert_zero(take_a * (out_scalar - s_a) + take_b * (local[COL_SCALAR].into() - s_b));
builder.assert_zero(take_both * (base_a - base_b));
let bnd_a = (is_combine.clone() + is_neg) * is_boundary.clone();
let bnd_b = is_combine * is_boundary;
let two16 = AB::Expr::from(Felt::from(TWO16));
let here_expr: AB::Expr = local[COL_EXPR_PTR].into();
let a_expr: AB::Expr = local[COL_A_EXPR].into();
let b_expr: AB::Expr = local[COL_B_EXPR].into();
let a_lo: AB::Expr = local[COL_A_DIFF_LO].into();
let a_hi: AB::Expr = local[COL_A_DIFF_HI].into();
let b_lo: AB::Expr = local[COL_B_DIFF_LO].into();
let b_hi: AB::Expr = local[COL_B_DIFF_HI].into();
builder.assert_zero(
bnd_a * (here_expr.clone() - a_expr - AB::Expr::ONE - a_lo - two16.clone() * a_hi),
);
builder.assert_zero(bnd_b * (here_expr - b_expr - AB::Expr::ONE - b_lo - two16 * b_hi));
let mut lb =
CyclicConstraintLookupBuilder::new(builder, self, self.preprocessed_width() > 0);
<Self as LookupAir<_>>::eval(self, &mut lb);
}
}
impl<LB> LookupAir<LB> for EcMsmAir
where
LB: LookupBuilder<F = Felt>,
{
fn num_columns(&self) -> usize {
NUM_LOGUP_COLS
}
fn column_shape(&self) -> &[usize] {
&COLUMN_SHAPE
}
fn max_message_width(&self) -> usize {
MAX_MESSAGE_WIDTH
}
fn num_bus_ids(&self) -> usize {
NUM_BUS_IDS
}
fn eval(&self, builder: &mut LB) {
let local: [LB::Var; NUM_MAIN_COLS] = current_main(builder.main(), 0);
let neg_mult: LB::Expr = LB::Expr::ZERO - local[COL_MULT].into();
let neg_claim_mult: LB::Expr = LB::Expr::ZERO - local[COL_CLAIM_MULT].into();
let is_boundary: LB::Expr = local[COL_IS_BOUNDARY].into();
let is_intro: LB::Expr = local[COL_IS_INTRO].into();
let is_combine: LB::Expr = local[COL_IS_COMBINE].into();
let is_neg: LB::Expr = local[COL_IS_NEG].into();
let expr_ptr: LB::Expr = local[COL_EXPR_PTR].into();
let group_ptr: LB::Expr = local[COL_GROUP_PTR].into();
let sbound_ptr: LB::Expr = local[COL_SBOUND_PTR].into();
let idx: LB::Expr = local[COL_IDX].into();
let base: LB::Expr = local[COL_BASE].into();
let scalar: LB::Expr = local[COL_SCALAR].into();
let val: LB::Expr = local[COL_VAL].into();
let a_expr: LB::Expr = local[COL_A_EXPR].into();
let b_expr: LB::Expr = local[COL_B_EXPR].into();
let i_cur: LB::Expr = local[COL_I].into();
let j_cur: LB::Expr = local[COL_J].into();
let take_a: LB::Expr = local[COL_TAKE_A].into();
let take_b: LB::Expr = local[COL_TAKE_B].into();
let take_both: LB::Expr = local[COL_TAKE_BOTH].into();
let base_a: LB::Expr = local[COL_BASE_A].into();
let s_a: LB::Expr = local[COL_S_A].into();
let base_b: LB::Expr = local[COL_BASE_B].into();
let s_b: LB::Expr = local[COL_S_B].into();
let val_a: LB::Expr = local[COL_VAL_A].into();
let val_b: LB::Expr = local[COL_VAL_B].into();
let a_ptr: LB::Expr = local[COL_A_PTR].into();
let b_ptr: LB::Expr = local[COL_B_PTR].into();
let bound_ptr: LB::Expr = local[COL_BOUND_PTR].into();
let a_lo: LB::Expr = local[COL_A_DIFF_LO].into();
let a_hi: LB::Expr = local[COL_A_DIFF_HI].into();
let b_lo: LB::Expr = local[COL_B_DIFF_LO].into();
let b_hi: LB::Expr = local[COL_B_DIFF_HI].into();
let neg_x: LB::Expr = local[COL_NEG_X].into();
let neg_ya: LB::Expr = local[COL_NEG_YA].into();
let neg_yr: LB::Expr = local[COL_NEG_YR].into();
let neg_minted: LB::Expr = local[COL_NEG_MINTED].into();
let adv_i = take_a + take_both.clone() + is_neg.clone();
let adv_j = take_b + take_both.clone();
let bnd_a = (is_combine.clone() + is_neg.clone()) * is_boundary.clone();
let bnd_b = is_combine * is_boundary.clone();
let bnd_neg = is_neg.clone() * is_boundary.clone();
let one_deg = Deg { v: 1, u: 1 };
let two_deg = Deg { v: 2, u: 1 };
let single_deg = Deg { v: 1, u: 2 };
let pair_deg = Deg { v: 3, u: 2 };
frac_col!(
builder,
"ec-msm-provide",
single_deg,
(
"provide-msmterm",
neg_mult.clone(),
MsmTermMsg {
expr_ptr: expr_ptr.clone(),
idx: idx.clone(),
base_ptr: base.clone(),
scalar_ptr: scalar.clone(),
},
one_deg
),
);
frac_col!(
builder,
"ec-msm-provide",
pair_deg,
(
"provide-msmexpr",
(neg_mult.clone() + neg_claim_mult.clone()) * is_boundary.clone(),
MsmExprMsg {
expr_ptr: expr_ptr.clone(),
group_ptr: group_ptr.clone(),
val_ptr: val.clone(),
k: idx.clone() + LB::Expr::ONE,
},
two_deg
),
(
"provide-msmclaimterm",
neg_claim_mult.clone(),
MsmClaimTermMsg {
expr_ptr: expr_ptr.clone(),
base_ptr: base.clone(),
scalar_ptr: scalar.clone(),
},
one_deg
),
);
frac_col!(
builder,
"ec-msm-provide",
pair_deg,
(
"provide-oncurvecert-neg",
LB::Expr::ZERO - neg_minted.clone() * is_boundary.clone(),
EcOnCurveCertMsg {
group_ptr: group_ptr.clone(),
r_ptr: val.clone()
},
two_deg
),
(
"consume-one",
is_intro.clone(),
UintValMsg {
ptr: scalar.clone(),
bound_ptr: sbound_ptr.clone(),
limbs: [
LB::Expr::ONE,
LB::Expr::ZERO,
LB::Expr::ZERO,
LB::Expr::ZERO,
LB::Expr::ZERO,
LB::Expr::ZERO,
LB::Expr::ZERO,
LB::Expr::ZERO,
],
},
one_deg
),
);
frac_col!(
builder,
"ec-msm-walk",
single_deg,
(
"consume-term-a",
adv_i.clone(),
MsmTermMsg {
expr_ptr: a_expr.clone(),
idx: i_cur.clone(),
base_ptr: base_a.clone(),
scalar_ptr: s_a.clone(),
},
one_deg
),
);
frac_col!(
builder,
"ec-msm-walk",
pair_deg,
(
"consume-term-b",
adv_j.clone(),
MsmTermMsg {
expr_ptr: b_expr.clone(),
idx: j_cur.clone(),
base_ptr: base_b.clone(),
scalar_ptr: s_b.clone(),
},
one_deg
),
(
"consume-uintadd",
take_both.clone(),
UintAddMsg {
bound_ptr: sbound_ptr.clone(),
a_ptr: s_a.clone(),
b_ptr: s_b.clone(),
c_ptr: scalar.clone(),
nz: LB::Expr::ZERO,
},
one_deg
),
);
frac_col!(
builder,
"ec-msm-walk",
pair_deg,
(
"consume-uintadd-neg",
is_neg.clone(),
UintAddMsg {
bound_ptr: sbound_ptr.clone(),
a_ptr: s_a.clone(),
b_ptr: scalar.clone(),
c_ptr: LB::Expr::ZERO,
nz: LB::Expr::ZERO,
},
one_deg
),
(
"consume-uintadd-neg-y",
bnd_neg.clone(),
UintAddMsg {
bound_ptr: bound_ptr.clone(),
a_ptr: neg_ya.clone(),
b_ptr: neg_yr.clone(),
c_ptr: LB::Expr::ZERO,
nz: LB::Expr::ZERO,
},
two_deg
),
);
frac_col!(
builder,
"ec-msm-heads",
pair_deg,
(
"consume-head-a",
bnd_a.clone(),
MsmExprMsg {
expr_ptr: a_expr.clone(),
group_ptr: group_ptr.clone(),
val_ptr: val_a.clone(),
k: i_cur.clone() + adv_i.clone(),
},
two_deg
),
(
"consume-head-b",
bnd_b.clone(),
MsmExprMsg {
expr_ptr: b_expr.clone(),
group_ptr: group_ptr.clone(),
val_ptr: val_b.clone(),
k: j_cur.clone() + adv_j.clone(),
},
two_deg
),
);
frac_col!(
builder,
"ec-msm-heads",
pair_deg,
(
"consume-ecgroupadd",
bnd_b.clone(),
EcGroupAddMsg {
group_ptr: group_ptr.clone(),
p_ptr: val_a.clone(),
q_ptr: val_b.clone(),
r_ptr: val.clone(),
},
two_deg
),
(
"consume-ecpoint-val-a",
bnd_neg.clone(),
EcPointMsg {
point_ptr: val_a.clone(),
group_ptr: group_ptr.clone(),
x_ptr: neg_x.clone(),
y_ptr: neg_ya.clone(),
is_pai: LB::Expr::ZERO,
},
two_deg
),
);
frac_col!(
builder,
"ec-msm-heads",
pair_deg,
(
"consume-ecpoint-neg",
bnd_neg.clone(),
EcPointMsg {
point_ptr: val.clone(),
group_ptr: group_ptr.clone(),
x_ptr: neg_x.clone(),
y_ptr: neg_yr.clone(),
is_pai: LB::Expr::ZERO,
},
two_deg
),
(
"consume-ecgroup",
bnd_a.clone(),
EcGroupMsg {
group_ptr: group_ptr.clone(),
a_ptr: a_ptr.clone(),
b_ptr: b_ptr.clone(),
bound_ptr: bound_ptr.clone(),
scalar_bound_ptr: sbound_ptr.clone(),
},
two_deg
),
);
frac_col!(
builder,
"ec-msm-order",
pair_deg,
("range-a-lo", bnd_a.clone(), Range16Msg { w: a_lo }, two_deg),
("range-a-hi", bnd_a, Range16Msg { w: a_hi }, two_deg),
);
frac_col!(
builder,
"ec-msm-order",
pair_deg,
("range-b-lo", bnd_b.clone(), Range16Msg { w: b_lo }, two_deg),
("range-b-hi", bnd_b, Range16Msg { w: b_hi }, two_deg),
);
}
}