use crate::common::{ProtocolSessionId, SecretSharingScheme, RBC};
use crate::honeybadger::input::InputError;
use crate::honeybadger::input::InputMessage;
use crate::honeybadger::robust_interpolate::robust_interpolate::RobustShare;
use crate::honeybadger::{ProtocolType, SessionId, WrappedMessage, MAX_MESSAGE_SIZE};
use ark_ff::FftField;
use ark_serialize::CanonicalSerialize;
use bincode::Options;
use std::collections::HashMap;
use std::sync::Arc;
use stoffelnet::network_utils::{ClientId, Network, PartyId};
use tokio::{
sync::{
watch::{channel, Receiver, Sender},
Mutex,
},
time::{timeout, Duration},
};
use tracing::{info, warn};
const MAX_INPUT_ELEMENTS: u64 = 65_536;
#[derive(PartialEq, Clone, Debug)]
pub enum InputType {
Empty,
RandomShares,
MaskedInputs,
InputShares,
}
#[derive(Clone, Debug)]
pub struct InputServer<F: FftField, R: RBC> {
pub id: usize,
pub n: usize,
pub rbc: R,
pub rbc_output: Arc<Mutex<tokio::sync::mpsc::Receiver<SessionId>>>,
status_sender: Sender<HashMap<ClientId, (InputType, Vec<RobustShare<F>>)>>,
pub status_receiver: Receiver<HashMap<ClientId, (InputType, Vec<RobustShare<F>>)>>,
}
fn calculate_input_shares<F: FftField>(
masked_inputs: &[RobustShare<F>],
random_shares: &Vec<RobustShare<F>>,
) -> Vec<RobustShare<F>> {
masked_inputs
.iter()
.zip(random_shares)
.map(|(masked_input, random_share)| {
RobustShare::new(
masked_input.share[0] - random_share.share[0],
random_share.id,
random_share.degree,
)
})
.collect()
}
impl<F: FftField, R: RBC<Id = SessionId>> InputServer<F, R> {
pub fn new(
id: usize,
n: usize,
t: usize,
input_ids: Vec<ClientId>,
) -> Result<Self, InputError> {
let (rbc_sender, rbc_receiver) = tokio::sync::mpsc::channel(200);
let rbc = R::new(
id,
n,
t,
t + 1,
rbc_sender,
Arc::new(WrappedMessage::rbc_wrap),
)?;
let (status_sender, status_receiver) = channel(
input_ids
.into_iter()
.map(|id| (id, (InputType::Empty, vec![])))
.collect(),
);
Ok(Self {
id,
n,
rbc,
rbc_output: Arc::new(Mutex::new(rbc_receiver)),
status_sender,
status_receiver,
})
}
pub async fn drain_rbc_output(&mut self) -> Result<(), InputError> {
loop {
let id = {
let mut rx = self.rbc_output.lock().await;
match rx.try_recv() {
Ok(id) => id,
Err(tokio::sync::mpsc::error::TryRecvError::Empty) => break,
Err(tokio::sync::mpsc::error::TryRecvError::Disconnected) => {
return Err(InputError::Abort);
}
}
};
let output = self.rbc.get_store(id).await?;
let msg: InputMessage = bincode::DefaultOptions::new()
.with_fixint_encoding()
.allow_trailing_bytes()
.with_limit(MAX_MESSAGE_SIZE)
.deserialize(&output)?;
let authenticated_sender = id.sub_id() as usize;
if msg.sender_id != authenticated_sender {
warn!(
"Dropping RBC output: inner sender_id {} does not match session sub_id {}",
msg.sender_id, authenticated_sender
);
continue;
}
match self.input_handler(authenticated_sender, msg.payload).await {
Ok(()) => {}
Err(e) => {
return Err(e);
}
}
}
Ok(())
}
pub async fn init<N: Network>(
&mut self,
client_id: usize,
shares: Vec<RobustShare<F>>,
input_len: usize,
net: Arc<N>,
) -> Result<(), InputError> {
if shares.len() != input_len {
return Err(InputError::InvalidInput(
"Incorrect number of shares".to_string(),
));
}
let mut send_over_network = false;
let mut already_rand_shares = false;
let mut unknown_client = false;
let mut invalid_length = false;
self.status_sender.send_if_modified(|status| {
match status.get(&client_id) {
Some((InputType::RandomShares | InputType::InputShares, _)) => {
already_rand_shares = true;
false
}
Some((InputType::MaskedInputs, masked_inputs)) => {
if masked_inputs.len() != shares.len() {
invalid_length = true;
return false;
}
let input_shares = calculate_input_shares(masked_inputs, &shares);
status.insert(client_id, (InputType::InputShares, input_shares));
info!("Calculated inputs for client {}", client_id);
true
}
Some((InputType::Empty, _)) => {
status.insert(client_id, (InputType::RandomShares, shares.clone()));
info!("Stored local mask shares for client {}", client_id);
send_over_network = true;
true
}
None => {
unknown_client = true;
false
}
}
});
if invalid_length {
return Err(InputError::InvalidInput(
"Mismatch in masked input and share length".to_string(),
));
}
if unknown_client {
return Err(InputError::InvalidInput(
"Unknown client {client_id}".to_string(),
));
}
if already_rand_shares {
return Err(InputError::Duplicate(format!(
"random shares already obtained for client {}",
client_id
)));
}
if send_over_network {
let mut payload = Vec::new();
shares.serialize_compressed(&mut payload)?;
let msg = InputMessage::new(self.id, payload);
let wrapped = WrappedMessage::Input(msg);
let bytes = bincode::serialize(&wrapped)?;
net.send_to_client(client_id, &bytes).await?;
info!("Server {} sent MaskShare to client {}", self.id, client_id);
}
Ok(())
}
pub async fn input_handler(
&mut self,
sender_id: PartyId,
payload: Vec<u8>,
) -> Result<(), InputError> {
info!(
"Server {} received MaskedInput from client {}",
self.id, sender_id
);
let masked_inputs_as_shares: Vec<RobustShare<F>> = {
if payload.len() < 8 {
return Err(InputError::InvalidInput("Payload too short".to_string()));
}
let declared_len = u64::from_le_bytes(payload[..8].try_into().unwrap());
if declared_len > MAX_INPUT_ELEMENTS {
return Err(InputError::InvalidInput(
"Declared input length exceeds maximum".to_string(),
));
}
let masked_inputs: Vec<F> =
ark_serialize::CanonicalDeserialize::deserialize_compressed(payload.as_slice())?;
masked_inputs
.iter()
.map(|m| RobustShare::new(*m, 0, 0))
.collect()
};
let mut unknown_client = false;
let mut already_masked_inputs = false;
let mut invalid_length = false;
self.status_sender
.send_if_modified(|status| match status.get(&sender_id) {
Some((InputType::MaskedInputs | InputType::InputShares, _)) => {
already_masked_inputs = true;
false
}
Some((InputType::RandomShares, random_shares)) => {
if masked_inputs_as_shares.len() != random_shares.len() {
invalid_length = true;
return false;
}
let input_shares =
calculate_input_shares(&masked_inputs_as_shares, random_shares);
status.insert(sender_id, (InputType::InputShares, input_shares));
info!(
"Server {} stored input shares from client {}",
self.id, sender_id
);
true
}
Some((InputType::Empty, _)) => {
status.insert(
sender_id,
(InputType::MaskedInputs, masked_inputs_as_shares),
);
info!(
"Server {} stored masked inputs from client {}",
self.id, sender_id
);
true
}
None => {
unknown_client = true;
false
}
});
if invalid_length {
return Err(InputError::InvalidInput(
"Mismatch in masked input and share length".to_string(),
));
}
if already_masked_inputs {
return Err(InputError::Duplicate(format!(
"Server {} already received masked inputs from {}",
self.id, sender_id
)));
}
if unknown_client {
return Err(InputError::InvalidInput(
"Unknown client {client_id}".to_string(),
));
}
Ok(())
}
pub async fn wait_for_all_inputs(
&mut self,
duration: Duration,
) -> Result<HashMap<ClientId, Vec<RobustShare<F>>>, InputError> {
let status_future = self.status_receiver.wait_for(|statuses| {
statuses
.iter()
.map(|(_, (status, _))| status)
.all(|status| *status == InputType::InputShares)
});
match timeout(duration, status_future).await {
Err(elapsed_err) => Err(InputError::Timeout(elapsed_err)),
Ok(Err(recv_err)) => Err(InputError::WaitingError(recv_err)),
Ok(Ok(statuses)) => {
info!("Server {} has inputs from all clients", self.id);
let input_shares = statuses
.iter()
.map(|(id, (_, shares))| (*id, shares.clone()))
.collect();
Ok(input_shares)
}
}
}
}
pub struct InputClientData<F: FftField, R: RBC> {
pub rbc: R,
pub inputs: Vec<F>,
pub rbc_done: bool,
pub received_shares: HashMap<usize, Vec<RobustShare<F>>>,
}
pub struct InputClient<F: FftField, R: RBC> {
pub client_id: usize,
pub n: usize,
pub t: usize,
pub instance_id: u32,
pub client_data: Arc<Mutex<InputClientData<F, R>>>,
}
impl<F: FftField, R: RBC> Clone for InputClient<F, R> {
fn clone(&self) -> Self {
Self {
client_id: self.client_id,
n: self.n,
t: self.t,
instance_id: self.instance_id,
client_data: Arc::clone(&self.client_data),
}
}
}
impl<F: FftField, R: RBC<Id = SessionId>> InputClient<F, R> {
pub fn new(
id: usize,
n: usize,
t: usize,
instance_id: u32,
inputs: Vec<F>,
) -> Result<Self, InputError> {
let (rbc_sender, _) = tokio::sync::mpsc::channel(200);
let rbc = R::new(
id,
n,
t,
t + 1,
rbc_sender,
Arc::new(WrappedMessage::rbc_wrap),
)?;
Ok(Self {
client_id: id,
n,
t,
instance_id,
client_data: Arc::new(Mutex::new(InputClientData::<F, R> {
rbc,
inputs,
received_shares: HashMap::new(),
rbc_done: false,
})),
})
}
pub async fn init_handler<N: Network + Send + Sync>(
&self,
msg: InputMessage,
net: Arc<N>,
) -> Result<(), InputError> {
let mut d = self.client_data.lock().await;
let input_len = d.inputs.len();
if msg.payload.len() < 8 {
return Err(InputError::InvalidInput("Payload too short".to_string()));
}
let declared_len = u64::from_le_bytes(msg.payload[..8].try_into().unwrap()) as usize;
if declared_len != input_len {
return Err(InputError::InvalidInput(
"Mismatch in input and share length".to_string(),
));
}
let mut shares: Vec<RobustShare<F>> =
ark_serialize::CanonicalDeserialize::deserialize_compressed(msg.payload.as_slice())?;
if !shares.iter().all(|s| s.id == msg.sender_id) {
return Err(InputError::InvalidInput(
"Share ID does not match authenticated sender".into(),
));
}
for share in &mut shares {
share.id = msg.sender_id;
if share.degree != self.t {
return Err(InputError::InvalidInput("Invalid share degree".to_string()));
}
}
if d.rbc_done {
return Ok(());
}
if d.received_shares.contains_key(&msg.sender_id) {
return Err(InputError::Duplicate(format!(
"Already random shares received from {}",
msg.sender_id
)));
}
if d.received_shares.len() == self.n {
return Err(InputError::InvalidInput(format!(
"Cannot receive from more than {} parties",
self.n
)));
}
d.received_shares.insert(msg.sender_id, shares.clone());
info!(
"Client {} received MaskShare from server {}",
self.client_id, msg.sender_id
);
let mut r_shares = vec![vec![]; input_len];
let mut masks = vec![];
let mut output = vec![];
if d.received_shares.len() >= 2 * self.t + 1 {
info!("Received enough shares to reconstruct");
for (_, r_share) in d.received_shares.iter() {
for i in 0..input_len {
r_shares[i].push(r_share[i].clone());
}
}
for recon in r_shares {
let secret = RobustShare::recover_secret(&recon, self.n, self.t)?;
masks.push(secret.1);
}
for (i, r) in masks.iter().enumerate() {
output.push(d.inputs[i] + r);
}
let mut payload = Vec::new();
output.serialize_compressed(&mut payload)?;
let msg = InputMessage::new(self.client_id, payload);
let bytes = bincode::serialize(&msg)?;
let sessionid = SessionId::new(
ProtocolType::Input,
SessionId::pack_slot(
0, self.client_id as u8,
0,
),
self.instance_id,
);
d.rbc.init(bytes, sessionid, net).await?;
d.rbc_done = true;
info!(
"Client {} initialized broadcasting of masked input to all servers",
self.client_id
);
}
Ok(())
}
pub async fn process<N: Network + Send + Sync>(
&mut self,
msg: InputMessage,
net: Arc<N>,
) -> Result<(), InputError> {
self.init_handler(msg, net).await
}
}
#[cfg(test)]
pub mod tests {
use super::*;
use crate::{
common::{rbc::rbc::Avid, SecretSharingScheme},
honeybadger::{robust_interpolate::robust_interpolate::RobustShare, WrappedMessage},
};
use ark_bls12_381::Fr;
use ark_std::test_rng;
use stoffelmpc_network::fake_network::{
FakeInnerNetwork, FakeNetwork, FakeNetworkConfig, SenderId,
};
use tokio::{
sync::mpsc,
time::{sleep, Duration},
};
pub fn fan_in_inboxes(
inboxes: Vec<(SenderId, tokio::sync::mpsc::Receiver<Vec<u8>>)>,
) -> tokio::sync::mpsc::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
}
#[tokio::test]
async fn test_init_before_input_handler() {
let n = 4;
let t = 1;
let clientid = 100;
let rand_secret = Fr::from(1);
let input = Fr::from(10);
let config = FakeNetworkConfig::new(500);
let (net, mut receivers, mut client_recv_map) =
FakeInnerNetwork::new(n, Some(vec![clientid]), config);
let client_inboxes = client_recv_map.remove(&clientid).unwrap();
let inbox: Vec<(SenderId, tokio::sync::mpsc::Receiver<Vec<u8>>)> = client_inboxes
.into_iter()
.enumerate()
.map(|(i, r)| (SenderId::Node(i), r))
.collect();
let mut client_recv = fan_in_inboxes(inbox);
let network: Vec<_> = (0..n)
.map(|id| Arc::new(FakeNetwork::new(id, net.clone())))
.collect();
let client_network: Arc<FakeNetwork> =
Arc::new(FakeNetwork::new_client(clientid, net.clone()));
let mut rng = test_rng();
let rand_shares = RobustShare::compute_shares(rand_secret, n, t, None, &mut rng).unwrap();
let mut client =
InputClient::<Fr, Avid<SessionId>>::new(clientid, n, t, 111, vec![input].clone())
.unwrap();
let mut nodes: Vec<_> = (0..n)
.map(|i| InputServer::<Fr, Avid<SessionId>>::new(i, n, t, vec![clientid]).unwrap())
.collect();
for i in 0..nodes.len() - 1 {
assert!(nodes[i]
.init(
clientid,
vec![rand_shares[i].clone()],
1,
network[i].clone()
)
.await
.is_ok());
let status = nodes[i].status_receiver.borrow();
let client_status = status.get(&clientid);
assert!(client_status.is_some() && client_status.unwrap().0 == InputType::RandomShares);
}
{
let status = nodes[3].status_receiver.borrow();
let client_status = status.get(&clientid);
assert!(client_status.is_some() && client_status.unwrap().0 == InputType::Empty);
}
for _ in 0..3 {
let (_, raw) = client_recv.recv().await.unwrap();
let wrapped: WrappedMessage =
bincode::deserialize(&raw).expect("deserialization error");
match wrapped {
WrappedMessage::Input(msg) => {
assert!(client.process(msg, client_network.clone()).await.is_ok());
}
_ => panic!("Unexpected message"),
}
}
for (i, node) in nodes.iter_mut().enumerate() {
let network = network.clone();
let mut node = node.clone();
let receiver = receivers.remove(0);
let inbox: Vec<(SenderId, tokio::sync::mpsc::Receiver<Vec<u8>>)> = receiver
.into_iter() .enumerate()
.map(|(i, r)| (SenderId::Node(i), r))
.collect();
let mut merged_rx = fan_in_inboxes(inbox);
tokio::spawn(async move {
while let Some(raw_msg) = merged_rx.recv().await {
let wrapped: WrappedMessage =
bincode::deserialize(&raw_msg.1).expect("deserialization error");
let _ = match wrapped {
WrappedMessage::Rbc(rbc_msg) => {
let _ = node.rbc.process(rbc_msg, network[i].clone()).await;
let _ = node.drain_rbc_output().await;
}
_ => {
panic!();
}
};
}
});
}
sleep(Duration::from_millis(200)).await;
{
let status = nodes[3].status_receiver.borrow();
let client_status = status.get(&clientid);
assert!(client_status.is_some() && client_status.unwrap().0 == InputType::MaskedInputs);
}
nodes[3]
.init(
clientid,
vec![rand_shares[3].clone()],
1,
network[3].clone(),
)
.await
.unwrap();
{
let status = nodes[3].status_receiver.borrow();
let client_status = status.get(&clientid);
assert!(client_status.is_some() && client_status.unwrap().0 == InputType::InputShares);
}
let mut recovered_shares = vec![];
for node in &mut nodes {
let shares = node
.wait_for_all_inputs(Duration::from_millis(1))
.await
.expect("input error");
let client_shares = shares.get(&clientid).unwrap();
recovered_shares.push(client_shares[0].clone());
}
let (_, r) = RobustShare::recover_secret(&recovered_shares, n, t).unwrap();
assert_eq!(r, input);
}
}