use super::{rbc_store::*, utils::*, RbcError};
use crate::common::{ProtocolSessionId, RbcWrapFn, RBC};
use async_trait::async_trait;
use bincode;
use std::{collections::HashMap, sync::Arc};
use stoffelnet::network_utils::Network;
use threshold_crypto::{
serde_impl::SerdeSecret, PublicKeySet, SecretKeySet, SecretKeyShare, SignatureShare,
};
use tokio::{
sync::{
mpsc::{self, Receiver, Sender},
Mutex, Notify, OnceCell,
},
time::Duration,
};
use tracing::{debug, error, info, warn};
#[derive(Clone)]
pub struct Bracha<Id: ProtocolSessionId + 'static> {
pub id: usize, pub n: usize, pub t: usize, pub k: usize, pub store: Arc<Mutex<HashMap<Id, (usize, Arc<Mutex<BrachaStore>>)>>>, pub output_sender: Sender<Id>,
pub wrapper: RbcWrapFn<Id>,
}
#[async_trait]
impl<Id> RBC for Bracha<Id>
where
Id: ProtocolSessionId + 'static,
{
type Id = Id;
fn new(
id: usize,
n: usize,
t: usize,
k: usize,
output_sender: Sender<Id>,
wrapper: RbcWrapFn<Id>,
) -> Result<Self, RbcError> {
if !(t < (n + 2) / 3) {
return Err(RbcError::InvalidThreshold(t, n));
}
Ok(Bracha {
id,
n,
t,
k,
store: Arc::new(Mutex::new(HashMap::new())),
output_sender,
wrapper,
})
}
fn id(&self) -> usize {
self.id
}
async fn get_store(&self, session_id: Id) -> Result<Vec<u8>, RbcError> {
let store = self.store.lock().await;
let (_, arc) = store
.get(&session_id)
.ok_or_else(|| RbcError::Internal("Session ID does not exist".to_string()))?;
let store_lock = arc.lock().await;
if !store_lock.ended {
return Err(RbcError::Internal("Rbc has not terminated".to_string()));
}
Ok(store_lock.output.clone())
}
async fn clear_store(&self) {
let mut store = self.store.lock().await;
store.clear();
}
async fn clear_session(&self, session_id: Id) {
let mut store = self.store.lock().await;
store.remove(&session_id);
}
async fn init<N: Network + Send + Sync>(
&self,
payload: Vec<u8>,
session_id: Id,
net: Arc<N>,
) -> Result<(), RbcError> {
let msg = Msg::new(
self.id,
session_id,
0,
payload.clone(),
vec![],
GenericMsgType::Bracha(MsgType::Init),
);
info!(
id = self.id,
session_id = session_id.as_u128(),
msg_type = "INIT",
"Broadcasting INIT message"
);
self.broadcast(msg, net).await?;
Ok(())
}
async fn process<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
match &msg.msg_type {
GenericMsgType::Bracha(msg_type) => match msg_type {
MsgType::Init => self.init_handler(msg, net).await?,
MsgType::Echo => self.echo_handler(msg, net).await?,
MsgType::Ready => self.ready_handler(msg, net).await?,
MsgType::Unknown(tag) => return Err(RbcError::UnknownMsgType(tag.clone())),
},
_ => return Err(RbcError::UnknownMsgType("non-Bracha".into())),
}
Ok(())
}
async fn broadcast<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
let encoded = (self.wrapper)(msg)?;
net.broadcast(&encoded).await?;
Ok(())
}
async fn send<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
recv: usize,
) -> Result<(), RbcError> {
let encoded = (self.wrapper)(msg)?;
net.send(recv, &encoded).await?;
Ok(())
}
}
impl<Id> Bracha<Id>
where
Id: ProtocolSessionId + 'static,
{
pub async fn init_handler<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
info!(
id = self.id,
session_id = msg.session_id.as_u128(),
sender = msg.sender_id,
msg_type = %msg.msg_type,
"Handling INIT message"
);
let session_store = match self
.get_or_create_store(msg.session_id, msg.sender_id)
.await
{
Some(s) => s,
None => {
debug!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Session already ended, ignoring"
);
return Ok(());
}
};
let mut store = session_store.lock().await;
if !store.echo {
let new_msg = Msg::new(
self.id,
msg.session_id,
msg.round_id,
msg.payload.clone(),
vec![],
GenericMsgType::Bracha(MsgType::Echo),
);
store.mark_echo(); info!(
id = self.id,
session_id = msg.session_id.as_u128(),
msg_type = "ECHO",
"Broadcasting ECHO in response to INIT"
);
drop(store);
self.broadcast(new_msg, net).await?;
}
Ok(())
}
pub async fn echo_handler<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
info!(
id = self.id,
session_id = msg.session_id.as_u128(),
sender = msg.sender_id,
msg_type = %msg.msg_type,
"Handling ECHO message"
);
let mut broadcast_ready: Option<Msg<Id>> = None;
let mut broadcast_echo: Option<Msg<Id>> = None;
let session_store = match self
.get_or_create_store(msg.session_id, msg.sender_id)
.await
{
Some(s) => s,
None => {
debug!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Session already ended, ignoring"
);
return Ok(());
}
};
let mut store = session_store.lock().await;
if store.ended {
debug!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Session already ended, ignoring ECHO"
);
return Ok(());
}
if !store.has_echo(msg.sender_id) {
store.set_echo_sent(msg.sender_id); store.increment_echo(&msg.payload); let count = store.get_echo_count(&msg.payload);
if count >= (self.n + self.t + 2) / 2 {
if !store.ready {
store.mark_ready(); broadcast_ready = Some(Msg::new(
self.id,
msg.session_id,
msg.round_id,
msg.payload.clone(),
vec![],
GenericMsgType::Bracha(MsgType::Ready),
));
}
if !store.echo {
store.mark_echo(); broadcast_echo = Some(Msg::new(
self.id,
msg.session_id,
msg.round_id,
msg.payload.clone(),
vec![],
GenericMsgType::Bracha(MsgType::Echo),
));
}
}
}
drop(store);
if let Some(m) = broadcast_ready {
info!(
id = self.id,
session_id = m.session_id.as_u128(),
msg_type = "READY",
"Broadcasting READY after ECHO threshold met"
);
self.broadcast(m, net.clone()).await?;
}
if let Some(m) = broadcast_echo {
info!(
id = self.id,
session_id = m.session_id.as_u128(),
msg_type = "ECHO",
"Re-broadcasting ECHO due to threshold"
);
self.broadcast(m, net).await?;
}
Ok(())
}
pub async fn ready_handler<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
info!(
id = self.id,
session_id = msg.session_id.as_u128(),
sender = msg.sender_id,
msg_type = %msg.msg_type,
"Handling READY message"
);
let mut broadcast_ready: Option<Msg<Id>> = None;
let mut broadcast_echo: Option<Msg<Id>> = None;
let mut send_output: Option<Vec<u8>> = None;
let session_store = match self
.get_or_create_store(msg.session_id, msg.sender_id)
.await
{
Some(s) => s,
None => {
debug!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Session already ended, ignoring"
);
return Ok(());
}
};
let mut store = session_store.lock().await;
if store.ended {
debug!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Session already ended, ignoring READY"
);
return Ok(());
}
if !store.has_ready(msg.sender_id) {
store.set_ready_sent(msg.sender_id); store.increment_ready(&msg.payload); let count = store.get_ready_count(&msg.payload);
if count >= self.t + 1 && count < 2 * self.t + 1 {
if !store.ready {
store.mark_ready(); broadcast_ready = Some(Msg::new(
self.id,
msg.session_id,
msg.round_id,
msg.payload.clone(),
vec![],
GenericMsgType::Bracha(MsgType::Ready),
));
}
if !store.echo {
store.mark_echo(); broadcast_echo = Some(Msg::new(
self.id,
msg.session_id,
msg.round_id,
msg.payload.clone(),
vec![],
GenericMsgType::Bracha(MsgType::Echo),
));
}
} else if count >= 2 * self.t + 1 {
store.mark_ended();
store.set_output(msg.payload.clone());
send_output = Some(msg.payload);
}
}
drop(store);
if let Some(m) = broadcast_ready {
info!(
id = self.id,
session_id = msg.session_id.as_u128(),
msg_type = "READY",
"Broadcasting READY after t+1 threshold"
);
self.broadcast(m, net.clone()).await?;
}
if let Some(m) = broadcast_echo {
info!(
id = self.id,
session_id = msg.session_id.as_u128(),
msg_type = "ECHO",
"Broadcasting ECHO along with READY"
);
self.broadcast(m, net.clone()).await?;
}
if let Some(m) = send_output {
info!(
id = self.id,
session_id = msg.session_id.as_u128(),
output = ?m,
"Consensus achieved; RBC instance ended"
);
self.output_sender
.send(msg.session_id)
.await
.map_err(|_| RbcError::SendError)?;
}
Ok(())
}
async fn get_or_create_store(
&self,
session_id: Id,
sender_id: usize,
) -> Option<Arc<Mutex<BrachaStore>>> {
let store_lock = {
let mut store = self.store.lock().await;
store
.entry(session_id)
.or_insert_with(|| (sender_id, Arc::new(Mutex::new(BrachaStore::default()))))
.1
.clone()
};
{
let store_guard = store_lock.lock().await;
if store_guard.ended {
return None;
}
}
Some(store_lock)
}
}
#[derive(Clone)]
pub struct Avid<Id: ProtocolSessionId> {
pub id: usize, pub n: usize, pub t: usize, pub k: usize, pub store: Arc<Mutex<HashMap<Id, (usize, Arc<Mutex<AvidStore>>)>>>, pub output_sender: Sender<Id>,
pub wrapper: RbcWrapFn<Id>,
}
#[async_trait]
impl<Id: ProtocolSessionId> RBC for Avid<Id> {
type Id = Id;
fn new(
id: usize,
n: usize,
t: usize,
k: usize,
output_sender: Sender<Id>,
wrapper: RbcWrapFn<Id>,
) -> Result<Self, RbcError> {
if !(t < (n + 2) / 3) {
return Err(RbcError::InvalidThreshold(t, n));
}
if !(t + 1 <= k && k <= n - 2 * t) {
return Err(RbcError::Internal(format!(
"Invalid k: must satisfy t + 1 <= k <= n - 2t (t={}, k={}, n={})",
t, k, n
)));
}
Ok(Avid {
id,
n,
t,
k,
store: Arc::new(Mutex::new(HashMap::new())),
output_sender,
wrapper,
})
}
fn id(&self) -> usize {
self.id
}
async fn clear_store(&self) {
let mut store = self.store.lock().await;
store.clear();
}
async fn clear_session(&self, session_id: Id) {
let mut store = self.store.lock().await;
store.remove(&session_id);
}
async fn get_store(&self, session_id: Id) -> Result<Vec<u8>, RbcError> {
let store = self.store.lock().await;
let (_, arc) = store
.get(&session_id)
.ok_or_else(|| RbcError::Internal("Session ID does not exist".to_string()))?;
let store_lock = arc.lock().await;
if !store_lock.ended {
return Err(RbcError::Internal("Rbc has not terminated".to_string()));
}
Ok(store_lock.output.clone())
}
async fn init<N: Network + Send + Sync>(
&self,
payload: Vec<u8>,
session_id: Id,
net: Arc<N>,
) -> Result<(), RbcError> {
info!(
id = self.id,
session_id = ?session_id,
msg_type = "SEND",
"Sending SEND message for AVID to all parties"
);
let shards = encode_rs(payload.clone(), self.k, self.n - self.k)?;
let tree = gen_merkletree(&shards);
let root = tree.root().ok_or_else(|| {
RbcError::Internal(format!("Merkle root missing for session {:?}", session_id))
})?;
for i in 0..self.n {
let fingerprint = tree.proof(&[i]).to_bytes();
let mut fp = Vec::with_capacity(root.len() + fingerprint.len());
fp.extend_from_slice(&root);
fp.extend_from_slice(&fingerprint);
let shard = shards[i].clone();
let msg = Msg::new(
self.id,
session_id,
0,
shard,
fp, GenericMsgType::Avid(MsgTypeAvid::Send),
);
if let Err(e) = self.send(msg, net.clone(), i).await {
warn!("Failed to send shard to party {}: {:?}", i, e);
}
}
Ok(())
}
async fn process<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
match &msg.msg_type {
GenericMsgType::Avid(msg_type) => match msg_type {
MsgTypeAvid::Send => self.send_handler(msg, net).await?,
MsgTypeAvid::Echo => self.echo_handler(msg, net).await?,
MsgTypeAvid::Ready => self.ready_handler(msg, net).await?,
MsgTypeAvid::Unknown(tag) => return Err(RbcError::UnknownMsgType(tag.clone())),
},
_ => return Err(RbcError::UnknownMsgType("non-Avid".into())),
}
Ok(())
}
async fn broadcast<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
let encoded = (self.wrapper)(msg)?;
net.broadcast(&encoded).await?;
Ok(())
}
async fn send<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
recv: usize,
) -> Result<(), RbcError> {
let encoded = (self.wrapper)(msg)?;
net.send(recv, &encoded).await?;
Ok(())
}
}
impl<Id: ProtocolSessionId> Avid<Id> {
pub async fn send_handler<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
info!(
id = self.id,
session_id = ?msg.session_id,
sender = msg.sender_id,
msg_type = %msg.msg_type,
"Handling SEND message"
);
if msg.metadata.len() < 32 {
return Err(RbcError::Internal("Incorrect message length".to_string()));
}
{
let store = self.store.lock().await;
if let Some((_, session_store)) = store.get(&msg.session_id) {
let session = session_store.lock().await;
if session.ended || session.echo {
return Ok(());
}
}
}
match verify_merkle(self.id, self.n, msg.metadata.clone(), msg.payload.clone()) {
Ok(true) => {}
Ok(false) => {
error!(
id = self.id,
session_id = msg.session_id.as_u128(),
sender = msg.sender_id,
"Merkle proof verification failed on SEND message"
);
return Err(RbcError::Internal("Merkle proof failed in SEND".into()));
}
Err(e) => {
error!(
id = self.id,
session_id = msg.session_id.as_u128(),
sender = msg.sender_id,
error = %e,
"Error during Merkle proof verification"
);
return Err(RbcError::ShardError(e));
}
}
let session_store = match self
.get_or_create_store(msg.session_id, msg.sender_id)
.await
{
Some(s) => s,
None => {
debug!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Session already ended, ignoring SEND"
);
return Ok(());
}
};
let mut store = session_store.lock().await;
if !store.echo {
let msg = Msg::new(
self.id,
msg.session_id,
msg.round_id,
msg.payload,
msg.metadata,
GenericMsgType::Avid(MsgTypeAvid::Echo),
);
store.mark_echo();
info!(
id = self.id,
session_id = ?msg.session_id,
msg_type = "ECHO",
"Broadcasting ECHO in response to SEND"
);
drop(store);
self.broadcast(msg, net).await?;
}
Ok(())
}
pub async fn echo_handler<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
info!(
id = self.id,
session_id = msg.session_id.as_u128(),
sender = msg.sender_id,
msg_type = %msg.msg_type,
"Handling ECHO message"
);
if msg.metadata.len() < 32 {
return Err(RbcError::Internal("Incorrect message length".to_string()));
}
{
let store = self.store.lock().await;
if let Some((_, session_store)) = store.get(&msg.session_id) {
let session = session_store.lock().await;
if session.ended || session.has_echo(msg.sender_id) {
return Ok(());
}
}
}
let root = msg.metadata[0..32].to_vec();
let proof_bytes = msg.metadata[32..].to_vec();
match verify_merkle(
msg.sender_id,
self.n,
msg.metadata.clone(),
msg.payload.clone(),
) {
Ok(true) => {}
Ok(false) => {
warn!(
id = self.id,
session_id = ?msg.session_id,
sender = msg.sender_id,
"Merkle verification failed for ECHO"
);
return Err(RbcError::Internal(
"Merkle verification failed in ECHO".into(),
));
}
Err(e) => {
warn!(
id = self.id,
session_id = ?msg.session_id,
sender = msg.sender_id,
error = %e,
"Merkle verification threw error"
);
return Err(RbcError::ShardError(e));
}
}
let session_store = match self
.get_or_create_store(msg.session_id, msg.sender_id)
.await
{
Some(s) => s,
None => {
debug!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Session already ended, ignoring"
);
return Ok(());
}
};
let mut store = session_store.lock().await;
if store.ended || store.has_echo(msg.sender_id) {
return Ok(());
}
store.insert_shard(root.clone(), msg.sender_id, msg.payload.clone(), self.k)?;
store.insert_fingerprint(root.clone(), msg.sender_id, proof_bytes);
store.increment_echo(&root);
store.set_echo_sent(msg.sender_id);
let echo_count = store.get_echo_count(&root);
let ready_count = store.get_ready_count(&root);
let threshold = usize::max((self.n + self.t + 2) / 2, self.k);
if echo_count == threshold && ready_count < self.k {
let shards_map = store.get_shards_for_root(&root);
drop(store);
self.send_ready(msg, shards_map, net).await?;
}
Ok(())
}
pub async fn ready_handler<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
info!(
id = self.id,
session_id = ?msg.session_id,
sender = msg.sender_id,
msg_type = %msg.msg_type,
"Handling READY message"
);
if msg.metadata.len() < 32 {
return Err(RbcError::Internal("Incorrect message length".to_string()));
}
{
let store = self.store.lock().await;
if let Some((_, session_store)) = store.get(&msg.session_id) {
let session = session_store.lock().await;
if session.ended || session.has_ready(msg.sender_id) {
return Ok(());
}
}
}
let root = msg.metadata[0..32].to_vec();
let proof_bytes = msg.metadata[32..].to_vec();
match verify_merkle(
msg.sender_id,
self.n,
msg.metadata.clone(),
msg.payload.clone(),
) {
Ok(true) => {}
Ok(false) => {
warn!(
id = self.id,
session_id = msg.session_id.as_u128(),
sender = msg.sender_id,
"Merkle verification failed in READY handler"
);
return Err(RbcError::Internal(
"Merkle verification failed in READY".into(),
));
}
Err(e) => {
warn!(
id = self.id,
session_id = msg.session_id.as_u128(),
sender = msg.sender_id,
error = %e,
"Error during Merkle verification in READY handler"
);
return Err(RbcError::ShardError(e));
}
}
let session_store = match self
.get_or_create_store(msg.session_id, msg.sender_id)
.await
{
Some(s) => s,
None => {
debug!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Session already ended, ignoring"
);
return Ok(());
}
};
let mut store = session_store.lock().await;
if store.ended || store.has_ready(msg.sender_id) {
return Ok(());
}
let mut send_shards_map: Option<HashMap<usize, Vec<u8>>> = None;
let mut send_output: Option<Vec<u8>> = None;
store.insert_shard(root.clone(), msg.sender_id, msg.payload.clone(), self.k)?;
store.insert_fingerprint(root.clone(), msg.sender_id, proof_bytes);
store.increment_ready(&root);
store.set_ready_sent(msg.sender_id);
let echo_count = store.get_echo_count(&root);
let ready_count = store.get_ready_count(&root);
let threshold = usize::max((self.n + self.t + 2) / 2, self.k);
if echo_count < threshold && ready_count == self.k {
let shards_map = store.get_shards_for_root(&root);
send_shards_map = Some(shards_map);
}
if ready_count >= (self.k + self.t) && !store.ended {
let shards = decode_rs(store.get_shards_for_root(&root), self.k, self.n - self.k)?;
let output = reconstruct_payload(shards, self.k)?;
store.mark_ended();
store.set_output(output.clone());
send_output = Some(output);
}
drop(store);
if let Some(m) = send_shards_map {
self.send_ready(msg.clone(), m, net.clone()).await?;
}
if let Some(m) = send_output {
info!(
id = self.id,
session_id = msg.session_id.as_u128(),
output = ?m,
"Consensus achieved; AVID instance ended"
);
self.output_sender
.send(msg.session_id)
.await
.map_err(|_| RbcError::SendError)?;
}
Ok(())
}
async fn send_ready<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
shards_map: HashMap<usize, Vec<u8>>,
net: Arc<N>,
) -> Result<(), RbcError> {
let root = &msg.metadata[0..32];
let handler_type = msg.msg_type;
let shards = decode_rs(shards_map, self.k, self.n - self.k)?;
let payload = shards[self.id].clone();
let mut fingerprint = root.to_vec();
match generate_merkle_proofs_map(&shards) {
Ok(proof_map) => {
let self_proof = proof_map.get(&(self.id)).cloned().unwrap_or_else(|| {
tracing::warn!(index = self.id, "Missing Merkle proof");
Vec::new()
});
fingerprint.extend(self_proof);
for (id, proof) in proof_map {
let mut fp = root.to_vec();
fp.extend(proof);
match verify_merkle(id, self.n, fp, shards[id].clone()) {
Ok(true) => {}
Ok(false) => {
error!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Merkle proof generation failed in {handler_type} handler. Aborting."
);
return Err(RbcError::Internal(format!(
"Merkle proof failed for id {}",
id
)));
}
Err(e) => {
error!(
id = self.id,
session_id = msg.session_id.as_u128(),
error = %e,
"Error during Merkle verification in {handler_type} handler"
);
return Err(RbcError::ShardError(e));
}
}
}
}
Err(e) => {
error!(
id = self.id,
session_id = msg.session_id.as_u128(),
error = %e,
"Failed to generate Merkle proof map in {handler_type} handler"
);
return Err(RbcError::ShardError(e));
}
}
let ready_msg = Msg::new(
self.id,
msg.session_id,
msg.round_id,
payload,
fingerprint,
GenericMsgType::Avid(MsgTypeAvid::Ready),
);
info!(
id = self.id,
session_id = msg.session_id.as_u128(),
msg_type = "READY",
"Broadcasting READY in response to a {handler_type}"
);
self.broadcast(ready_msg, net).await?;
Ok(())
}
async fn get_or_create_store(
&self,
session_id: Id,
sender_id: usize,
) -> Option<Arc<Mutex<AvidStore>>> {
let store_lock = {
let mut store = self.store.lock().await;
store
.entry(session_id)
.or_insert_with(|| (sender_id, Arc::new(Mutex::new(AvidStore::default()))))
.1
.clone()
};
{
let store_guard = store_lock.lock().await;
if store_guard.ended {
return None;
}
}
Some(store_lock)
}
}
#[derive(Clone)]
pub struct ABA<Id: ProtocolSessionId> {
pub id: usize, pub n: usize, pub t: usize, pub k: usize, pub skshare: Arc<OnceCell<Vec<u8>>>, pub pkset: Arc<OnceCell<Vec<u8>>>, pub store: Arc<Mutex<HashMap<Id, Arc<Mutex<AbaStore>>>>>, pub coin: Arc<Mutex<HashMap<Id, Arc<Mutex<CoinStore>>>>>, pub output_sender: Sender<Id>,
pub wrapper: RbcWrapFn<Id>,
}
#[async_trait]
impl<Id: ProtocolSessionId + 'static> RBC for ABA<Id> {
type Id = Id;
fn new(
id: usize,
n: usize,
t: usize,
k: usize,
output_sender: Sender<Id>,
wrapper: RbcWrapFn<Id>,
) -> Result<Self, RbcError> {
if !(t < (n + 2) / 3) {
return Err(RbcError::InvalidThreshold(t, n));
}
if !(t + 1 <= k && k <= n - 2 * t) {
return Err(RbcError::Internal(format!(
"Invalid k: must satisfy t + 1 <= k <= n - 2t (t={}, k={}, n={})",
t, k, n
)));
}
if id >= n {
return Err(RbcError::Internal(
"PartyID is greater than party size".to_string(),
));
}
Ok(ABA {
id,
n,
t,
k,
skshare: Arc::new(OnceCell::new()),
pkset: Arc::new(OnceCell::new()),
store: Arc::new(Mutex::new(HashMap::new())),
coin: Arc::new(Mutex::new(HashMap::new())),
output_sender,
wrapper,
})
}
fn id(&self) -> usize {
self.id
}
async fn clear_store(&self) {
let mut store = self.store.lock().await;
let mut coin_store = self.coin.lock().await;
store.clear();
coin_store.clear();
}
async fn clear_session(&self, session_id: Id) {
let mut store = self.store.lock().await;
store.remove(&session_id);
}
async fn get_store(&self, session_id: Id) -> Result<Vec<u8>, RbcError> {
let store = self.store.lock().await;
let output_store = store
.get(&session_id)
.ok_or_else(|| RbcError::Internal("Session ID does not exist".to_string()))?;
let store_lock = output_store.lock().await;
if !store_lock.ended {
return Err(RbcError::Internal("Rbc has not terminated".to_string()));
}
Ok(vec![store_lock.output as u8])
}
async fn init<N: Network + Send + Sync>(
&self,
payload: Vec<u8>,
session_id: Id,
net: Arc<N>,
) -> Result<(), RbcError> {
info!(
id = self.id,
session_id = session_id.as_u128(),
msg_type = "EST",
"Broadcasting EST message"
);
let v_r = match get_value_round(&payload) {
Some(r) => r,
None => {
error!(
id = self.id,
session_id = session_id.as_u128(),
"Error while getting roundid at init"
);
return Err(RbcError::Internal(
"Error while getting roundid at init".to_string(),
));
}
};
let msg = Msg::new(
self.id,
session_id,
v_r.1 as usize, vec![v_r.0 as u8], vec![],
GenericMsgType::ABA(MsgTypeAba::Est),
);
let session_store = self.get_or_create_store(msg.session_id).await;
let mut store = session_store.lock().await;
store.mark_est(msg.round_id, v_r.0);
drop(store);
self.broadcast(msg, net).await?;
Ok(())
}
async fn process<N: Network + Send + Sync + 'static>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
match &msg.msg_type {
GenericMsgType::ABA(msg_type) => match msg_type {
MsgTypeAba::Est => self.est_handler(msg, net).await?,
MsgTypeAba::Aux => self.aux_handler(msg, net).await?,
MsgTypeAba::Key => self.key_handler(msg)?,
MsgTypeAba::Coin => self.coin_handler(msg).await?,
MsgTypeAba::Unknown(tag) => return Err(RbcError::UnknownMsgType(tag.clone())),
},
_ => return Err(RbcError::UnknownMsgType("non-ABA".into())),
}
Ok(())
}
async fn broadcast<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
let encoded = (self.wrapper)(msg)?;
net.broadcast(&encoded).await?;
Ok(())
}
async fn send<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
recv: usize,
) -> Result<(), RbcError> {
let encoded = (self.wrapper)(msg)?;
net.send(recv, &encoded).await?;
Ok(())
}
}
impl<Id: ProtocolSessionId + 'static> ABA<Id> {
pub async fn est_handler<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
info!(
session_id = ?msg.session_id.as_u128(),
id = self.id,
sender = msg.sender_id,
msg_type = %msg.msg_type,
"Handling EST message for round {}",msg.round_id,
);
let session_store = self.get_or_create_store(msg.session_id).await;
let mut store = session_store.lock().await;
if store.ended {
debug!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Session already ended, ignoring est"
);
return Ok(());
}
let value = match get_value(&msg.payload) {
Some(v) => v,
None => {
warn!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Error while getting value at est handler"
);
return Err(RbcError::Internal(
"Error while getting value at est handler".to_string(),
));
}
};
if !store.has_sent_est(msg.round_id, msg.sender_id, value) {
store.set_est_sent(msg.round_id, msg.sender_id, value); store.increment_est(msg.round_id, value); let count = store.get_est_count(msg.round_id)[value as usize];
if count >= self.t + 1 && !store.get_est(msg.round_id, value) {
store.mark_est(msg.round_id, value); let new_msg = Msg::new(
self.id,
msg.session_id,
msg.round_id,
msg.payload.clone(),
msg.metadata.clone(),
GenericMsgType::ABA(MsgTypeAba::Est),
);
self.broadcast(new_msg, net.clone()).await?;
}
if count >= 2 * self.t + 1 {
store.insert_bin_value(msg.round_id, value);
if !store.get_aux(msg.round_id, value) {
store.mark_aux(msg.round_id, value); let new_msg = Msg::new(
self.id,
msg.session_id,
msg.round_id,
msg.payload,
msg.metadata,
GenericMsgType::ABA(MsgTypeAba::Aux),
);
drop(store);
self.broadcast(new_msg, net).await?;
}
}
}
Ok(())
}
pub async fn aux_handler<N: Network + Send + Sync + 'static>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
info!(
id = self.id,
session_id = msg.session_id.as_u128(),
sender = msg.sender_id,
msg_type = %msg.msg_type,
"Handling AUX message for round {}",msg.round_id
);
let session_store = self.get_or_create_store(msg.session_id).await;
let mut store = session_store.lock().await;
let value = match get_value(&msg.payload) {
Some(v) => v,
None => {
warn!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Error while getting value at aux handler"
);
return Err(RbcError::Internal(
"Error while getting value at aux handler".to_string(),
));
}
};
if store.ended {
debug!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Session already ended, ignoring aux"
);
return Ok(());
}
let bin_val = store.get_bin_values(msg.round_id);
if !bin_val.contains(&value) {
return Err(RbcError::Internal(
"Not accepting AUX message for values not in bin_values".to_string(),
));
}
if !store.has_sent_aux(msg.round_id, msg.sender_id, value) {
store.set_aux_sent(msg.round_id, msg.sender_id, value); store.insert_values(msg.round_id, msg.sender_id, value);
let count = store.get_sender_count(msg.round_id);
let bin_val = store.get_bin_values(msg.round_id);
let values = store.get_all_values(msg.round_id);
drop(store);
if count >= self.n - self.t && values.is_subset(&bin_val) {
{
let coin_store = self.get_or_create_coinstore(msg.session_id).await;
let mut store = coin_store.lock().await;
if store.get_start(msg.round_id) {
return Ok(());
}
store.set_start(msg.round_id);
}
self.init_coin(msg.clone(), net.clone()).await?;
let cloned_msg = msg.clone();
let cloned_net = net.clone();
let cloned_self = self.clone();
tokio::spawn(async move {
let coin_opt = cloned_self
.wait_for_coin(cloned_msg.session_id, cloned_msg.round_id, 1000)
.await;
let coin_value = match coin_opt {
Some(coin) => coin,
None => {
error!(
id = cloned_self.id,
session_id = cloned_msg.session_id.as_u128(),
round = cloned_msg.round_id,
"Failed to get coin value in time"
);
return;
}
};
let session_store =
cloned_self.get_or_create_store(cloned_msg.session_id).await;
let mut store = session_store.lock().await;
if store.ended {
debug!(
id = cloned_self.id,
session_id = cloned_msg.session_id.as_u128(),
"Session already ended, ignoring coin result"
);
return;
}
if values.len() == 1 {
let v = match values.iter().next().copied() {
Some(v) => v,
None => {
error!(
id = cloned_self.id,
session_id = cloned_msg.session_id.as_u128(),
round = cloned_msg.round_id,
"Could not get the value from values set"
);
return;
}
};
if v == coin_value {
store.mark_ended();
store.set_output(v);
info!(
id = cloned_self.id,
session_id = cloned_msg.session_id.as_u128(),
output = ?cloned_msg.payload,
"Binary agreement achieved; ABA instance ended at round {}",msg.round_id
);
} else {
info!(
id = cloned_self.id,
session_id = cloned_msg.session_id.as_u128(),
"Entering round {} with value",
msg.round_id + 1
);
let _ = cloned_self
.send_est_for_next_round(
&cloned_msg,
cloned_msg.round_id + 1,
v,
cloned_net,
)
.await
.map_err(|err| {
error!(
id = cloned_self.id,
session_id = cloned_msg.session_id.as_u128(),
error = ?err,
"Starting next round failed"
);
});
}
} else {
info!(
id = cloned_self.id,
session_id = cloned_msg.session_id.as_u128(),
"Entering round {} with coin",
msg.round_id + 1
);
let _ = cloned_self
.send_est_for_next_round(
&cloned_msg,
cloned_msg.round_id + 1,
coin_value,
cloned_net,
)
.await
.map_err(|err| {
error!(
id = cloned_self.id,
session_id = cloned_msg.session_id.as_u128(),
error = ?err,
"Starting next round failed"
);
});
}
});
}
}
Ok(())
}
async fn wait_for_coin(
&self,
session_id: Id,
round_id: usize,
timeout_ms: u64,
) -> Option<bool> {
let coin_store = self.get_or_create_coinstore(session_id).await;
let notify = {
let mut store = coin_store.lock().await;
if let Some(coin) = store.coin(round_id) {
return Some(coin);
}
let entry = store
.notifiers
.entry(round_id)
.or_insert_with(|| Arc::new(Notify::new()));
Arc::clone(entry)
};
let timeout = Duration::from_millis(timeout_ms);
tokio::select! {
_ = notify.notified() => {
let store = coin_store.lock().await;
store.coins.get(&round_id).copied()
}
_ = tokio::time::sleep(timeout) => {
warn!(
"Timed out waiting for coin for session {} round {}",
session_id.as_u128(), round_id
);
None
}
}
}
async fn get_or_create_store(&self, session_id: Id) -> Arc<Mutex<AbaStore>> {
let mut store = self.store.lock().await;
store
.entry(session_id)
.or_insert_with(|| Arc::new(Mutex::new(AbaStore::default())))
.clone()
}
async fn get_or_create_coinstore(&self, session_id: Id) -> Arc<Mutex<CoinStore>> {
let mut store = self.coin.lock().await;
store
.entry(session_id)
.or_insert_with(|| Arc::new(Mutex::new(CoinStore::default())))
.clone()
}
async fn send_est_for_next_round<N: Network + Send + Sync>(
&self,
msg: &Msg<Id>,
round: usize,
value: bool,
net: Arc<N>,
) -> Result<(), RbcError> {
let msg = Msg::new(
self.id,
msg.session_id,
round,
vec![value as u8],
msg.metadata.clone(),
GenericMsgType::ABA(MsgTypeAba::Est),
);
self.broadcast(msg, net).await?;
Ok(())
}
fn key_handler(&self, msg: Msg<Id>) -> Result<(), RbcError> {
info!(
id = self.id,
session_id = msg.session_id.as_u128(),
sender = msg.sender_id,
msg_type = %msg.msg_type,
"Handling Key setup message"
);
let _ = self
.skshare
.set(msg.payload)
.map_err(|e| RbcError::Internal(e.to_string()))?;
let _ = self
.pkset
.set(msg.metadata)
.map_err(|e| RbcError::Internal(e.to_string()))?;
Ok(())
}
pub async fn init_coin<N: Network + Send + Sync>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
info!(
id = self.id,
session_id = msg.session_id.as_u128(),
sender = msg.sender_id,
msg_type = %msg.msg_type,
round = msg.round_id,
"Initialising common coin"
);
let sk_payload = match self.skshare.get() {
Some(sk) => sk,
None => {
error!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Error while getting secret key share"
);
return Err(RbcError::Internal(
"Error while getting secret key share".to_string(),
));
}
};
let skshare: SerdeSecret<SecretKeyShare> =
bincode::deserialize(&sk_payload).map_err(|e| RbcError::SerializationError(e))?;
let signshare = skshare.sign(msg.round_id.to_be_bytes());
let new_msg = Msg::new(
self.id,
msg.session_id,
msg.round_id,
signshare.to_bytes().to_vec(),
vec![],
GenericMsgType::ABA(MsgTypeAba::Coin),
);
info!(
session_id = msg.session_id.as_u128(),
id = self.id,
sender = msg.sender_id,
msg_type = %msg.msg_type,
"Broadcasting signature shares"
);
self.broadcast(new_msg, net).await?;
Ok(())
}
async fn coin_handler(&self, msg: Msg<Id>) -> Result<(), RbcError> {
info!(
session_id = msg.session_id.as_u128(),
id = self.id,
sender = msg.sender_id,
msg_type = %msg.msg_type,
"At coin handler"
);
let coin_store = self.get_or_create_coinstore(msg.session_id).await;
let mut store = coin_store.lock().await;
if !store.has_sent_sign(msg.round_id, msg.sender_id) {
let sigshare_bytes: &[u8; 96] = match msg.payload.as_slice().try_into() {
Ok(bytes) => bytes,
Err(_) => {
warn!("Invalid signature share size from {}", msg.sender_id);
return Err(RbcError::Internal(
"Invalid signature share size".to_string(),
));
}
};
let sigshare = match SignatureShare::from_bytes(sigshare_bytes) {
Ok(share) => share,
Err(e) => {
warn!(
"Failed to deserialize signature share from {}",
msg.sender_id
);
return Err(RbcError::Internal(format!(
"Failed to deserialize signature share from {} : {e}",
msg.sender_id
)));
}
};
let pkset_bytes = match self.pkset.get() {
Some(pk) => pk,
None => {
error!(
id = self.id,
session_id = msg.session_id.as_u128(),
"Error while getting pk key set"
);
return Err(RbcError::Internal(
"Error while getting pk key set".to_string(),
));
}
};
let pkset: PublicKeySet = match bincode::deserialize(&pkset_bytes) {
Ok(pk) => pk,
Err(e) => {
warn!("Failed to deserialize PublicKeySet");
return Err(RbcError::SerializationError(e));
}
};
if !pkset
.public_key_share(msg.sender_id as i32)
.verify(&sigshare, msg.round_id.to_be_bytes())
{
warn!("Invalid signature share from {}", msg.sender_id);
return Err(RbcError::Internal(format!(
"Invalid signature share from {}",
msg.sender_id
)));
}
store.set_sign_sent(msg.round_id, msg.sender_id);
store.insert_share(msg.round_id, msg.sender_id, msg.payload); store.increment_sign(msg.round_id); let count = store.get_sign_count(msg.round_id);
if count == self.t + 1 {
let shares = store.get_shares_map(msg.round_id);
if let Some(shares_map) = shares {
let sig_shares: Vec<(usize, SignatureShare)> = shares_map
.iter()
.filter_map(|(&sender_id, bytes)| {
let array: &[u8; 96] = bytes.as_slice().try_into().ok()?;
SignatureShare::from_bytes(array)
.ok()
.map(|s| (sender_id, s))
})
.collect();
match pkset.combine_signatures(sig_shares.iter().map(|(i, s)| (*i, s))) {
Ok(signature) => {
if pkset
.public_key()
.verify(&signature, msg.round_id.to_be_bytes())
{
let coin_bit = signature.to_bytes()[0] & 1 == 1;
store.set_coin(msg.round_id, coin_bit);
info!(
session_id = msg.session_id.as_u128(),
id = self.id,
"Successfully combined and verified signature for round {} with coin = {}",
msg.round_id,
coin_bit
);
return Ok(());
} else {
warn!("Combined signature failed verification");
return Err(RbcError::Internal(
"Combined signature failed verification".to_string(),
));
}
}
Err(err) => {
error!("Failed to combine signature shares: {:?}", err);
return Err(RbcError::Internal(format!(
"Failed to combine signature shares: {err}"
)));
}
}
}
}
}
Ok(())
}
}
pub struct Dealer {
n: usize,
t: usize,
}
impl Dealer {
pub fn new(n: usize, t: usize) -> Self {
Dealer { n, t }
}
pub async fn distribute_keys<N: Network, Id: ProtocolSessionId>(
&self,
msg: Msg<Id>,
net: Arc<N>,
wrapper: RbcWrapFn<Id>,
) -> Result<(), RbcError> {
let mut rng = rand::thread_rng();
let skset = SecretKeySet::random(self.t, &mut rng);
let pkset = skset.public_keys();
let pkset_serial = bincode::serialize(&pkset).expect("Failed to serialize pkset");
for i in 0..self.n {
let skshare = SerdeSecret(skset.secret_key_share(i as i32));
let serialized_share =
bincode::serialize(&skshare).map_err(|e| RbcError::SerializationError(e))?;
let key_msg = Msg::new(
msg.sender_id,
msg.session_id,
msg.round_id,
serialized_share,
pkset_serial.clone(),
msg.msg_type.clone(),
);
let encoded = (wrapper)(key_msg)?;
net.send(i, &encoded).await?;
}
Ok(())
}
}
#[derive(Clone)]
pub struct ACS<Id: ProtocolSessionId> {
pub id: usize, pub n: usize, pub t: usize, pub k: usize, pub store: Arc<Mutex<AcsStore>>, pub aba: ABA<Id>, pub aba_output: Arc<Mutex<Receiver<Id>>>,
pub acs_output_sender: Sender<Id>,
}
impl<Id: ProtocolSessionId + 'static> ACS<Id> {
pub fn new(
id: usize,
n: usize,
t: usize,
k: usize,
acs_output_sender: Sender<Id>,
wrapper: RbcWrapFn<Id>,
) -> Result<Self, RbcError> {
if !(t < (n + 2) / 3) {
return Err(RbcError::InvalidThreshold(t, n));
}
if !(t + 1 <= k && k <= n - 2 * t) {
return Err(RbcError::Internal(format!(
"Invalid k: must satisfy t + 1 <= k <= n - 2t (t={}, k={}, n={})",
t, k, n
)));
}
if id >= n {
return Err(RbcError::Internal(
"PartyID is greater than party size".to_string(),
));
}
let (s, r) = mpsc::channel(256);
let aba = ABA::new(id, n, t, k, s, wrapper)?;
Ok(ACS {
id,
n,
t,
k,
store: Arc::new(Mutex::new(AcsStore::default())),
aba: aba,
aba_output: Arc::new(Mutex::new(r)),
acs_output_sender,
})
}
pub async fn init<N: Network + Send + Sync + 'static>(
&self,
msg: Msg<Id>,
net: Arc<N>,
) -> Result<(), RbcError> {
info!(
id = self.id,
session_id = msg.session_id.as_u128(),
sender = msg.sender_id,
msg_type = %msg.msg_type,
"Initiating common subset"
);
assert!(msg.session_id.calling_protocol().is_some());
let mut store = self.store.lock().await;
if !store.has_aba_input(msg.session_id.slot().into()) {
store.set_aba_input(msg.session_id.slot().into(), true);
store.set_rbc_output(msg.session_id.slot().into(), msg.payload);
let payload = set_value_round(true, 0);
self.aba.init(payload, msg.session_id, net.clone()).await?;
let aba_store = self.aba.get_or_create_store(msg.session_id).await;
let aba_store_clone = aba_store.clone();
let self_clone = self.clone();
let net_clone = net.clone();
let store_clone = self.store.clone();
drop(store);
tokio::spawn(async move {
let notify = {
let aba = aba_store_clone.lock().await;
aba.notify.clone()
};
notify.notified().await;
let output: bool = {
let aba = aba_store_clone.lock().await;
aba.output
};
let mut store = store_clone.lock().await;
store.set_aba_output(msg.session_id.slot().into(), output);
let true_count = store.get_aba_output_one_count();
if true_count >= self_clone.n - self_clone.t {
let uninitiated = (0..self_clone.n)
.filter(|sid| !store.has_aba_input(*sid as u128))
.collect::<Vec<_>>();
if uninitiated.len() == 0 {
let store_clone2 = store_clone.clone();
let self_clone2 = self_clone.clone();
tokio::spawn(async move {
self_clone2.check_and_finalize_output(store_clone2).await;
});
return;
} else {
for sid in uninitiated {
let payload = set_value_round(false, 0);
let sessionid = Id::new(
msg.session_id.calling_protocol().unwrap(),
sid as u128,
msg.session_id.instance_id(),
);
let _ = self_clone
.aba
.init(payload, sessionid, net_clone.clone())
.await
.map_err(|err| {
error!(
id = self_clone.id,
session_id = sid,
error = ?err,
"ABA init failed"
);
});
store.set_aba_input(sid as u128, false);
let aba_store = self_clone.aba.get_or_create_store(sessionid).await;
let aba_store_clone = aba_store.clone();
let self_clone2 = self_clone.clone();
let store_clone2 = store_clone.clone();
tokio::spawn(async move {
let notify = {
let aba = aba_store_clone.lock().await;
aba.notify.clone()
};
notify.notified().await;
let output = {
let aba = aba_store_clone.lock().await;
aba.output
};
{
let mut store = self_clone2.store.lock().await;
store.set_aba_output(sid as u128, output);
}
self_clone2.check_and_finalize_output(store_clone2).await;
});
}
}
}
});
} else if store.get_rbc_output(msg.session_id.slot().into()).is_none() {
store.set_rbc_output(msg.session_id.slot().into(), msg.payload);
let store_clone = self.store.clone();
let self_clone = self.clone();
drop(store);
info!(id = self.id, "Collect rbc");
tokio::spawn(async move {
self_clone.check_and_finalize_output(store_clone).await;
});
}
Ok(())
}
async fn check_and_finalize_output(&self, session_store: Arc<Mutex<AcsStore>>) {
let mut store = session_store.lock().await;
if store.aba_output.len() < self.n {
return;
}
let mut consensus_indices: Vec<u128> = store
.aba_output
.iter()
.filter(|(_, &v)| v)
.map(|(&id, _)| id)
.collect();
consensus_indices.sort();
let mut values = Vec::new();
for &j in &consensus_indices {
if let Some(value) = store.get_rbc_output(j) {
values.push(value.clone());
} else {
return;
}
}
info!(
id = self.id,
"ACS output finalized with {} values from {:?}",
values.len(),
consensus_indices
);
store.set_acs(values);
store.mark_ended();
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::honeybadger::{SessionId, WrappedMessage};
fn default_hb_rbc_wrap(msg: Msg<SessionId>) -> Result<Vec<u8>, RbcError> {
let wrapped = WrappedMessage::Rbc(msg);
Ok(bincode::serialize(&wrapped)?)
}
#[test]
fn test_bracha_avid_valid_params() {
let (s, _) = mpsc::channel(256);
let bracha = Bracha::new(0, 4, 1, 2, s.clone(), Arc::new(default_hb_rbc_wrap));
assert!(
bracha.is_ok(),
"Expected valid parameters for Bracha to succeed"
);
let avid = Avid::new(0, 6, 1, 2, s.clone(), Arc::new(default_hb_rbc_wrap));
assert!(
avid.is_ok(),
"Expected valid parameters for Avid to succeed"
);
let avid = Avid::new(1, 9, 2, 4, s.clone(), Arc::new(default_hb_rbc_wrap));
assert!(
avid.is_ok(),
"Expected valid parameters for Avid to succeed"
);
}
#[test]
fn test_bracha_avid_invalid_t() {
let (s, _) = mpsc::channel(256);
let bracha = Bracha::new(0, 4, 2, 2, s.clone(), Arc::new(default_hb_rbc_wrap)); assert!(bracha.is_err(), "Expected invalid t to fail for Bracha");
if let Err(msg) = bracha {
assert!(
msg.to_string().contains("t"),
"Expected error message to mention t for Bracha"
);
}
let avid = Avid::new(0, 6, 2, 2, s.clone(), Arc::new(default_hb_rbc_wrap)); assert!(avid.is_err(), "Expected invalid t to fail for Avid");
if let Err(msg) = avid {
assert!(
msg.to_string().contains("t"),
"Expected error message to mention t for Avid"
);
}
let avid = Avid::new(1, 9, 4, 4, s, Arc::new(default_hb_rbc_wrap)); assert!(avid.is_err(), "Expected invalid t to fail for Avid");
}
#[test]
fn test_bracha_avid_invalid_k() {
let (s, _) = mpsc::channel(256);
let avid = Avid::new(0, 6, 1, 0, s.clone(), Arc::new(default_hb_rbc_wrap)); assert!(avid.is_err(), "Expected invalid k to fail for Avid");
if let Err(msg) = avid {
assert!(
msg.to_string().contains("k"),
"Expected error message to mention k for Avid"
);
}
let avid = Avid::new(1, 9, 2, 7, s.clone(), Arc::new(default_hb_rbc_wrap)); assert!(avid.is_err(), "Expected invalid k to fail for Avid");
let bracha = Bracha::new(0, 5, 1, 3, s, Arc::new(default_hb_rbc_wrap)); assert!(bracha.is_ok(), "Expected valid parameters for Bracha");
}
#[test]
fn test_bracha_avid_edge_cases() {
let (s, _) = mpsc::channel(256);
let bracha = Bracha::new(0, 5, 1, 2, s.clone(), Arc::new(default_hb_rbc_wrap));
assert!(bracha.is_ok(), "Expected valid parameters for Bracha");
let avid = Avid::new(0, 5, 1, 2, s.clone(), Arc::new(default_hb_rbc_wrap));
assert!(avid.is_ok(), "Expected valid parameters for Avid");
let avid_invalid = Avid::new(0, 5, 2, 2, s.clone(), Arc::new(default_hb_rbc_wrap));
assert!(avid_invalid.is_err(), "Expected invalid k to fail for Avid");
let bracha_invalid = Bracha::new(0, 5, 2, 2, s, Arc::new(default_hb_rbc_wrap));
assert!(
bracha_invalid.is_err(),
"Expected invalid t to fail for Bracha"
);
}
#[test]
fn test_bracha_avid_zero_t() {
let (s, _) = mpsc::channel(256);
let bracha = Bracha::new(2, 3, 0, 1, s.clone(), Arc::new(default_hb_rbc_wrap));
assert!(bracha.is_ok(), "Expected t = 0 to be valid for Bracha");
let avid = Avid::new(2, 5, 0, 1, s.clone(), Arc::new(default_hb_rbc_wrap));
assert!(avid.is_ok(), "Expected t = 0 to be valid for Avid");
let bracha = Bracha::new(2, 3, 0, 2, s.clone(), Arc::new(default_hb_rbc_wrap));
assert!(bracha.is_ok(), "Expected valid parameters for Bracha");
let avid = Avid::new(3, 3, 0, 2, s.clone(), Arc::new(default_hb_rbc_wrap));
assert!(avid.is_ok(), "Expected valid parameters for Avid");
}
}