use crate::{
common::{
share::{shamir::NonRobustShare, ShareError},
utils::deser_bounded_vec,
ProtocolSessionId, SecretSharingScheme,
},
honeybadger::{
double_share::{DouShaError, DouShaMessage, DouShaPayload, DouShaStorage},
SessionId, WrappedMessage,
},
};
use ark_ff::FftField;
use ark_serialize::{CanonicalDeserialize, CanonicalSerialize};
use ark_std::rand::Rng;
use itertools::izip;
use std::{collections::BTreeMap, sync::Arc};
use stoffelnet::network_utils::{Network, PartyId};
use tokio::{
sync::Mutex,
time::{timeout, Duration},
};
use tracing::{info, warn};
use super::DoubleShamirShare;
#[derive(Clone, PartialEq, Debug)]
pub enum ProtocolState {
Initialized,
Finished,
NotInitialized,
}
#[derive(Clone, Debug)]
pub struct DoubleShareNode<F>
where
F: FftField,
{
pub id: PartyId,
pub n_parties: usize,
pub threshold: usize,
pub storage: Arc<Mutex<BTreeMap<SessionId, (usize, Arc<Mutex<DouShaStorage<F>>>)>>>,
}
pub static MAX_DOUSHA_SESSIONS: usize = 256;
impl<F> DoubleShareNode<F>
where
F: FftField,
{
pub async fn process(&mut self, message: DouShaMessage) -> Result<(), DouShaError> {
self.receive_double_shares_handler(message).await?;
Ok(())
}
pub fn new(id: PartyId, n_parties: usize, threshold: usize) -> Self {
Self {
id,
n_parties,
threshold,
storage: Arc::new(Mutex::new(BTreeMap::new())),
}
}
pub async fn get_or_create_store(
&mut self,
session_id: SessionId,
initiator_id: usize,
) -> Result<Arc<Mutex<DouShaStorage<F>>>, DouShaError> {
let mut storage = self.storage.lock().await;
Ok(storage
.entry(session_id)
.or_insert((
initiator_id,
Arc::new(Mutex::new(DouShaStorage::empty(self.n_parties))),
))
.1
.clone())
}
pub async fn clear_store(&self, session_id: SessionId) -> bool {
let mut store = self.storage.lock().await;
store.remove(&session_id).is_some()
}
pub async fn store_len(&self) -> usize {
self.storage.lock().await.len()
}
pub async fn wait_for_result(
&self,
session_id: SessionId,
duration: Duration,
) -> Result<Vec<DoubleShamirShare<F>>, DouShaError> {
let output_receiver = {
let storage = self.storage.lock().await;
let storage_bind = match storage.get(&session_id) {
Some((_, arc)) => arc,
None => return Err(DouShaError::NoSuchSessionId(session_id)),
};
let mut storage = storage_bind.lock().await;
storage
.output_receiver
.take()
.ok_or(DouShaError::ResultAlreadyReceived(session_id))?
};
match timeout(duration, output_receiver).await {
Err(_) => Err(DouShaError::Timeout(session_id)),
Ok(Err(_)) => Err(DouShaError::ReceiveError(session_id)),
Ok(Ok(shares)) => Ok(shares),
}
}
pub async fn init<N, R>(
&mut self,
session_id: SessionId,
rng: &mut R,
network: Arc<N>,
) -> Result<(), DouShaError>
where
N: Network,
R: Rng,
{
self.init_batch(session_id, 1, rng, network).await
}
pub async fn init_batch<N, R>(
&mut self,
session_id: SessionId,
batch_size: usize,
rng: &mut R,
network: Arc<N>,
) -> Result<(), DouShaError>
where
N: Network,
R: Rng,
{
info!("Receiving init for faulty double share from {0:?}", self.id);
let batch_size = batch_size.max(1);
let mut shares_by_recipient = vec![Vec::with_capacity(batch_size); self.n_parties];
for _ in 0..batch_size {
let secret = F::rand(rng);
let shares_deg_t =
NonRobustShare::compute_shares(secret, self.n_parties, self.threshold, None, rng)?;
let shares_deg_2t = NonRobustShare::compute_shares(
secret,
self.n_parties,
2 * self.threshold,
None,
rng,
)?;
for (recipient_id, (share_t, share_2t)) in
izip!(shares_deg_t, shares_deg_2t).enumerate()
{
shares_by_recipient[recipient_id].push(DoubleShamirShare::new(share_t, share_2t));
}
}
for (recipient_id, double_shares) in shares_by_recipient.into_iter().enumerate() {
let mut payload = Vec::new();
let payload = if batch_size == 1 {
double_shares[0].serialize_compressed(&mut payload)?;
DouShaPayload::Share(payload)
} else {
double_shares.serialize_compressed(&mut payload)?;
DouShaPayload::Shares(payload)
};
let generic_message =
WrappedMessage::Dousha(DouShaMessage::new(self.id, session_id, payload));
let bytes_generic_msg = bincode::serialize(&generic_message)?;
info!(
"sending double shares from {:?} to {:?}",
self.id, recipient_id
);
network.send(recipient_id, &bytes_generic_msg).await?;
}
let storage_access = self.get_or_create_store(session_id, self.id).await?;
let mut storage = storage_access.lock().await;
storage.batch_size = batch_size;
storage.state = ProtocolState::Initialized;
Ok(())
}
pub async fn receive_double_shares_handler(
&mut self,
recv_message: DouShaMessage,
) -> Result<(), DouShaError> {
let double_shares: Vec<DoubleShamirShare<F>> = match recv_message.payload {
DouShaPayload::Share(payload) => {
vec![CanonicalDeserialize::deserialize_compressed(
payload.as_slice(),
)?]
}
DouShaPayload::Shares(payload) => {
deser_bounded_vec(&mut payload.as_slice(), payload.len())?
}
};
for double_share in &double_shares {
if double_share.degree_t.id != self.id || double_share.degree_2t.id != self.id {
return Err(ShareError::IdMismatch.into());
}
if double_share.degree_t.degree != self.threshold {
return Err(ShareError::DegreeMismatch.into());
}
if double_share.degree_2t.degree != 2 * self.threshold {
return Err(ShareError::DegreeMismatch.into());
}
}
let binding = self
.get_or_create_store(recv_message.session_id, recv_message.sender_id)
.await?;
let mut dousha_storage = binding.lock().await;
if dousha_storage.share.is_empty() {
dousha_storage.batch_size = double_shares.len();
} else if dousha_storage.batch_size != double_shares.len() {
return Err(DouShaError::ShareError(ShareError::DegreeMismatch));
}
if dousha_storage.state == ProtocolState::Finished {
return Ok(());
}
if dousha_storage.share.contains_key(&recv_message.sender_id) {
warn!(
session_id = recv_message.session_id.as_u128(),
"Duplicate double share received from party {:?}, ignoring.",
recv_message.sender_id
);
return Ok(()); }
if recv_message.sender_id >= self.n_parties {
return Err(DouShaError::InvalidPartyId);
}
dousha_storage
.share
.insert(recv_message.sender_id, double_shares);
info!(
session_id = recv_message.session_id.as_u128(),
"party {:?} received double shares from {:?}", self.id, recv_message.sender_id,
);
dousha_storage.reception_tracker[recv_message.sender_id] = true;
if dousha_storage
.reception_tracker
.iter()
.all(|&received| received)
{
let mut output =
Vec::with_capacity(dousha_storage.batch_size * dousha_storage.share.len());
for batch_index in 0..dousha_storage.batch_size {
for shares in dousha_storage.share.values() {
output.push(shares[batch_index].clone());
}
}
dousha_storage.protocol_output = output.clone();
dousha_storage.state = ProtocolState::Finished;
let taken_output_sender = dousha_storage.output_sender.take().unwrap();
taken_output_sender
.send(output)
.map_err(|_| DouShaError::SendError(recv_message.session_id))?;
}
Ok(())
}
}