Skip to main content

core_utils/preprocessing/
iterator.rs

1use primitives::{
2    algebra::field::binary::Gf2_128,
3    correlated_randomness::{
4        bundler::{errors::BundlerError, BundleIterator},
5        dabits::DaBit,
6        singlets::Singlet,
7        stream::{Next, NextVec, NextVecIterator},
8    },
9    sharing::{Malicious, ThreatModel},
10    types::identifiers::Named,
11};
12
13use crate::{
14    circuit::preprocessing::CircuitPreprocessing,
15    config::{BaseFieldOf, MpcConfig, MpcFieldOf, ScalarFieldOf},
16    errors::AbortError,
17    preprocessing::{NextDaBit, NextDaBits, NextSinglet, NextSinglets, PreprocessingKind},
18};
19
20/// An iterator containing preprocessing futures for every gate in a circuit.
21pub struct PreprocessingIterator<C: MpcConfig, M: ThreatModel = Malicious> {
22    // Base field iterators
23    pub base_field_dabits: NextVecIterator<DaBit<BaseFieldOf<C>>, AbortError>,
24    pub base_field_singlets: NextVecIterator<Singlet<BaseFieldOf<C>>, AbortError>,
25    pub base_field_triples: NextVecIterator<M::Triple<BaseFieldOf<C>>, AbortError>,
26
27    // Binary field iterators
28    pub binary_singlets: NextVecIterator<Singlet<Gf2_128>, AbortError>,
29    pub binary_triples: NextVecIterator<M::Triple<Gf2_128>, AbortError>,
30
31    // MpcFieldOf<C> iterators
32    pub mpc_field_dabits: NextVecIterator<DaBit<MpcFieldOf<C>>, AbortError>,
33    pub mpc_field_singlets: NextVecIterator<Singlet<MpcFieldOf<C>>, AbortError>,
34    pub mpc_field_triples: NextVecIterator<M::Triple<MpcFieldOf<C>>, AbortError>,
35
36    // Scalar field iterators
37    pub scalar_dabits: NextVecIterator<DaBit<ScalarFieldOf<C>>, AbortError>,
38    pub scalar_singlets: NextVecIterator<Singlet<ScalarFieldOf<C>>, AbortError>,
39    pub scalar_triples: NextVecIterator<M::Triple<ScalarFieldOf<C>>, AbortError>,
40}
41
42impl<C: MpcConfig, M: ThreatModel> BundleIterator for PreprocessingIterator<C, M> {
43    type Size = CircuitPreprocessing;
44    type Error = AbortError;
45
46    fn len(&self) -> CircuitPreprocessing {
47        let mut len = CircuitPreprocessing::default();
48        len[PreprocessingKind::BitSinglets] = self.binary_singlets.len();
49        len[PreprocessingKind::BitTriples] = self.binary_triples.len();
50        len[PreprocessingKind::BaseFieldSinglets] = self.base_field_singlets.len();
51        len[PreprocessingKind::BaseFieldTriples] = self.base_field_triples.len();
52        len[PreprocessingKind::BaseFieldDaBits] = self.base_field_dabits.len();
53        len[PreprocessingKind::ScalarSinglets] = self.scalar_singlets.len();
54        len[PreprocessingKind::ScalarTriples] = self.scalar_triples.len();
55        len[PreprocessingKind::ScalarDaBits] = self.scalar_dabits.len();
56        len[PreprocessingKind::MpcFieldSinglets] = self.mpc_field_singlets.len();
57        len[PreprocessingKind::MpcFieldTriples] = self.mpc_field_triples.len();
58        len[PreprocessingKind::MpcFieldDaBits] = self.mpc_field_dabits.len();
59        len
60    }
61}
62
63macro_rules! next_preprocessing_item {
64    ($( ($fn:ident, $iter:ident, $ret:ty, $ty:ty, $kind:literal) ),* $(,)?) => {
65        $(
66            pub fn $fn(&mut self) -> Result<$ret, BundlerError> {
67                self.$iter.next().ok_or_else(|| {
68                    BundlerError::InsufficientPreprocessing(
69                        format!("{} {}", <$ty>::get_name(), $kind),
70                    )
71                })
72            }
73        )*
74    };
75}
76
77macro_rules! next_n_preprocessing_items {
78    ($( ($fn:ident, $iter:ident, $ret:ty, $ty:ty, $kind:literal) ),* $(,)?) => {
79        $(
80            pub fn $fn(&mut self, n: usize) -> Result<$ret, BundlerError> {
81                self.$iter.next_n(n).ok_or_else(|| {
82                    BundlerError::InsufficientPreprocessing(
83                        format!("{} {} (requested {}, available {})", <$ty>::get_name(), $kind, n, self.$iter.len()),
84                    )
85                })
86            }
87        )*
88    };
89}
90
91impl<C: MpcConfig, M: ThreatModel> PreprocessingIterator<C, M> {
92    next_preprocessing_item!(
93        (next_base_field_dabit, base_field_dabits, NextDaBit<BaseFieldOf<C>>, BaseFieldOf<C>, "dabits"),
94        (next_base_field_singlet, base_field_singlets, NextSinglet<BaseFieldOf<C>>, BaseFieldOf<C>, "singlets"),
95        (next_base_field_triple, base_field_triples, Next<M::Triple<BaseFieldOf<C>>, AbortError>, BaseFieldOf<C>, "triples"),
96        (next_bit_singlet, binary_singlets, NextSinglet<Gf2_128>, Gf2_128, "singlets"),
97        (next_bit_triple, binary_triples, Next<M::Triple<Gf2_128>, AbortError>, Gf2_128, "triples"),
98        (next_mpc_field_dabit, mpc_field_dabits, NextDaBit<MpcFieldOf<C>>, MpcFieldOf<C>, "dabits"),
99        (next_mpc_field_singlet, mpc_field_singlets, NextSinglet<MpcFieldOf<C>>, MpcFieldOf<C>, "singlets"),
100        (next_mpc_field_triple, mpc_field_triples, Next<M::Triple<MpcFieldOf<C>>, AbortError>, MpcFieldOf<C>, "triples"),
101        (next_scalar_dabit, scalar_dabits, NextDaBit<ScalarFieldOf<C>>, ScalarFieldOf<C>, "dabits"),
102        (next_scalar_singlet, scalar_singlets, NextSinglet<ScalarFieldOf<C>>, ScalarFieldOf<C>, "singlets"),
103        (next_scalar_triple, scalar_triples, Next<M::Triple<ScalarFieldOf<C>>, AbortError>, ScalarFieldOf<C>, "triples"),
104    );
105
106    next_n_preprocessing_items!(
107        (next_n_base_field_dabits, base_field_dabits, NextDaBits<BaseFieldOf<C>>, BaseFieldOf<C>, "dabits"),
108        (next_n_base_field_singlets, base_field_singlets, NextSinglets<BaseFieldOf<C>>, BaseFieldOf<C>, "singlets"),
109        (next_n_base_field_triples, base_field_triples, NextVec<M::Triple<BaseFieldOf<C>>, AbortError>, BaseFieldOf<C>, "triples"),
110        (next_n_bit_singlets, binary_singlets, NextSinglets<Gf2_128>, Gf2_128, "singlets"),
111        (next_n_bit_triples, binary_triples, NextVec<M::Triple<Gf2_128>, AbortError>, Gf2_128, "triples"),
112        (next_n_mpc_field_dabits, mpc_field_dabits, NextDaBits<MpcFieldOf<C>>, MpcFieldOf<C>, "dabits"),
113        (next_n_mpc_field_singlets, mpc_field_singlets, NextSinglets<MpcFieldOf<C>>, MpcFieldOf<C>, "singlets"),
114        (next_n_mpc_field_triples, mpc_field_triples, NextVec<M::Triple<MpcFieldOf<C>>, AbortError>, MpcFieldOf<C>, "triples"),
115        (next_n_scalar_dabits, scalar_dabits, NextDaBits<ScalarFieldOf<C>>, ScalarFieldOf<C>, "dabits"),
116        (next_n_scalar_singlets, scalar_singlets, NextSinglets<ScalarFieldOf<C>>, ScalarFieldOf<C>, "singlets"),
117        (next_n_scalar_triples, scalar_triples, NextVec<M::Triple<ScalarFieldOf<C>>, AbortError>, ScalarFieldOf<C>, "triples"),
118    );
119}
120
121impl<C: MpcConfig, M: ThreatModel> std::fmt::Debug for PreprocessingIterator<C, M> {
122    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
123        write!(f, "PreprocessingIterator with len: {:?}", self.len())
124    }
125}