use crate::{
avss_mpc::{
share_gen::{RanShaAvssError, RanShaAvssStore},
AvssSessionId, AvssWrappedMessage,
},
common::{
share::{
apply_vandermonde,
avss::{AvssError, AvssNode},
feldman::FeldmanShamirShare,
make_vandermonde,
},
ProtocolSessionId, ShamirShare, RBC,
},
};
use ark_ec::CurveGroup;
use ark_ff::FftField;
use ark_std::rand::Rng;
use std::{collections::HashMap, sync::Arc};
use stoffelnet::network_utils::{Network, PartyId};
use tokio::{
sync::{
mpsc::{self, Receiver},
Mutex,
},
time::{timeout, Duration},
};
use tracing::info;
#[derive(Clone, Debug)]
pub struct RanShaAvssNode<F: FftField, R: RBC, G: CurveGroup<ScalarField = F>> {
pub id: usize,
pub n_parties: usize,
pub threshold: usize,
pub store: Arc<Mutex<HashMap<AvssSessionId, (usize, Arc<Mutex<RanShaAvssStore<F, G>>>)>>>,
pub avss: AvssNode<F, R, G, AvssSessionId>,
pub avss_output: Arc<Mutex<Receiver<AvssSessionId>>>,
}
impl<F, R, C> RanShaAvssNode<F, R, C>
where
F: FftField,
R: RBC<Id = AvssSessionId>,
C: CurveGroup<ScalarField = F> + Send + Sync,
{
pub fn new(
id: PartyId,
n_parties: usize,
threshold: usize,
sk_i: F,
pk_map: Arc<Vec<C>>,
) -> Result<Self, RanShaAvssError> {
let (avss_sender, avss_receiver) = mpsc::channel(128);
let avss = AvssNode::new(
id,
n_parties,
(1..=n_parties).collect(),
threshold,
sk_i,
pk_map,
avss_sender,
Arc::new(AvssWrappedMessage::rbc_wrap),
Arc::new(AvssWrappedMessage::avss_wrap),
)?;
Ok(Self {
id,
n_parties,
threshold,
store: Arc::new(Mutex::new(HashMap::new())),
avss,
avss_output: Arc::new(Mutex::new(avss_receiver)),
})
}
pub async fn get_or_create_store(
&mut self,
session_id: AvssSessionId,
initiator_id: usize,
) -> Result<Arc<Mutex<RanShaAvssStore<F, C>>>, RanShaAvssError> {
let mut storage = self.store.lock().await;
Ok(storage
.entry(session_id)
.or_insert((
initiator_id,
Arc::new(Mutex::new(RanShaAvssStore::empty(self.n_parties))),
))
.1
.clone())
}
pub async fn wait_for_result(
&self,
session_id: AvssSessionId,
duration: Duration,
) -> Result<Vec<FeldmanShamirShare<F, C>>, RanShaAvssError> {
let output_receiver = {
let storage = self.store.lock().await;
let storage_bind = match storage.get(&session_id) {
Some((_, arc)) => arc,
None => return Err(RanShaAvssError::NoSuchSessionId(session_id)),
};
let mut storage = storage_bind.lock().await;
storage
.output_receiver
.take()
.ok_or(RanShaAvssError::ResultAlreadyReceived(session_id))?
};
match timeout(duration, output_receiver).await {
Err(_) => Err(RanShaAvssError::Timeout(session_id)),
Ok(Err(_)) => Err(RanShaAvssError::ReceiveError(session_id)),
Ok(Ok(shares)) => Ok(shares),
}
}
pub async fn init<N, G>(
&mut self,
session_id: AvssSessionId,
rng: &mut G,
network: Arc<N>,
) -> Result<(), RanShaAvssError>
where
N: Network + Send + Sync,
G: Rng + Send,
{
info!("Receiving init for share from {0:?}", self.id);
let secret = F::rand(rng);
let avss_sessionid = AvssSessionId::new(
session_id.calling_protocol().unwrap(),
AvssSessionId::pack_slot(session_id.exec_id(), self.id as u8, session_id.round_id()),
session_id.instance_id(),
);
self.avss
.init(vec![secret], avss_sessionid, rng, network.clone())
.await?;
while let Some(id) = {
let mut rx = self.avss_output.lock().await;
rx.recv().await
} {
if id.calling_protocol().unwrap() == session_id.calling_protocol().unwrap()
&& id.exec_id() == session_id.exec_id()
&& id.round_id() == session_id.round_id()
&& id.instance_id() == session_id.instance_id()
{
let mut store = self.avss.shares.lock().await;
let avss_share = store.remove(&id).unwrap().unwrap();
drop(store);
let binding = self.get_or_create_store(session_id, self.id).await?;
let mut ransha_storage = binding.lock().await;
let sender_id = id.sub_id();
if usize::from(sender_id) >= self.n_parties {
return Err(RanShaAvssError::InvalidPartyId);
}
ransha_storage.initial_shares.insert(
sender_id.into(),
avss_share
.first()
.ok_or(RanShaAvssError::InvalidPartyId)?
.clone(),
);
ransha_storage.reception_tracker[sender_id as usize] = true;
if ransha_storage
.reception_tracker
.iter()
.all(|&received| received)
{
let mut shares_deg_t: Vec<(usize, FeldmanShamirShare<F, C>)> = ransha_storage
.initial_shares
.iter()
.map(|(sid, s)| (*sid, s.clone()))
.collect();
drop(ransha_storage);
shares_deg_t.sort_by_key(|(sid, _)| *sid);
let shares_deg_t: Vec<FeldmanShamirShare<F, C>> =
shares_deg_t.into_iter().map(|(_, s)| s).collect();
self.ransha_gen(shares_deg_t, session_id).await?;
break;
}
}
}
Ok(())
}
pub async fn ransha_gen(
&mut self,
shares_deg_t: Vec<FeldmanShamirShare<F, C>>,
session_id: AvssSessionId,
) -> Result<(), RanShaAvssError> {
info!(
"party {:?} received shares for Random sharing generation",
self.id
);
let n = self.n_parties;
let t = self.threshold;
let shares: Vec<ShamirShare<_, 1, _>> = shares_deg_t
.iter()
.map(|s| s.feldmanshare.clone())
.collect();
let vandermonde_matrix = make_vandermonde(n, n - 1)?;
let r_deg_t = apply_vandermonde(&vandermonde_matrix, &shares)?;
let mut r_commitments: Vec<Vec<C>> = Vec::with_capacity(n);
for k in 0..n {
let mut ck = vec![C::zero(); t + 1];
for i in 0..n {
let a_ki = vandermonde_matrix[k][i]; let ci = &shares_deg_t[i].commitments;
if ci.len() != t + 1 {
return Err(RanShaAvssError::AvssError(
AvssError::InvalidCommitmentLength,
));
}
for j in 0..=t {
ck[j] += ci[j].mul(a_ki);
}
}
r_commitments.push(ck);
}
let bind_store = self.get_or_create_store(session_id, self.id).await?;
let mut store = bind_store.lock().await;
store.computed_r_shares = (0..n)
.map(|k| FeldmanShamirShare {
feldmanshare: r_deg_t[k].clone(),
commitments: r_commitments[k].clone(),
})
.collect();
let output = store.computed_r_shares[2 * t..].to_vec();
store.protocol_output = output.clone();
if let Some(sender) = store.output_sender.take() {
sender
.send(output)
.map_err(|_| RanShaAvssError::SendError(session_id))?;
}
Ok(())
}
}