use alloc::{vec, vec::Vec};
use miden_precompiles::glv_decompose;
use crate::{
ec::msm::trace::EcExprPtr,
math::{U256, from_limbs32, to_limbs32},
session::{EcNode, Session},
};
pub fn straus(session: &mut Session, terms: &[(EcNode, U256)]) -> EcExprPtr {
let k = terms.len();
assert!(k >= 1, "an MSM needs at least one base");
assert!(k <= 16, "Straus' 2ᵏ table is impractical past k = 16");
let intros: Vec<EcExprPtr> = terms.iter().map(|(p, _)| session.msm_intro(p)).collect();
let scalars: Vec<U256> = terms.iter().map(|(_, s)| *s).collect();
let bits = scalars.iter().map(U256::bit_len).max().unwrap_or(0);
let mut table: Vec<Option<EcExprPtr>> = vec![None; 1usize << k];
for (j, &intro) in intros.iter().enumerate() {
table[1usize << j] = Some(intro);
}
for mask in 1usize..(1usize << k) {
if mask.count_ones() > 1 {
let lsb = mask & mask.wrapping_neg(); let rest = mask ^ lsb;
let combined = session.msm_combine(table[rest].unwrap(), table[lsb].unwrap());
table[mask] = Some(combined);
}
}
let mut acc: Option<EcExprPtr> = None;
for i in (0..bits).rev() {
let mask: usize = (0..k).filter(|&j| scalars[j].bit(i)).map(|j| 1usize << j).sum();
let sel = table[mask];
acc = match acc {
None => sel,
Some(a) => {
let doubled = session.msm_combine(a, a);
Some(match sel {
Some(t) => session.msm_combine(doubled, t),
None => doubled,
})
},
};
}
acc.expect("at least one scalar must be nonzero")
}
pub fn joint_naf(session: &mut Session, terms: &[(EcNode, U256)]) -> EcExprPtr {
assert_eq!(terms.len(), 2, "joint_naf is a 2-base strategy");
let p = session.msm_intro(&terms[0].0);
let q = session.msm_intro(&terms[1].0);
let pq = session.msm_combine(p, q); let nq = session.msm_neg(q); let pmq = session.msm_combine(p, nq); let np = session.msm_neg(p); let npq = session.msm_neg(pq); let npmq = session.msm_neg(pmq); let sel = |d0: i8, d1: i8| -> Option<EcExprPtr> {
match (d0, d1) {
(0, 0) => None,
(1, 0) => Some(p),
(-1, 0) => Some(np),
(0, 1) => Some(q),
(0, -1) => Some(nq),
(1, 1) => Some(pq),
(-1, -1) => Some(npq),
(1, -1) => Some(pmq),
(-1, 1) => Some(npmq),
_ => unreachable!("NAF digits are in {{-1, 0, 1}}"),
}
};
let naf0 = naf(terms[0].1);
let naf1 = naf(terms[1].1);
let len = naf0.len().max(naf1.len());
let digit = |d: &[i8], i: usize| -> i8 { d.get(i).copied().unwrap_or(0) };
let mut acc: Option<EcExprPtr> = None;
for i in (0..len).rev() {
let entry = sel(digit(&naf0, i), digit(&naf1, i));
acc = match acc {
None => entry,
Some(a) => {
let doubled = session.msm_combine(a, a);
Some(match entry {
Some(t) => session.msm_combine(doubled, t),
None => doubled,
})
},
};
}
acc.expect("at least one scalar must be nonzero")
}
pub struct WnafTable {
odds: Vec<EcExprPtr>,
w: usize,
}
pub fn wnaf_table(session: &mut Session, base: &EcNode, w: usize) -> WnafTable {
let p1 = session.msm_intro(base); wnaf_table_from_seed(session, p1, w)
}
pub fn wnaf_table_endo(session: &mut Session, base: &EcNode, w: usize) -> WnafTable {
let e1 = session.msm_intro_endo(base); wnaf_table_from_seed(session, e1, w)
}
fn wnaf_table_from_seed(session: &mut Session, seed: EcExprPtr, w: usize) -> WnafTable {
assert!((2..=8).contains(&w), "wNAF window w ∈ [2, 8]");
let two_p = session.msm_combine(seed, seed); let n_odds = 1usize << (w - 2);
let mut odds = Vec::with_capacity(n_odds);
odds.push(seed);
let mut cur = seed;
for _ in 1..n_odds {
cur = session.msm_combine(cur, two_p); odds.push(cur);
}
WnafTable { odds, w }
}
pub fn wnaf_scalarmul(session: &mut Session, table: &WnafTable, k: U256) -> EcExprPtr {
let digits = wnaf(k, table.w);
let mut acc: Option<EcExprPtr> = None;
for i in (0..digits.len()).rev() {
let d = digits[i];
let entry = (d != 0).then(|| {
let pos = table.odds[(d.unsigned_abs() as usize - 1) / 2];
if d > 0 { pos } else { session.msm_neg(pos) }
});
acc = match acc {
None => entry, Some(a) => {
let doubled = session.msm_combine(a, a);
Some(match entry {
Some(t) => session.msm_combine(doubled, t),
None => doubled,
})
},
};
}
acc.expect("wnaf_scalarmul needs k > 0")
}
pub fn wnaf_msm(session: &mut Session, terms: &[(&WnafTable, U256)]) -> EcExprPtr {
let mut acc: Option<EcExprPtr> = None;
for &(table, k) in terms {
let term = wnaf_scalarmul(session, table, k);
acc = Some(match acc {
None => term,
Some(a) => session.msm_combine(a, term),
});
}
acc.expect("an MSM needs at least one term")
}
pub fn joint_wnaf(session: &mut Session, terms: &[(EcNode, U256)], w: usize) -> EcExprPtr {
let tables: Vec<WnafTable> = terms.iter().map(|(p, _)| wnaf_table(session, p, w)).collect();
let table_terms: Vec<(&WnafTable, U256)> =
tables.iter().zip(terms).map(|(table, &(_, k))| (table, k)).collect();
joint_wnaf_with_tables(session, &table_terms)
}
pub fn joint_wnaf_with_tables(session: &mut Session, terms: &[(&WnafTable, U256)]) -> EcExprPtr {
let signed_terms: Vec<(&WnafTable, U256, bool)> =
terms.iter().map(|&(table, k)| (table, k, false)).collect();
joint_wnaf_with_signed_tables(session, &signed_terms)
}
pub fn joint_wnaf_with_signed_tables(
session: &mut Session,
terms: &[(&WnafTable, U256, bool)],
) -> EcExprPtr {
let digits: Vec<Vec<i8>> = terms.iter().map(|(table, k, _)| wnaf(*k, table.w)).collect();
let len = digits.iter().map(Vec::len).max().unwrap_or(0);
let mut acc: Option<EcExprPtr> = None;
let mut column = Vec::with_capacity(terms.len());
for i in (0..len).rev() {
if let Some(a) = acc {
acc = Some(session.msm_combine(a, a));
}
column.clear();
for ((table, _, negate), base_digits) in terms.iter().zip(&digits) {
let d = base_digits.get(i).copied().unwrap_or(0);
if d != 0 {
let pos = table.odds[(d.unsigned_abs() as usize - 1) / 2];
let want_neg = (d < 0) ^ *negate;
column.push(if want_neg { session.msm_neg(pos) } else { pos });
}
}
while column.len() > 1 {
let len = column.len();
for pair in 0..len / 2 {
column[pair] = session.msm_combine(column[2 * pair], column[2 * pair + 1]);
}
if len % 2 != 0 {
column[len / 2] = column[len - 1];
}
column.truncate(len.div_ceil(2));
}
if let Some(sum) = column.pop() {
acc = Some(match acc {
None => sum,
Some(a) => session.msm_combine(a, sum),
});
}
}
acc.expect("joint_wnaf needs a nonzero scalar")
}
pub fn glv_joint_wnaf_with_tables(
session: &mut Session,
terms: &[(&WnafTable, Option<&WnafTable>, U256)],
) -> EcExprPtr {
assert!(!terms.is_empty(), "an MSM needs at least one base");
let mut signed_terms: Vec<(&WnafTable, U256, bool)> = Vec::with_capacity(terms.len() * 2);
for &(plain, endo, k) in terms {
assert_ne!(k, U256::ZERO, "glv_joint_wnaf_with_tables terms must have a nonzero scalar");
match endo {
Some(endo) => {
let [(a_neg, a_mag), (b_neg, b_mag)] = glv_decompose(to_limbs32(k));
signed_terms.push((plain, from_limbs32(&a_mag), a_neg));
signed_terms.push((endo, from_limbs32(&b_mag), b_neg));
},
None => signed_terms.push((plain, k, false)),
}
}
joint_wnaf_with_signed_tables(session, &signed_terms)
}
fn naf(mut k: U256) -> Vec<i8> {
let one = U256::from(1u64);
let mut out = Vec::new();
while k > U256::ZERO {
if k.bit(0) {
let d: i8 = if k.bit(1) { -1 } else { 1 }; out.push(d);
k = if d == 1 { k - one } else { k + one };
} else {
out.push(0);
}
k >>= 1usize;
}
out
}
fn wnaf(mut k: U256, w: usize) -> Vec<i8> {
let half = 1u64 << (w - 1); let modulus = 1u64 << w; let mut out = Vec::new();
while k > U256::ZERO {
let d: i64 = if k.bit(0) {
let low = (0..w).filter(|&j| k.bit(j)).fold(0u64, |a, j| a | (1u64 << j));
let d = if low >= half {
low as i64 - modulus as i64
} else {
low as i64
};
k = if d >= 0 {
k - U256::from(d as u64)
} else {
k + U256::from((-d) as u64)
};
d
} else {
0
};
out.push(d as i8);
k >>= 1usize;
}
out
}