1use binius_utils::serialization::{DeserializeBytes, SerializationError, SerializeBytes};
3use bytes::{Buf, BufMut};
4
5use super::{ValueIndex, ValueVec};
6use crate::word::Word;
7
8#[derive(Copy, Clone, Debug, Hash, PartialEq, Eq, PartialOrd, Ord)]
14#[repr(u8)]
15pub enum ShiftVariant {
16 Sll = 0,
18 Slr = 1,
20 Sar = 2,
25 Rotr = 3,
29 Sll32 = 4,
34 Srl32 = 5,
39 Sra32 = 6,
45 Rotr32 = 7,
51}
52
53impl ShiftVariant {
54 #[inline]
59 pub const fn from_u8(byte: u8) -> Option<Self> {
60 match byte {
61 0 => Some(ShiftVariant::Sll),
62 1 => Some(ShiftVariant::Slr),
63 2 => Some(ShiftVariant::Sar),
64 3 => Some(ShiftVariant::Rotr),
65 4 => Some(ShiftVariant::Sll32),
66 5 => Some(ShiftVariant::Srl32),
67 6 => Some(ShiftVariant::Sra32),
68 7 => Some(ShiftVariant::Rotr32),
69 _ => None,
70 }
71 }
72
73 #[inline]
79 pub const fn is_half_word(self) -> bool {
80 matches!(
81 self,
82 ShiftVariant::Sll32 | ShiftVariant::Srl32 | ShiftVariant::Sra32 | ShiftVariant::Rotr32
83 )
84 }
85
86 #[inline]
94 pub const fn max_amount(self) -> usize {
95 if self.is_half_word() { 32 } else { 64 }
96 }
97
98 #[inline]
108 pub fn apply(self, word: Word, amount: usize) -> Word {
109 let amount = amount as u32;
111 match self {
117 ShiftVariant::Sll => word << amount,
118 ShiftVariant::Slr => word >> amount,
119 ShiftVariant::Sar => word.sar(amount),
120 ShiftVariant::Rotr => word.rotr(amount),
121 ShiftVariant::Sll32 => word.sll32(amount),
122 ShiftVariant::Srl32 => word.srl32(amount),
123 ShiftVariant::Sra32 => word.sra32(amount),
124 ShiftVariant::Rotr32 => word.rotr32(amount),
125 }
126 }
127}
128
129impl SerializeBytes for ShiftVariant {
130 fn serialize(&self, write_buf: impl BufMut) -> Result<(), SerializationError> {
131 (*self as u8).serialize(write_buf)
132 }
133}
134
135impl DeserializeBytes for ShiftVariant {
136 fn deserialize(read_buf: impl Buf) -> Result<Self, SerializationError>
137 where
138 Self: Sized,
139 {
140 let index = u8::deserialize(read_buf)?;
141 match index {
142 0 => Ok(ShiftVariant::Sll),
143 1 => Ok(ShiftVariant::Slr),
144 2 => Ok(ShiftVariant::Sar),
145 3 => Ok(ShiftVariant::Rotr),
146 4 => Ok(ShiftVariant::Sll32),
147 5 => Ok(ShiftVariant::Srl32),
148 6 => Ok(ShiftVariant::Sra32),
149 7 => Ok(ShiftVariant::Rotr32),
150 _ => Err(SerializationError::UnknownEnumVariant {
151 name: "ShiftVariant",
152 index,
153 }),
154 }
155 }
156}
157
158#[derive(Copy, Clone, Debug, Hash, Eq, PartialEq, Ord, PartialOrd)]
165pub struct ShiftedValueIndex {
166 pub value_index: ValueIndex,
168 pub shift_variant: ShiftVariant,
170 pub amount: u8,
175}
176
177impl ShiftedValueIndex {
178 pub const fn plain(value_index: ValueIndex) -> Self {
181 Self {
182 value_index,
183 shift_variant: ShiftVariant::Sll,
184 amount: 0,
185 }
186 }
187
188 pub fn sll(value_index: ValueIndex, amount: usize) -> Self {
193 assert!(amount < 64, "shift amount n={amount} out of range");
194 Self {
195 value_index,
196 shift_variant: ShiftVariant::Sll,
197 amount: amount as u8,
198 }
199 }
200
201 pub fn srl(value_index: ValueIndex, amount: usize) -> Self {
206 assert!(amount < 64, "shift amount n={amount} out of range");
207 Self {
208 value_index,
209 shift_variant: ShiftVariant::Slr,
210 amount: amount as u8,
211 }
212 }
213
214 pub fn sar(value_index: ValueIndex, amount: usize) -> Self {
222 assert!(amount < 64, "shift amount n={amount} out of range");
223 Self {
224 value_index,
225 shift_variant: ShiftVariant::Sar,
226 amount: amount as u8,
227 }
228 }
229
230 pub fn rotr(value_index: ValueIndex, amount: usize) -> Self {
237 assert!(amount < 64, "shift amount n={amount} out of range");
238 Self {
239 value_index,
240 shift_variant: ShiftVariant::Rotr,
241 amount: amount as u8,
242 }
243 }
244
245 pub fn sll32(value_index: ValueIndex, amount: usize) -> Self {
253 assert!(amount < 32, "shift amount n={amount} out of range for 32-bit shift");
254 Self {
255 value_index,
256 shift_variant: ShiftVariant::Sll32,
257 amount: amount as u8,
258 }
259 }
260
261 pub fn srl32(value_index: ValueIndex, amount: usize) -> Self {
269 assert!(amount < 32, "shift amount n={amount} out of range for 32-bit shift");
270 Self {
271 value_index,
272 shift_variant: ShiftVariant::Srl32,
273 amount: amount as u8,
274 }
275 }
276
277 pub fn sra32(value_index: ValueIndex, amount: usize) -> Self {
286 assert!(amount < 32, "shift amount n={amount} out of range for 32-bit shift");
287 Self {
288 value_index,
289 shift_variant: ShiftVariant::Sra32,
290 amount: amount as u8,
291 }
292 }
293
294 pub fn rotr32(value_index: ValueIndex, amount: usize) -> Self {
303 assert!(amount < 32, "shift amount n={amount} out of range for 32-bit rotate");
304 Self {
305 value_index,
306 shift_variant: ShiftVariant::Rotr32,
307 amount: amount as u8,
308 }
309 }
310
311 #[inline]
316 pub fn eval(&self, witness: &ValueVec) -> Word {
317 self.shift_variant
319 .apply(witness[self.value_index], self.amount as usize)
320 }
321}
322
323impl SerializeBytes for ShiftedValueIndex {
324 fn serialize(&self, mut write_buf: impl BufMut) -> Result<(), SerializationError> {
325 self.value_index.serialize(&mut write_buf)?;
326 self.shift_variant.serialize(&mut write_buf)?;
327 (self.amount as usize).serialize(write_buf)
329 }
330}
331
332impl DeserializeBytes for ShiftedValueIndex {
333 fn deserialize(mut read_buf: impl Buf) -> Result<Self, SerializationError>
334 where
335 Self: Sized,
336 {
337 let value_index = ValueIndex::deserialize(&mut read_buf)?;
338 let shift_variant = ShiftVariant::deserialize(&mut read_buf)?;
339 let amount = usize::deserialize(read_buf)?;
340
341 if amount >= shift_variant.max_amount() {
346 return Err(SerializationError::InvalidConstruction {
347 name: "ShiftedValueIndex::amount",
348 });
349 }
350
351 Ok(ShiftedValueIndex {
352 value_index,
353 shift_variant,
354 amount: amount as u8,
355 })
356 }
357}
358
359#[cfg(test)]
360mod tests {
361 use super::*;
362
363 #[test]
364 fn test_shift_variant_serialization_round_trip() {
365 let variants = [
366 ShiftVariant::Sll,
367 ShiftVariant::Slr,
368 ShiftVariant::Sar,
369 ShiftVariant::Rotr,
370 ];
371
372 for variant in variants {
373 let mut buf = Vec::new();
374 variant.serialize(&mut buf).unwrap();
375
376 let deserialized = ShiftVariant::deserialize(&mut buf.as_slice()).unwrap();
377 match (variant, deserialized) {
378 (ShiftVariant::Sll, ShiftVariant::Sll)
379 | (ShiftVariant::Slr, ShiftVariant::Slr)
380 | (ShiftVariant::Sar, ShiftVariant::Sar)
381 | (ShiftVariant::Rotr, ShiftVariant::Rotr) => {}
382 _ => panic!("ShiftVariant round trip failed: {:?} != {:?}", variant, deserialized),
383 }
384 }
385 }
386
387 #[test]
388 fn test_shift_variant_unknown_variant() {
389 let mut buf = Vec::new();
391 255u8.serialize(&mut buf).unwrap();
392
393 let result = ShiftVariant::deserialize(&mut buf.as_slice());
394 assert!(result.is_err());
395 match result.unwrap_err() {
396 SerializationError::UnknownEnumVariant { name, index } => {
397 assert_eq!(name, "ShiftVariant");
398 assert_eq!(index, 255);
399 }
400 _ => panic!("Expected UnknownEnumVariant error"),
401 }
402 }
403
404 #[test]
405 fn test_shifted_value_index_serialization_round_trip() {
406 let shifted_value_index = ShiftedValueIndex::srl(ValueIndex(42), 23);
407
408 let mut buf = Vec::new();
409 shifted_value_index.serialize(&mut buf).unwrap();
410
411 let deserialized = ShiftedValueIndex::deserialize(&mut buf.as_slice()).unwrap();
412 assert_eq!(shifted_value_index.value_index, deserialized.value_index);
413 assert_eq!(shifted_value_index.amount, deserialized.amount);
414 match (shifted_value_index.shift_variant, deserialized.shift_variant) {
415 (ShiftVariant::Slr, ShiftVariant::Slr) => {}
416 _ => panic!("ShiftVariant mismatch"),
417 }
418 }
419
420 #[test]
421 fn test_shifted_value_index_invalid_amount() {
422 let mut buf = Vec::new();
424 ValueIndex(0).serialize(&mut buf).unwrap();
425 ShiftVariant::Sll.serialize(&mut buf).unwrap();
426 64usize.serialize(&mut buf).unwrap(); let result = ShiftedValueIndex::deserialize(&mut buf.as_slice());
429 assert!(result.is_err());
430 match result.unwrap_err() {
431 SerializationError::InvalidConstruction { name } => {
432 assert_eq!(name, "ShiftedValueIndex::amount");
433 }
434 _ => panic!("Expected InvalidConstruction error"),
435 }
436 }
437
438 #[test]
439 fn test_max_amount_and_is_half_word() {
440 for variant in [
442 ShiftVariant::Sll,
443 ShiftVariant::Slr,
444 ShiftVariant::Sar,
445 ShiftVariant::Rotr,
446 ] {
447 assert!(!variant.is_half_word());
448 assert_eq!(variant.max_amount(), 64);
449 }
450 for variant in [
452 ShiftVariant::Sll32,
453 ShiftVariant::Srl32,
454 ShiftVariant::Sra32,
455 ShiftVariant::Rotr32,
456 ] {
457 assert!(variant.is_half_word());
458 assert_eq!(variant.max_amount(), 32);
459 }
460 }
461
462 fn deserialize_amount(
465 shift_variant: ShiftVariant,
466 amount: usize,
467 ) -> Result<ShiftedValueIndex, SerializationError> {
468 let mut buf = Vec::new();
469 ValueIndex(0).serialize(&mut buf).unwrap();
470 shift_variant.serialize(&mut buf).unwrap();
471 amount.serialize(&mut buf).unwrap();
472 ShiftedValueIndex::deserialize(&mut buf.as_slice())
473 }
474
475 #[test]
476 fn test_deserialize_rejects_half_word_amount_at_or_above_32() {
477 assert_eq!(
479 deserialize_amount(ShiftVariant::Sll32, 31).unwrap(),
480 ShiftedValueIndex {
481 value_index: ValueIndex(0),
482 shift_variant: ShiftVariant::Sll32,
483 amount: 31,
484 }
485 );
486 match deserialize_amount(ShiftVariant::Sll32, 32).unwrap_err() {
488 SerializationError::InvalidConstruction { name } => {
489 assert_eq!(name, "ShiftedValueIndex::amount");
490 }
491 other => panic!("Expected InvalidConstruction, got: {other:?}"),
492 }
493 assert_eq!(
495 deserialize_amount(ShiftVariant::Sll, 32).unwrap(),
496 ShiftedValueIndex {
497 value_index: ValueIndex(0),
498 shift_variant: ShiftVariant::Sll,
499 amount: 32,
500 }
501 );
502 assert_eq!(
503 deserialize_amount(ShiftVariant::Sll, 63).unwrap(),
504 ShiftedValueIndex {
505 value_index: ValueIndex(0),
506 shift_variant: ShiftVariant::Sll,
507 amount: 63,
508 }
509 );
510 }
511
512 #[test]
513 fn shifted_value_index_fits_in_a_word() {
514 assert_eq!(size_of::<ShiftedValueIndex>(), 8);
518 }
519
520 #[test]
521 fn test_shift_variant_from_u8_round_trip() {
522 let variants = [
524 ShiftVariant::Sll,
525 ShiftVariant::Slr,
526 ShiftVariant::Sar,
527 ShiftVariant::Rotr,
528 ShiftVariant::Sll32,
529 ShiftVariant::Srl32,
530 ShiftVariant::Sra32,
531 ShiftVariant::Rotr32,
532 ];
533 for variant in variants {
534 assert_eq!(ShiftVariant::from_u8(variant as u8), Some(variant));
535 }
536 assert_eq!(ShiftVariant::from_u8(8), None);
538 assert_eq!(ShiftVariant::from_u8(255), None);
539 }
540}