Skip to main content

core_utils/preprocessing/
iterator.rs

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
22/// An iterator containing preprocessing futures for every gate in a circuit.
23pub struct PreprocessingIterator<C: Curve> {
24    // Base field iterators
25    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    // Binary field iterators
30    pub binary_singlets: NextVecIterator<Singlet<Gf2_128>, AbortError>,
31    pub binary_triples: NextVecIterator<Triple<Gf2_128>, AbortError>,
32
33    // Mersenne107 iterators
34    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    // Scalar field iterators
39    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}