arcium-core-utils 0.8.0

Arcium core utils
Documentation
use primitives::{
    algebra::field::binary::Gf2_128,
    correlated_randomness::{
        bundler::{errors::BundlerError, BundleIterator},
        dabits::DaBit,
        singlets::Singlet,
        stream::{Next, NextVec, NextVecIterator},
    },
    sharing::{Malicious, ThreatModel},
    types::identifiers::Named,
};

use crate::{
    circuit::preprocessing::{CircuitPreprocessing, FieldCircuitPreprocessing},
    config::{BaseFieldOf, MpcConfig, MpcFieldOf, ScalarFieldOf},
    errors::AbortError,
    preprocessing::{NextDaBit, NextDaBits, NextSinglet, NextSinglets},
};

/// An iterator containing preprocessing futures for every gate in a circuit.
pub struct PreprocessingIterator<C: MpcConfig, M: ThreatModel = Malicious> {
    // Base field iterators
    pub base_field_dabits: NextVecIterator<DaBit<BaseFieldOf<C>>, AbortError>,
    pub base_field_singlets: NextVecIterator<Singlet<BaseFieldOf<C>>, AbortError>,
    pub base_field_triples: NextVecIterator<M::Triple<BaseFieldOf<C>>, AbortError>,

    // Binary field iterators
    pub binary_singlets: NextVecIterator<Singlet<Gf2_128>, AbortError>,
    pub binary_triples: NextVecIterator<M::Triple<Gf2_128>, AbortError>,

    // MpcFieldOf<C> iterators
    pub mpc_field_dabits: NextVecIterator<DaBit<MpcFieldOf<C>>, AbortError>,
    pub mpc_field_singlets: NextVecIterator<Singlet<MpcFieldOf<C>>, AbortError>,
    pub mpc_field_triples: NextVecIterator<M::Triple<MpcFieldOf<C>>, AbortError>,

    // Scalar field iterators
    pub scalar_dabits: NextVecIterator<DaBit<ScalarFieldOf<C>>, AbortError>,
    pub scalar_singlets: NextVecIterator<Singlet<ScalarFieldOf<C>>, AbortError>,
    pub scalar_triples: NextVecIterator<M::Triple<ScalarFieldOf<C>>, AbortError>,
}

impl<C: MpcConfig, M: ThreatModel> BundleIterator for PreprocessingIterator<C, M> {
    type Size = CircuitPreprocessing;

    fn len(&self) -> CircuitPreprocessing {
        CircuitPreprocessing {
            bit_singlets: self.binary_singlets.len(),
            bit_triples: self.binary_triples.len(),
            base_field: FieldCircuitPreprocessing {
                singlets: self.base_field_singlets.len(),
                triples: self.base_field_triples.len(),
                dabits: self.base_field_dabits.len(),
            },
            scalar: FieldCircuitPreprocessing {
                singlets: self.scalar_singlets.len(),
                triples: self.scalar_triples.len(),
                dabits: self.scalar_dabits.len(),
            },
            mpc_field: FieldCircuitPreprocessing {
                singlets: self.mpc_field_singlets.len(),
                triples: self.mpc_field_triples.len(),
                dabits: self.mpc_field_dabits.len(),
            },
        }
    }
    // `is_empty` uses the `BundleIterator` default (`len() == Size::default()`).
}

macro_rules! next_preprocessing_item {
    ($( ($fn:ident, $iter:ident, $ret:ty, $ty:ty, $kind:literal) ),* $(,)?) => {
        $(
            pub fn $fn(&mut self) -> Result<$ret, BundlerError> {
                self.$iter.next().ok_or_else(|| {
                    BundlerError::InsufficientPreprocessing(
                        format!("{} {}", <$ty>::get_name(), $kind),
                    )
                })
            }
        )*
    };
}

macro_rules! next_n_preprocessing_items {
    ($( ($fn:ident, $iter:ident, $ret:ty, $ty:ty, $kind:literal) ),* $(,)?) => {
        $(
            pub fn $fn(&mut self, n: usize) -> Result<$ret, BundlerError> {
                self.$iter.next_n(n).ok_or_else(|| {
                    BundlerError::InsufficientPreprocessing(
                        format!("{} {} (requested {}, available {})", <$ty>::get_name(), $kind, n, self.$iter.len()),
                    )
                })
            }
        )*
    };
}

