Skip to main content

binius_circuits/
multiplexer.rs

1// Copyright 2025 Irreducible Inc.
2use binius_core::word::Word;
3use binius_frontend::{CircuitBuilder, Wire};
4
5/// Creates a multiplexer circuit that selects a group of wires from multiple groups based on a
6/// selector value.
7///
8/// This circuit validates that the output contains the group at position `sel` from the input
9/// groups. Each group must have the same number of wires.
10///
11/// # Arguments
12/// * `b` - Circuit builder
13/// * `inputs` - Slice of wire groups, where each group has the same number of wires
14/// * `sel` - Selector value (only ceil(log2(N)) LSB bits are used, where N is the number of groups)
15///
16/// # Returns
17/// A vector of wires representing the selected group
18///
19/// # Implementation Details
20/// For each wire position across all groups, builds a separate multiplexer tree using select gates.
21/// If inputs is `[[a1,a2], [b1,b2], [c1,c2]]` and sel=1, output is `[b1,b2]`.
22///
23/// # Panics
24/// * If inputs is empty
25/// * If any group is empty
26/// * If groups have different lengths
27pub 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	// Assert all groups have the same length
34	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	// For each position in the groups, build a multiplexer
46	(0..group_size)
47		.map(|position| {
48			// Collect all wires at this position across groups
49			let wires_at_position: Vec<Wire> = inputs.iter().map(|group| group[position]).collect();
50			// Build multiplexer for this position
51			single_wire_multiplex(b, &wires_at_position, sel)
52		})
53		.collect()
54}
55
56/// Creates a single-wire multiplexer circuit that selects an element from a vector based on a
57/// selector value.
58///
59/// This circuit validates that the output contains the element at position `sel` from the input
60/// vector `inputs`. The selection is done using a binary tree of 2-to-1 select gates.
61///
62/// # Arguments
63/// * `b` - Circuit builder
64/// * `inputs` - Input vector of N elements (N can be any positive number)
65/// * `sel` - Selector value (only ceil(log2(N)) LSB bits are used)
66///
67/// # Returns
68/// The output wire containing the selected element
69///
70/// # Implementation Details
71/// - Builds a binary tree of 2-to-1 select gates, level by level
72/// - Binary tree has ceil(log2(N)) levels for N inputs
73/// - For non-power-of-two inputs, unpaired wires are carried forward to the next level
74/// - Each level uses a different bit from the selector
75/// - The final output is the single wire remaining after all levels
76///
77/// # Panics
78/// * If inputs.len() is 0
79pub 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	// Calculate number of selector bits needed
86	let num_sel_bits = log2_ceil_usize(n);
87
88	// Build MUX tree from bottom to top using level-by-level approach
89	// This creates an optimal tree with exactly N-1 MUX gates
90	let mut current_level = inputs.to_vec();
91
92	// Process level by level until we have a single output
93	for bit_level in 0..num_sel_bits {
94		let sel_bit = b.shl(sel, (Word::BITS - 1 - bit_level) as u32);
95
96		// Process pairs of wires at the current level
97		let next_level = current_level
98			.chunks(2)
99			.map(|pair| {
100				if let Ok([lhs, rhs]) = TryInto::<[Wire; 2]>::try_into(pair) {
101					// We have a pair - create a MUX gate
102					// Use the current bit level for selection
103					b.select(sel_bit, rhs, lhs)
104				} else {
105					// Odd wire out - carry it forward to the next level
106					pair[0]
107				}
108			})
109			.collect();
110
111		current_level = next_level;
112	}
113
114	// The final wire is our output
115	current_level[0]
116}
117
118/// Rotates `words` left by a dynamic word count, returning the first `n_out` positions.
119///
120/// `out[i]` is `words[(shift + i) % words.len()]`. A caller that only needs a prefix reads the
121/// wrap-around at positions past its own validity bound, so it must treat those as unconstrained.
122///
123/// # Arguments
124/// * `b` - Circuit builder
125/// * `words` - The array to rotate
126/// * `shift` - Rotate amount in words (only `ceil(log2(len))` LSB bits are used)
127/// * `n_out` - Number of leading positions to return
128///
129/// # Implementation Details
130///
131/// One barrel stage per bit of `shift`, largest stage first so the working window shrinks as the
132/// stages are applied. That costs one select per position per stage, where a multiplexer per
133/// output position would instead cost `words.len() - 1` selects each.
134///
135/// # Panics
136/// * If `words` is empty
137/// * If `n_out` exceeds `words.len()`
138pub 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	// Width the window must have before each stage, derived back from `n_out`: a stage rotating by
151	// `2^bit` reads positions `i` and `i + 2^bit`, and indices wrap, so the width caps at `n`.
152	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	/// Exhaustively checks `rotate_left_dynamic` against `words[(shift + i) % n]` for every shift
185	/// the rotate can express, at sizes on both sides of a power of two.
186	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		// The rotate reads `ceil(log2(n))` bits of the shift, so that range covers every distinct
198		// permutation it can produce.
199		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				// Distinct per position so a wrong index cannot pass by coincidence.
204				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		// Powers of two, just under, just over, and a prefix much shorter than the array.
223		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	/// Helper function to verify single-wire multiplexer behavior
256	/// Takes input values and test cases as (selector, expected_output) pairs
257	fn verify_single_wire_multiplex(values: &[u64], test_cases: &[(u64, u64)]) {
258		let n = values.len();
259		let builder = CircuitBuilder::new();
260
261		// Create input wires
262		let inputs: Vec<Wire> = (0..n).map(|_| builder.add_inout()).collect();
263		let sel = builder.add_inout();
264
265		// Create multiplexer circuit
266		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		// Test each case
273		for &(selector, expected_val) in test_cases {
274			let mut w = built.new_witness_filler();
275
276			// Set input values
277			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			// Populate witness
284			built.populate_wire_witness(&mut w).unwrap();
285
286			// Verify constraints
287			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		// Test with 4 elements (common power-of-two case)
295		verify_single_wire_multiplex(
296			&[13, 7, 25, 100],
297			&[
298				(0, 13),  // Select index 0
299				(1, 7),   // Select index 1
300				(2, 25),  // Select index 2
301				(3, 100), // Select index 3
302			],
303		);
304
305		// Test with 8 elements (larger power-of-two)
306		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		// Test with 3 elements (creates asymmetric tree)
314		verify_single_wire_multiplex(
315			&[10, 20, 30],
316			&[
317				(0, 10), // Select index 0
318				(1, 20), // Select index 1
319				(2, 30), // Select index 2
320				(3, 30), // Index 3 wraps in a specific way due to tree structure
321			],
322		);
323
324		// Test with 5 elements
325		verify_single_wire_multiplex(
326			&[100, 200, 300, 400, 500],
327			&[
328				(0, 100), // Select index 0
329				(2, 300), // Select index 2
330				(4, 500), // Select index 4
331			],
332		);
333
334		// Test with 7 elements
335		let values = [11, 22, 33, 44, 55, 66, 77];
336		verify_single_wire_multiplex(
337			&values,
338			&[
339				(0, 11), // Select index 0
340				(3, 44), // Select index 3
341				(6, 77), // Select index 6
342				(7, 77), // Index 7 wraps to 6 in the tree structure
343			],
344		);
345	}
346
347	#[test]
348	fn test_single_element() {
349		// Edge case: single input always returns that input regardless of selector
350		verify_single_wire_multiplex(
351			&[42],
352			&[
353				(0, 42),   // Selector 0
354				(1, 42),   // Selector 1 (ignored)
355				(100, 42), // Large selector (ignored)
356			],
357		);
358	}
359
360	#[test]
361	fn test_out_of_bounds_selector() {
362		// Test selector wrapping behavior with power-of-two size
363		verify_single_wire_multiplex(
364			&[10, 20, 30, 40],
365			&[
366				(4, 10),   // 4 & 3 = 0
367				(5, 20),   // 5 & 3 = 1
368				(6, 30),   // 6 & 3 = 2
369				(7, 40),   // 7 & 3 = 3
370				(15, 40),  // 15 & 3 = 3
371				(100, 10), // 100 & 3 = 0
372			],
373		);
374
375		// Test with non-power-of-two (behavior depends on tree structure)
376		verify_single_wire_multiplex(
377			&[1, 2, 3],
378			&[
379				(3, 3), // Out of bounds wraps based on tree structure
380				(4, 1), // Wraps around
381				(5, 2), // Wraps around
382			],
383		);
384	}
385
386	/// Helper function to verify multi-wire multiplexer behavior
387	/// Takes groups of values and test cases as (selector, expected_group_index) pairs
388	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		// Create input wire groups
394		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		// Convert to the format needed by multi_wire_multiplex
400		let input_refs: Vec<&[Wire]> = input_groups.iter().map(|g| g.as_slice()).collect();
401
402		// Create multiplexer circuit
403		let outputs = multi_wire_multiplex(&builder, &input_refs, sel);
404
405		// Create expected output wires
406		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		// Test each case
414		for &(selector, expected_group_idx) in test_cases {
415			let mut w = built.new_witness_filler();
416
417			// Set input values
418			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			// Set expected values
426			for (i, &val) in groups[expected_group_idx].iter().enumerate() {
427				w[expected[i]] = Word(val);
428			}
429
430			// Populate witness
431			built.populate_wire_witness(&mut w).unwrap();
432
433			// Verify constraints
434			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		// Test with 2-wire groups
442		let groups = vec![
443			vec![10, 11], // Group 0
444			vec![20, 21], // Group 1
445			vec![30, 31], // Group 2
446			vec![40, 41], // Group 3
447		];
448
449		verify_multi_wire_multiplex(
450			&groups,
451			&[
452				(0, 0), // Select group 0
453				(1, 1), // Select group 1
454				(2, 2), // Select group 2
455				(3, 3), // Select group 3
456				(4, 0), // Wraps to group 0
457				(7, 3), // Wraps to group 3
458			],
459		);
460	}
461
462	#[test]
463	fn test_multi_wire_three_wire_groups() {
464		// Test with 3-wire groups
465		let groups = vec![
466			vec![100, 101, 102], // Group 0
467			vec![200, 201, 202], // Group 1
468			vec![300, 301, 302], // Group 2
469		];
470
471		verify_multi_wire_multiplex(
472			&groups,
473			&[
474				(0, 0), // Select group 0
475				(1, 1), // Select group 1
476				(2, 2), // Select group 2
477				(3, 2), // Wraps based on tree structure
478			],
479		);
480	}
481
482	#[test]
483	fn test_multi_wire_single_group() {
484		// Edge case: single group with multiple wires
485		let groups = vec![
486			vec![50, 51, 52, 53], // Only group
487		];
488
489		verify_multi_wire_multiplex(
490			&groups,
491			&[
492				(0, 0),   // Select the only group
493				(5, 0),   // Any selector returns the only group
494				(100, 0), // Any selector returns the only group
495			],
496		);
497	}
498
499	#[test]
500	fn test_multi_wire_single_wire_per_group() {
501		// Edge case: multiple groups but each has only one wire
502		// This should behave identically to single_wire_multiplex
503		let groups = vec![
504			vec![10], // Group 0
505			vec![20], // Group 1
506			vec![30], // Group 2
507			vec![40], // Group 3
508		];
509
510		verify_multi_wire_multiplex(
511			&groups,
512			&[
513				(0, 0), // Select group 0
514				(1, 1), // Select group 1
515				(2, 2), // Select group 2
516				(3, 3), // Select group 3
517			],
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		// Create mismatched groups
527		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		// This should panic
534		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		// This should panic
546		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		// This should panic
559		multi_wire_multiplex(&builder, &inputs, sel);
560	}
561}