arcium-core-utils 0.7.1

Arcium core utils
Documentation
use std::marker::PhantomData;

use primitives::{
    algebra::field::binary::Gf2_128,
    correlated_randomness::{
        bundler::{errors::BundlerError, Bundler},
        dabits::DaBit,
        singlets::Singlet,
        stream::{
            Buffer,
            CorrelatedStream,
            CorrelatedStreamError,
            PrefetchHandle,
            SharedBufferConfig,
        },
        triples::Triple,
    },
};

use crate::{
    circuit::preprocessing::CircuitPreprocessing,
    config::{BaseFieldOf, MpcConfig, MpcFieldOf, ScalarFieldOf},
    errors::AbortError,
    preprocessing::iterator::PreprocessingIterator,
};

// ── Fan-out over the bundler's 11 streams ───────────────────────────────────────────────────── //
// Single source of truth pairing each stream with its iterator field and `CircuitPreprocessing`
// size path. `for_each_stream!(cb!(args,))` forwards `args` plus the 11 triples
// `(iterator_field, stream_field, size_path)` to the callback macro `cb!`.
macro_rules! for_each_stream {
    ($cb:ident ! ( $($pre:tt)* )) => {
        $cb!($($pre)*
            (base_field_dabits,    basefield_dabit_stream,     base_field.dabits),
            (base_field_singlets,  basefield_singlet_stream,   base_field.singlets),
            (base_field_triples,   basefield_triple_stream,    base_field.triples),
            (binary_singlets,      binary_singlet_stream,      bit_singlets),
            (binary_triples,       binary_triple_stream,       bit_triples),
            (mpc_field_dabits,   mpc_field_dabit_stream,   mpc_field.dabits),
            (mpc_field_singlets, mpc_field_singlet_stream, mpc_field.singlets),
            (mpc_field_triples,  mpc_field_triple_stream,  mpc_field.triples),
            (scalar_dabits,        scalar_dabit_stream,        scalar.dabits),
            (scalar_singlets,      scalar_singlet_stream,      scalar.singlets),
            (scalar_triples,       scalar_triple_stream,       scalar.triples),
        )
    };
}
macro_rules! build_iterator {
    ($self:ident, $req:ident, $(($it:ident, $s:ident, $($p:tt)+)),+ $(,)?) => {
        PreprocessingIterator { $( $it: $self.$s.next_n($req.$($p)+)?.into_iter() ),+ }
    };
}
macro_rules! assign_positions {
    ($self:ident, $pos:ident, $(($it:ident, $s:ident, $($p:tt)+)),+ $(,)?) => {
        $( $pos.$($p)+ = $self.$s.position() as usize; )+
    };
}
macro_rules! check_no_rewind {
    ($cur:ident, $tgt:ident, $(($it:ident, $s:ident, $($p:tt)+)),+ $(,)?) => {
        $( if $tgt.$($p)+ < $cur.$($p)+ {
            return Err(CorrelatedStreamError::ResyncRewind {
                current: $cur.$($p)+ as u64,
                target: $tgt.$($p)+ as u64,
            }.into());
        } )+
    };
}
macro_rules! resync_handles {
    ($self:ident, $tgt:ident, $(($it:ident, $s:ident, $($p:tt)+)),+ $(,)?) => {
        [ $( $self.$s.resync($tgt.$($p)+ as u64) ),+ ]
    };
}
macro_rules! prefetch_handles {
    ($self:ident, $req:ident, $(($it:ident, $s:ident, $($p:tt)+)),+ $(,)?) => {
        [ $( $self.$s.prefetch_n($req.$($p)+) ),+ ]
    };
}
macro_rules! collect_configs {
    ($self:ident, $(($it:ident, $s:ident, $($p:tt)+)),+ $(,)?) => {
        [ $( $self.$s.config().clone() ),+ ]
    };
}

/// Stream bundler, holding one stream per preprocessing type, to provide all preprocessing for a
/// circuit via its streams.
pub struct StreamBundler<
    C: MpcConfig,
    BFDS: CorrelatedStream<DaBit<BaseFieldOf<C>>>,
    BFSS: CorrelatedStream<Singlet<BaseFieldOf<C>>>,
    BFTS: CorrelatedStream<Triple<BaseFieldOf<C>>>,
    BSS: CorrelatedStream<Singlet<Gf2_128>>,
    BTS: CorrelatedStream<Triple<Gf2_128>>,
    MDS: CorrelatedStream<DaBit<MpcFieldOf<C>>>,
    MSS: CorrelatedStream<Singlet<MpcFieldOf<C>>>,
    MTS: CorrelatedStream<Triple<MpcFieldOf<C>>>,
    SDS: CorrelatedStream<DaBit<ScalarFieldOf<C>>>,
    SSS: CorrelatedStream<Singlet<ScalarFieldOf<C>>>,
    STS: CorrelatedStream<Triple<ScalarFieldOf<C>>>,
