use crate::common::SecretSharingScheme;
use crate::honeybadger::output::{OutputError, OutputMessage};
use crate::honeybadger::robust_interpolate::robust_interpolate::RobustShare;
use crate::honeybadger::WrappedMessage;
use ark_ff::FftField;
use ark_serialize::CanonicalSerialize;
use std::collections::HashMap;
use std::sync::Arc;
use stoffelnet::network_utils::Network;
use tokio::sync::watch::{channel, Receiver, Sender};
use tokio::time::{timeout, Duration};
use tracing::info;
#[derive(Clone, Debug)]
pub struct OutputServer {
pub id: usize,
pub n: usize,
}
impl OutputServer {
pub fn new(id: usize, n: usize) -> Result<Self, OutputError> {
Ok(Self { id, n })
}
pub async fn init<N: Network, F: FftField>(
&self,
client_id: usize,
shares: Vec<RobustShare<F>>,
input_len: usize,
net: Arc<N>,
) -> Result<(), OutputError> {
if shares.len() != input_len {
return Err(OutputError::InvalidInput(
"Incorrect number of shares".to_string(),
));
}
let mut payload = Vec::new();
shares.serialize_compressed(&mut payload)?;
let msg = OutputMessage::new(self.id, payload);
let wrapped = WrappedMessage::Output(msg);
let bytes = bincode::serialize(&wrapped)?;
net.send_to_client(client_id, &bytes).await?;
info!(
"Server {} sent output share to client {}",
self.id, client_id
);
Ok(())
}
}
pub struct OutputClientData<F: FftField> {
pub output: Option<Vec<F>>,
pub output_shares: HashMap<usize, Vec<RobustShare<F>>>,
}
#[derive(Clone)]
pub struct OutputClient<F: FftField> {
pub client_id: usize,
pub n: usize,
pub t: usize,
pub input_len: usize,
pub output_sender: Sender<OutputClientData<F>>,
pub output_receiver: Receiver<OutputClientData<F>>,
}
impl<F: FftField> OutputClient<F> {
pub fn new(id: usize, n: usize, t: usize, input_len: usize) -> Result<Self, OutputError> {
let (output_sender, output_receiver) = channel(OutputClientData::<F> {
output: None,
output_shares: HashMap::new(),
});
Ok(Self {
client_id: id,
n,
t,
input_len,
output_sender,
output_receiver,
})
}
pub async fn output_handler(&mut self, msg: OutputMessage) -> Result<(), OutputError> {
if msg.payload.len() < 8 {
return Err(OutputError::InvalidInput("Payload too short".to_string()));
}
let declared_len = u64::from_le_bytes(msg.payload[..8].try_into().unwrap()) as usize;
if declared_len != self.input_len {
return Err(OutputError::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(OutputError::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(OutputError::InvalidInput(
"Invalid share degree".to_string(),
));
}
}
if shares.len() != self.input_len {
return Err(OutputError::InvalidInput(
"Mismatch in input and share length".to_string(),
));
}
let mut already_recvd = false;
let mut recovery_err = None;
self.output_sender.send_if_modified(|output_data| {
let share_store = &mut output_data.output_shares;
if share_store.contains_key(&msg.sender_id) {
already_recvd = true;
return false;
}
share_store.insert(msg.sender_id, shares.clone());
info!("Received Output share from server {}", msg.sender_id);
let mut r_shares = vec![vec![]; self.input_len];
if output_data.output.is_none() && share_store.len() >= 2 * self.t + 1 {
info!("Received enough shares to reconstruct");
for (_, r_share) in share_store.iter() {
for i in 0..self.input_len {
r_shares[i].push(r_share[i].clone());
}
}
let mut output = Vec::new();
for output_elem in r_shares {
let secret = match RobustShare::recover_secret(&output_elem, self.n, self.t) {
Ok(secret) => secret,
Err(e) => {
recovery_err = Some(e);
return false;
}
};
output.push(secret.1);
}
output_data.output = Some(output);
return true;
}
false
});
if already_recvd {
return Err(OutputError::Duplicate(format!(
"Already received from {}",
msg.sender_id
)));
}
if let Some(err) = recovery_err {
return Err(OutputError::InterpolateError(err));
}
Ok(())
}
pub async fn wait_for_output(&mut self, duration: Duration) -> Result<Vec<F>, OutputError> {
let output_future = self
.output_receiver
.wait_for(|output_data| output_data.output.is_some());
match timeout(duration, output_future).await {
Err(elapsed_err) => Err(OutputError::Timeout(elapsed_err)),
Ok(Err(recv_err)) => Err(OutputError::WaitingError(recv_err)),
Ok(Ok(output_data)) => Ok(output_data.output.as_ref().unwrap().clone()),
}
}
pub fn get_output(&self) -> Option<Vec<F>> {
let output_data = self.output_receiver.borrow();
output_data.output.clone()
}
pub async fn process(&mut self, msg: OutputMessage) -> Result<(), OutputError> {
self.output_handler(msg).await?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::honeybadger::output::OutputMessage;
use crate::honeybadger::robust_interpolate::robust_interpolate::RobustShare;
use ark_bls12_381::Fr;
use ark_ff::UniformRand;
use ark_std::test_rng;
use tokio::time::Duration;
#[tokio::test]
async fn test_get_output() {
let n = 5;
let t = 1;
let input_len = 1;
let client_id = 7;
let mut rng = test_rng();
let mut client = OutputClient::<Fr>::new(client_id, n, t, input_len).unwrap();
let secret = Fr::rand(&mut rng);
let shares_vec = RobustShare::compute_shares(secret, n, t, None, &mut rng).unwrap();
for i in 0..2 {
let mut payload = Vec::new();
vec![shares_vec[i].clone()]
.serialize_compressed(&mut payload)
.unwrap();
let msg = OutputMessage::new(i, payload);
client.output_handler(msg).await.unwrap();
}
assert_eq!(client.get_output(), None);
let mut payload = Vec::new();
vec![shares_vec[2].clone()]
.serialize_compressed(&mut payload)
.unwrap();
let msg = OutputMessage::new(2, payload);
client.output_handler(msg).await.unwrap();
assert_eq!(client.get_output(), Some(vec![secret]));
}
#[tokio::test]
async fn test_wait_for_output() {
let n = 5;
let t = 1;
let input_len = 1;
let client_id = 7;
let mut rng = test_rng();
let mut client = OutputClient::<Fr>::new(client_id, n, t, input_len).unwrap();
let secret = Fr::rand(&mut rng);
let shares_vec = RobustShare::compute_shares(secret, n, t, None, &mut rng).unwrap();
for i in 0..2 {
let mut payload = Vec::new();
vec![shares_vec[i].clone()]
.serialize_compressed(&mut payload)
.unwrap();
let msg = OutputMessage::new(i, payload);
client.output_handler(msg).await.unwrap();
}
let result = client.wait_for_output(Duration::from_millis(10)).await;
assert!(
result.is_err(),
"Expected timeout error when only 2 shares are sent"
);
let mut payload = Vec::new();
vec![shares_vec[2].clone()]
.serialize_compressed(&mut payload)
.unwrap();
let msg = OutputMessage::new(2, payload);
client.output_handler(msg).await.unwrap();
let result2 = client.wait_for_output(Duration::from_millis(10)).await;
assert!(
result2.is_ok(),
"Expected output to be reconstructed after enough shares"
);
assert_eq!(result2.unwrap(), vec![secret]);
}
}