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::{
13 Buffer,
14 CorrelatedStream,
15 CorrelatedStreamError,
16 PrefetchHandle,
17 SharedBufferConfig,
18 },
19 triples::Triple,
20 },
21};
22
23use crate::{
24 circuit::preprocessing::CircuitPreprocessing,
25 errors::AbortError,
26 preprocessing::iterator::PreprocessingIterator,
27};
28
29macro_rules! for_each_stream {
34 ($cb:ident ! ( $($pre:tt)* )) => {
35 $cb!($($pre)*
36 (base_field_dabits, basefield_dabit_stream, base_field.dabits),
37 (base_field_singlets, basefield_singlet_stream, base_field.singlets),
38 (base_field_triples, basefield_triple_stream, base_field.triples),
39 (binary_singlets, binary_singlet_stream, bit_singlets),
40 (binary_triples, binary_triple_stream, bit_triples),
41 (mersenne107_dabits, mersenne107_dabit_stream, mersenne107.dabits),
42 (mersenne107_singlets, mersenne107_singlet_stream, mersenne107.singlets),
43 (mersenne107_triples, mersenne107_triple_stream, mersenne107.triples),
44 (scalar_dabits, scalar_dabit_stream, scalar.dabits),
45 (scalar_singlets, scalar_singlet_stream, scalar.singlets),
46 (scalar_triples, scalar_triple_stream, scalar.triples),
47 )
48 };
49}
50macro_rules! build_iterator {
51 ($self:ident, $req:ident, $(($it:ident, $s:ident, $($p:tt)+)),+ $(,)?) => {
52 PreprocessingIterator { $( $it: $self.$s.next_n($req.$($p)+)?.into_iter() ),+ }
53 };
54}
55macro_rules! assign_positions {
56 ($self:ident, $pos:ident, $(($it:ident, $s:ident, $($p:tt)+)),+ $(,)?) => {
57 $( $pos.$($p)+ = $self.$s.position() as usize; )+
58 };
59}
60macro_rules! check_no_rewind {
61 ($cur:ident, $tgt:ident, $(($it:ident, $s:ident, $($p:tt)+)),+ $(,)?) => {
62 $( if $tgt.$($p)+ < $cur.$($p)+ {
63 return Err(CorrelatedStreamError::ResyncRewind {
64 current: $cur.$($p)+ as u64,
65 target: $tgt.$($p)+ as u64,
66 }.into());
67 } )+
68 };
69}
70macro_rules! resync_handles {
71 ($self:ident, $tgt:ident, $(($it:ident, $s:ident, $($p:tt)+)),+ $(,)?) => {
72 [ $( $self.$s.resync($tgt.$($p)+ as u64) ),+ ]
73 };
74}
75macro_rules! prefetch_handles {
76 ($self:ident, $req:ident, $(($it:ident, $s:ident, $($p:tt)+)),+ $(,)?) => {
77 [ $( $self.$s.prefetch_n($req.$($p)+) ),+ ]
78 };
79}
80macro_rules! collect_configs {
81 ($self:ident, $(($it:ident, $s:ident, $($p:tt)+)),+ $(,)?) => {
82 [ $( $self.$s.config().clone() ),+ ]
83 };
84}
85
86pub struct StreamBundler<
89 C: Curve,
90 BFDS: CorrelatedStream<DaBit<C::BaseField>>,
91 BFSS: CorrelatedStream<Singlet<C::BaseField>>,
92 BFTS: CorrelatedStream<Triple<C::BaseField>>,
93 BSS: CorrelatedStream<Singlet<Gf2_128>>,
94 BTS: CorrelatedStream<Triple<Gf2_128>>,
95 MDS: CorrelatedStream<DaBit<Mersenne107>>,
96 MSS: CorrelatedStream<Singlet<Mersenne107>>,
97 MTS: CorrelatedStream<Triple<Mersenne107>>,
98 SDS: CorrelatedStream<DaBit<C::Scalar>>,
99 SSS: CorrelatedStream<Singlet<C::Scalar>>,
100 STS: CorrelatedStream<Triple<C::Scalar>>,
101> {
102 pub basefield_dabit_stream: BFDS,
104 pub basefield_singlet_stream: BFSS,
105 pub basefield_triple_stream: BFTS,
106 pub binary_singlet_stream: BSS,
108 pub binary_triple_stream: BTS,
109 pub mersenne107_dabit_stream: MDS,
111 pub mersenne107_singlet_stream: MSS,
112 pub mersenne107_triple_stream: MTS,
113 pub scalar_dabit_stream: SDS,
115 pub scalar_singlet_stream: SSS,
116 pub scalar_triple_stream: STS,
117
118 pub _c: PhantomData<C>,
119}
120
121impl<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
122 StreamBundler<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
123where
124 C: Curve,
125 BFDS: CorrelatedStream<DaBit<C::BaseField>>,
126 BFSS: CorrelatedStream<Singlet<C::BaseField>>,
127 BFTS: CorrelatedStream<Triple<C::BaseField>>,
128 BSS: CorrelatedStream<Singlet<Gf2_128>>,
129 BTS: CorrelatedStream<Triple<Gf2_128>>,
130 MDS: CorrelatedStream<DaBit<Mersenne107>>,
131 MSS: CorrelatedStream<Singlet<Mersenne107>>,
132 MTS: CorrelatedStream<Triple<Mersenne107>>,
133 SDS: CorrelatedStream<DaBit<C::Scalar>>,
134 SSS: CorrelatedStream<Singlet<C::Scalar>>,
135 STS: CorrelatedStream<Triple<C::Scalar>>,
136{
137 #[allow(clippy::too_many_arguments)]
139 pub fn new(
140 basefield_dabit_stream: BFDS,
141 basefield_singlet_stream: BFSS,
142 basefield_triple_stream: BFTS,
143 binary_singlet_stream: BSS,
144 binary_triple_stream: BTS,
145 mersenne107_dabit_stream: MDS,
146 mersenne107_singlet_stream: MSS,
147 mersenne107_triple_stream: MTS,
148 scalar_dabit_stream: SDS,
149 scalar_singlet_stream: SSS,
150 scalar_triple_stream: STS,
151 ) -> Self {
152 Self {
153 basefield_dabit_stream,
154 basefield_singlet_stream,
155 basefield_triple_stream,
156 binary_singlet_stream,
157 binary_triple_stream,
158 mersenne107_dabit_stream,
159 mersenne107_singlet_stream,
160 mersenne107_triple_stream,
161 scalar_dabit_stream,
162 scalar_singlet_stream,
163 scalar_triple_stream,
164 _c: PhantomData,
165 }
166 }
167}
168
169impl<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS> Bundler
172 for StreamBundler<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
173where
174 C: Curve,
175 BFDS: CorrelatedStream<DaBit<C::BaseField>, Error = AbortError>,
176 BFSS: CorrelatedStream<Singlet<C::BaseField>, Error = AbortError>,
177 BFTS: CorrelatedStream<Triple<C::BaseField>, Error = AbortError>,
178 BSS: CorrelatedStream<Singlet<Gf2_128>, Error = AbortError>,
179 BTS: CorrelatedStream<Triple<Gf2_128>, Error = AbortError>,
180 MDS: CorrelatedStream<DaBit<Mersenne107>, Error = AbortError>,
181 MSS: CorrelatedStream<Singlet<Mersenne107>, Error = AbortError>,
182 MTS: CorrelatedStream<Triple<Mersenne107>, Error = AbortError>,
183 SDS: CorrelatedStream<DaBit<C::Scalar>, Error = AbortError>,
184 SSS: CorrelatedStream<Singlet<C::Scalar>, Error = AbortError>,
185 STS: CorrelatedStream<Triple<C::Scalar>, Error = AbortError>,
186{
187 type Iterator = PreprocessingIterator<C>;
188 fn fetch(
189 &mut self,
190 req: &CircuitPreprocessing,
191 ) -> Result<PreprocessingIterator<C>, BundlerError> {
192 Ok(for_each_stream!(build_iterator!(self, req,)))
193 }
194}
195
196impl<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
199 StreamBundler<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
200where
201 C: Curve,
202 BFDS: CorrelatedStream<DaBit<C::BaseField>, Error = AbortError>,
203 BFSS: CorrelatedStream<Singlet<C::BaseField>, Error = AbortError>,
204 BFTS: CorrelatedStream<Triple<C::BaseField>, Error = AbortError>,
205 BSS: CorrelatedStream<Singlet<Gf2_128>, Error = AbortError>,
206 BTS: CorrelatedStream<Triple<Gf2_128>, Error = AbortError>,
207 MDS: CorrelatedStream<DaBit<Mersenne107>, Error = AbortError>,
208 MSS: CorrelatedStream<Singlet<Mersenne107>, Error = AbortError>,
209 MTS: CorrelatedStream<Triple<Mersenne107>, Error = AbortError>,
210 SDS: CorrelatedStream<DaBit<C::Scalar>, Error = AbortError>,
211 SSS: CorrelatedStream<Singlet<C::Scalar>, Error = AbortError>,
212 STS: CorrelatedStream<Triple<C::Scalar>, Error = AbortError>,
213{
214 pub fn prefetch(&self, req: &CircuitPreprocessing) -> PrefetchHandle<AbortError> {
218 let handles = for_each_stream!(prefetch_handles!(self, req,));
219 PrefetchHandle::from_future(async move {
220 let mut first_err = None;
221 for h in handles {
222 if let Err(e) = h.await {
223 first_err.get_or_insert(e);
224 }
225 }
226 first_err.map_or(Ok(()), Err)
227 })
228 }
229
230 pub fn positions(&self) -> CircuitPreprocessing {
234 let mut pos = CircuitPreprocessing::default();
235 for_each_stream!(assign_positions!(self, pos,));
236 pos
237 }
238
239 pub async fn resync(&self, targets: &CircuitPreprocessing) -> Result<(), AbortError> {
245 let current = self.positions();
246 for_each_stream!(check_no_rewind!(current, targets,));
247 let handles = for_each_stream!(resync_handles!(self, targets,));
248 let mut first_err = None;
250 for handle in handles {
251 if let Err(e) = handle.await {
252 first_err.get_or_insert(e);
253 }
254 }
255 first_err.map_or(Ok(()), Err)
256 }
257}
258
259impl<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
262 StreamBundler<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
263where
264 C: Curve,
265 BFDS: CorrelatedStream<DaBit<C::BaseField>> + Buffer,
266 BFSS: CorrelatedStream<Singlet<C::BaseField>> + Buffer,
267 BFTS: CorrelatedStream<Triple<C::BaseField>> + Buffer,
268 BSS: CorrelatedStream<Singlet<Gf2_128>> + Buffer,
269 BTS: CorrelatedStream<Triple<Gf2_128>> + Buffer,
270 MDS: CorrelatedStream<DaBit<Mersenne107>> + Buffer,
271 MSS: CorrelatedStream<Singlet<Mersenne107>> + Buffer,
272 MTS: CorrelatedStream<Triple<Mersenne107>> + Buffer,
273 SDS: CorrelatedStream<DaBit<C::Scalar>> + Buffer,
274 SSS: CorrelatedStream<Singlet<C::Scalar>> + Buffer,
275 STS: CorrelatedStream<Triple<C::Scalar>> + Buffer,
276{
277 pub fn buffer_configs(&self) -> [SharedBufferConfig; 11] {
281 for_each_stream!(collect_configs!(self,))
282 }
283}