> {
    // Base field
    pub basefield_dabit_stream: BFDS,
    pub basefield_singlet_stream: BFSS,
    pub basefield_triple_stream: BFTS,
    // Binary (Gf2_128)
    pub binary_singlet_stream: BSS,
    pub binary_triple_stream: BTS,
    // MPC field
    pub mpc_field_dabit_stream: MDS,
    pub mpc_field_singlet_stream: MSS,
    pub mpc_field_triple_stream: MTS,
    // Scalar
    pub scalar_dabit_stream: SDS,
    pub scalar_singlet_stream: SSS,
    pub scalar_triple_stream: STS,

    pub _c: PhantomData<C>,
}

impl<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
    StreamBundler<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
where
    C: MpcConfig,
    BFDS: CorrelatedStream<DaBit<BaseFieldOf<C>>>,
    BFSS: CorrelatedStream<Singlet<BaseFieldOf<C>>>,
    BFTS: CorrelatedStream<Triple<BaseFieldOf<C>>>,
    BSS: CorrelatedStream<Singlet<Gf2_128>>,
    BTS: CorrelatedStream<Triple<Gf2_128>>,
    MDS: CorrelatedStream<DaBit<MpcFieldOf<C>>>,
    MSS: CorrelatedStream<Singlet<MpcFieldOf<C>>>,
    MTS: CorrelatedStream<Triple<MpcFieldOf<C>>>,
    SDS: CorrelatedStream<DaBit<ScalarFieldOf<C>>>,
    SSS: CorrelatedStream<Singlet<ScalarFieldOf<C>>>,
    STS: CorrelatedStream<Triple<ScalarFieldOf<C>>>,
{
    /// Creates a new bundler from the given streams.
    #[allow(clippy::too_many_arguments)]
    pub fn new(
        basefield_dabit_stream: BFDS,
        basefield_singlet_stream: BFSS,
        basefield_triple_stream: BFTS,
        binary_singlet_stream: BSS,
        binary_triple_stream: BTS,
        mpc_field_dabit_stream: MDS,
        mpc_field_singlet_stream: MSS,
        mpc_field_triple_stream: MTS,
        scalar_dabit_stream: SDS,
        scalar_singlet_stream: SSS,
        scalar_triple_stream: STS,
    ) -> Self {
        Self {
            basefield_dabit_stream,
            basefield_singlet_stream,
            basefield_triple_stream,
            binary_singlet_stream,
            binary_triple_stream,
            mpc_field_dabit_stream,
            mpc_field_singlet_stream,
            mpc_field_triple_stream,
            scalar_dabit_stream,
            scalar_singlet_stream,
            scalar_triple_stream,
            _c: PhantomData,
        }
    }
}

// ──────────────────────── PreprocessingBundler impl ──────────────────────── //

impl<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS> Bundler
    for StreamBundler<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
where
    C: MpcConfig,
    BFDS: CorrelatedStream<DaBit<BaseFieldOf<C>>, Error = AbortError>,
    BFSS: CorrelatedStream<Singlet<BaseFieldOf<C>>, Error = AbortError>,
    BFTS: CorrelatedStream<Triple<BaseFieldOf<C>>, Error = AbortError>,
    BSS: CorrelatedStream<Singlet<Gf2_128>, Error = AbortError>,
    BTS: CorrelatedStream<Triple<Gf2_128>, Error = AbortError>,
    MDS: CorrelatedStream<DaBit<MpcFieldOf<C>>, Error = AbortError>,
    MSS: CorrelatedStream<Singlet<MpcFieldOf<C>>, Error = AbortError>,
    MTS: CorrelatedStream<Triple<MpcFieldOf<C>>, Error = AbortError>,
    SDS: CorrelatedStream<DaBit<ScalarFieldOf<C>>, Error = AbortError>,
    SSS: CorrelatedStream<Singlet<ScalarFieldOf<C>>, Error = AbortError>,
    STS: CorrelatedStream<Triple<ScalarFieldOf<C>>, Error = AbortError>,
{
    type Iterator = PreprocessingIterator<C>;
    fn fetch(
        &mut self,
        req: &CircuitPreprocessing,
    ) -> Result<PreprocessingIterator<C>, BundlerError> {
        Ok(for_each_stream!(build_iterator!(self, req,)))
    }
}

// ──────────────────────── Resynchronization ──────────────────────── //

impl<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
    StreamBundler<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
