Skip to main content

binius_circuits/
shift.rs

1// Copyright 2025 Irreducible Inc.
2//! Variable-amount shift gadgets.
3//!
4//! `CircuitBuilder` only exposes shifts by a compile-time-constant amount. This module provides
5//! barrel-shifter gadgets that shift by a runtime [`Wire`] amount:
6//!
7//! - [`var_sll`], [`var_srl`], [`var_sra`] — shift by the low 6 bits of `shift` (full bit range for
8//!   a 64-bit word).
9//! - [`var_sll_blocks`], [`var_srl_blocks`], [`var_sra_blocks`] — most general: the actual
10//!   bit-shift is `shift * 2^block_bits`, with `shift` treated as a `shift_bits`-bit unsigned
11//!   integer. Useful when the caller knows the shift is a multiple of a fixed power of two, or when
12//!   only a narrow range of shifts is possible.
13//! - [`var_sll_bytes`], [`var_srl_bytes`], [`var_sra_bytes`] — byte-granularity shifts (shift in
14//!   `0..8` bytes).
15//!
16//! ## Precondition
17//!
18//! For `*_blocks`, the shift amount is treated as a `shift_bits`-bit unsigned integer; bits above
19//! position `shift_bits - 1` are ignored. Callers must ensure `shift < 2^shift_bits`; this is
20//! not checked by the gadget. The other variants impose the same precondition with `shift_bits`
21//! fixed (6 for the base variants, 3 for the byte variants).
22//!
23//! ## Cost
24//!
25//! `*_blocks`: `3 * shift_bits` AND constraints. Base variants: 18. Byte variants: 9.
26
27use binius_frontend::{CircuitBuilder, Wire};
28
29/// Variable-amount logical left shift.
30///
31/// Returns `x << shift`, reading the low 6 bits of `shift` as the shift amount.
32pub fn var_sll(b: &CircuitBuilder, x: Wire, shift: Wire) -> Wire {
33	var_sll_blocks(b, x, shift, 6, 0)
34}
35
36/// Variable-amount logical right shift.
37///
38/// Returns `x >> shift`, reading the low 6 bits of `shift` as the shift amount.
39pub fn var_srl(b: &CircuitBuilder, x: Wire, shift: Wire) -> Wire {
40	var_srl_blocks(b, x, shift, 6, 0)
41}
42
43/// Variable-amount arithmetic right shift.
44///
45/// Returns `x SAR shift` (sign-extending), reading the low 6 bits of `shift` as the shift amount.
46pub fn var_sra(b: &CircuitBuilder, x: Wire, shift: Wire) -> Wire {
47	var_sra_blocks(b, x, shift, 6, 0)
48}
49
50/// Variable-amount logical left shift by blocks of `2^block_bits` bits.
51///
52/// Returns `x << (shift * 2^block_bits)`, where the low `shift_bits` bits of `shift` are the
53/// block count.
54///
55/// # Panics
56///
57/// Panics if `shift_bits + block_bits > 6`.
58pub fn var_sll_blocks(
59	b: &CircuitBuilder,
60	x: Wire,
61	shift: Wire,
62	shift_bits: usize,
63	block_bits: usize,
64) -> Wire {
65	var_shift_blocks(b, x, shift, shift_bits, block_bits, CircuitBuilder::shl)
66}
67
68/// Variable-amount logical right shift by blocks of `2^block_bits` bits.
69///
70/// Returns `x >> (shift * 2^block_bits)`, where the low `shift_bits` bits of `shift` are the
71/// block count.
72///
73/// # Panics
74///
75/// Panics if `shift_bits + block_bits > 6`.
76pub fn var_srl_blocks(
77	b: &CircuitBuilder,
78	x: Wire,
79	shift: Wire,
80	shift_bits: usize,
81	block_bits: usize,
82) -> Wire {
83	var_shift_blocks(b, x, shift, shift_bits, block_bits, CircuitBuilder::shr)
84}
85
86/// Variable-amount arithmetic right shift by blocks of `2^block_bits` bits.
87///
88/// Returns `x SAR (shift * 2^block_bits)` (sign-extending), where the low `shift_bits` bits of
89/// `shift` are the block count.
90///
91/// # Panics
92///
93/// Panics if `shift_bits + block_bits > 6`.
94pub fn var_sra_blocks(
95	b: &CircuitBuilder,
96	x: Wire,
97	shift: Wire,
98	shift_bits: usize,
99	block_bits: usize,
100) -> Wire {
101	var_shift_blocks(b, x, shift, shift_bits, block_bits, CircuitBuilder::sar)
102}
103
104/// Variable-amount logical left shift by whole bytes.
105///
106/// Returns `x << (shift * 8)`. The low 3 bits of `shift` are the byte count (range `0..8`).
107pub fn var_sll_bytes(b: &CircuitBuilder, x: Wire, shift: Wire) -> Wire {
108	var_sll_blocks(b, x, shift, 3, 3)
109}
110
111/// Variable-amount logical right shift by whole bytes.
112///
113/// Returns `x >> (shift * 8)`. The low 3 bits of `shift` are the byte count (range `0..8`).
114pub fn var_srl_bytes(b: &CircuitBuilder, x: Wire, shift: Wire) -> Wire {
115	var_srl_blocks(b, x, shift, 3, 3)
116}
117
118/// Variable-amount arithmetic right shift by whole bytes.
119///
120/// Returns `x SAR (shift * 8)` (sign-extending). The low 3 bits of `shift` are the byte count
121/// (range `0..8`).
122pub fn var_sra_bytes(b: &CircuitBuilder, x: Wire, shift: Wire) -> Wire {
123	var_sra_blocks(b, x, shift, 3, 3)
124}
125
126fn var_shift_blocks(
127	b: &CircuitBuilder,
128	x: Wire,
129	shift: Wire,
130	shift_bits: usize,
131	block_bits: usize,
132	step: impl Fn(&CircuitBuilder, Wire, u32) -> Wire,
133) -> Wire {
134	assert!(
135		shift_bits + block_bits <= 6,
136		"shift_bits={shift_bits} + block_bits={block_bits} > 6 (max for 64-bit word)"
137	);
138	let mut result = x;
139	for i in 0..shift_bits {
140		// Move bit i of `shift` into the MSB position so `select` reads it as the condition.
141		let cond = b.shl(shift, 63 - i as u32);
142		let shifted = step(b, result, 1u32 << (i + block_bits));
143		result = b.select(cond, shifted, result);
144	}
145	result
146}
147
148#[cfg(test)]
149mod tests {
150	use binius_core::word::Word;
151	use binius_frontend::Circuit;
152	use proptest::prelude::*;
153
154	use super::*;
155
156	type Gadget = fn(&CircuitBuilder, Wire, Wire) -> Wire;
157	type BlocksGadget = fn(&CircuitBuilder, Wire, Wire, usize, usize) -> Wire;
158
159	fn build_circuit(gadget: Gadget) -> (Circuit, Wire, Wire, Wire) {
160		let builder = CircuitBuilder::new();
161		let x = builder.add_witness();
162		let shift = builder.add_witness();
163		let output = builder.add_witness();
164		let computed = gadget(&builder, x, shift);
165		builder.assert_eq("var_shift_result", computed, output);
166		let circuit = builder.build();
167		(circuit, x, shift, output)
168	}
169
170	fn build_blocks_circuit(
171		gadget: BlocksGadget,
172		shift_bits: usize,
173		block_bits: usize,
174	) -> (Circuit, Wire, Wire, Wire) {
175		let builder = CircuitBuilder::new();
176		let x = builder.add_witness();
177		let shift = builder.add_witness();
178		let output = builder.add_witness();
179		let computed = gadget(&builder, x, shift, shift_bits, block_bits);
180		builder.assert_eq("var_shift_result", computed, output);
181		let circuit = builder.build();
182		(circuit, x, shift, output)
183	}
184
185	fn check_ok(gadget: Gadget, x_val: u64, shift_val: u64, expected: u64) {
186		let (circuit, x, shift, output) = build_circuit(gadget);
187		fill_and_check(&circuit, x, shift, output, x_val, shift_val, expected);
188	}
189
190	fn check_ok_blocks(
191		gadget: BlocksGadget,
192		shift_bits: usize,
193		block_bits: usize,
194		x_val: u64,
195		shift_val: u64,
196		expected: u64,
197	) {
198		let (circuit, x, shift, output) = build_blocks_circuit(gadget, shift_bits, block_bits);
199		fill_and_check(&circuit, x, shift, output, x_val, shift_val, expected);
200	}
201
202	fn fill_and_check(
203		circuit: &Circuit,
204		x: Wire,
205		shift: Wire,
206		output: Wire,
207		x_val: u64,
208		shift_val: u64,
209		expected: u64,
210	) {
211		let mut w = circuit.new_witness_filler();
212		w[x] = Word(x_val);
213		w[shift] = Word(shift_val);
214		w[output] = Word(expected);
215		circuit.populate_wire_witness(&mut w).unwrap_or_else(|e| {
216			panic!(
217				"populate failed: x=0x{x_val:016x} shift={shift_val} expected=0x{expected:016x}: {e:?}"
218			)
219		});
220	}
221
222	fn ref_sll(x: u64, s: u64) -> u64 {
223		if s >= 64 { 0 } else { x << s }
224	}
225	fn ref_srl(x: u64, s: u64) -> u64 {
226		if s >= 64 { 0 } else { x >> s }
227	}
228	fn ref_sra(x: u64, s: u64) -> u64 {
229		let s = s.min(63);
230		((x as i64) >> s) as u64
231	}
232
233	// Static edge-case fixtures.
234	const X_FIXTURES: &[u64] = &[
235		0,
236		1,
237		0xFFFF_FFFF_FFFF_FFFF,
238		0x8000_0000_0000_0000, // MSB only — important for sra
239		0x0123_4567_89AB_CDEF,
240		0xDEAD_BEEF_CAFE_F00D,
241		0x5555_5555_5555_5555,
242	];
243
244	#[test]
245	fn var_sll_fixtures() {
246		for &x_val in X_FIXTURES {
247			for shift_val in 0u64..64 {
248				check_ok(var_sll, x_val, shift_val, ref_sll(x_val, shift_val));
249			}
250		}
251	}
252
253	#[test]
254	fn var_srl_fixtures() {
255		for &x_val in X_FIXTURES {
256			for shift_val in 0u64..64 {
257				check_ok(var_srl, x_val, shift_val, ref_srl(x_val, shift_val));
258			}
259		}
260	}
261
262	#[test]
263	fn var_sra_fixtures() {
264		for &x_val in X_FIXTURES {
265			for shift_val in 0u64..64 {
266				check_ok(var_sra, x_val, shift_val, ref_sra(x_val, shift_val));
267			}
268		}
269	}
270
271	const BLOCK_CONFIGS: &[(usize, usize)] = &[
272		(0, 0),
273		(0, 6),
274		(3, 0),
275		(3, 3),
276		(6, 0),
277		(1, 5),
278		(2, 4),
279		(4, 2),
280	];
281
282	#[test]
283	fn var_sll_blocks_fixtures() {
284		for &(shift_bits, block_bits) in BLOCK_CONFIGS {
285			let max_shift = if shift_bits == 0 {
286				1
287			} else {
288				1u64 << shift_bits
289			};
290			for &x_val in X_FIXTURES {
291				for shift_val in 0..max_shift {
292					let effective = shift_val << block_bits;
293					check_ok_blocks(
294						var_sll_blocks,
295						shift_bits,
296						block_bits,
297						x_val,
298						shift_val,
299						ref_sll(x_val, effective),
300					);
301				}
302			}
303		}
304	}
305
306	#[test]
307	fn var_srl_blocks_fixtures() {
308		for &(shift_bits, block_bits) in BLOCK_CONFIGS {
309			let max_shift = if shift_bits == 0 {
310				1
311			} else {
312				1u64 << shift_bits
313			};
314			for &x_val in X_FIXTURES {
315				for shift_val in 0..max_shift {
316					let effective = shift_val << block_bits;
317					check_ok_blocks(
318						var_srl_blocks,
319						shift_bits,
320						block_bits,
321						x_val,
322						shift_val,
323						ref_srl(x_val, effective),
324					);
325				}
326			}
327		}
328	}
329
330	#[test]
331	fn var_sra_blocks_fixtures() {
332		for &(shift_bits, block_bits) in BLOCK_CONFIGS {
333			let max_shift = if shift_bits == 0 {
334				1
335			} else {
336				1u64 << shift_bits
337			};
338			for &x_val in X_FIXTURES {
339				for shift_val in 0..max_shift {
340					let effective = shift_val << block_bits;
341					check_ok_blocks(
342						var_sra_blocks,
343						shift_bits,
344						block_bits,
345						x_val,
346						shift_val,
347						ref_sra(x_val, effective),
348					);
349				}
350			}
351		}
352	}
353
354	#[test]
355	fn var_sll_bytes_fixtures() {
356		for &x_val in X_FIXTURES {
357			for shift_val in 0u64..8 {
358				check_ok(var_sll_bytes, x_val, shift_val, ref_sll(x_val, shift_val * 8));
359			}
360		}
361	}
362
363	#[test]
364	fn var_srl_bytes_fixtures() {
365		for &x_val in X_FIXTURES {
366			for shift_val in 0u64..8 {
367				check_ok(var_srl_bytes, x_val, shift_val, ref_srl(x_val, shift_val * 8));
368			}
369		}
370	}
371
372	#[test]
373	fn var_sra_bytes_fixtures() {
374		for &x_val in X_FIXTURES {
375			for shift_val in 0u64..8 {
376				check_ok(var_sra_bytes, x_val, shift_val, ref_sra(x_val, shift_val * 8));
377			}
378		}
379	}
380
381	proptest! {
382		#![proptest_config(ProptestConfig::with_cases(64))]
383
384		#[test]
385		fn var_sll_random(x_val in any::<u64>(), shift_val in 0u64..64) {
386			check_ok(var_sll, x_val, shift_val, ref_sll(x_val, shift_val));
387		}
388
389		#[test]
390		fn var_srl_random(x_val in any::<u64>(), shift_val in 0u64..64) {
391			check_ok(var_srl, x_val, shift_val, ref_srl(x_val, shift_val));
392		}
393
394		#[test]
395		fn var_sra_random(x_val in any::<u64>(), shift_val in 0u64..64) {
396			check_ok(var_sra, x_val, shift_val, ref_sra(x_val, shift_val));
397		}
398
399		#[test]
400		fn var_sll_bytes_random(x_val in any::<u64>(), shift_val in 0u64..8) {
401			check_ok(var_sll_bytes, x_val, shift_val, ref_sll(x_val, shift_val * 8));
402		}
403
404		#[test]
405		fn var_srl_bytes_random(x_val in any::<u64>(), shift_val in 0u64..8) {
406			check_ok(var_srl_bytes, x_val, shift_val, ref_srl(x_val, shift_val * 8));
407		}
408
409		#[test]
410		fn var_sra_bytes_random(x_val in any::<u64>(), shift_val in 0u64..8) {
411			check_ok(var_sra_bytes, x_val, shift_val, ref_sra(x_val, shift_val * 8));
412		}
413	}
414
415	#[test]
416	fn rejects_incorrect_output() {
417		let (circuit, x, shift, output) = build_circuit(var_sll);
418		let mut w = circuit.new_witness_filler();
419		w[x] = Word(0x1);
420		w[shift] = Word(4);
421		w[output] = Word(0x20); // wrong — should be 0x10
422		assert!(circuit.populate_wire_witness(&mut w).is_err());
423	}
424
425	#[test]
426	#[should_panic(expected = "shift_bits=4 + block_bits=3")]
427	fn rejects_excessive_total_bits() {
428		let builder = CircuitBuilder::new();
429		let x = builder.add_witness();
430		let shift = builder.add_witness();
431		let _ = var_sll_blocks(&builder, x, shift, 4, 3);
432	}
433}