Skip to main content

core_utils/preprocessing/
bundler.rs

1use std::marker::PhantomData;
2
3use primitives::{
4    algebra::{
5        elliptic_curve::Curve,
6        field::{binary::Gf2_128, mersenne::Mersenne107},
7    },
8    correlated_randomness::{
9        bundler::{errors::BundlerError, Bundler},
10        dabits::DaBit,
11        singlets::Singlet,
12        stream::CorrelatedStream,
13        triples::Triple,
14    },
15};
16
17use crate::{
18    circuit::preprocessing::CircuitPreprocessing,
19    errors::AbortError,
20    preprocessing::iterator::PreprocessingIterator,
21};
22
23/// Stream bundler, holding one stream per preprocessing type, to provide all preprocessing for a
24/// circuit via its streams.
25pub struct StreamBundler<
26    C: Curve,
27    BFDS: CorrelatedStream<DaBit<C::BaseField>>,
28    BFSS: CorrelatedStream<Singlet<C::BaseField>>,
29    BFTS: CorrelatedStream<Triple<C::BaseField>>,
30    BSS: CorrelatedStream<Singlet<Gf2_128>>,
31    BTS: CorrelatedStream<Triple<Gf2_128>>,
32    MDS: CorrelatedStream<DaBit<Mersenne107>>,
33    MSS: CorrelatedStream<Singlet<Mersenne107>>,
34    MTS: CorrelatedStream<Triple<Mersenne107>>,
35    SDS: CorrelatedStream<DaBit<C::Scalar>>,
36    SSS: CorrelatedStream<Singlet<C::Scalar>>,
37    STS: CorrelatedStream<Triple<C::Scalar>>,
38> {
39    // Base field
40    pub basefield_dabit_stream: BFDS,
41    pub basefield_singlet_stream: BFSS,
42    pub basefield_triple_stream: BFTS,
43    // Binary (Gf2_128)
44    pub binary_singlet_stream: BSS,
45    pub binary_triple_stream: BTS,
46    // Mersenne107
47    pub mersenne107_dabit_stream: MDS,
48    pub mersenne107_singlet_stream: MSS,
49    pub mersenne107_triple_stream: MTS,
50    // Scalar
51    pub scalar_dabit_stream: SDS,
52    pub scalar_singlet_stream: SSS,
53    pub scalar_triple_stream: STS,
54
55    pub _c: PhantomData<C>,
56}
57
58impl<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
59    StreamBundler<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
60where
61    C: Curve,
62    BFDS: CorrelatedStream<DaBit<C::BaseField>>,
63    BFSS: CorrelatedStream<Singlet<C::BaseField>>,
64    BFTS: CorrelatedStream<Triple<C::BaseField>>,
65    BSS: CorrelatedStream<Singlet<Gf2_128>>,
66    BTS: CorrelatedStream<Triple<Gf2_128>>,
67    MDS: CorrelatedStream<DaBit<Mersenne107>>,
68    MSS: CorrelatedStream<Singlet<Mersenne107>>,
69    MTS: CorrelatedStream<Triple<Mersenne107>>,
70    SDS: CorrelatedStream<DaBit<C::Scalar>>,
71    SSS: CorrelatedStream<Singlet<C::Scalar>>,
72    STS: CorrelatedStream<Triple<C::Scalar>>,
73{
74    /// Creates a new bundler from the given streams.
75    #[allow(clippy::too_many_arguments)]
76    pub fn new(
77        basefield_dabit_stream: BFDS,
78        basefield_singlet_stream: BFSS,
79        basefield_triple_stream: BFTS,
80        binary_singlet_stream: BSS,
81        binary_triple_stream: BTS,
82        mersenne107_dabit_stream: MDS,
83        mersenne107_singlet_stream: MSS,
84        mersenne107_triple_stream: MTS,
85        scalar_dabit_stream: SDS,
86        scalar_singlet_stream: SSS,
87        scalar_triple_stream: STS,
88    ) -> Self {
89        Self {
90            basefield_dabit_stream,
91            basefield_singlet_stream,
92            basefield_triple_stream,
93            binary_singlet_stream,
94            binary_triple_stream,
95            mersenne107_dabit_stream,
96            mersenne107_singlet_stream,
97            mersenne107_triple_stream,
98            scalar_dabit_stream,
99            scalar_singlet_stream,
100            scalar_triple_stream,
101            _c: PhantomData,
102        }
103    }
104}
105
106// ──────────────────────── PreprocessingBundler impl ──────────────────────── //
107
108impl<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS> Bundler
109    for StreamBundler<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
110where
111    C: Curve,
112    BFDS: CorrelatedStream<DaBit<C::BaseField>, Error = AbortError>,
113    BFSS: CorrelatedStream<Singlet<C::BaseField>, Error = AbortError>,
114    BFTS: CorrelatedStream<Triple<C::BaseField>, Error = AbortError>,
115    BSS: CorrelatedStream<Singlet<Gf2_128>, Error = AbortError>,
116    BTS: CorrelatedStream<Triple<Gf2_128>, Error = AbortError>,
117    MDS: CorrelatedStream<DaBit<Mersenne107>, Error = AbortError>,
118    MSS: CorrelatedStream<Singlet<Mersenne107>, Error = AbortError>,
119    MTS: CorrelatedStream<Triple<Mersenne107>, Error = AbortError>,
120    SDS: CorrelatedStream<DaBit<C::Scalar>, Error = AbortError>,
121    SSS: CorrelatedStream<Singlet<C::Scalar>, Error = AbortError>,
122    STS: CorrelatedStream<Triple<C::Scalar>, Error = AbortError>,
123{
124    type Iterator = PreprocessingIterator<C>;
125    fn fetch(
126        &mut self,
127        req: &CircuitPreprocessing,
128    ) -> Result<PreprocessingIterator<C>, BundlerError> {
129        Ok(PreprocessingIterator {
130            base_field_dabits: self
131                .basefield_dabit_stream
132                .next_n(req.base_field.dabits)?
133                .into_iter(),
134            base_field_singlets: self
135                .basefield_singlet_stream
136                .next_n(req.base_field.singlets)?
137                .into_iter(),
138            base_field_triples: self
139                .basefield_triple_stream
140                .next_n(req.base_field.triples)?
141                .into_iter(),
142            binary_singlets: self
143                .binary_singlet_stream
144                .next_n(req.bit_singlets)?
145                .into_iter(),
146            binary_triples: self
147                .binary_triple_stream
148                .next_n(req.bit_triples)?
149                .into_iter(),
150            mersenne107_dabits: self
151                .mersenne107_dabit_stream
152                .next_n(req.mersenne107.dabits)?
153                .into_iter(),
154            mersenne107_singlets: self
155                .mersenne107_singlet_stream
156                .next_n(req.mersenne107.singlets)?
157                .into_iter(),
158            mersenne107_triples: self
159                .mersenne107_triple_stream
160                .next_n(req.mersenne107.triples)?
161                .into_iter(),
162            scalar_dabits: self
163                .scalar_dabit_stream
164                .next_n(req.scalar.dabits)?
165                .into_iter(),
166            scalar_singlets: self
167                .scalar_singlet_stream
168                .next_n(req.scalar.singlets)?
169                .into_iter(),
170            scalar_triples: self
171                .scalar_triple_stream
172                .next_n(req.scalar.triples)?
173                .into_iter(),
174        })
175    }
176}