where
    C: MpcConfig,
    BFDS: CorrelatedStream<DaBit<BaseFieldOf<C>>, Error = AbortError>,
    BFSS: CorrelatedStream<Singlet<BaseFieldOf<C>>, Error = AbortError>,
    BFTS: CorrelatedStream<Triple<BaseFieldOf<C>>, Error = AbortError>,
    BSS: CorrelatedStream<Singlet<Gf2_128>, Error = AbortError>,
    BTS: CorrelatedStream<Triple<Gf2_128>, Error = AbortError>,
    MDS: CorrelatedStream<DaBit<MpcFieldOf<C>>, Error = AbortError>,
    MSS: CorrelatedStream<Singlet<MpcFieldOf<C>>, Error = AbortError>,
    MTS: CorrelatedStream<Triple<MpcFieldOf<C>>, Error = AbortError>,
    SDS: CorrelatedStream<DaBit<ScalarFieldOf<C>>, Error = AbortError>,
    SSS: CorrelatedStream<Singlet<ScalarFieldOf<C>>, Error = AbortError>,
    STS: CorrelatedStream<Triple<ScalarFieldOf<C>>, Error = AbortError>,
{
    /// Prefetches the given amount into every stream and returns a single [`PrefetchHandle`] that
    /// resolves once all of them complete (first error wins). The prefetches run concurrently in
    /// the background; the handle can be awaited or dropped.
    pub fn prefetch(&self, req: &CircuitPreprocessing) -> PrefetchHandle<AbortError> {
        let handles = for_each_stream!(prefetch_handles!(self, req,));
        PrefetchHandle::from_future(async move {
            let mut first_err = None;
            for h in handles {
                if let Err(e) = h.await {
                    first_err.get_or_insert(e);
                }
            }
            first_err.map_or(Ok(()), Err)
        })
    }

    /// The logical position (elements delivered) of every stream, per type. Stays in sync across
    /// parties; take the per-type maximum to agree on a resync target, then pass it to
    /// [`resync`](Self::resync).
    pub fn positions(&self) -> CircuitPreprocessing {
        let mut pos = CircuitPreprocessing::default();
        for_each_stream!(assign_positions!(self, pos,));
        pos
    }

    /// Advances every stream to its per-type `target`, realigning all parties on the same prefix.
    ///
    /// All-or-nothing on rewinds: any target behind a stream's current position is rejected before
    /// touching any stream. Otherwise all resyncs are dispatched together, every handle awaited,
    /// and the first error (if any) returned — no short-circuiting mid-fan-out.
    pub async fn resync(&self, targets: &CircuitPreprocessing) -> Result<(), AbortError> {
        let current = self.positions();
        for_each_stream!(check_no_rewind!(current, targets,));
        let handles = for_each_stream!(resync_handles!(self, targets,));
        // Await every handle (no early `?`) so all streams advance, then surface the first error.
        let mut first_err = None;
        for handle in handles {
            if let Err(e) = handle.await {
                first_err.get_or_insert(e);
            }
        }
        first_err.map_or(Ok(()), Err)
    }
}

// ──────────────────────── Buffer configuration ──────────────────────── //

impl<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
    StreamBundler<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS>
where
    C: MpcConfig,
    BFDS: CorrelatedStream<DaBit<BaseFieldOf<C>>> + Buffer,
    BFSS: CorrelatedStream<Singlet<BaseFieldOf<C>>> + Buffer,
    BFTS: CorrelatedStream<Triple<BaseFieldOf<C>>> + Buffer,
    BSS: CorrelatedStream<Singlet<Gf2_128>> + Buffer,
    BTS: CorrelatedStream<Triple<Gf2_128>> + Buffer,
    MDS: CorrelatedStream<DaBit<MpcFieldOf<C>>> + Buffer,
    MSS: CorrelatedStream<Singlet<MpcFieldOf<C>>> + Buffer,
    MTS: CorrelatedStream<Triple<MpcFieldOf<C>>> + Buffer,
    SDS: CorrelatedStream<DaBit<ScalarFieldOf<C>>> + Buffer,
    SSS: CorrelatedStream<Singlet<ScalarFieldOf<C>>> + Buffer,
    STS: CorrelatedStream<Triple<ScalarFieldOf<C>>> + Buffer,
{
    /// The shared buffer config of every stream. The returned array is a [`Buffer`] slice, so
    /// the whole bundle can be tuned in one call: `bundler.buffer_configs().set_capacity(n)`
    /// writes all streams, while `.capacity()` etc. read the first.
    pub fn buffer_configs(&self) -> [SharedBufferConfig; 11] {
        for_each_stream!(collect_configs!(self,))
    }
}