1use primitives::{
2 algebra::{
3 elliptic_curve::Curve,
4 field::{binary::Gf2_128, mersenne::Mersenne107},
5 },
6 correlated_randomness::{
7 bundler::{errors::BundlerError, BundleIterator},
8 dabits::DaBit,
9 singlets::Singlet,
10 stream::NextVecIterator,
11 triples::Triple,
12 },
13 types::identifiers::Named,
14};
15
16use crate::{
17 circuit::preprocessing::{CircuitPreprocessing, FieldCircuitPreprocessing},
18 errors::AbortError,
19 preprocessing::{NextDaBit, NextDaBits, NextSinglet, NextSinglets, NextTriple, NextTriples},
20};
21
22pub struct PreprocessingIterator<C: Curve> {
24 pub base_field_dabits: NextVecIterator<DaBit<C::BaseField>, AbortError>,
26 pub base_field_singlets: NextVecIterator<Singlet<C::BaseField>, AbortError>,
27 pub base_field_triples: NextVecIterator<Triple<C::BaseField>, AbortError>,
28
29 pub binary_singlets: NextVecIterator<Singlet<Gf2_128>, AbortError>,
31 pub binary_triples: NextVecIterator<Triple<Gf2_128>, AbortError>,
32
33 pub mersenne107_dabits: NextVecIterator<DaBit<Mersenne107>, AbortError>,
35 pub mersenne107_singlets: NextVecIterator<Singlet<Mersenne107>, AbortError>,
36 pub mersenne107_triples: NextVecIterator<Triple<Mersenne107>, AbortError>,
37
38 pub scalar_dabits: NextVecIterator<DaBit<C::Scalar>, AbortError>,
40 pub scalar_singlets: NextVecIterator<Singlet<C::Scalar>, AbortError>,
41 pub scalar_triples: NextVecIterator<Triple<C::Scalar>, AbortError>,
42}
43
44impl<C: Curve> BundleIterator for PreprocessingIterator<C> {
45 type Size = CircuitPreprocessing;
46
47 fn len(&self) -> CircuitPreprocessing {
48 CircuitPreprocessing {
49 bit_singlets: self.binary_singlets.len(),
50 bit_triples: self.binary_triples.len(),
51 base_field: FieldCircuitPreprocessing {
52 singlets: self.base_field_singlets.len(),
53 triples: self.base_field_triples.len(),
54 dabits: self.base_field_dabits.len(),
55 },
56 scalar: FieldCircuitPreprocessing {
57 singlets: self.scalar_singlets.len(),
58 triples: self.scalar_triples.len(),
59 dabits: self.scalar_dabits.len(),
60 },
61 mersenne107: FieldCircuitPreprocessing {
62 singlets: self.mersenne107_singlets.len(),
63 triples: self.mersenne107_triples.len(),
64 dabits: self.mersenne107_dabits.len(),
65 },
66 }
67 }
68 fn is_empty(&self) -> bool {
69 self.base_field_dabits.len() == 0
70 && self.base_field_singlets.len() == 0
71 && self.base_field_triples.len() == 0
72 && self.binary_singlets.len() == 0
73 && self.binary_triples.len() == 0
74 && self.mersenne107_dabits.len() == 0
75 && self.mersenne107_singlets.len() == 0
76 && self.mersenne107_triples.len() == 0
77 && self.scalar_dabits.len() == 0
78 && self.scalar_singlets.len() == 0
79 && self.scalar_triples.len() == 0
80 }
81}
82
83macro_rules! next_preprocessing_item {
84 ($( ($fn:ident, $iter:ident, $ret:ty, $ty:ty, $kind:literal) ),* $(,)?) => {
85 $(
86 pub fn $fn(&mut self) -> Result<$ret, BundlerError> {
87 self.$iter.next().ok_or_else(|| {
88 BundlerError::InsufficientPreprocessing(
89 format!("{} {}", <$ty>::get_name(), $kind),
90 )
91 })
92 }
93 )*
94 };
95}
96
97macro_rules! next_n_preprocessing_items {
98 ($( ($fn:ident, $iter:ident, $ret:ty, $ty:ty, $kind:literal) ),* $(,)?) => {
99 $(
100 pub fn $fn(&mut self, n: usize) -> Result<$ret, BundlerError> {
101 self.$iter.next_n(n).ok_or_else(|| {
102 BundlerError::InsufficientPreprocessing(
103 format!("{} {} (requested {}, available {})", <$ty>::get_name(), $kind, n, self.$iter.len()),
104 )
105 })
106 }
107 )*
108 };
109}
110
111impl<C: Curve> PreprocessingIterator<C> {
112 next_preprocessing_item!(
113 (next_base_field_dabit, base_field_dabits, NextDaBit<C::BaseField>, C::BaseField, "dabits"),
114 (next_base_field_singlet, base_field_singlets, NextSinglet<C::BaseField>, C::BaseField, "singlets"),
115 (next_base_field_triple, base_field_triples, NextTriple<C::BaseField>, C::BaseField, "triples"),
116 (next_bit_singlet, binary_singlets, NextSinglet<Gf2_128>, Gf2_128, "singlets"),
117 (next_bit_triple, binary_triples, NextTriple<Gf2_128>, Gf2_128, "triples"),
118 (next_mersenne107_dabit, mersenne107_dabits, NextDaBit<Mersenne107>, Mersenne107, "dabits"),
119 (next_mersenne107_singlet, mersenne107_singlets, NextSinglet<Mersenne107>, Mersenne107, "singlets"),
120 (next_mersenne107_triple, mersenne107_triples, NextTriple<Mersenne107>, Mersenne107, "triples"),
121 (next_scalar_dabit, scalar_dabits, NextDaBit<C::Scalar>, C::Scalar, "dabits"),
122 (next_scalar_singlet, scalar_singlets, NextSinglet<C::Scalar>, C::Scalar, "singlets"),
123 (next_scalar_triple, scalar_triples, NextTriple<C::Scalar>, C::Scalar, "triples"),
124 );
125
126 next_n_preprocessing_items!(
127 (next_n_base_field_dabits, base_field_dabits, NextDaBits<C::BaseField>, C::BaseField, "dabits"),
128 (next_n_base_field_singlets, base_field_singlets, NextSinglets<C::BaseField>, C::BaseField, "singlets"),
129 (next_n_base_field_triples, base_field_triples, NextTriples<C::BaseField>, C::BaseField, "triples"),
130 (next_n_bit_singlets, binary_singlets, NextSinglets<Gf2_128>, Gf2_128, "singlets"),
131 (next_n_bit_triples, binary_triples, NextTriples<Gf2_128>, Gf2_128, "triples"),
132 (next_n_mersenne107_dabits, mersenne107_dabits, NextDaBits<Mersenne107>, Mersenne107, "dabits"),
133 (next_n_mersenne107_singlets, mersenne107_singlets, NextSinglets<Mersenne107>, Mersenne107, "singlets"),
134 (next_n_mersenne107_triples, mersenne107_triples, NextTriples<Mersenne107>, Mersenne107, "triples"),
135 (next_n_scalar_dabits, scalar_dabits, NextDaBits<C::Scalar>, C::Scalar, "dabits"),
136 (next_n_scalar_singlets, scalar_singlets, NextSinglets<C::Scalar>, C::Scalar, "singlets"),
137 (next_n_scalar_triples, scalar_triples, NextTriples<C::Scalar>, C::Scalar, "triples"),
138 );
139}
140
141impl<C: Curve> std::fmt::Debug for PreprocessingIterator<C> {
142 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
143 write!(f, "PreprocessingIterator with len: {:?}", self.len())
144 }
145}