1use binius_core::word::Word;
4use binius_frontend::{CircuitBuilder, Wire};
5
6use crate::{
7 bignum::{BigUint, assert_eq, select as select_biguint},
8 multiplexer::multi_wire_multiplex,
9 secp256k1::{
10 N_LIMBS, Secp256k1, Secp256k1Affine, Secp256k1EndosplitHint, coord_lambda, coord_zero,
11 },
12};
13
14pub const MSM_WINDOW: usize = 4;
17
18pub fn scalar_mul(
33 b: &CircuitBuilder,
34 curve: &Secp256k1,
35 scalar: &BigUint,
36 point: Secp256k1Affine,
37) -> Secp256k1Affine {
38 msm_strauss_endo(b, curve, MSM_WINDOW, std::slice::from_ref(scalar), &[point])
39}
40
41fn check_endomorphism_split(
44 b: &CircuitBuilder,
45 curve: &Secp256k1,
46 k1_neg: Wire,
47 k2_neg: Wire,
48 k1_abs: [Wire; 2],
49 k2_abs: [Wire; 2],
50 k: &BigUint,
51) {
52 assert_eq!(k.limbs.len(), N_LIMBS);
53
54 let k1_abs = BigUint {
55 limbs: k1_abs.to_vec(),
56 }
57 .zero_extend(b, N_LIMBS);
58 let k2_abs = BigUint {
59 limbs: k2_abs.to_vec(),
60 }
61 .zero_extend(b, N_LIMBS);
62
63 let f_scalar = curve.f_scalar();
64 let k1 = select_biguint(b, k1_neg, &f_scalar.sub(b, &coord_zero(b), &k1_abs), &k1_abs);
65 let k2 = select_biguint(b, k2_neg, &f_scalar.sub(b, &coord_zero(b), &k2_abs), &k2_abs);
66
67 assert_eq(
68 b,
69 "endomorphism split k1 + λk2 = k (mod n)",
70 k,
71 &f_scalar.add(b, &k1, &f_scalar.mul(b, &k2, &coord_lambda(b))),
72 );
73}
74
75pub fn msm_strauss_endo(
108 b: &CircuitBuilder,
109 curve: &Secp256k1,
110 window: usize,
111 scalars: &[BigUint],
112 points: &[Secp256k1Affine],
113) -> Secp256k1Affine {
114 let n = points.len();
115 assert_eq!(scalars.len(), n, "scalars and points must have the same length");
116 assert!(n >= 1, "MSM requires at least one point");
117 assert!(0 < window && window < Word::BITS, "window must be in 1..Word::BITS");
118
119 let mut tables = Vec::with_capacity(2 * n);
123 let mut subscalars = Vec::with_capacity(2 * n);
124 for (scalar, point) in scalars.iter().zip(points) {
125 assert_eq!(scalar.limbs.len(), N_LIMBS);
126
127 let (k1_neg, k2_neg, k1_abs, k2_abs) = Secp256k1EndosplitHint::call(b, &scalar.limbs);
128 check_endomorphism_split(b, curve, k1_neg, k2_neg, k1_abs, k2_abs, scalar);
129
130 let base = curve.negate_if(b, k1_neg, point);
132 let table_p = build_strauss_table(b, curve, &base, window);
133
134 let rel_neg = b.bxor(k1_neg, k2_neg);
137 let table_phi = table_p
138 .iter()
139 .map(|q| {
140 let phi = curve.endomorphism(b, q);
141 curve.negate_if(b, rel_neg, &phi)
142 })
143 .collect::<Vec<_>>();
144
145 tables.push(table_p);
146 subscalars.push(k1_abs);
147 tables.push(table_phi);
148 subscalars.push(k2_abs);
149 }
150
151 let subscalar_refs = subscalars
152 .iter()
153 .map(<[Wire; 2]>::as_slice)
154 .collect::<Vec<_>>();
155 strauss_accumulate(b, curve, &tables, &subscalar_refs, window, 128)
157}
158
159fn build_strauss_table(
164 b: &CircuitBuilder,
165 curve: &Secp256k1,
166 point: &Secp256k1Affine,
167 window: usize,
168) -> Vec<Secp256k1Affine> {
169 let mut table = Vec::with_capacity(1 << window);
170 table.push(Secp256k1Affine::point_at_infinity(b));
171 table.push(point.clone());
172 for x in 2..1 << window {
173 let multiple = if x == 2 {
174 curve.double(b, point)
175 } else {
176 curve.add_incomplete(b, &table[x - 1], point)
177 };
178 table.push(multiple);
179 }
180 table
181}
182
183fn strauss_accumulate(
192 b: &CircuitBuilder,
193 curve: &Secp256k1,
194 tables: &[Vec<Secp256k1Affine>],
195 subscalars: &[&[Wire]],
196 window: usize,
197 exponent_bits: usize,
198) -> Secp256k1Affine {
199 assert_eq!(subscalars.len(), tables.len(), "one subscalar per base point");
200
201 let tables_flat: Vec<Vec<Vec<Wire>>> = tables
203 .iter()
204 .map(|table| table.iter().map(Secp256k1Affine::to_wires).collect())
205 .collect();
206 let table_refs: Vec<Vec<&[Wire]>> = tables_flat
207 .iter()
208 .map(|table| table.iter().map(Vec::as_slice).collect())
209 .collect();
210
211 let n_windows = exponent_bits.div_ceil(window);
212 let mut acc = Secp256k1Affine::point_at_infinity(b);
213
214 for w_idx in (0..n_windows).rev() {
215 if w_idx != n_windows - 1 {
218 for _ in 0..window {
219 acc = curve.double(b, &acc);
220 }
221 }
222
223 let base_bit = w_idx * window;
224 for (point_idx, subscalar) in subscalars.iter().enumerate() {
225 let n_bits = (base_bit + window).min(exponent_bits) - base_bit;
230 let mask = b.add_constant_64((1u64 << n_bits) - 1);
231 let offset = (base_bit % Word::BITS) as u32;
232 let lo = base_bit / Word::BITS;
233 let hi = (base_bit + n_bits - 1) / Word::BITS;
234 let sel = if lo == hi {
235 b.band(b.shr(subscalar[lo], offset), mask)
236 } else {
237 let low = b.shr(subscalar[lo], offset);
239 let high = b.shl(subscalar[hi], Word::BITS as u32 - offset);
243 b.band(b.bxor(low, high), mask)
244 };
245
246 let selected =
247 Secp256k1Affine::from_wires(&multi_wire_multiplex(b, &table_refs[point_idx], sel));
248 acc = curve.add_incomplete(b, &acc, &selected);
249 }
250 }
251
252 acc
253}
254
255#[cfg(test)]
256mod tests {
257 use binius_core::word::Word;
258 use binius_frontend::CircuitBuilder;
259 use k256::{
260 ProjectivePoint, Scalar, U256,
261 elliptic_curve::{scalar::FromUintUnchecked, sec1::ToSec1Point},
262 };
263 use rand::prelude::*;
264
265 use super::*;
266 use crate::{
267 bignum::{BigUint, assert_eq},
268 secp256k1::{Secp256k1, Secp256k1Affine},
269 };
270
271 #[test]
272 fn test_scalar_mul_with_endomorphism() {
273 let builder = CircuitBuilder::new();
274 let curve = Secp256k1::new(&builder);
275
276 let mut rng = StdRng::seed_from_u64(0);
278 let mut scalar_bytes = [0u8; 32];
279 rng.fill(&mut scalar_bytes);
280
281 let k256_uint = U256::from_be_slice(&scalar_bytes);
283 let k256_scalar = Scalar::from_uint_unchecked(k256_uint);
284 let scalar_bigint = num_bigint::BigUint::from_bytes_be(&scalar_bytes);
285 let scalar = BigUint::new_constant(&builder, &scalar_bigint).zero_extend(&builder, N_LIMBS);
286
287 let k256_point = ProjectivePoint::mul_by_generator(&k256_scalar).to_affine();
289
290 let point_bytes = k256_point.to_sec1_point(false).to_bytes();
292 let x_coord = num_bigint::BigUint::from_bytes_be(&point_bytes[1..33]);
293 let y_coord = num_bigint::BigUint::from_bytes_be(&point_bytes[33..65]);
294
295 let expected_x = BigUint::new_constant(&builder, &x_coord);
297 let expected_y = BigUint::new_constant(&builder, &y_coord);
298
299 let generator = Secp256k1Affine::generator(&builder);
301
302 let result = scalar_mul(&builder, &curve, &scalar, generator);
304
305 assert_eq(&builder, "result_x", &result.x, &expected_x);
307 assert_eq(&builder, "result_y", &result.y, &expected_y);
308
309 builder.force_commit(result.is_point_at_infinity);
311
312 let cs = builder.build();
314 let mut w = cs.new_witness_filler();
315 assert!(cs.populate_wire_witness(&mut w).is_ok());
316
317 assert_eq!(w[result.is_point_at_infinity], Word::ZERO);
319 }
320
321 type StraussFn =
322 fn(&CircuitBuilder, &Secp256k1, usize, &[BigUint], &[Secp256k1Affine]) -> Secp256k1Affine;
323
324 fn check_msm_strauss(msm_fn: StraussFn, window: usize, n: usize, seed: u64) {
328 let builder = CircuitBuilder::new();
329 let curve = Secp256k1::new(&builder);
330 let mut rng = StdRng::seed_from_u64(seed);
331
332 let mut scalars = Vec::with_capacity(n);
333 let mut points = Vec::with_capacity(n);
334 let mut expected = ProjectivePoint::IDENTITY;
335
336 for _ in 0..n {
337 let mut scalar_bytes = [0u8; 32];
338 rng.fill(&mut scalar_bytes);
339 let k256_scalar = Scalar::from_uint_unchecked(U256::from_be_slice(&scalar_bytes));
340
341 let mut point_seed = [0u8; 32];
343 rng.fill(&mut point_seed);
344 let r = Scalar::from_uint_unchecked(U256::from_be_slice(&point_seed));
345 let point = ProjectivePoint::mul_by_generator(&r);
346 expected += point * k256_scalar;
347
348 let scalar_bigint = num_bigint::BigUint::from_bytes_be(&scalar_bytes);
349 scalars.push(
350 BigUint::new_constant(&builder, &scalar_bigint).zero_extend(&builder, N_LIMBS),
351 );
352
353 let point_bytes = point.to_affine().to_sec1_point(false).to_bytes();
354 let x_coord = num_bigint::BigUint::from_bytes_be(&point_bytes[1..33]);
355 let y_coord = num_bigint::BigUint::from_bytes_be(&point_bytes[33..65]);
356 points.push(Secp256k1Affine {
357 x: BigUint::new_constant(&builder, &x_coord).zero_extend(&builder, N_LIMBS),
358 y: BigUint::new_constant(&builder, &y_coord).zero_extend(&builder, N_LIMBS),
359 is_point_at_infinity: builder.add_constant(Word::ZERO),
360 });
361 }
362
363 let result = msm_fn(&builder, &curve, window, &scalars, &points);
364
365 let expected_bytes = expected.to_affine().to_sec1_point(false).to_bytes();
366 let expected_x = BigUint::new_constant(
367 &builder,
368 &num_bigint::BigUint::from_bytes_be(&expected_bytes[1..33]),
369 );
370 let expected_y = BigUint::new_constant(
371 &builder,
372 &num_bigint::BigUint::from_bytes_be(&expected_bytes[33..65]),
373 );
374
375 assert_eq(&builder, "msm_x", &result.x, &expected_x);
376 assert_eq(&builder, "msm_y", &result.y, &expected_y);
377
378 builder.force_commit(result.is_point_at_infinity);
380
381 let cs = builder.build();
382 let mut w = cs.new_witness_filler();
383 assert!(cs.populate_wire_witness(&mut w).is_ok());
384 assert_eq!(w[result.is_point_at_infinity], Word::ZERO);
385 }
386
387 #[test]
389 fn test_msm_strauss_endo_window1() {
390 check_msm_strauss(msm_strauss_endo, 1, 1, 6);
391 check_msm_strauss(msm_strauss_endo, 1, 2, 0);
392 }
393
394 #[test]
396 fn test_msm_strauss_endo_window2() {
397 check_msm_strauss(msm_strauss_endo, 2, 1, 7);
398 check_msm_strauss(msm_strauss_endo, 2, 2, 8);
399 check_msm_strauss(msm_strauss_endo, 2, 3, 9);
400 }
401
402 #[test]
405 fn test_msm_strauss_endo_window3() {
406 check_msm_strauss(msm_strauss_endo, 3, 1, 10);
407 check_msm_strauss(msm_strauss_endo, 3, 2, 11);
408 check_msm_strauss(msm_strauss_endo, 3, 3, 12);
409 }
410}