stoffelcrypto 0.1.0

Asynchronous HoneyBadgerMPC protocols, preprocessing, and arithmetic for Stoffel.
Documentation
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,
}

/// Node participating in a non-robust double share protocol.
#[derive(Clone, Debug)]
pub struct DoubleShareNode<F>
where
    F: FftField,
{
    /// ID of the party.
    pub id: PartyId,
    /// Number of parties participating in the protocol.
    pub n_parties: usize,
    /// Threshold for the corrupted parties.
    pub threshold: usize,
    /// Storage of the party.
    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(())
    }

    /// Creates a new node for the faulty double share protocol.
    pub fn new(id: PartyId, n_parties: usize, threshold: usize) -> Self {
        Self {
            id,
            n_parties,
            threshold,
            storage: Arc::new(Mutex::new(BTreeMap::new())),
        }
    }

    /// Returns the storage for a node in the Random Double Sharing protocol. If the storage has
    /// not been created yet, the function will create an empty storage and return it.
    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;

        // TODO: restore session limits
        // if !storage.contains_key(&session_id) {
        //     if storage.len() >= MAX_DOUSHA_SESSIONS {
        //         warn!("DouSha session limit reached");
        //         return Err(DouShaError::LimitError);
        //     }
        //     let per_peer_limit = MAX_DOUSHA_SESSIONS / self.n_parties;
        //     let peer_count = storage
        //         .values()
        //         .filter(|(id, _)| *id == initiator_id)
        //         .count();
        //     if peer_count >= per_peer_limit {
        //         warn!("DouSha per-peer session limit reached");
        //         return Err(DouShaError::LimitError);
        //     }
        // }
        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() {
            // Create and serialize the payload.
            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)
            };

            // Create and serialize the generic message.
            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?;
        }

        // Update the state of the protocol to Initialized.
        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(());
        }
        //todo: Better handle duplicate messages from a sender, check the shares
        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(()); // ignore duplicate message
        }

        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;

        // Check if the protocol has reached an end
        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(())
    }
}