1use std::ops::Deref;
11
12use binius_field::{BinaryField, BinaryField1b};
13
14#[derive(Debug, Clone, PartialEq, Eq)]
19pub struct BinarySubspace<F, Data: Deref<Target = [F]> = Vec<F>> {
20 basis: Data,
22}
23
24impl<F: BinaryField, Data: Deref<Target = [F]>> BinarySubspace<F, Data> {
25 pub const fn new_unchecked(basis: Data) -> Self {
29 Self { basis }
30 }
31
32 pub fn isomorphic<FIso>(&self) -> BinarySubspace<FIso>
36 where
37 FIso: BinaryField + From<F>,
38 {
39 BinarySubspace {
40 basis: self.basis.iter().copied().map(Into::into).collect(),
42 }
43 }
44
45 pub fn dim(&self) -> usize {
47 self.basis.len()
48 }
49
50 pub fn basis(&self) -> &[F] {
52 &self.basis
53 }
54
55 pub fn get(&self, index: usize) -> F {
67 assert!(
70 self.dim() >= usize::BITS as usize || index < 1 << self.dim(),
71 "precondition: index must be less than 2^dim"
72 );
73
74 element_at(&self.basis, index)
75 }
76
77 pub fn iter(&self) -> BinarySubspaceIterator<'_, F> {
83 BinarySubspaceIterator::new(&self.basis)
84 }
85}
86
87impl<F: BinaryField> BinarySubspace<F> {
88 pub fn with_dim(dim: usize) -> Self {
96 assert!(dim <= F::DEGREE, "precondition: dim must be at most F::DEGREE");
97
98 let basis = (0..dim).map(|i| F::basis(i)).collect();
100 Self { basis }
101 }
102
103 pub fn reduce_dim(&self, dim: usize) -> Self {
108 assert!(dim <= self.dim(), "precondition: dim must be at most this subspace's dimension");
109
110 Self {
111 basis: self.basis[..dim].to_vec(),
112 }
113 }
114}
115
116fn element_at<F: BinaryField>(basis: &[F], index: usize) -> F {
118 basis
119 .iter()
120 .take(usize::BITS as usize)
123 .enumerate()
124 .map(|(i, &basis_i)| basis_i * BinaryField1b::from((index >> i) & 1 == 1))
127 .sum()
128}
129
130#[derive(Debug, Clone)]
136pub struct BinarySubspaceIterator<'a, F> {
137 basis: &'a [F],
139 index: usize,
141 next: Option<F>,
143}
144
145impl<'a, F: BinaryField> BinarySubspaceIterator<'a, F> {
146 pub fn new(basis: &'a [F]) -> Self {
152 assert!(basis.len() < usize::BITS as usize);
153
154 Self {
156 basis,
157 index: 0,
158 next: Some(F::ZERO),
159 }
160 }
161}
162
163impl<'a, F: BinaryField> Iterator for BinarySubspaceIterator<'a, F> {
164 type Item = F;
165
166 #[inline]
176 fn next(&mut self) -> Option<Self::Item> {
177 let ret = self.next?;
178
179 let ones = self.index.trailing_ones() as usize;
182
183 let mut next = ret;
185 for &basis_i in &self.basis[..ones] {
186 next -= basis_i;
187 }
188
189 self.next = self.basis.get(ones).map(|&basis_i| next + basis_i);
192
193 self.index += 1;
194 Some(ret)
195 }
196
197 fn size_hint(&self) -> (usize, Option<usize>) {
198 let last = 1 << self.basis.len();
201 let remaining = last - self.index;
202 (remaining, Some(remaining))
203 }
204
205 fn nth(&mut self, n: usize) -> Option<Self::Item> {
211 match self.index.checked_add(n) {
212 Some(new_index) if new_index < 1 << self.basis.len() => {
214 self.index = new_index;
215 self.next = Some(element_at(self.basis, new_index));
216 }
217 _ => {
219 self.index = 1 << self.basis.len();
220 self.next = None;
221 }
222 }
223
224 self.next()
225 }
226}
227
228impl<'a, F: BinaryField> ExactSizeIterator for BinarySubspaceIterator<'a, F> {
229 fn len(&self) -> usize {
230 let last = 1 << self.basis.len();
232 last - self.index
233 }
234}
235
236impl<'a, F: BinaryField> std::iter::FusedIterator for BinarySubspaceIterator<'a, F> {}
237
238impl<F: BinaryField> Default for BinarySubspace<F> {
239 fn default() -> Self {
241 let basis = (0..F::DEGREE).map(|i| F::basis(i)).collect();
243 Self { basis }
244 }
245}
246
247#[cfg(test)]
248mod tests {
249 use binius_field::{ExtensionField, Field, Ghash128b as B128, Rijndael8b as B8};
250
251 use super::*;
252
253 #[test]
254 fn test_default_binary_subspace_iterates_elements() {
255 let subspace = BinarySubspace::<B8>::default();
258 for i in 0..=255 {
259 assert_eq!(subspace.get(i), B8::new(i as u8));
260 }
261 }
262
263 #[test]
264 fn test_get_on_a_subspace_wider_than_a_usize() {
265 let basis = <B128 as ExtensionField<BinaryField1b>>::basis;
268 let subspace = BinarySubspace::<B128>::default();
269 assert_eq!(subspace.dim(), 128);
270 assert_eq!(subspace.get(0), B128::ZERO);
271 assert_eq!(subspace.get(1), basis(0));
272 assert_eq!(subspace.get(5), basis(0) + basis(2));
273 let low_bits: B128 = (0..usize::BITS as usize).map(basis).sum();
275 assert_eq!(subspace.get(usize::MAX), low_bits);
276 }
277
278 #[test]
279 #[should_panic(expected = "precondition")]
280 fn test_binary_subspace_range_error() {
281 let subspace = BinarySubspace::<B8>::default();
283 let _ = subspace.get(256);
284 }
285
286 #[test]
287 fn test_default_binary_subspace() {
288 let subspace = BinarySubspace::<B8>::default();
289 assert_eq!(subspace.dim(), 8);
290 assert_eq!(subspace.basis().len(), 8);
291
292 assert_eq!(
294 subspace.basis(),
295 [
296 B8::new(0b00000001),
297 B8::new(0b00000010),
298 B8::new(0b00000100),
299 B8::new(0b00001000),
300 B8::new(0b00010000),
301 B8::new(0b00100000),
302 B8::new(0b01000000),
303 B8::new(0b10000000)
304 ]
305 );
306
307 let expected_elements: [u8; 256] = (0..=255).collect::<Vec<_>>().try_into().unwrap();
309
310 for (i, &expected) in expected_elements.iter().enumerate() {
311 assert_eq!(subspace.get(i), B8::new(expected));
312 }
313 }
314
315 #[test]
316 fn test_with_dim_valid() {
317 let subspace = BinarySubspace::<B8>::with_dim(3);
319 assert_eq!(subspace.dim(), 3);
320 assert_eq!(subspace.basis().len(), 3);
321
322 assert_eq!(subspace.basis(), [B8::new(0b001), B8::new(0b010), B8::new(0b100)]);
323
324 let expected_elements: [u8; 8] = [0b000, 0b001, 0b010, 0b011, 0b100, 0b101, 0b110, 0b111];
326
327 for (i, &expected) in expected_elements.iter().enumerate() {
328 assert_eq!(subspace.get(i), B8::new(expected));
329 }
330 }
331
332 #[test]
333 #[should_panic(expected = "precondition")]
334 fn test_with_dim_invalid() {
335 let _ = BinarySubspace::<B8>::with_dim(10);
337 }
338
339 #[test]
340 fn test_reduce_dim_valid() {
341 let subspace = BinarySubspace::<B8>::with_dim(6);
343 let reduced = subspace.reduce_dim(4);
344 assert_eq!(reduced.dim(), 4);
345 assert_eq!(reduced.basis().len(), 4);
346
347 assert_eq!(
349 reduced.basis(),
350 [
351 B8::new(0b0001),
352 B8::new(0b0010),
353 B8::new(0b0100),
354 B8::new(0b1000)
355 ]
356 );
357
358 let expected_elements: [u8; 16] = (0..16).collect::<Vec<_>>().try_into().unwrap();
359
360 for (i, &expected) in expected_elements.iter().enumerate() {
361 assert_eq!(reduced.get(i), B8::new(expected));
362 }
363 }
364
365 #[test]
366 #[should_panic(expected = "precondition")]
367 fn test_reduce_dim_invalid() {
368 let subspace = BinarySubspace::<B8>::with_dim(4);
370 let _ = subspace.reduce_dim(6);
371 }
372
373 #[test]
374 fn test_isomorphic_conversion() {
375 let subspace = BinarySubspace::<B8>::with_dim(3);
376 let iso_subspace: BinarySubspace<B128> = subspace.isomorphic();
378 assert_eq!(iso_subspace.dim(), 3);
379 assert_eq!(iso_subspace.basis().len(), 3);
380
381 assert_eq!(
383 iso_subspace.basis(),
384 [
385 B128::from(B8::new(0b001)),
386 B128::from(B8::new(0b010)),
387 B128::from(B8::new(0b100)),
388 ]
389 );
390 }
391
392 #[test]
393 fn test_iterate_subspace() {
394 let subspace = BinarySubspace::<B8>::with_dim(3);
395 let elements: Vec<_> = subspace.iter().collect();
397 assert_eq!(elements.len(), 8);
398
399 let expected_elements: [u8; 8] = [0b000, 0b001, 0b010, 0b011, 0b100, 0b101, 0b110, 0b111];
400
401 for (i, &expected) in expected_elements.iter().enumerate() {
402 assert_eq!(elements[i], B8::new(expected));
403 }
404 }
405
406 #[test]
407 fn test_iterator_matches_get() {
408 let subspace = BinarySubspace::<B8>::with_dim(5);
409
410 for (i, elem) in subspace.iter().enumerate() {
412 assert_eq!(elem, subspace.get(i), "Mismatch at index {}", i);
413 }
414 }
415
416 #[test]
417 #[allow(clippy::iter_nth_zero)]
418 fn test_iterator_nth() {
419 let subspace = BinarySubspace::<B8>::with_dim(4);
420
421 let mut iter = subspace.iter();
423 assert_eq!(iter.nth(0), Some(subspace.get(0)));
424 assert_eq!(iter.nth(0), Some(subspace.get(1)));
425 assert_eq!(iter.nth(2), Some(subspace.get(4)));
427 assert_eq!(iter.nth(5), Some(subspace.get(10)));
428
429 let mut iter = subspace.iter();
431 assert_eq!(iter.nth(15), Some(subspace.get(15)));
432 assert_eq!(iter.nth(0), None);
434 }
435
436 #[test]
437 fn test_iterator_nth_skips_efficiently() {
438 let subspace = BinarySubspace::<B8>::with_dim(6);
439
440 let mut iter = subspace.iter();
442 assert_eq!(iter.nth(30), Some(subspace.get(30)));
443 assert_eq!(iter.next(), Some(subspace.get(31)));
445
446 let mut iter = subspace.iter();
448 assert_eq!(iter.nth(50), Some(subspace.get(50)));
449 }
450
451 #[test]
452 fn test_iterator_size_hint() {
453 let subspace = BinarySubspace::<B8>::with_dim(3);
454 let mut iter = subspace.iter();
455
456 assert_eq!(iter.size_hint(), (8, Some(8)));
458 iter.next();
459 assert_eq!(iter.size_hint(), (7, Some(7)));
460 iter.nth(3);
462 assert_eq!(iter.size_hint(), (3, Some(3)));
463 }
464
465 #[test]
466 fn test_iterator_exact_size() {
467 let subspace = BinarySubspace::<B8>::with_dim(4);
468 let mut iter = subspace.iter();
469
470 assert_eq!(iter.len(), 16);
471 iter.next();
472 assert_eq!(iter.len(), 15);
473 iter.nth(5);
474 assert_eq!(iter.len(), 9);
475 }
476
477 #[test]
478 fn test_iterator_empty_subspace() {
479 let subspace = BinarySubspace::<B8>::with_dim(0);
481 let mut iter = subspace.iter();
482
483 assert_eq!(iter.len(), 1);
484 assert_eq!(iter.next(), Some(B8::ZERO));
485 assert_eq!(iter.next(), None);
486 }
487
488 #[test]
489 fn test_iterator_full_iteration() {
490 let subspace = BinarySubspace::<B8>::default();
493 let collected: Vec<_> = subspace.iter().collect();
494
495 assert_eq!(collected.len(), 256);
496 for (i, elem) in collected.iter().enumerate() {
497 assert_eq!(*elem, subspace.get(i));
498 }
499 }
500
501 #[test]
502 fn test_iterator_partial_then_nth() {
503 let subspace = BinarySubspace::<B8>::with_dim(5);
504 let mut iter = subspace.iter();
505
506 assert_eq!(iter.next(), Some(subspace.get(0)));
508 assert_eq!(iter.next(), Some(subspace.get(1)));
509 assert_eq!(iter.next(), Some(subspace.get(2)));
510
511 assert_eq!(iter.nth(5), Some(subspace.get(8)));
513 assert_eq!(iter.next(), Some(subspace.get(9)));
514 }
515
516 #[test]
517 fn test_iterator_clone() {
518 let subspace = BinarySubspace::<B8>::with_dim(3);
519 let mut iter1 = subspace.iter();
520
521 iter1.next();
522 iter1.next();
523
524 let mut iter2 = iter1.clone();
526
527 assert_eq!(iter1.next(), iter2.next());
528 assert_eq!(iter1.collect::<Vec<_>>(), iter2.collect::<Vec<_>>());
529 }
530}