use alloc::{vec, vec::Vec};
use crate::{
ec::msm::trace::EcExprPtr,
math::U256,
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 {
assert!((2..=8).contains(&w), "wNAF window w ∈ [2, 8]");
let p1 = session.msm_intro(base); let two_p = session.msm_combine(p1, p1); let n_odds = 1usize << (w - 2);
let mut odds = Vec::with_capacity(n_odds);
odds.push(p1);
let mut cur = p1;
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 digits: Vec<Vec<i8>> = terms.iter().map(|(_, k)| wnaf(*k, w)).collect();
let len = digits.iter().map(Vec::len).max().unwrap_or(0);
let mut acc: Option<EcExprPtr> = None;
for i in (0..len).rev() {
if let Some(a) = acc {
acc = Some(session.msm_combine(a, a));
}
for (table, base_digits) in tables.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 entry = if d > 0 { pos } else { session.msm_neg(pos) };
acc = Some(match acc {
None => entry, Some(a) => session.msm_combine(a, entry),
});
}
}
}
acc.expect("joint_wnaf needs a nonzero scalar")
}
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
}