Skip to main content

binius_prover/protocols/shift/key_collection/
dense_shift_encoding.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use 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/// A dense re-encoding of the shift sequences occurring in a key segment.
12/// A key names a sequence of two shifts, each slot drawn from a fixed alphabet of 512 spellings.
13///
14/// - The sequence space is therefore 512^2 = 262,144 entries.
15/// - A constraint system uses only a few dozen of them.
16/// - The segment's own sequences are re-encoded as a contiguous range.
17/// - A per-sequence table is sized by the sequences present, not the whole space.
18/// - The space is too large for a lookup table, so a sequence is located by binary search.
19#[derive(Debug, Clone, Default)]
20pub struct DenseShiftEncoding {
21	/// The shift sequence each dense index encodes, in ascending sequence order.
22	///
23	/// Invariant: sorted and distinct.
24	/// This is what makes finding a sequence's index a binary search.
25	/// A deserialized encoding is checked for both properties.
26	shifts: Vec<[Shift; 2]>,
27}
28
29impl DenseShiftEncoding {
30	/// Builds the encoding of the shift sequences in an iterator, neither sorted nor distinct.
31	///
32	/// # Panics
33	///
34	/// Panics if the sequences do not fit the `u16` a key addresses them with.
35	/// That needs more than 65,536 distinct shift sequences in one segment.
36	/// No real system reaches that many.
37	pub(super) fn new(shifts: impl IntoIterator<Item = [Shift; 2]>) -> Self {
38		// A sorted set dedupes and orders in one pass, over the sequences present.
39		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	/// The number of distinct shift sequences the segment uses.
51	pub const fn len(&self) -> usize {
52		self.shifts.len()
53	}
54
55	/// Whether the segment uses no shifted values at all.
56	pub const fn is_empty(&self) -> bool {
57		self.shifts.is_empty()
58	}
59
60	/// The shift sequences the segment uses, in dense index order.
61	pub fn iter(&self) -> impl Iterator<Item = [Shift; 2]> + '_ {
62		self.shifts.iter().copied()
63	}
64
65	/// Where every sequence the segment uses sits in the space two shift slots span.
66	///
67	/// - The space is addressed outer-major: the outer slot's index sits above the inner slot's.
68	/// - This matches the reduction's round order, which peels the outer shift first.
69	/// - Distinct sequences land on distinct indices, so two segments' encodings can merge.
70	/// - The indices do not come out ascending, since sequences are sorted by inner slot first.
71	/// - Nothing needs them ascending.
72	///
73	/// ```text
74	/// index = (outer_index << LOG_SHIFT_COUNT) | inner_index
75	/// ```
76	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	/// The dense index of one shift sequence, for lookup while the keys are built.
83	/// The sequences are sorted, so this is a binary search over the ones the segment uses.
84	/// A lookup table over the whole sequence space would instead be 262,144 entries wide.
85	///
86	/// # Panics
87	///
88	/// Panics if the sequence is not one this encoding covers.
89	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		// `new` bounds the length, so every index it yields fits.
95		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		// Half-word (*32) variants cap at 32, full-width ones at 64.
127		// An amount past its variant's bound denotes no shift at all.
128		let amounts_in_range = shifts
129			.iter()
130			.flatten()
131			.all(|shift| (shift.amount as usize) < shift.variant.max_amount());
132		// A dense index is only meaningful against a list of distinct sequences, which the strictly
133		// ascending order this writes them in also gives.
134		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		// A key addresses a sequence with a `u16`, so a longer list could not be indexed.
141		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	/// A shift sequence carrying one shift, which the canonical form places in the inner slot.
158	fn single(shift: Shift) -> [Shift; 2] {
159		[shift, Shift::IDENTITY]
160	}
161
162	// Serializes an encoding built raw, bypassing `new`'s sorting and deduplication, so that a
163	// malformed list reaches the deserializer.
164	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		// A dense index means nothing against an unsorted list: `dense_idx` binary-searches it.
173		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		// Ascending order is checked strictly, so a repeat is rejected along with a swap: two equal
184		// sequences would give one shift sequence two dense indices.
185		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		// The bound is the variant's own: a half-word (*32) variant caps at 32, not at 64.
196		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		// Both slots are checked, so an outer amount its variant cannot represent is rejected too.
211		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		// Indexing inverts the dense index directly: no production caller needs a decode-by-index
226		// method, so the inverse is checked here against the private field instead of one.
227		// The input is unsorted and repeats one sequence, which also pins the sort and the dedupe.
228		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}