impl<C: MpcConfig, M: ThreatModel> PreprocessingIterator<C, M> {
    next_preprocessing_item!(
        (next_base_field_dabit, base_field_dabits, NextDaBit<BaseFieldOf<C>>, BaseFieldOf<C>, "dabits"),
        (next_base_field_singlet, base_field_singlets, NextSinglet<BaseFieldOf<C>>, BaseFieldOf<C>, "singlets"),
        (next_base_field_triple, base_field_triples, Next<M::Triple<BaseFieldOf<C>>, AbortError>, BaseFieldOf<C>, "triples"),
        (next_bit_singlet, binary_singlets, NextSinglet<Gf2_128>, Gf2_128, "singlets"),
        (next_bit_triple, binary_triples, Next<M::Triple<Gf2_128>, AbortError>, Gf2_128, "triples"),
        (next_mpc_field_dabit, mpc_field_dabits, NextDaBit<MpcFieldOf<C>>, MpcFieldOf<C>, "dabits"),
        (next_mpc_field_singlet, mpc_field_singlets, NextSinglet<MpcFieldOf<C>>, MpcFieldOf<C>, "singlets"),
        (next_mpc_field_triple, mpc_field_triples, Next<M::Triple<MpcFieldOf<C>>, AbortError>, MpcFieldOf<C>, "triples"),
        (next_scalar_dabit, scalar_dabits, NextDaBit<ScalarFieldOf<C>>, ScalarFieldOf<C>, "dabits"),
        (next_scalar_singlet, scalar_singlets, NextSinglet<ScalarFieldOf<C>>, ScalarFieldOf<C>, "singlets"),
        (next_scalar_triple, scalar_triples, Next<M::Triple<ScalarFieldOf<C>>, AbortError>, ScalarFieldOf<C>, "triples"),
    );

    next_n_preprocessing_items!(
        (next_n_base_field_dabits, base_field_dabits, NextDaBits<BaseFieldOf<C>>, BaseFieldOf<C>, "dabits"),
        (next_n_base_field_singlets, base_field_singlets, NextSinglets<BaseFieldOf<C>>, BaseFieldOf<C>, "singlets"),
        (next_n_base_field_triples, base_field_triples, NextVec<M::Triple<BaseFieldOf<C>>, AbortError>, BaseFieldOf<C>, "triples"),
        (next_n_bit_singlets, binary_singlets, NextSinglets<Gf2_128>, Gf2_128, "singlets"),
        (next_n_bit_triples, binary_triples, NextVec<M::Triple<Gf2_128>, AbortError>, Gf2_128, "triples"),
        (next_n_mpc_field_dabits, mpc_field_dabits, NextDaBits<MpcFieldOf<C>>, MpcFieldOf<C>, "dabits"),
        (next_n_mpc_field_singlets, mpc_field_singlets, NextSinglets<MpcFieldOf<C>>, MpcFieldOf<C>, "singlets"),
        (next_n_mpc_field_triples, mpc_field_triples, NextVec<M::Triple<MpcFieldOf<C>>, AbortError>, MpcFieldOf<C>, "triples"),
        (next_n_scalar_dabits, scalar_dabits, NextDaBits<ScalarFieldOf<C>>, ScalarFieldOf<C>, "dabits"),
        (next_n_scalar_singlets, scalar_singlets, NextSinglets<ScalarFieldOf<C>>, ScalarFieldOf<C>, "singlets"),
        (next_n_scalar_triples, scalar_triples, NextVec<M::Triple<ScalarFieldOf<C>>, AbortError>, ScalarFieldOf<C>, "triples"),
    );
}

impl<C: MpcConfig, M: ThreatModel> std::fmt::Debug for PreprocessingIterator<C, M> {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        write!(f, "PreprocessingIterator with len: {:?}", self.len())
    }
}