use crate::utils::test_utils::fan_in_inboxes;
use ark_ff::FftField;
use ark_std::test_rng;
use std::sync::Arc;
use std::time::Duration;
use stoffelcrypto::common::{SecretSharingScheme, RBC};
use stoffelcrypto::honeybadger::fpmul::rand_bit::RandBit;
use stoffelcrypto::honeybadger::robust_interpolate::robust_interpolate::RobustShare;
use stoffelcrypto::honeybadger::triple_gen::ShamirBeaverTriple;
use stoffelcrypto::honeybadger::{SessionId, WrappedMessage};
use stoffelmpc_network::fake_network::{FakeNetwork, SenderId};
use tokio::sync::mpsc::Receiver;
use tokio::task::JoinSet;
pub fn create_rand_bit_input<F>(
n_parties: usize,
threshold: usize,
batch_size: usize,
) -> (Vec<Vec<RobustShare<F>>>, Vec<Vec<ShamirBeaverTriple<F>>>)
where
F: FftField,
{
let mut a_shares = vec![vec![]; n_parties];
let mut mult_triples = vec![vec![]; n_parties];
let mut rng = test_rng();
for _ in 0..batch_size {
let a = F::rand(&mut rng);
let shares_a =
RobustShare::<F>::compute_shares(a, n_parties, threshold, None, &mut rng).unwrap();
let x = F::rand(&mut rng);
let y = F::rand(&mut rng);
let mult = x * y;
let x_shares =
RobustShare::<F>::compute_shares(x, n_parties, threshold, None, &mut rng).unwrap();
let y_shares =
RobustShare::<F>::compute_shares(y, n_parties, threshold, None, &mut rng).unwrap();
let mult_shares =
RobustShare::<F>::compute_shares(mult, n_parties, threshold, None, &mut rng).unwrap();
for party_id in 0..n_parties {
a_shares[party_id].push(shares_a[party_id].clone());
let mult_triple = ShamirBeaverTriple::new(
x_shares[party_id].clone(),
y_shares[party_id].clone(),
mult_shares[party_id].clone(),
);
mult_triples[party_id].push(mult_triple);
}
}
(a_shares, mult_triples)
}
pub async fn spawn_receiver_tasks<F, R>(
num_parties: usize,
mut receivers: Vec<Vec<Receiver<Vec<u8>>>>,
nodes: Vec<RandBit<F, R>>,
network: Vec<Arc<FakeNetwork>>,
) -> JoinSet<()>
where
F: FftField,
R: RBC<Id = SessionId> + Clone + 'static,
{
let mut set = JoinSet::new();
for i in 0..num_parties {
let mut node = nodes[i].clone();
let receiver = receivers.remove(0);
let net = network[i].clone();
let inbox: Vec<(SenderId, Receiver<Vec<u8>>)> = receiver
.into_iter() .enumerate()
.map(|(i, r)| (SenderId::Node(i), r))
.collect();
let mut merge_rx = fan_in_inboxes(inbox);
set.spawn(async move {
while let Some((_, bytes)) = merge_rx.recv().await {
let wrapped: WrappedMessage = bincode::deserialize(&bytes).unwrap();
match wrapped {
WrappedMessage::BatchRecon(msg) => {
if msg.session_id.round_id() == 0 {
let _ = node.batch_recon.process(msg, net.clone()).await;
node.drain_batch_recon_output().await.unwrap();
} else {
let _ = node.mult_node.batch_recon.process(msg, net.clone()).await;
node.mult_node.drain_batch_recon_output().await.unwrap();
}
}
_ => {}
}
}
});
}
set
}
pub async fn initialize_nodes<F, R>(
num_parties: usize,
a_shares: Vec<Vec<RobustShare<F>>>,
mult_triples: Vec<Vec<ShamirBeaverTriple<F>>>,
session_id: SessionId,
nodes: Vec<RandBit<F, R>>,
network: Arc<FakeNetwork>,
duration: Duration,
) -> JoinSet<()>
where
F: FftField,
R: RBC<Id = SessionId> + Clone + 'static,
{
let mut init_set = JoinSet::new();
for i in 0..num_parties {
let mut node = nodes[i].clone();
let net = network.clone();
let a = a_shares[i].clone();
let triples = mult_triples[i].clone();
init_set.spawn(async move {
node.init(a, triples, session_id, duration, net)
.await
.unwrap();
});
}
init_set
}