1use binius_frontend::{CircuitBuilder, Wire};
28
29pub fn var_sll(b: &CircuitBuilder, x: Wire, shift: Wire) -> Wire {
33 var_sll_blocks(b, x, shift, 6, 0)
34}
35
36pub fn var_srl(b: &CircuitBuilder, x: Wire, shift: Wire) -> Wire {
40 var_srl_blocks(b, x, shift, 6, 0)
41}
42
43pub fn var_sra(b: &CircuitBuilder, x: Wire, shift: Wire) -> Wire {
47 var_sra_blocks(b, x, shift, 6, 0)
48}
49
50pub 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
68pub 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
86pub 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
104pub fn var_sll_bytes(b: &CircuitBuilder, x: Wire, shift: Wire) -> Wire {
108 var_sll_blocks(b, x, shift, 3, 3)
109}
110
111pub fn var_srl_bytes(b: &CircuitBuilder, x: Wire, shift: Wire) -> Wire {
115 var_srl_blocks(b, x, shift, 3, 3)
116}
117
118pub 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 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 const X_FIXTURES: &[u64] = &[
235 0,
236 1,
237 0xFFFF_FFFF_FFFF_FFFF,
238 0x8000_0000_0000_0000, 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); 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}