use ark_bls12_381::Fr;
use ark_ff::{FftField, PrimeField, UniformRand};
use ark_std::rand::rngs::{OsRng, StdRng};
use ark_std::rand::SeedableRng;
use ark_std::test_rng;
use once_cell::sync::Lazy;
use std::collections::HashMap;
use std::marker::PhantomData;
use std::time::Duration;
use std::{sync::atomic::AtomicUsize, sync::atomic::Ordering, sync::Arc, vec};
use stoffelcrypto::common::rbc::rbc::Avid;
use stoffelcrypto::common::rbc::RbcError;
use stoffelcrypto::common::share::shamir::NonRobustShare;
use stoffelcrypto::common::{MPCProtocol, SecretSharingScheme, RBC};
use stoffelcrypto::honeybadger::ran_dou_sha::{RanDouShaError, RanDouShaNode};
use stoffelcrypto::honeybadger::robust_interpolate::robust_interpolate::RobustShare;
use stoffelcrypto::honeybadger::share_gen::RanShaError;
use stoffelcrypto::honeybadger::triple_gen::ShamirBeaverTriple;
use stoffelcrypto::honeybadger::{
HoneyBadgerMPCClient, HoneyBadgerMPCNode, HoneyBadgerMPCNodeOpts, SessionId, WrappedMessage,
};
use stoffelmpc_network::fake_network::{
FakeInnerNetwork, FakeNetwork, FakeNetworkConfig, SenderId,
};
use stoffelnet::network_utils::{ClientId, Network, NetworkError};
use tokio::sync::mpsc::{self, Receiver};
use tokio::sync::Mutex;
use tokio::task::JoinSet;
use tracing::{info, warn};
use tracing_subscriber::EnvFilter;
use tracing_subscriber::FmtSubscriber;
pub async fn setup_network_and_parties<T: RBC<Id = SessionId>, N: Network>(
n: usize,
t: usize,
k: usize,
buffer_size: usize,
) -> Result<(Vec<T>, Vec<Arc<FakeNetwork>>, Vec<Vec<Receiver<Vec<u8>>>>), RbcError> {
let config = FakeNetworkConfig::new(buffer_size);
let (inner, receivers, _) = FakeInnerNetwork::new(n as usize, None, config);
let net: Vec<_> = (0..n)
.map(|id| Arc::new(FakeNetwork::new(id, inner.clone())))
.collect();
let mut parties = Vec::with_capacity(n as usize);
for i in 0..n {
let (rbc_sender, _) = mpsc::channel(200);
let rbc = T::new(i, n, t, k, rbc_sender, Arc::new(WrappedMessage::rbc_wrap))?; parties.push(rbc);
}
Ok((parties, net, receivers))
}
pub async fn spawn_parties<T, N>(
parties: &[T],
receivers: Vec<Vec<mpsc::Receiver<Vec<u8>>>>,
net: Vec<Arc<N>>,
) where
T: RBC<Id = SessionId> + Clone + Send + Sync + 'static,
N: Network + Send + Sync + 'static,
{
for (rbc, rx) in parties.iter().cloned().zip(receivers.into_iter()) {
let id = rbc.id();
let net_clone = net[id].clone();
let inbox: Vec<(SenderId, Receiver<Vec<u8>>)> = rx
.into_iter() .enumerate()
.map(|(i, r)| (SenderId::Node(i), r))
.collect();
let mut merge_rx = fan_in_inboxes(inbox);
tokio::spawn(async move {
while let Some(msg) = merge_rx.recv().await {
let wrapped: WrappedMessage = match bincode::deserialize(&msg.1) {
Ok(m) => m,
Err(_) => {
warn!("Malformed or unrecognized message format.");
continue;
}
};
match wrapped {
WrappedMessage::RanDouSha(_) => todo!(),
WrappedMessage::Rbc(msg) => {
if let Err(e) = rbc.process(msg, Arc::clone(&net_clone)).await {
warn!(error = %e, "Message processing failed");
}
}
_ => todo!(),
}
}
});
}
}
pub fn test_setup(
n: usize,
clientid: Vec<ClientId>,
) -> (
Vec<Arc<FakeNetwork>>,
Vec<Vec<Receiver<Vec<u8>>>>,
HashMap<usize, Arc<FakeNetwork>>,
HashMap<usize, Vec<Receiver<Vec<u8>>>>,
) {
let config = FakeNetworkConfig::new(500);
let (inner, receivers, client_recv) = FakeInnerNetwork::new(n, Some(clientid.clone()), config);
let network: Vec<_> = (0..n)
.map(|id| Arc::new(FakeNetwork::new(id, inner.clone())))
.collect();
let client_networks: HashMap<ClientId, Arc<FakeNetwork>> = clientid
.into_iter()
.map(|client_id| {
(
client_id,
Arc::new(FakeNetwork::new_client(client_id, inner.clone())),
)
})
.collect();
(network, receivers, client_networks, client_recv)
}
pub fn get_reconstruct_input(
n: usize,
degree_t: usize,
) -> (Fr, Vec<NonRobustShare<Fr>>, Vec<NonRobustShare<Fr>>) {
let mut rng = test_rng();
let secret = Fr::rand(&mut rng);
let shares_si_t = NonRobustShare::compute_shares(secret, n, degree_t, None, &mut rng).unwrap();
let shares_si_2t =
NonRobustShare::compute_shares(secret, n, degree_t * 2, None, &mut rng).unwrap();
(secret, shares_si_t, shares_si_2t)
}
pub fn construct_e2e_input(
n: usize,
degree_t: usize,
) -> (
Vec<Fr>,
Vec<Vec<NonRobustShare<Fr>>>,
Vec<Vec<NonRobustShare<Fr>>>,
) {
let mut n_shares_t = vec![vec![]; n];
let mut n_shares_2t = vec![vec![]; n];
let mut secrets = Vec::new();
let mut rng = test_rng();
for _ in 0..n {
let secret = Fr::rand(&mut rng);
secrets.push(secret);
let shares_si_t =
NonRobustShare::compute_shares(secret, n, degree_t, None, &mut rng).unwrap();
let shares_si_2t =
NonRobustShare::compute_shares(secret, n, degree_t * 2, None, &mut rng).unwrap();
for j in 0..n {
n_shares_t[j].push(shares_si_t[j].clone());
n_shares_2t[j].push(shares_si_2t[j].clone());
}
}
return (secrets, n_shares_t, n_shares_2t);
}
pub fn initialize_node(
node_id: usize,
n: usize,
t: usize,
k: usize,
) -> RanDouShaNode<Fr, Avid<SessionId>> {
RanDouShaNode::new(node_id, n, t, k).unwrap()
}
pub fn create_nodes(
n_parties: usize,
t: usize,
k: usize,
) -> Vec<Arc<Mutex<RanDouShaNode<Fr, Avid<SessionId>>>>> {
(0..n_parties)
.map(|id| Arc::new(Mutex::new(initialize_node(id, n_parties, t, k))))
.collect()
}
pub async fn initialize_all_nodes(
nodes: &[Arc<Mutex<RanDouShaNode<Fr, Avid<SessionId>>>>],
n_shares_t: &[Vec<NonRobustShare<Fr>>],
n_shares_2t: &[Vec<NonRobustShare<Fr>>],
session_id: SessionId,
network: Vec<Arc<FakeNetwork>>,
) {
assert!(nodes.len() == n_shares_t.len());
assert!(nodes.len() == n_shares_2t.len());
for node in nodes {
let node_locked = &mut node.lock().await;
let node_id = node_locked.id;
match node_locked
.init(
n_shares_t[node_id].clone(),
n_shares_2t[node_id].clone(),
session_id,
network[node_id].clone(),
)
.await
{
Ok(()) => (),
Err(e) => {
if let RanDouShaError::NetworkError(NetworkError::SendError) = e {
eprintln!(
"Test: Init handler for node {} got expected SendError: {:?}",
node_locked.id, e
);
} else {
panic!(
"Test: Unexpected error during init_handler for node {}: {:?}",
node_locked.id, e
);
}
}
}
}
}
pub fn spawn_receiver_tasks(
nodes: Vec<Arc<Mutex<RanDouShaNode<Fr, Avid<SessionId>>>>>,
mut receivers: Vec<Vec<Receiver<Vec<u8>>>>,
network: Vec<Arc<FakeNetwork>>,
abort_counter: Option<Arc<AtomicUsize>>,
) -> JoinSet<()> {
let mut set = JoinSet::new();
for (i, node) in nodes.iter().enumerate() {
let randousha_node = Arc::clone(&node);
let receiver = receivers.remove(0);
let net_clone = network[i].clone();
let abort_count = abort_counter.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(msg_bytes) = merge_rx.recv().await {
let wrapped: WrappedMessage = match bincode::deserialize(&msg_bytes.1) {
Ok(m) => m,
Err(_) => {
warn!("Malformed or unrecognized message format.");
continue;
}
};
match &wrapped {
WrappedMessage::RanDouSha(rds) => {
let result = randousha_node
.lock()
.await
.process(rds.clone(), Arc::clone(&net_clone))
.await;
match result {
Ok(()) => {}
Err(RanDouShaError::Abort) => {
let id = randousha_node.lock().await.id;
println!("RanDouSha aborted by node {id}");
if let Some(c) = abort_count {
c.fetch_add(1, Ordering::SeqCst);
}
break;
}
Err(RanDouShaError::NetworkError(NetworkError::SendError)) => {
eprintln!(
"Party {} encountered SendError (ignored)",
randousha_node.lock().await.id
);
continue;
}
Err(e) => {
panic!(
"Node {} encountered unexpected error: {e}",
randousha_node.lock().await.id
);
}
}
}
WrappedMessage::Rbc(msg) => {
if let Err(e) = randousha_node
.lock()
.await
.rbc
.process(msg.clone(), Arc::clone(&net_clone))
.await
{
warn!("Rbc processing error: {e}");
}
match randousha_node.lock().await.drain_rbc_output().await {
Ok(()) => {}
Err(RanDouShaError::Abort) => {
info!("RanDouSha aborted");
if let Some(c) = abort_count.clone() {
c.fetch_add(1, Ordering::SeqCst);
}
break;
}
Err(e) => {
warn!("RBC output handling error: {e}");
}
}
}
_ => todo!(),
}
}
});
}
set
}
static TRACING_INIT: Lazy<()> = Lazy::new(|| {
let subscriber = FmtSubscriber::builder()
.with_env_filter(EnvFilter::from_default_env().add_directive("info".parse().unwrap()))
.pretty()
.finish();
tracing::subscriber::set_global_default(subscriber).expect("setting default subscriber failed");
let old_hook = std::panic::take_hook();
std::panic::set_hook(Box::new(move |info| {
old_hook(info);
tracing::error!("{}", info);
std::process::exit(1);
}));
});
static QUIET_TRACING_INIT: Lazy<()> = Lazy::new(|| {
let subscriber = FmtSubscriber::builder()
.with_env_filter(EnvFilter::new(
"warn,stoffelcrypto::honeybadger=info,stoffelcrypto::honeybadger::batch_recon=warn,stoffelcrypto::honeybadger::mul=warn,stoffelcrypto::honeybadger::share_gen=warn,stoffelcrypto::honeybadger::triple_gen=warn,stoffelcrypto::honeybadger::ran_dou_sha=warn,stoffelcrypto::honeybadger::double_share=warn,stoffelcrypto::common::rbc=warn",
))
.pretty()
.finish();
tracing::subscriber::set_global_default(subscriber).expect("setting default subscriber failed");
let old_hook = std::panic::take_hook();
std::panic::set_hook(Box::new(move |info| {
old_hook(info);
tracing::error!("{}", info);
std::process::exit(1);
}));
});
pub fn setup_tracing() {
Lazy::force(&TRACING_INIT);
}
pub fn setup_quiet_tracing() {
Lazy::force(&QUIET_TRACING_INIT);
}
pub fn generate_independent_shares<F: FftField>(
secrets: &[F],
t: usize,
n: usize,
) -> Vec<Vec<RobustShare<F>>> {
let mut rng = test_rng();
let mut shares = vec![
vec![
RobustShare {
share: [F::zero()],
id: 0,
degree: t,
_sharetype: PhantomData
};
secrets.len()
];
n
];
for (j, secret) in secrets.iter().enumerate() {
let secret_shares = RobustShare::compute_shares(*secret, n, t, None, &mut rng).unwrap();
for i in 0..n {
shares[i][j] = secret_shares[i].clone(); }
}
shares
}
pub fn fan_in_inboxes(
inboxes: Vec<(SenderId, Receiver<Vec<u8>>)>,
) -> Receiver<(SenderId, Vec<u8>)> {
let (tx, rx) = mpsc::channel(300);
for (sender, mut rx_i) in inboxes {
let tx_i = tx.clone();
tokio::spawn(async move {
while let Some(msg) = rx_i.recv().await {
let _ = tx_i.send((sender, msg)).await;
}
});
}
rx
}
pub fn receive<F, R, S, N>(
mut receivers: Vec<Vec<Receiver<Vec<u8>>>>, mut nodes: Vec<HoneyBadgerMPCNode<F, R>>,
net: Vec<Arc<N>>,
client_ids: Option<Vec<ClientId>>,
) where
F: PrimeField,
R: RBC + 'static,
N: Network + Send + Sync + 'static,
S: SecretSharingScheme<F>,
HoneyBadgerMPCNode<F, R>: MPCProtocol<F, S, N>,
{
assert_eq!(
receivers.len(),
nodes.len(),
"Each node must have a receiver"
);
let n_len = nodes.len();
for i in 0..n_len {
let inbox_row = receivers.remove(0);
let mut node = nodes.remove(0);
let net_clone = net[i].clone();
let mut labeled_inboxes = Vec::with_capacity(inbox_row.len());
for (idx, rx) in inbox_row.into_iter().enumerate() {
if idx < n_len {
labeled_inboxes.push((SenderId::Node(idx), rx));
} else if let Some(ref clients) = client_ids {
let client_idx = idx - n_len;
labeled_inboxes.push((SenderId::Client(clients[client_idx]), rx));
}
}
let mut merged_rx = fan_in_inboxes(labeled_inboxes);
tokio::spawn(async move {
while let Some((sender, raw_msg)) = merged_rx.recv().await {
let id = match sender {
SenderId::Node(i) => i,
SenderId::Client(i) => i,
};
if let Err(e) = node.process(id, raw_msg, net_clone.clone()).await {
tracing::error!(
"Node {:?} failed to process message from {:?}: {:?}",
i,
sender,
e
);
}
}
tracing::info!("Receiver task for node {:?} ended", i);
});
}
}
pub fn create_global_nodes<F: PrimeField, R: RBC + 'static, S, N>(
n_parties: usize,
t: usize,
n_triples: usize,
n_random_shares: usize,
instance_id: u32,
n_prandbit: usize,
n_prandint: usize,
l: usize,
k: usize,
timeout: Duration,
input_ids: Vec<ClientId>,
) -> Vec<HoneyBadgerMPCNode<F, R>>
where
N: Network + Send + Sync + 'static,
S: SecretSharingScheme<F>,
HoneyBadgerMPCNode<F, R>: MPCProtocol<F, S, N, MPCOpts = HoneyBadgerMPCNodeOpts>,
{
let parameters = HoneyBadgerMPCNodeOpts::new(
n_parties,
t,
n_triples,
n_random_shares,
instance_id,
n_prandbit,
n_prandint,
l,
k,
timeout,
)
.unwrap();
(0..n_parties)
.map(|id| HoneyBadgerMPCNode::setup(id, parameters.clone(), input_ids.clone()).unwrap())
.collect()
}
pub async fn initialize_global_nodes_randousha<F, R, N>(
nodes: Vec<HoneyBadgerMPCNode<F, R>>,
n_shares_t: &[Vec<NonRobustShare<F>>],
n_shares_2t: &[Vec<NonRobustShare<F>>],
session_id: SessionId,
network: Vec<Arc<N>>,
) where
F: PrimeField,
R: RBC<Id = SessionId> + 'static,
N: Network + Send + Sync + 'static,
{
assert!(nodes.len() == n_shares_t.len());
assert!(nodes.len() == n_shares_2t.len());
for node in nodes {
let mut node_rds = node.preprocess.ran_dou_sha;
let node_id = node_rds.id;
match node_rds
.init(
n_shares_t[node_id].clone(),
n_shares_2t[node_id].clone(),
session_id,
network[node.id].clone(),
)
.await
{
Ok(()) => (),
Err(e) => {
if let RanDouShaError::NetworkError(NetworkError::SendError) = e {
eprintln!(
"Test: Init handler for node {} got expected SendError: {:?}",
node_id, e
);
} else {
panic!(
"Test: Unexpected error during init_handler for node {}: {:?}",
node_id, e
);
}
}
}
}
}
pub fn construct_e2e_input_ransha(
n: usize,
degree_t: usize,
) -> (Vec<Fr>, Vec<Vec<RobustShare<Fr>>>) {
let mut n_shares_t = vec![vec![]; n];
let mut secrets = Vec::new();
let mut rng = test_rng();
for _ in 0..n {
let secret = Fr::rand(&mut rng);
secrets.push(secret);
let shares_si_t = RobustShare::compute_shares(secret, n, degree_t, None, &mut rng).unwrap();
for j in 0..n {
n_shares_t[j].push(shares_si_t[j].clone());
}
}
return (secrets, n_shares_t);
}
pub async fn initialize_global_nodes_ransha<F, R, N>(
nodes: Vec<HoneyBadgerMPCNode<F, R>>,
session_id: SessionId,
network: Vec<Arc<N>>,
) where
F: PrimeField,
R: RBC<Id = SessionId> + 'static,
N: Network + Send + Sync + 'static,
{
let mut rng = StdRng::from_rng(OsRng).unwrap();
for node in nodes {
let mut node_rds = node.preprocess.share_gen;
let node_id = node_rds.id;
match node_rds
.init(session_id, &mut rng, network[node_id].clone())
.await
{
Ok(()) => (),
Err(e) => {
if let RanShaError::NetworkError(NetworkError::SendError) = e {
eprintln!(
"Test: Init handler for node {} got expected SendError: {:?}",
node_id, e
);
} else {
panic!(
"Test: Unexpected error during init_handler for node {}: {:?}",
node_id, e
);
}
}
}
}
}
pub fn construct_e2e_input_mul<F: PrimeField + FftField>(
n_parties: usize,
n_triples: usize,
threshold: usize,
) -> ((Vec<F>, Vec<F>, Vec<F>), Vec<Vec<ShamirBeaverTriple<F>>>) {
let mut rng = test_rng();
let mut secrets_a = Vec::new();
let mut secrets_b = Vec::new();
let mut secrets_c = Vec::new();
let mut per_party_triples: Vec<Vec<ShamirBeaverTriple<F>>> = vec![Vec::new(); n_parties];
for _i in 0..n_triples {
let a_secret = F::rand(&mut rng);
let b_secret = F::rand(&mut rng);
let c_secret = a_secret * b_secret;
let shares_a = RobustShare::compute_shares(a_secret, n_parties, threshold, None, &mut rng)
.expect("share a creation failed");
let shares_b = RobustShare::compute_shares(b_secret, n_parties, threshold, None, &mut rng)
.expect("share b creation failed");
let shares_c = RobustShare::compute_shares(c_secret, n_parties, threshold, None, &mut rng)
.expect("share c creation failed");
secrets_a.push(a_secret);
secrets_b.push(b_secret);
secrets_c.push(c_secret);
for pid in 0..n_parties {
let triple = ShamirBeaverTriple {
a: shares_a[pid].clone(),
b: shares_b[pid].clone(),
mult: shares_c[pid].clone(),
};
per_party_triples[pid].push(triple);
}
}
((secrets_a, secrets_b, secrets_c), per_party_triples)
}
pub fn create_clients<F: FftField, R: RBC<Id = SessionId> + 'static>(
client_ids: Vec<ClientId>,
n_parties: usize,
t: usize,
instance_id: u32,
inputs: Vec<F>,
input_len: usize,
) -> HashMap<ClientId, HoneyBadgerMPCClient<F, R>> {
client_ids
.into_iter()
.map(|id| {
let client =
HoneyBadgerMPCClient::new(id, n_parties, t, instance_id, inputs.clone(), input_len)
.unwrap();
(id, client)
})
.collect()
}
pub fn receive_client<F, R, N>(
mut receivers: HashMap<ClientId, Vec<Receiver<Vec<u8>>>>,
clients: HashMap<ClientId, HoneyBadgerMPCClient<F, R>>,
mut net: HashMap<usize, Arc<N>>,
) where
F: FftField + 'static,
R: RBC<Id = SessionId> + 'static,
N: Network + Send + Sync + 'static,
{
assert_eq!(
receivers.len(),
clients.len(),
"Each node must have a receiver"
);
for (clientid, inbox_vec) in receivers.drain() {
let mut client = clients[&clientid].clone();
let net_clone = net.remove(&clientid).unwrap();
let inbox: Vec<(SenderId, tokio::sync::mpsc::Receiver<Vec<u8>>)> = inbox_vec
.into_iter()
.enumerate()
.map(|(from, r)| (SenderId::Node(from), r))
.collect();
let merged_rx = fan_in_inboxes(inbox);
tokio::spawn(async move {
let mut merged_rx = merged_rx;
while let Some((sender, received)) = merged_rx.recv().await {
if let SenderId::Node(s) = sender {
if let Err(e) = client.process(s, received, net_clone.clone()).await {
tracing::error!(
"Client {} failed processing from {:?}: {:?}",
clientid,
sender,
e
);
}
}
}
tracing::info!("Receiver task for client {clientid} ended");
});
}
}