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::{
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
29// ── Fan-out over the bundler's 11 streams ───────────────────────────────────────────────────── //
30// Single source of truth pairing each stream with its iterator field and `CircuitPreprocessing`
31// size path. `for_each_stream!(cb!(args,))` forwards `args` plus the 11 triples
32// `(iterator_field, stream_field, size_path)` to the callback macro `cb!`.
33macro_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
86/// Stream bundler, holding one stream per preprocessing type, to provide all preprocessing for a
87/// circuit via its streams.
88pub 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    // Base field
103    pub basefield_dabit_stream: BFDS,
104    pub basefield_singlet_stream: BFSS,
105    pub basefield_triple_stream: BFTS,
106    // Binary (Gf2_128)
107    pub binary_singlet_stream: BSS,
108    pub binary_triple_stream: BTS,
109    // Mersenne107
110    pub mersenne107_dabit_stream: MDS,
111    pub mersenne107_singlet_stream: MSS,
112    pub mersenne107_triple_stream: MTS,
113    // Scalar
114    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    /// Creates a new bundler from the given streams.
138    #[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
169// ──────────────────────── PreprocessingBundler impl ──────────────────────── //
170
171impl<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
196// ──────────────────────── Resynchronization ──────────────────────── //
197
198impl<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    /// Prefetches the given amount into every stream and returns a single [`PrefetchHandle`] that
215    /// resolves once all of them complete (first error wins). The prefetches run concurrently in
216    /// the background; the handle can be awaited or dropped.
217    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    /// The logical position (elements delivered) of every stream, per type. Stays in sync across
231    /// parties; take the per-type maximum to agree on a resync target, then pass it to
232    /// [`resync`](Self::resync).
233    pub fn positions(&self) -> CircuitPreprocessing {
234        let mut pos = CircuitPreprocessing::default();
235        for_each_stream!(assign_positions!(self, pos,));
236        pos
237    }
238
239    /// Advances every stream to its per-type `target`, realigning all parties on the same prefix.
240    ///
241    /// All-or-nothing on rewinds: any target behind a stream's current position is rejected before
242    /// touching any stream. Otherwise all resyncs are dispatched together, every handle awaited,
243    /// and the first error (if any) returned — no short-circuiting mid-fan-out.
244    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        // Await every handle (no early `?`) so all streams advance, then surface the first error.
249        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
259// ──────────────────────── Buffer configuration ──────────────────────── //
260
261impl<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    /// The shared buffer config of every stream. The returned array is a [`Buffer`] slice, so
278    /// the whole bundle can be tuned in one call: `bundler.buffer_configs().set_capacity(n)`
279    /// writes all streams, while `.capacity()` etc. read the first.
280    pub fn buffer_configs(&self) -> [SharedBufferConfig; 11] {
281        for_each_stream!(collect_configs!(self,))
282    }
283}