use alloc::vec::Vec;
use crate::{
ec::{
EcStores,
msm::trace::{CombineRow, EcExprPtr, EcMsmRequires, NegRow},
trace::EcPointPtr,
},
math::from_hex,
uint::{UintRequire, UintStores, trace::UintPtr},
};
pub fn intro(
msm: &mut EcMsmRequires,
ec: &mut EcStores,
uint: &mut UintStores,
base: EcPointPtr,
) -> EcExprPtr {
if let Some(e) = msm.lookup_intro(base) {
return e; }
let group = ec.store.point_params(base).0;
let sbound = ec.store.group_sbound(group);
let one = uint.require().intern(from_hex("1"), sbound);
msm.intro(group, sbound, base, one)
}
pub fn intro_zero(
msm: &mut EcMsmRequires,
ec: &mut EcStores,
uint: &mut UintStores,
base: EcPointPtr,
) -> EcExprPtr {
if let Some(e) = msm.lookup_intro_zero(base) {
return e; }
let group = ec.store.point_params(base).0;
let (a_ptr, b_ptr, bound_ptr) = ec.store.group_params(group);
let (beta_ptr, lambda_ptr) = ec.store.group_glv_params(group);
let sbound = ec.store.group_sbound(group);
let zero = uint.require().intern(from_hex("0"), sbound);
let (base_x, base_y) =
ec.store.point_params(base).1.expect("intro_zero of the point at infinity");
let val = ec.store.group_pai(group);
ec.store.require_ecpoint(base);
ec.store.require_ecpoint(val);
ec.store.require_ecgroup(group);
msm.intro_zero(
group, sbound, a_ptr, b_ptr, bound_ptr, beta_ptr, lambda_ptr, base, base_x, base_y, zero,
val,
)
}
pub fn intro_endo(
msm: &mut EcMsmRequires,
ec: &mut EcStores,
uint: &mut UintStores,
base: EcPointPtr,
) -> EcExprPtr {
if let Some(e) = msm.lookup_intro_endo(base) {
return e; }
let group = ec.store.point_params(base).0;
let (beta_ptr, lambda_ptr) = ec.store.group_glv_params(group);
assert_ne!(beta_ptr.addr(), 0, "intro_endo requires a group with a GLV endomorphism");
let (a_ptr, b_ptr, bound_ptr) = ec.store.group_params(group);
let sbound = ec.store.group_sbound(group);
let (px, py) = ec.store.point_params(base).1.expect("intro_endo of the point at infinity");
let phi_x = uint.require().mac(1, beta_ptr, px, 0, bound_ptr);
let (val, minted) = ec.store.add_point_cert(group, phi_x, py);
ec.store.require_ecpoint(base);
ec.store.require_ecpoint(val);
ec.store.require_ecgroup(group);
msm.intro_endo(
group, sbound, a_ptr, b_ptr, bound_ptr, beta_ptr, lambda_ptr, base, val, px, py, phi_x,
minted,
)
}
pub fn combine(
msm: &mut EcMsmRequires,
ec: &mut EcStores,
uint: &mut UintStores,
a: EcExprPtr,
b: EcExprPtr,
) -> EcExprPtr {
if let Some(e) = msm.lookup_combine(a, b) {
return e; }
let group = msm.group(a);
let sbound = msm.sbound(a);
let a_terms = msm.terms(a);
let b_terms = msm.terms(b);
let val_a = msm.value(a);
let val_b = msm.value(b);
let (a_ptr, b_ptr, bound_ptr) = ec.store.group_params(group);
let (beta_ptr, lambda_ptr) = ec.store.group_glv_params(group);
let rows = merge_terms(&a_terms, &b_terms, &mut uint.require());
let val = ec.require(uint.require()).add(val_a, val_b, 1);
ec.store.require_ecgroup(group);
let c = msm.combine(
group, sbound, a_ptr, b_ptr, bound_ptr, beta_ptr, lambda_ptr, a, b, val_a, val_b, val, rows,
);
msm.consume_op(a, 1);
msm.consume_op(b, 1);
c
}
pub fn combine_terms_preserving(
msm: &mut EcMsmRequires,
ec: &mut EcStores,
uint: &mut UintStores,
a: EcExprPtr,
b: EcExprPtr,
) -> EcExprPtr {
if let Some(e) = msm.lookup_concat_combine(a, b) {
return e; }
let group = msm.group(a);
let sbound = msm.sbound(a);
let a_terms = msm.terms(a);
let b_terms = msm.terms(b);
let val_a = msm.value(a);
let val_b = msm.value(b);
let (a_ptr, b_ptr, bound_ptr) = ec.store.group_params(group);
let (beta_ptr, lambda_ptr) = ec.store.group_glv_params(group);
let rows = concat_terms(&a_terms, &b_terms);
let val = ec.require(uint.require()).add(val_a, val_b, 1);
ec.store.require_ecgroup(group);
let c = msm.combine_concat(
group, sbound, a_ptr, b_ptr, bound_ptr, beta_ptr, lambda_ptr, a, b, val_a, val_b, val, rows,
);
msm.consume_op(a, 1);
msm.consume_op(b, 1);
c
}
pub fn neg(
msm: &mut EcMsmRequires,
ec: &mut EcStores,
uint: &mut UintStores,
a: EcExprPtr,
) -> EcExprPtr {
if let Some(e) = msm.lookup_neg(a) {
return e; }
let group = msm.group(a);
let sbound = msm.sbound(a);
let a_terms = msm.terms(a);
let val_a = msm.value(a);
let (a_ptr, b_ptr, bound_ptr) = ec.store.group_params(group);
let (beta_ptr, lambda_ptr) = ec.store.group_glv_params(group);
let mut rows = Vec::with_capacity(a_terms.len());
for (i, (base, s)) in a_terms.iter().enumerate() {
let out_scalar = uint.require().neg(*s);
rows.push(NegRow {
i: i as u32,
base: *base,
s_a: *s,
out_scalar,
});
}
let (px, py) = ec.store.point_params(val_a).1.expect("neg of the point at infinity");
let neg_py = uint.require().neg(py);
let (val, minted) = ec.store.add_point_cert(group, px, neg_py);
ec.store.require_ecpoint(val_a); ec.store.require_ecpoint(val); ec.store.require_ecgroup(group);
let c = msm.neg(
group, sbound, a_ptr, b_ptr, bound_ptr, beta_ptr, lambda_ptr, a, val_a, val, px, py,
neg_py, minted, rows,
);
msm.consume_op(a, 1);
c
}
pub fn merge_terms(
a_terms: &[(EcPointPtr, UintPtr)],
b_terms: &[(EcPointPtr, UintPtr)],
uint: &mut UintRequire<'_>,
) -> Vec<CombineRow> {
let zero_pt = EcPointPtr::from_addr(0);
let zero_u = UintPtr::from_addr(0);
let (mut i, mut j) = (0usize, 0usize);
let mut rows = Vec::new();
while i < a_terms.len() || j < b_terms.len() {
let (ci, cj) = (i as u32, j as u32);
let a_first =
j >= b_terms.len() || (i < a_terms.len() && a_terms[i].0.addr() < b_terms[j].0.addr());
let b_first =
i >= a_terms.len() || (j < b_terms.len() && b_terms[j].0.addr() < a_terms[i].0.addr());
if a_first {
let (base, s) = a_terms[i];
rows.push(CombineRow {
take_a: true,
take_b: false,
take_both: false,
i: ci,
j: cj,
base_a: base,
s_a: s,
base_b: zero_pt,
s_b: zero_u,
out_base: base,
out_scalar: s,
});
i += 1;
} else if b_first {
let (base, s) = b_terms[j];
rows.push(CombineRow {
take_a: false,
take_b: true,
take_both: false,
i: ci,
j: cj,
base_a: zero_pt,
s_a: zero_u,
base_b: base,
s_b: s,
out_base: base,
out_scalar: s,
});
j += 1;
} else {
let (base, sa) = a_terms[i];
let (_, sb) = b_terms[j];
let s_out = uint.add(sa, sb);
rows.push(CombineRow {
take_a: false,
take_b: false,
take_both: true,
i: ci,
j: cj,
base_a: base,
s_a: sa,
base_b: base,
s_b: sb,
out_base: base,
out_scalar: s_out,
});
i += 1;
j += 1;
}
}
rows
}
pub fn concat_terms(
a_terms: &[(EcPointPtr, UintPtr)],
b_terms: &[(EcPointPtr, UintPtr)],
) -> Vec<CombineRow> {
let zero_pt = EcPointPtr::from_addr(0);
let zero_u = UintPtr::from_addr(0);
let mut rows = Vec::with_capacity(a_terms.len() + b_terms.len());
for (i, &(base, s)) in a_terms.iter().enumerate() {
rows.push(CombineRow {
take_a: true,
take_b: false,
take_both: false,
i: i as u32,
j: 0,
base_a: base,
s_a: s,
base_b: zero_pt,
s_b: zero_u,
out_base: base,
out_scalar: s,
});
}
for (j, &(base, s)) in b_terms.iter().enumerate() {
rows.push(CombineRow {
take_a: false,
take_b: true,
take_both: false,
i: a_terms.len() as u32,
j: j as u32,
base_a: zero_pt,
s_a: zero_u,
base_b: base,
s_b: s,
out_base: base,
out_scalar: s,
});
}
rows
}