binius_prover/protocols/shift/key_collection/
dense_shift_encoding.rs1use std::collections::BTreeSet;
5
6use binius_core::{ShiftVariant, constraint_system::Shift};
7use binius_utils::serialization::{DeserializeBytes, SerializationError, SerializeBytes};
8use binius_verifier::protocols::shift::LOG_SHIFT_COUNT;
9use bytes::{Buf, BufMut};
10
11#[derive(Debug, Clone, Default)]
20pub struct DenseShiftEncoding {
21 shifts: Vec<[Shift; 2]>,
27}
28
29impl DenseShiftEncoding {
30 pub(super) fn new(shifts: impl IntoIterator<Item = [Shift; 2]>) -> Self {
38 let shifts = shifts.into_iter().collect::<BTreeSet<_>>();
40 assert!(
41 shifts.len() <= u16::MAX as usize + 1,
42 "a key segment uses {} distinct shift sequences, more than the u16 dense index addresses",
43 shifts.len()
44 );
45 Self {
46 shifts: shifts.into_iter().collect(),
47 }
48 }
49
50 pub const fn len(&self) -> usize {
52 self.shifts.len()
53 }
54
55 pub const fn is_empty(&self) -> bool {
57 self.shifts.is_empty()
58 }
59
60 pub fn iter(&self) -> impl Iterator<Item = [Shift; 2]> + '_ {
62 self.shifts.iter().copied()
63 }
64
65 pub fn shift_indices(&self) -> impl Iterator<Item = usize> + '_ {
77 self.shifts
78 .iter()
79 .map(|&[inner, outer]| outer.index() << LOG_SHIFT_COUNT | inner.index())
80 }
81
82 pub(super) fn dense_idx(&self, shift_seq: [Shift; 2]) -> u16 {
90 let index = self
91 .shifts
92 .binary_search(&shift_seq)
93 .expect("the encoding covers every shift sequence its segment's keys name");
94 index as u16
96 }
97}
98
99impl SerializeBytes for DenseShiftEncoding {
100 fn serialize(&self, mut write_buf: impl BufMut) -> Result<(), SerializationError> {
101 (self.shifts.len() as u32).serialize(&mut write_buf)?;
102 for shift_seq in &self.shifts {
103 for shift in shift_seq {
104 shift.variant.serialize(&mut write_buf)?;
105 shift.amount.serialize(&mut write_buf)?;
106 }
107 }
108 Ok(())
109 }
110}
111
112impl DeserializeBytes for DenseShiftEncoding {
113 fn deserialize(mut read_buf: impl Buf) -> Result<Self, SerializationError> {
114 let len = u32::deserialize(&mut read_buf)? as usize;
115 let mut shifts = Vec::with_capacity(len);
116 for _ in 0..len {
117 let mut shift_seq = [Shift::IDENTITY; 2];
118 for shift in &mut shift_seq {
119 let variant = ShiftVariant::deserialize(&mut read_buf)?;
120 let amount = u8::deserialize(&mut read_buf)?;
121 *shift = Shift { variant, amount };
122 }
123 shifts.push(shift_seq);
124 }
125
126 let amounts_in_range = shifts
129 .iter()
130 .flatten()
131 .all(|shift| (shift.amount as usize) < shift.variant.max_amount());
132 let strictly_ascending = shifts.windows(2).all(|window| window[0] < window[1]);
135 if !amounts_in_range || !strictly_ascending {
136 return Err(SerializationError::InvalidConstruction {
137 name: "DenseShiftEncoding::shifts",
138 });
139 }
140 if len > u16::MAX as usize + 1 {
142 return Err(SerializationError::InvalidConstruction {
143 name: "DenseShiftEncoding::shifts",
144 });
145 }
146
147 Ok(DenseShiftEncoding { shifts })
148 }
149}
150
151#[cfg(test)]
152mod tests {
153 use binius_core::word::Word;
154
155 use super::*;
156
157 fn single(shift: Shift) -> [Shift; 2] {
159 [shift, Shift::IDENTITY]
160 }
161
162 fn deserialize_raw(shifts: Vec<[Shift; 2]>) -> Result<DenseShiftEncoding, SerializationError> {
165 let mut buf = Vec::new();
166 DenseShiftEncoding { shifts }.serialize(&mut buf).unwrap();
167 DenseShiftEncoding::deserialize(buf.as_slice())
168 }
169
170 #[test]
171 fn dense_shift_encoding_rejects_an_unordered_serialization() {
172 match deserialize_raw(vec![single(Shift::srl(3)), single(Shift::IDENTITY)]).unwrap_err() {
174 SerializationError::InvalidConstruction { name } => {
175 assert_eq!(name, "DenseShiftEncoding::shifts");
176 }
177 other => panic!("Expected InvalidConstruction, got: {other:?}"),
178 }
179 }
180
181 #[test]
182 fn dense_shift_encoding_rejects_a_repeated_sequence() {
183 match deserialize_raw(vec![single(Shift::srl(3)), single(Shift::srl(3))]).unwrap_err() {
186 SerializationError::InvalidConstruction { name } => {
187 assert_eq!(name, "DenseShiftEncoding::shifts");
188 }
189 other => panic!("Expected InvalidConstruction, got: {other:?}"),
190 }
191 }
192
193 #[test]
194 fn dense_shift_encoding_rejects_an_out_of_range_shift_amount() {
195 let out_of_range = Shift {
197 variant: ShiftVariant::Sll32,
198 amount: 32,
199 };
200 match deserialize_raw(vec![single(out_of_range)]).unwrap_err() {
201 SerializationError::InvalidConstruction { name } => {
202 assert_eq!(name, "DenseShiftEncoding::shifts");
203 }
204 other => panic!("Expected InvalidConstruction, got: {other:?}"),
205 }
206 }
207
208 #[test]
209 fn dense_shift_encoding_rejects_an_out_of_range_outer_shift_amount() {
210 let out_of_range = Shift {
212 variant: ShiftVariant::Sll,
213 amount: Word::BITS as u8,
214 };
215 match deserialize_raw(vec![[Shift::srl(3), out_of_range]]).unwrap_err() {
216 SerializationError::InvalidConstruction { name } => {
217 assert_eq!(name, "DenseShiftEncoding::shifts");
218 }
219 other => panic!("Expected InvalidConstruction, got: {other:?}"),
220 }
221 }
222
223 #[test]
224 fn dense_shift_encoding_indexes_every_sequence_it_covers() {
225 let sequences = [
229 [Shift::srl(3), Shift::sll(3)],
230 single(Shift::rotr(1)),
231 single(Shift::IDENTITY),
232 single(Shift::rotr(1)),
233 ];
234 let encoding = DenseShiftEncoding::new(sequences);
235
236 assert_eq!(encoding.len(), 3);
237 for dense_idx in 0..encoding.len() {
238 let sequence = encoding.shifts[dense_idx];
239 assert_eq!(encoding.dense_idx(sequence) as usize, dense_idx);
240 }
241 }
242}