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, BundlerPositions},
};
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! assign_buffered {
($self:ident, $buf:ident, $(($it:ident, $s:ident, $($p:tt)+)),+ $(,)?) => {
$( $buf.$($p)+ = $self.$s.buffered() 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() ),+ ]
};
}
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>>>,
> {
pub basefield_dabit_stream: BFDS,
pub basefield_singlet_stream: BFSS,
pub basefield_triple_stream: BFTS,
pub binary_singlet_stream: BSS,
pub binary_triple_stream: BTS,
pub mpc_field_dabit_stream: MDS,
pub mpc_field_singlet_stream: MSS,
pub mpc_field_triple_stream: MTS,
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>>>,
{
#[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,
}
}
}
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,)))
}
}
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>,
{
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)
})
}
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,));
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)
}
}
impl<C, BFDS, BFSS, BFTS, BSS, BTS, MDS, MSS, MTS, SDS, SSS, STS> BundlerPositions
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>,
{
fn positions(&self) -> CircuitPreprocessing {
let mut pos = CircuitPreprocessing::default();
for_each_stream!(assign_positions!(self, pos,));
pos
}
fn buffered(&self) -> CircuitPreprocessing {
let mut buf = CircuitPreprocessing::default();
for_each_stream!(assign_buffered!(self, buf,));
buf
}
async fn resync(&self, targets: &CircuitPreprocessing) -> Result<(), AbortError> {
self.resync(targets).await
}
}
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,
{
pub fn buffer_configs(&self) -> [SharedBufferConfig; 11] {
for_each_stream!(collect_configs!(self,))
}
}