1use binius_core::word::Word;
3use binius_frontend::{CircuitBuilder, Wire};
4
5pub fn multi_wire_multiplex(b: &CircuitBuilder, inputs: &[&[Wire]], sel: Wire) -> Vec<Wire> {
28 assert!(!inputs.is_empty(), "Input groups must not be empty");
29
30 let group_size = inputs[0].len();
31 assert!(group_size > 0, "Groups must not be empty");
32
33 for (i, group) in inputs.iter().enumerate() {
35 assert_eq!(
36 group.len(),
37 group_size,
38 "All groups must have the same length. Group {} has length {}, expected {}",
39 i,
40 group.len(),
41 group_size
42 );
43 }
44
45 (0..group_size)
47 .map(|position| {
48 let wires_at_position: Vec<Wire> = inputs.iter().map(|group| group[position]).collect();
50 single_wire_multiplex(b, &wires_at_position, sel)
52 })
53 .collect()
54}
55
56pub fn single_wire_multiplex(b: &CircuitBuilder, inputs: &[Wire], sel: Wire) -> Wire {
80 let n = inputs.len();
81 if n == 0 {
82 return b.add_constant(Word::ZERO);
83 }
84
85 let num_sel_bits = log2_ceil_usize(n);
87
88 let mut current_level = inputs.to_vec();
91
92 for bit_level in 0..num_sel_bits {
94 let sel_bit = b.shl(sel, (Word::BITS - 1 - bit_level) as u32);
95
96 let next_level = current_level
98 .chunks(2)
99 .map(|pair| {
100 if let Ok([lhs, rhs]) = TryInto::<[Wire; 2]>::try_into(pair) {
101 b.select(sel_bit, rhs, lhs)
104 } else {
105 pair[0]
107 }
108 })
109 .collect();
110
111 current_level = next_level;
112 }
113
114 current_level[0]
116}
117
118pub fn rotate_left_dynamic(
139 b: &CircuitBuilder,
140 words: &[Wire],
141 shift: Wire,
142 n_out: usize,
143) -> Vec<Wire> {
144 let n = words.len();
145 assert!(n > 0, "words must not be empty");
146 assert!(n_out <= n, "n_out ({n_out}) must not exceed words.len() ({n})");
147
148 let n_bits = log2_ceil_usize(n);
149
150 let mut widths = vec![n_out];
153 for bit in 0..n_bits {
154 widths.push(n.min(widths[widths.len() - 1] + (1 << bit)));
155 }
156
157 let mut current = words[..widths[n_bits]].to_vec();
158 for bit in (0..n_bits).rev() {
159 let shift_bit = b.shl(shift, (Word::BITS - 1 - bit) as u32);
160 let amount = 1usize << bit;
161 current = (0..widths[bit])
162 .map(|i| b.select(shift_bit, current[(i + amount) % current.len()], current[i]))
163 .collect();
164 }
165
166 current
167}
168
169#[inline]
170const fn log2_ceil_usize(n: usize) -> usize {
171 if n <= 1 {
172 0
173 } else {
174 (usize::BITS as usize) - ((n - 1).leading_zeros() as usize)
175 }
176}
177
178#[cfg(test)]
179mod tests {
180 use binius_core::word::Word;
181
182 use super::*;
183
184 fn verify_rotate(n: usize, n_out: usize) {
187 let builder = CircuitBuilder::new();
188 let input: Vec<Wire> = (0..n).map(|_| builder.add_inout()).collect();
189 let shift = builder.add_inout();
190 let out = rotate_left_dynamic(&builder, &input, shift, n_out);
191 let expected: Vec<Wire> = (0..n_out).map(|_| builder.add_inout()).collect();
192 for (i, (got, want)) in out.iter().zip(&expected).enumerate() {
193 builder.assert_eq(format!("rot[{i}]"), *got, *want);
194 }
195 let circuit = builder.build();
196
197 let n_shifts = 1usize << log2_ceil_usize(n);
200 for sh in 0..n_shifts {
201 let mut w = circuit.new_witness_filler();
202 for (i, wire) in input.iter().enumerate() {
203 w[*wire] = Word(0x1000 + i as u64);
205 }
206 w[shift] = Word(sh as u64);
207 for (i, wire) in expected.iter().enumerate() {
208 w[*wire] = Word(0x1000 + ((sh + i) % n) as u64);
209 }
210 circuit
211 .populate_wire_witness(&mut w)
212 .unwrap_or_else(|e| panic!("n={n} n_out={n_out} shift={sh}: {e}"));
213 circuit
214 .constraint_system()
215 .verify(&w.into_value_vec())
216 .unwrap_or_else(|e| panic!("n={n} n_out={n_out} shift={sh}: {e}"));
217 }
218 }
219
220 #[test]
221 fn rotate_matches_modular_index() {
222 for (n, n_out) in [
224 (8, 8),
225 (8, 3),
226 (7, 7),
227 (9, 9),
228 (9, 2),
229 (16, 16),
230 (5, 5),
231 (1, 1),
232 (2, 2),
233 ] {
234 verify_rotate(n, n_out);
235 }
236 }
237
238 #[test]
239 #[should_panic(expected = "words must not be empty")]
240 fn rotate_rejects_empty() {
241 let builder = CircuitBuilder::new();
242 let shift = builder.add_inout();
243 rotate_left_dynamic(&builder, &[], shift, 0);
244 }
245
246 #[test]
247 #[should_panic(expected = "must not exceed")]
248 fn rotate_rejects_oversized_prefix() {
249 let builder = CircuitBuilder::new();
250 let input: Vec<Wire> = (0..4).map(|_| builder.add_inout()).collect();
251 let shift = builder.add_inout();
252 rotate_left_dynamic(&builder, &input, shift, 5);
253 }
254
255 fn verify_single_wire_multiplex(values: &[u64], test_cases: &[(u64, u64)]) {
258 let n = values.len();
259 let builder = CircuitBuilder::new();
260
261 let inputs: Vec<Wire> = (0..n).map(|_| builder.add_inout()).collect();
263 let sel = builder.add_inout();
264
265 let output = single_wire_multiplex(&builder, &inputs, sel);
267 let expected = builder.add_inout();
268 builder.assert_eq("single_wire_multiplex_output", output, expected);
269
270 let built = builder.build();
271
272 for &(selector, expected_val) in test_cases {
274 let mut w = built.new_witness_filler();
275
276 for (i, &val) in values.iter().enumerate() {
278 w[inputs[i]] = Word(val);
279 }
280 w[sel] = Word(selector);
281 w[expected] = Word(expected_val);
282
283 built.populate_wire_witness(&mut w).unwrap();
285
286 let cs = built.constraint_system();
288 cs.verify(&w.into_value_vec()).unwrap();
289 }
290 }
291
292 #[test]
293 fn test_power_of_two_size() {
294 verify_single_wire_multiplex(
296 &[13, 7, 25, 100],
297 &[
298 (0, 13), (1, 7), (2, 25), (3, 100), ],
303 );
304
305 let values: Vec<u64> = (10..18).collect();
307 let test_cases: Vec<_> = (0..8).map(|i| (i, values[i as usize])).collect();
308 verify_single_wire_multiplex(&values, &test_cases);
309 }
310
311 #[test]
312 fn test_non_power_of_two() {
313 verify_single_wire_multiplex(
315 &[10, 20, 30],
316 &[
317 (0, 10), (1, 20), (2, 30), (3, 30), ],
322 );
323
324 verify_single_wire_multiplex(
326 &[100, 200, 300, 400, 500],
327 &[
328 (0, 100), (2, 300), (4, 500), ],
332 );
333
334 let values = [11, 22, 33, 44, 55, 66, 77];
336 verify_single_wire_multiplex(
337 &values,
338 &[
339 (0, 11), (3, 44), (6, 77), (7, 77), ],
344 );
345 }
346
347 #[test]
348 fn test_single_element() {
349 verify_single_wire_multiplex(
351 &[42],
352 &[
353 (0, 42), (1, 42), (100, 42), ],
357 );
358 }
359
360 #[test]
361 fn test_out_of_bounds_selector() {
362 verify_single_wire_multiplex(
364 &[10, 20, 30, 40],
365 &[
366 (4, 10), (5, 20), (6, 30), (7, 40), (15, 40), (100, 10), ],
373 );
374
375 verify_single_wire_multiplex(
377 &[1, 2, 3],
378 &[
379 (3, 3), (4, 1), (5, 2), ],
383 );
384 }
385
386 fn verify_multi_wire_multiplex(groups: &[Vec<u64>], test_cases: &[(u64, usize)]) {
389 let num_groups = groups.len();
390 let group_size = groups[0].len();
391 let builder = CircuitBuilder::new();
392
393 let input_groups: Vec<Vec<Wire>> = (0..num_groups)
395 .map(|_| (0..group_size).map(|_| builder.add_inout()).collect())
396 .collect();
397 let sel = builder.add_inout();
398
399 let input_refs: Vec<&[Wire]> = input_groups.iter().map(|g| g.as_slice()).collect();
401
402 let outputs = multi_wire_multiplex(&builder, &input_refs, sel);
404
405 let expected: Vec<Wire> = (0..group_size).map(|_| builder.add_inout()).collect();
407 for (i, &output) in outputs.iter().enumerate() {
408 builder.assert_eq(format!("multi_wire_output_{i}"), output, expected[i]);
409 }
410
411 let built = builder.build();
412
413 for &(selector, expected_group_idx) in test_cases {
415 let mut w = built.new_witness_filler();
416
417 for (group_idx, group) in groups.iter().enumerate() {
419 for (wire_idx, &val) in group.iter().enumerate() {
420 w[input_groups[group_idx][wire_idx]] = Word(val);
421 }
422 }
423 w[sel] = Word(selector);
424
425 for (i, &val) in groups[expected_group_idx].iter().enumerate() {
427 w[expected[i]] = Word(val);
428 }
429
430 built.populate_wire_witness(&mut w).unwrap();
432
433 let cs = built.constraint_system();
435 cs.verify(&w.into_value_vec()).unwrap();
436 }
437 }
438
439 #[test]
440 fn test_multi_wire_two_wire_groups() {
441 let groups = vec![
443 vec![10, 11], vec![20, 21], vec![30, 31], vec![40, 41], ];
448
449 verify_multi_wire_multiplex(
450 &groups,
451 &[
452 (0, 0), (1, 1), (2, 2), (3, 3), (4, 0), (7, 3), ],
459 );
460 }
461
462 #[test]
463 fn test_multi_wire_three_wire_groups() {
464 let groups = vec![
466 vec![100, 101, 102], vec![200, 201, 202], vec![300, 301, 302], ];
470
471 verify_multi_wire_multiplex(
472 &groups,
473 &[
474 (0, 0), (1, 1), (2, 2), (3, 2), ],
479 );
480 }
481
482 #[test]
483 fn test_multi_wire_single_group() {
484 let groups = vec![
486 vec![50, 51, 52, 53], ];
488
489 verify_multi_wire_multiplex(
490 &groups,
491 &[
492 (0, 0), (5, 0), (100, 0), ],
496 );
497 }
498
499 #[test]
500 fn test_multi_wire_single_wire_per_group() {
501 let groups = vec![
504 vec![10], vec![20], vec![30], vec![40], ];
509
510 verify_multi_wire_multiplex(
511 &groups,
512 &[
513 (0, 0), (1, 1), (2, 2), (3, 3), ],
518 );
519 }
520
521 #[test]
522 #[should_panic(expected = "All groups must have the same length")]
523 fn test_multi_wire_mismatched_group_sizes() {
524 let builder = CircuitBuilder::new();
525
526 let group1: Vec<Wire> = (0..2).map(|_| builder.add_inout()).collect();
528 let group2: Vec<Wire> = (0..3).map(|_| builder.add_inout()).collect();
529 let sel = builder.add_inout();
530
531 let inputs = vec![group1.as_slice(), group2.as_slice()];
532
533 multi_wire_multiplex(&builder, &inputs, sel);
535 }
536
537 #[test]
538 #[should_panic(expected = "Input groups must not be empty")]
539 fn test_multi_wire_empty_inputs() {
540 let builder = CircuitBuilder::new();
541 let sel = builder.add_inout();
542
543 let inputs: Vec<&[Wire]> = vec![];
544
545 multi_wire_multiplex(&builder, &inputs, sel);
547 }
548
549 #[test]
550 #[should_panic(expected = "Groups must not be empty")]
551 fn test_multi_wire_empty_group() {
552 let builder = CircuitBuilder::new();
553 let sel = builder.add_inout();
554
555 let empty_group: &[Wire] = &[];
556 let inputs = vec![empty_group];
557
558 multi_wire_multiplex(&builder, &inputs, sel);
560 }
561}