pub mod robust_interpolate;
pub mod batch_recon;
pub mod ran_dou_sha;
pub mod double_share;
pub mod triple_gen;
pub mod fpdiv;
pub mod fpmul;
pub mod input;
pub mod mul;
pub mod output;
pub mod preprocessing;
pub mod share_gen;
use crate::{
common::{
math::goldilocks::GoldilocksField,
rbc::{rbc_store::Msg, RbcError},
types::{
fixed::{ClearFixedPoint, SecretFixedPoint},
integer::{ClearInt, SecretInt},
TypeError,
},
MPCProtocol, MPCTypeOps, PreprocessingMPCProtocol, ProtocolSessionId, ProtocolTag,
ShamirShare, RBC,
},
honeybadger::{
batch_recon::{BatchReconError, BatchReconMsg},
double_share::{double_share_generation, DouShaError, DouShaMessage, DoubleShamirShare},
fpdiv::fpdiv_const::{FPDivConstError, FPDivConstNode},
fpmul::{
fpmul::{FPError, FPMulNode},
prandbitd::PRandBitDNode,
rand_bit::RandBit,
PRandBitDMessage, PRandError, RandBitError, TruncPrError,
},
input::{
input::{InputClient, InputServer},
InputError, InputMessage,
},
mul::{multiplication::Multiply, MulError},
output::{
output::{OutputClient, OutputServer},
OutputError, OutputMessage,
},
preprocessing::HoneyBadgerMPCNodePreprocMaterial,
ran_dou_sha::messages::RanDouShaMessage,
robust_interpolate::robust_interpolate::Robust,
share_gen::{share_gen::RanShaNode, RanShaError, RanShaMessage},
triple_gen::TripleGenError,
},
};
use ark_ff::{FftField, PrimeField};
use ark_std::rand::rngs::{OsRng, StdRng};
use ark_std::rand::{Rng, SeedableRng};
use async_trait::async_trait;
use bincode::{ErrorKind, Options};
use double_share_generation::DoubleShareNode;
use ran_dou_sha::{RanDouShaError, RanDouShaNode};
use robust_interpolate::robust_interpolate::RobustShare;
use serde::{Deserialize, Serialize};
use std::{fmt, sync::Arc, time::Instant};
use stoffelnet::network_utils::{ClientId, Network, NetworkError, PartyId};
use thiserror::Error;
use tokio::{sync::Mutex, time::Duration};
use tracing::{info, warn};
use triple_gen::triple_generation::TripleGenNode;
const MAX_MESSAGE_SIZE: u64 = 10 * 1024 * 1024;
fn preprocessing_trace_enabled() -> bool {
std::env::var("HMPC_PREPROCESSING_TRACE")
.map(|value| matches!(value.as_str(), "1" | "true" | "TRUE" | "yes" | "YES"))
.unwrap_or(false)
}
fn trace_preprocessing_phase(party_id: PartyId, phase: &str, items: usize, started: Instant) {
if preprocessing_trace_enabled() {
eprintln!(
"[hmpc preprocessing] party={} phase={} items={} elapsed_ms={}",
party_id,
phase,
items,
started.elapsed().as_millis()
);
}
}
fn triple_batch_groups_limit() -> usize {
std::env::var("HMPC_TRIPLE_BATCH_GROUPS")
.ok()
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| *value > 0)
.unwrap_or(4096)
}
fn ran_dou_sha_batch_columns_limit() -> usize {
std::env::var("HMPC_RANDOUSHA_BATCH_COLUMNS")
.ok()
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| *value > 0)
.unwrap_or(1536)
}
#[derive(Error, Debug)]
pub enum HoneyBadgerError {
#[error("network error: {0:?}")]
NetworkError(#[from] NetworkError),
#[error("error in share generation: {0:?}")]
RanShaError(#[from] RanShaError),
#[error("error in Input share generation: {0:?}")]
InputError(#[from] InputError),
#[error("error in faulty double share generation: {0:?}")]
DouShaError(#[from] DouShaError),
#[error("error in random double share generation: {0:?}")]
RanDouShaError(#[from] RanDouShaError),
#[error("there is not enough preprocessing to complete the protocol")]
NotEnoughPreprocessing,
#[error("error in triple generation protocol: {0:?}")]
TripleGenError(#[from] TripleGenError),
#[error("error in the RBC: {0:?}")]
RbcError(#[from] RbcError),
#[error("error in the Mul: {0:?}")]
MulError(#[from] MulError),
#[error("error in the Output server: {0:?}")]
OutputError(#[from] OutputError),
#[error("error in the Batch Reconstruction: {0:?}")]
BatchReconError(#[from] BatchReconError),
#[error("error in random bit generation: {0:?}")]
RandBitError(#[from] RandBitError),
#[error("error in Prand bit generation: {0:?}")]
PRandError(#[from] PRandError),
#[error("error in FPMul: {0:?}")]
FPError(#[from] FPError),
#[error("error in FPDiv_Const: {0:?}")]
FPDivConstError(#[from] FPDivConstError),
#[error("error in Truncation: {0:?}")]
TruncPrError(#[from] TruncPrError),
#[error("error in types: {0:?}")]
TypeError(#[from] TypeError),
#[error("Already reserved batch")]
AlreadyReserved,
#[error("error during the serialization using bincode: {0:?}")]
BincodeSerializationError(#[from] Box<ErrorKind>),
#[error("failed to join spawned task")]
JoinError,
#[error("instance ID {0:?} is incorrect")]
InstanceIdError(u32),
#[error("output channel closed before result was received")]
ChannelClosed,
#[error("Invalid threshold t={0} for n={1}, must satisfy t < ceil(n / 3)")]
InvalidThreshold(usize, usize),
#[error("Party size is too large")]
InvalidPartySize,
#[error("Party Id is out of bounds")]
InvalidPartyId,
#[error("the protocol cannot be executed any more")]
LimitError,
}
pub struct HoneyBadgerMPCClient<F: FftField, R: RBC> {
pub id: usize,
pub input: InputClient<F, R>,
pub output: OutputClient<F>,
}
impl<F, R> Clone for HoneyBadgerMPCClient<F, R>
where
F: FftField,
R: RBC,
{
fn clone(&self) -> Self {
Self {
id: self.id,
input: self.input.clone(),
output: self.output.clone(),
}
}
}
impl<F: FftField, R: RBC<Id = SessionId>> HoneyBadgerMPCClient<F, R> {
pub fn new(
id: usize,
n: usize,
t: usize,
instance_id: u32,
inputs: Vec<F>,
input_len: usize,
) -> Result<Self, HoneyBadgerError> {
let input = InputClient::new(id, n, t, instance_id, inputs)?;
let output = OutputClient::new(id, n, t, input_len)?;
Ok(Self { id, input, output })
}
pub async fn process<N: Network + Send + Sync>(
&mut self,
sender_id: ClientId,
raw_msg: Vec<u8>,
net: Arc<N>,
) -> Result<(), HoneyBadgerError> {
let wrapped: WrappedMessage = bincode::DefaultOptions::new()
.with_fixint_encoding()
.allow_trailing_bytes()
.with_limit(MAX_MESSAGE_SIZE)
.deserialize(&raw_msg)?;
match wrapped {
WrappedMessage::Input(input_msg) => {
if sender_id != input_msg.sender_id {
return Err(HoneyBadgerError::InvalidPartyId);
}
self.input.process(input_msg, net).await?;
}
WrappedMessage::Output(output_msg) => {
if sender_id != output_msg.sender_id {
return Err(HoneyBadgerError::InvalidPartyId);
}
self.output.process(output_msg).await?
}
_ => warn!("Incorrect message type recieved at input"),
}
Ok(())
}
}
#[derive(Clone, Debug)]
pub struct HoneyBadgerMPCNode<F: PrimeField, R: RBC> {
pub id: PartyId,
pub preprocessing_material: Arc<Mutex<HoneyBadgerMPCNodePreprocMaterial<F>>>,
pub params: HoneyBadgerMPCNodeOpts,
pub preprocess: PreprocessNodes<F, R>,
pub operations: Operation<F, R>,
pub type_ops: TypeOperations<F, R>,
pub output: OutputServer,
pub counters: SubProtocolCounters,
}
impl<F, R> HoneyBadgerMPCNode<F, R>
where
F: PrimeField,
R: RBC<Id = SessionId>,
{
pub async fn debug_store_sizes(&self) -> String {
let len = self.preprocessing_material.lock().await.length();
let triples = len.beaver_triples;
let random_shares = len.random_shr;
let prandbit = len.prandbit;
let prandint = len.prandint;
format!(
"material=(triples:{triples},random:{random_shares},prandbit:{prandbit},prandint:{prandint}) \
stores=(share_gen:{},dou_sha:{},ran_dou_sha:{},triple:{},triple_batch_recon:{},mul:{},rand_bit:{},rand_bit_mul:{},rand_bit_batch_recon:{},prand_bit:{},prand_bit_batch_recon:{},fpmul_mul:{},fpmul_trunc:{})",
self.preprocess.share_gen.store_len().await,
self.preprocess.dou_sha.store_len().await,
self.preprocess.ran_dou_sha.store_len().await,
self.preprocess.triple_gen.store_len().await,
self.preprocess.triple_gen.batch_recon_node.store_len().await,
self.operations.mul.store_len().await,
self.preprocess.small_field_preproc.rand_bit.store_len().await,
self.preprocess.small_field_preproc.rand_bit.mult_node.store_len().await,
self.preprocess.small_field_preproc.rand_bit.batch_recon.store_len().await,
self.preprocess.prand_bit.store_len().await,
self.preprocess.prand_bit.batch_recon.store_len().await,
self.type_ops.fpmul.mult_node.store_len().await,
self.type_ops.fpmul.trunc_node.store_len().await,
)
}
}
#[derive(Clone, Debug)]
pub struct Operation<F: FftField, R: RBC> {
pub mul: Multiply<F, R>,
}
#[derive(Clone, Debug)]
pub struct TypeOperations<F: PrimeField, R: RBC> {
pub fpmul: FPMulNode<F, R>,
pub fpdiv_const: FPDivConstNode<F, R>,
}
#[derive(Clone, Debug)]
pub struct PreprocessNodes<F: PrimeField, R: RBC> {
pub input: InputServer<F, R>,
pub share_gen: RanShaNode<F, R>,
pub dou_sha: DoubleShareNode<F>,
pub ran_dou_sha: RanDouShaNode<F, R>,
pub triple_gen: TripleGenNode<F>,
pub prand_bit: PRandBitDNode<GoldilocksField, F>,
pub small_field_preproc: PreprocNodesSmallField<R>,
}
#[derive(Clone, Debug)]
pub struct PreprocNodesSmallField<R: RBC> {
pub share_gen: RanShaNode<GoldilocksField, R>,
pub triple_gen: TripleGenNode<GoldilocksField>,
pub rand_bit: RandBit<GoldilocksField, R>,
pub ran_dou_sha: RanDouShaNode<GoldilocksField, R>,
pub dou_sha: DoubleShareNode<GoldilocksField>,
}
#[derive(Clone, Debug)]
pub struct SubProtocolCounter(Arc<Mutex<Option<u64>>>);
trait GetNext<T> {
async fn get_next(&self) -> Result<T, HoneyBadgerError>;
}
impl GetNext<u64> for SubProtocolCounter {
async fn get_next(&self) -> Result<u64, HoneyBadgerError> {
let mut counter = self.0.lock().await;
match &mut *counter {
None => Err(HoneyBadgerError::LimitError),
Some(value) => {
let current = *value;
if *value == u64::MAX {
*counter = None;
} else {
*value += 1;
}
Ok(current)
}
}
}
}
#[derive(Clone, Debug)]
pub struct SubProtocolCounters {
pub ran_dou_sha_counter: SubProtocolCounter,
pub ran_sha_counter: SubProtocolCounter,
pub triple_counter: SubProtocolCounter,
pub batch_recon_counter: SubProtocolCounter,
pub dou_sha_counter: SubProtocolCounter,
pub mul_counter: SubProtocolCounter,
pub rand_bit_counter: SubProtocolCounter,
pub prand_bit_counter: SubProtocolCounter,
pub prand_int_counter: SubProtocolCounter,
pub fpmul_counter: SubProtocolCounter,
pub fpdiv_const_counter: SubProtocolCounter,
pub ran_sha_small_field_counter: SubProtocolCounter,
pub triple_small_field_counter: SubProtocolCounter,
pub rand_bit_small_field_counter: SubProtocolCounter,
pub dou_sha_small_field_counter: SubProtocolCounter,
pub ran_dou_sha_small_field_counter: SubProtocolCounter,
}
impl SubProtocolCounters {
pub fn new() -> Self {
Self {
ran_dou_sha_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
ran_sha_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
triple_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
batch_recon_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
dou_sha_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
mul_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
rand_bit_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
prand_bit_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
prand_int_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
fpmul_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
fpdiv_const_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
ran_sha_small_field_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
triple_small_field_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
rand_bit_small_field_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
dou_sha_small_field_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
ran_dou_sha_small_field_counter: SubProtocolCounter(Arc::new(Mutex::new(Some(0)))),
}
}
}
#[derive(Clone, Debug)]
pub struct HoneyBadgerMPCNodeOpts {
pub n_parties: usize,
pub threshold: usize,
pub n_triples: usize,
pub n_random_shares: usize,
pub instance_id: u32,
pub n_prandbit: usize,
pub n_prandint: usize,
pub k: usize,
pub l: usize,
pub timeout: Duration,
}
impl HoneyBadgerMPCNodeOpts {
pub fn new(
n_parties: usize,
threshold: usize,
n_triples: usize,
n_random_shares: usize,
instance_id: u32,
n_prandbit: usize,
n_prandint: usize,
l: usize,
k: usize,
timeout: Duration,
) -> Result<Self, HoneyBadgerError> {
if n_parties > 255 {
return Err(HoneyBadgerError::InvalidPartySize);
}
if !(threshold < (n_parties + 2) / 3) {
return Err(HoneyBadgerError::InvalidThreshold(threshold, n_parties));
}
Ok(Self {
n_parties,
threshold,
n_triples,
n_random_shares,
instance_id,
n_prandbit,
n_prandint,
k,
l,
timeout,
})
}
pub fn set_timeout(&mut self, secs: u64) {
self.timeout = Duration::from_secs(secs)
}
}
#[async_trait]
impl<F, R, N> MPCProtocol<F, RobustShare<F>, N> for HoneyBadgerMPCNode<F, R>
where
N: Network + Send + Sync + 'static,
F: PrimeField,
R: RBC<Id = SessionId>,
{
type MPCOpts = HoneyBadgerMPCNodeOpts;
type Error = HoneyBadgerError;
fn setup(
id: PartyId,
params: Self::MPCOpts,
input_ids: Vec<ClientId>,
) -> Result<Self, HoneyBadgerError> {
if id >= params.n_parties {
return Err(HoneyBadgerError::InvalidPartyId);
}
let dousha_node = DoubleShareNode::new(id, params.n_parties, params.threshold);
let prand_bit_node = PRandBitDNode::new(id, params.n_parties, params.threshold)?;
let ran_dou_sha_node =
RanDouShaNode::new(id, params.n_parties, params.threshold, params.threshold + 1)?;
let triple_gen_node = TripleGenNode::new(id, params.n_parties, params.threshold)?;
let mul_node = Multiply::new(id, params.n_parties, params.threshold)?;
let share_gen =
RanShaNode::new(id, params.n_parties, params.threshold, params.threshold + 1)?;
let fpmul_node = FPMulNode::new(id, params.n_parties, params.threshold)?;
let fpdiv_const_node = FPDivConstNode::new(id, params.n_parties, params.threshold)?;
let input = InputServer::new(id, params.n_parties, params.threshold, input_ids)?;
let output = OutputServer::new(id, params.n_parties)?;
let triple_gen_small_field_node =
TripleGenNode::new(id, params.n_parties, params.threshold)?;
let share_gen_small_field =
RanShaNode::new(id, params.n_parties, params.threshold, params.threshold + 1)?;
let rand_bit_node = RandBit::new(id, params.n_parties, params.threshold)?;
let ran_dou_sha_small_field =
RanDouShaNode::new(id, params.n_parties, params.threshold, params.threshold + 1)?;
let dousha_node_small_field = DoubleShareNode::new(id, params.n_parties, params.threshold);
let small_field_preproc = PreprocNodesSmallField {
triple_gen: triple_gen_small_field_node,
rand_bit: rand_bit_node,
share_gen: share_gen_small_field,
ran_dou_sha: ran_dou_sha_small_field,
dou_sha: dousha_node_small_field,
};
Ok(Self {
id,
preprocessing_material: Arc::new(
Mutex::new(HoneyBadgerMPCNodePreprocMaterial::empty()),
),
params,
preprocess: PreprocessNodes {
input,
share_gen,
dou_sha: dousha_node,
ran_dou_sha: ran_dou_sha_node,
triple_gen: triple_gen_node,
prand_bit: prand_bit_node,
small_field_preproc,
},
operations: Operation { mul: mul_node },
type_ops: TypeOperations {
fpmul: fpmul_node,
fpdiv_const: fpdiv_const_node,
},
output,
counters: SubProtocolCounters::new(),
})
}
async fn mul(
&mut self,
x: Vec<RobustShare<F>>,
y: Vec<RobustShare<F>>,
network: Arc<N>,
) -> Result<Vec<RobustShare<F>>, Self::Error> {
assert_eq!(x.len(), y.len());
if x.is_empty() {
return Ok(Vec::new());
}
let no_triples = {
let store = self.preprocessing_material.lock().await;
store.length().beaver_triples
};
if no_triples < x.len() {
let mut rng = StdRng::from_rng(OsRng).unwrap();
self.run_preprocessing(network.clone(), &mut rng).await?;
}
let max_pairs_per_session = max_mul_pairs_per_session(self.params.threshold);
let mut result = Vec::with_capacity(x.len());
let mut session_ids = Vec::new();
for (x_chunk, y_chunk) in x
.chunks(max_pairs_per_session)
.zip(y.chunks(max_pairs_per_session))
{
let beaver_triples = self
.preprocessing_material
.lock()
.await
.take_beaver_triples(x_chunk.len())?;
let session_id = SessionId::new(
ProtocolType::Mul,
SessionId::pack_slot(self.counters.mul_counter.get_next().await?, 0, 0),
self.params.instance_id,
);
self.operations
.mul
.init(
session_id,
x_chunk.to_vec(),
y_chunk.to_vec(),
beaver_triples,
network.clone(),
)
.await?;
session_ids.push(session_id);
}
for session_id in &session_ids {
let mut chunk_result = self
.operations
.mul
.wait_for_result(*session_id, self.params.timeout)
.await
.map_err(HoneyBadgerError::from)?;
result.append(&mut chunk_result);
}
for session_id in &session_ids {
if let Err(error) = self.operations.mul.clear_store(*session_id).await {
warn!(
?session_id,
?error,
"failed to clear completed multiplication protocol state"
);
}
}
Ok(result)
}
async fn rand(&mut self, network: Arc<N>) -> Result<RobustShare<F>, Self::Error> {
let no_rand = {
let store = self.preprocessing_material.lock().await;
store.length().random_shr
};
if no_rand == 0 {
let mut rng = StdRng::from_rng(OsRng).unwrap();
self.run_preprocessing(network.clone(), &mut rng).await?;
}
let rand_value = self
.preprocessing_material
.lock()
.await
.take_random_shares(1)?;
Ok(rand_value[0].clone())
}
async fn process(
&mut self,
sender_id: PartyId,
raw_msg: Vec<u8>,
net: Arc<N>,
) -> Result<(), Self::Error> {
let wrapped: WrappedMessage = bincode::DefaultOptions::new()
.with_fixint_encoding()
.allow_trailing_bytes()
.with_limit(MAX_MESSAGE_SIZE)
.deserialize(&raw_msg)?;
match wrapped {
WrappedMessage::Rbc(rbc_msg) => {
if sender_id != rbc_msg.sender_id {
return Err(HoneyBadgerError::InvalidPartyId);
}
if rbc_msg.session_id.instance_id() != self.params.instance_id {
return Err(HoneyBadgerError::InstanceIdError(
rbc_msg.session_id.instance_id(),
));
}
if rbc_msg.msg_type.is_dealer_message() {
let expected_dealer = rbc_msg.session_id.sub_id() as usize;
if rbc_msg.sender_id != expected_dealer {
warn!(
"Rejecting dealer message: sender {} is not expected dealer {} for session {:?}",
rbc_msg.sender_id, expected_dealer, rbc_msg.session_id
);
return Err(HoneyBadgerError::InvalidPartyId);
}
}
match rbc_msg.session_id.calling_protocol() {
Some(ProtocolType::Randousha) => {
self.preprocess
.ran_dou_sha
.rbc
.process(rbc_msg, net)
.await?;
self.preprocess.ran_dou_sha.drain_rbc_output().await?;
}
Some(ProtocolType::RanDouShaSmallField) => {
self.preprocess
.small_field_preproc
.ran_dou_sha
.rbc
.process(rbc_msg, net)
.await?;
self.preprocess
.small_field_preproc
.ran_dou_sha
.drain_rbc_output()
.await?;
}
Some(ProtocolType::Ransha) => {
self.preprocess.share_gen.rbc.process(rbc_msg, net).await?;
self.preprocess.share_gen.drain_rbc_output().await?;
}
Some(ProtocolType::RanShaSmallField) => {
self.preprocess
.small_field_preproc
.share_gen
.rbc
.process(rbc_msg, net)
.await?;
self.preprocess
.small_field_preproc
.share_gen
.drain_rbc_output()
.await?;
}
Some(ProtocolType::Input) => {
self.preprocess.input.rbc.process(rbc_msg, net).await?;
self.preprocess.input.drain_rbc_output().await?;
}
Some(ProtocolType::Mul) => {
self.operations.mul.rbc.process(rbc_msg, net).await?;
self.operations.mul.drain_rbc_output().await?;
}
Some(ProtocolType::RandBit) => {
self.preprocess
.small_field_preproc
.rand_bit
.mult_node
.rbc
.process(rbc_msg, net)
.await?;
self.preprocess
.small_field_preproc
.rand_bit
.mult_node
.drain_rbc_output()
.await?;
}
Some(ProtocolType::FpMul) => {
if rbc_msg.session_id.round_id() == 0 {
self.type_ops
.fpmul
.trunc_node
.rbc
.process(rbc_msg, net)
.await?;
self.type_ops.fpmul.trunc_node.drain_rbc_output().await?;
} else {
self.type_ops
.fpmul
.mult_node
.rbc
.process(rbc_msg, net)
.await?;
self.type_ops.fpmul.mult_node.drain_rbc_output().await?;
}
}
Some(ProtocolType::FpDivConst) => {
self.type_ops
.fpdiv_const
.trunc_node
.rbc
.process(rbc_msg, net)
.await?;
self.type_ops
.fpdiv_const
.trunc_node
.drain_rbc_output()
.await?;
}
_ => {
warn!(
"Unknown protocol ID in session ID: {:?} in RBC",
rbc_msg.session_id
);
}
}
}
WrappedMessage::RanSha(rs_msg) => {
if sender_id != rs_msg.sender_id {
return Err(HoneyBadgerError::InvalidPartyId);
}
if rs_msg.session_id.instance_id() != self.params.instance_id {
return Err(HoneyBadgerError::InstanceIdError(
rs_msg.session_id.instance_id(),
));
}
if let Some(ProtocolType::RanShaSmallField) = rs_msg.session_id.calling_protocol() {
self.preprocess
.small_field_preproc
.share_gen
.process(rs_msg, net)
.await?;
} else {
self.preprocess.share_gen.process(rs_msg, net).await?;
}
}
WrappedMessage::Dousha(ds_msg) => {
if sender_id != ds_msg.sender_id {
return Err(HoneyBadgerError::InvalidPartyId);
}
if ds_msg.session_id.instance_id() != self.params.instance_id {
return Err(HoneyBadgerError::InstanceIdError(
ds_msg.session_id.instance_id(),
));
}
if let Some(ProtocolType::DouShaSmallField) = ds_msg.session_id.calling_protocol() {
self.preprocess
.small_field_preproc
.dou_sha
.process(ds_msg)
.await?;
} else {
self.preprocess.dou_sha.process(ds_msg).await?;
}
}
WrappedMessage::RanDouSha(rds_msg) => {
if sender_id != rds_msg.sender_id {
return Err(HoneyBadgerError::InvalidPartyId);
}
if rds_msg.session_id.instance_id() != self.params.instance_id {
return Err(HoneyBadgerError::InstanceIdError(
rds_msg.session_id.instance_id(),
));
}
if let Some(ProtocolType::RanDouShaSmallField) =
rds_msg.session_id.calling_protocol()
{
self.preprocess
.small_field_preproc
.ran_dou_sha
.process(rds_msg, net)
.await?;
} else {
self.preprocess.ran_dou_sha.process(rds_msg, net).await?;
}
}
WrappedMessage::BatchRecon(batch_msg) => {
if sender_id != batch_msg.sender_id {
return Err(HoneyBadgerError::InvalidPartyId);
}
if batch_msg.session_id.instance_id() != self.params.instance_id {
return Err(HoneyBadgerError::InstanceIdError(
batch_msg.session_id.instance_id(),
));
}
match batch_msg.session_id.calling_protocol() {
Some(ProtocolType::Mul) => {
self.operations
.mul
.batch_recon
.process(batch_msg, net)
.await?;
self.operations.mul.drain_batch_recon_output().await?
}
Some(ProtocolType::Triple) => {
self.preprocess
.triple_gen
.batch_recon_node
.process(batch_msg, net)
.await?;
self.preprocess
.triple_gen
.drain_batch_recon_output()
.await?
}
Some(ProtocolType::TripleSmallField) => {
self.preprocess
.small_field_preproc
.triple_gen
.batch_recon_node
.process(batch_msg, net)
.await?;
self.preprocess
.small_field_preproc
.triple_gen
.drain_batch_recon_output()
.await?
}
Some(ProtocolType::RandBit) => {
if batch_msg.session_id.round_id() == 0 {
self.preprocess
.small_field_preproc
.rand_bit
.batch_recon
.process(batch_msg, net)
.await?;
self.preprocess
.small_field_preproc
.rand_bit
.drain_batch_recon_output()
.await?;
} else {
self.preprocess
.small_field_preproc
.rand_bit
.mult_node
.batch_recon
.process(batch_msg, net)
.await?;
self.preprocess
.small_field_preproc
.rand_bit
.mult_node
.drain_batch_recon_output()
.await?;
}
}
Some(ProtocolType::PRandBit) => {
self.preprocess
.prand_bit
.batch_recon
.process(batch_msg, net)
.await?;
self.preprocess.prand_bit.drain_batch_recon_output().await?;
}
Some(ProtocolType::FpMul) => {
self.type_ops
.fpmul
.mult_node
.batch_recon
.process(batch_msg, net)
.await?;
self.type_ops
.fpmul
.mult_node
.drain_batch_recon_output()
.await?;
}
_ => {
warn!(
"Unknown protocol ID in session ID: {:?} at Batch reconstruction",
batch_msg.session_id
);
}
}
}
WrappedMessage::PRandBitD(prand_message) => {
if sender_id != prand_message.sender_id {
return Err(HoneyBadgerError::InvalidPartyId);
}
if prand_message.session_id.instance_id() != self.params.instance_id {
return Err(HoneyBadgerError::InstanceIdError(
prand_message.session_id.instance_id(),
));
}
self.preprocess
.prand_bit
.process(prand_message, net)
.await?;
}
WrappedMessage::Input(_) => warn!("Incorrect message recieved at process function"),
WrappedMessage::Output(_) => warn!("Incorrect message recieved at process function"),
}
Ok(())
}
}
#[async_trait]
impl<F, N, R> MPCTypeOps<F, RobustShare<F>, N> for HoneyBadgerMPCNode<F, R>
where
F: PrimeField,
N: Network + Send + Sync + 'static,
R: RBC<Id = SessionId>,
{
type Error = HoneyBadgerError;
type Sfix = SecretFixedPoint<F, RobustShare<F>>;
type Sint = SecretInt<F, RobustShare<F>>;
type Cfix = ClearFixedPoint<F>;
type Cint = ClearInt<F>;
async fn add_fixed(
&self,
x: Vec<Self::Sfix>,
y: Vec<Self::Sfix>,
) -> Result<Vec<Self::Sfix>, Self::Error> {
if x.len() != y.len() {
return Err(HoneyBadgerError::FPError(FPError::IncompatiblePrecision));
}
Ok(x.into_iter()
.zip(y)
.map(|(a, b)| a + b)
.collect::<Result<Vec<_>, _>>()?)
}
async fn sub_fixed(
&self,
x: Vec<Self::Sfix>,
y: Vec<Self::Sfix>,
) -> Result<Vec<Self::Sfix>, Self::Error> {
if x.len() != y.len() {
return Err(HoneyBadgerError::FPError(FPError::IncompatiblePrecision));
}
Ok(x.into_iter()
.zip(y)
.map(|(a, b)| a - b)
.collect::<Result<Vec<_>, _>>()?)
}
async fn mul_fixed(
&mut self,
x: SecretFixedPoint<F, RobustShare<F>>,
y: SecretFixedPoint<F, RobustShare<F>>,
net: Arc<N>,
) -> Result<SecretFixedPoint<F, RobustShare<F>>, Self::Error> {
if x.precision() != y.precision() {
return Err(HoneyBadgerError::FPError(FPError::IncompatiblePrecision));
}
let (no_rand_bit, no_rand_int) = {
let store = self.preprocessing_material.lock().await;
(store.length().prandbit, store.length().prandint)
};
if no_rand_bit < x.precision().f() || no_rand_int == 0 {
let mut rng = StdRng::from_rng(OsRng).unwrap();
self.run_preprocessing(net.clone(), &mut rng).await?;
}
let beaver_triples = self
.preprocessing_material
.lock()
.await
.take_beaver_triples(1)?;
let r_bits_vec = self
.preprocessing_material
.lock()
.await
.take_prandbit_shares(x.precision().f())?;
let r_int = self
.preprocessing_material
.lock()
.await
.take_prandint_shares(1)?;
let session_id = SessionId::new(
ProtocolType::FpMul,
SessionId::pack_slot(self.counters.fpmul_counter.get_next().await?, 0, 0),
self.params.instance_id,
);
let r_bits = r_bits_vec.iter().map(|(a, _)| a.clone()).collect();
self.type_ops
.fpmul
.init(
x,
y,
beaver_triples[0].clone(),
r_bits,
r_int[0].clone(),
self.params.timeout,
session_id,
net,
)
.await
.map_err(HoneyBadgerError::from)
}
async fn div_with_const_fixed(
&mut self,
x: SecretFixedPoint<F, RobustShare<F>>,
y: ClearFixedPoint<F>,
net: Arc<N>,
) -> Result<SecretFixedPoint<F, RobustShare<F>>, Self::Error> {
if x.precision() != y.precision() {
return Err(HoneyBadgerError::FPDivConstError(
FPDivConstError::IncompatiblePrecision,
));
}
let (no_rand_bit, no_rand_int) = {
let store = self.preprocessing_material.lock().await;
(store.length().prandbit, store.length().prandint)
};
if no_rand_bit < x.precision().f() || no_rand_int == 0 {
let mut rng = StdRng::from_rng(OsRng).unwrap();
self.run_preprocessing(net.clone(), &mut rng).await?;
}
let r_bits_vec = self
.preprocessing_material
.lock()
.await
.take_prandbit_shares(x.precision().f())?;
let r_int = self
.preprocessing_material
.lock()
.await
.take_prandint_shares(1)?;
let r_bits_only = r_bits_vec
.iter()
.map(|(a, _)| a.clone())
.collect::<Vec<_>>();
let session_id = SessionId::new(
ProtocolType::FpDivConst,
SessionId::pack_slot(self.counters.fpdiv_const_counter.get_next().await?, 0, 0),
self.params.instance_id,
);
self.type_ops
.fpdiv_const
.init(
x,
y,
r_bits_only,
r_int[0].clone(),
self.params.timeout,
session_id,
net.clone(),
)
.await
.map_err(HoneyBadgerError::from)
}
async fn add_int(
&self,
x: Vec<Self::Sint>,
y: Vec<Self::Sint>,
) -> Result<Vec<Self::Sint>, Self::Error> {
if x.len() != y.len() {
return Err(HoneyBadgerError::FPError(FPError::IncompatiblePrecision));
}
let mut out = Vec::with_capacity(x.len());
for (a, b) in x.into_iter().zip(y.into_iter()) {
let sum = (a + b)?;
out.push(sum);
}
Ok(out)
}
async fn sub_int(
&self,
x: Vec<Self::Sint>,
y: Vec<Self::Sint>,
) -> Result<Vec<Self::Sint>, Self::Error> {
if x.len() != y.len() {
return Err(HoneyBadgerError::FPError(FPError::IncompatiblePrecision));
}
let mut out = Vec::with_capacity(x.len());
for (a, b) in x.into_iter().zip(y.into_iter()) {
let sum = (a - b)?;
out.push(sum);
}
Ok(out)
}
async fn mul_int(
&mut self,
x: Vec<Self::Sint>,
y: Vec<Self::Sint>,
net: Arc<N>,
) -> Result<Vec<Self::Sint>, Self::Error> {
if x.len() != y.len() {
return Err(HoneyBadgerError::FPError(FPError::IncompatiblePrecision));
}
let bitlen_x = x
.first()
.map(|v| v.bit_length())
.ok_or(HoneyBadgerError::FPError(FPError::IncompatiblePrecision))?;
let x_ok = x.iter().all(|v| v.bit_length() == bitlen_x);
if !x_ok {
return Err(HoneyBadgerError::FPError(FPError::IncompatiblePrecision));
}
let bitlen_y = y
.first()
.map(|v| v.bit_length())
.ok_or(HoneyBadgerError::FPError(FPError::IncompatiblePrecision))?;
let y_ok = y.iter().all(|v| v.bit_length() == bitlen_y);
if !y_ok {
return Err(HoneyBadgerError::FPError(FPError::IncompatiblePrecision));
}
if bitlen_x != bitlen_y {
return Err(HoneyBadgerError::FPError(FPError::IncompatiblePrecision));
}
let bitlen = bitlen_x;
let a: Vec<ShamirShare<F, 1, Robust>> = x.iter().map(|s| s.share().clone()).collect();
let b: Vec<ShamirShare<F, 1, Robust>> = y.iter().map(|s| s.share().clone()).collect();
let result = self.mul(a, b, net).await?;
let output = result
.into_iter()
.map(|share| SecretInt::new(share, bitlen))
.collect();
Ok(output)
}
}
#[async_trait]
impl<F, R, N> PreprocessingMPCProtocol<F, RobustShare<F>, N> for HoneyBadgerMPCNode<F, R>
where
N: Network + Send + Sync + 'static,
F: PrimeField,
R: RBC<Id = SessionId>,
{
async fn run_preprocessing<G>(
&mut self,
network: Arc<N>,
rng: &mut G,
) -> Result<(), Self::Error>
where
N: 'async_trait,
G: Rng + Send,
{
let (no_of_triples_avail, no_of_random_shares_avail) = {
let store = self.preprocessing_material.lock().await;
(store.length().beaver_triples, store.length().random_shr)
};
let mut no_of_triples = self.params.n_triples;
let mut no_of_random_shares = self.params.n_random_shares;
let group_size = 2 * self.params.threshold + 1;
let total_triples_to_generate = if no_of_triples_avail >= no_of_triples {
no_of_triples = 0;
0
} else {
((no_of_triples - no_of_triples_avail + group_size - 1) / group_size) * group_size
};
let total_random_shares_to_generate = if total_triples_to_generate > 0 {
let baseline = if no_of_random_shares_avail < no_of_random_shares {
no_of_random_shares - no_of_random_shares_avail
} else {
no_of_random_shares = 0;
0
};
baseline + 2 * total_triples_to_generate
} else if no_of_random_shares_avail < no_of_random_shares {
no_of_random_shares - no_of_random_shares_avail
} else {
no_of_random_shares = 0;
0
};
if no_of_triples == 0 && no_of_random_shares == 0 {
info!("There are enough Random shares and Beaver triples");
} else {
let mut triple_counter = self.counters.triple_counter.get_next().await?;
let phase_start = Instant::now();
self.ensure_random_shares(network.clone(), rng, total_random_shares_to_generate)
.await?;
trace_preprocessing_phase(
self.id,
"random_shares",
total_random_shares_to_generate,
phase_start,
);
info!("Random share generation done");
let phase_start = Instant::now();
let ran_dou_sha_pair = self
.ensure_ran_dou_sha_pair(network.clone(), rng, total_triples_to_generate)
.await?;
trace_preprocessing_phase(self.id, "randousha", total_triples_to_generate, phase_start);
info!("Randousha pair generation done");
let random_shares_a = self
.preprocessing_material
.lock()
.await
.take_random_shares(total_triples_to_generate)?;
let random_shares_b = self
.preprocessing_material
.lock()
.await
.take_random_shares(total_triples_to_generate)?;
let mut round_id = 0u8;
let mut group_index = 0;
let total_groups = total_triples_to_generate / group_size;
let phase_start = Instant::now();
let max_batch_groups = triple_batch_groups_limit();
let mut sessions: Vec<(SessionId, usize, usize)> = Vec::new();
while group_index < total_groups {
let batch_groups = (total_groups - group_index).min(max_batch_groups);
let share_start = group_index * group_size;
let share_end = share_start + batch_groups * group_size;
let sessionid = SessionId::new(
ProtocolType::Triple,
SessionId::pack_slot(triple_counter, 0, round_id),
self.params.instance_id,
);
sessions.push((sessionid, share_start, share_end));
if round_id == 255 {
triple_counter = self.counters.triple_counter.get_next().await?;
round_id = 0;
} else {
round_id += 1;
}
group_index += batch_groups;
}
for (sessionid, share_start, share_end) in &sessions {
self.preprocess
.triple_gen
.init_batch(
random_shares_a[*share_start..*share_end].to_vec(),
random_shares_b[*share_start..*share_end].to_vec(),
ran_dou_sha_pair[*share_start..*share_end].to_vec(),
*sessionid,
network.clone(),
)
.await?;
}
for (sessionid, _, _) in &sessions {
let triples = self
.preprocess
.triple_gen
.wait_for_result(*sessionid, self.params.timeout)
.await?;
self.preprocessing_material.lock().await.add(
Some(triples),
None,
None,
None,
None,
None,
);
assert!(self.preprocess.triple_gen.clear_store(*sessionid).await);
}
trace_preprocessing_phase(self.id, "triples", total_triples_to_generate, phase_start);
}
let phase_start = Instant::now();
self.ensure_prandbit_shares(rng, network.clone()).await?;
trace_preprocessing_phase(self.id, "prandbit", self.params.n_prandbit, phase_start);
info!("PrandBit share generation done");
let phase_start = Instant::now();
self.ensure_prandint_shares(network.clone()).await?;
trace_preprocessing_phase(self.id, "prandint", self.params.n_prandint, phase_start);
info!("PrandInt share generation done");
Ok(())
}
}
impl<F, R> HoneyBadgerMPCNode<F, R>
where
F: PrimeField,
R: RBC<Id = SessionId>,
{
async fn ensure_random_shares<G, N>(
&mut self,
network: Arc<N>,
rng: &mut G,
needed: usize,
) -> Result<(), HoneyBadgerError>
where
N: Network + Send + Sync + 'static,
G: Rng + Send,
{
let output_per_column = self.params.n_parties - 2 * self.params.threshold;
let columns_needed = (needed + output_per_column - 1) / output_per_column;
let max_columns_per_run = 2048usize;
let run = (columns_needed + max_columns_per_run - 1) / max_columns_per_run;
let mut round_id = 0u8;
let mut ran_sha_counter = self.counters.ran_sha_counter.get_next().await?;
let mut sessions: Vec<(SessionId, usize)> = Vec::with_capacity(run);
for i in 0..run {
info!("Random share generation run {}", i);
let columns_remaining = columns_needed - i * max_columns_per_run;
let batch_size = columns_remaining.min(max_columns_per_run);
let sessionid = SessionId::new(
ProtocolType::Ransha,
SessionId::pack_slot(ran_sha_counter, 0, round_id),
self.params.instance_id,
);
sessions.push((sessionid, batch_size));
if round_id == 255 {
ran_sha_counter = self.counters.ran_sha_counter.get_next().await.unwrap();
round_id = 0;
} else {
round_id += 1;
}
}
for (sessionid, batch_size) in &sessions {
self.preprocess
.share_gen
.init_batch(*sessionid, *batch_size, rng, network.clone())
.await?;
}
for (sessionid, _) in &sessions {
let output = self
.preprocess
.share_gen
.wait_for_result(*sessionid, self.params.timeout)
.await?;
self.preprocessing_material.lock().await.add(
None,
None,
Some(output),
None,
None,
None,
);
assert!(self.preprocess.share_gen.clear_store(*sessionid).await);
}
Ok(())
}
async fn ensure_ran_dou_sha_pair<G, N>(
&mut self,
network: Arc<N>,
rng: &mut G,
needed: usize,
) -> Result<Vec<DoubleShamirShare<F>>, HoneyBadgerError>
where
N: Network + Send + Sync + 'static,
G: Rng + Send,
{
let mut pair = Vec::new();
let output_per_column = self.params.threshold + 1;
let columns_needed = (needed + output_per_column - 1) / output_per_column;
let max_columns_per_run = ran_dou_sha_batch_columns_limit();
let run = (columns_needed + max_columns_per_run - 1) / max_columns_per_run;
let mut round_id = 0u8;
let mut ran_dou_sha_counter = self.counters.ran_dou_sha_counter.get_next().await?;
let mut sessions: Vec<(SessionId, usize)> = Vec::with_capacity(run);
for i in 0..run {
let columns_remaining = columns_needed - i * max_columns_per_run;
let batch_size = columns_remaining.min(max_columns_per_run);
let sessionid = SessionId::new(
ProtocolType::Randousha,
SessionId::pack_slot(ran_dou_sha_counter, 0, round_id),
self.params.instance_id,
);
sessions.push((sessionid, batch_size));
if round_id == 255 {
ran_dou_sha_counter = self.counters.ran_dou_sha_counter.get_next().await.unwrap();
round_id = 0;
} else {
round_id += 1;
}
}
for (sessionid, batch_size) in &sessions {
self.preprocess
.dou_sha
.init_batch(*sessionid, *batch_size, rng, network.clone())
.await?;
}
let mut all_double_shares: Vec<Vec<DoubleShamirShare<F>>> =
Vec::with_capacity(sessions.len());
for (sessionid, _) in &sessions {
let double_shares = self
.preprocess
.dou_sha
.wait_for_result(*sessionid, self.params.timeout)
.await?;
assert!(self.preprocess.dou_sha.clear_store(*sessionid).await);
all_double_shares.push(double_shares);
}
let mut rds_sessions: Vec<SessionId> = Vec::with_capacity(sessions.len());
for ((sessionid, batch_size), double_shares) in
sessions.iter().zip(all_double_shares.into_iter())
{
let mut shares_deg_t_by_batch = Vec::with_capacity(*batch_size);
let mut shares_deg_2t_by_batch = Vec::with_capacity(*batch_size);
for double_share_batch in double_shares.chunks_exact(self.params.n_parties) {
let (shares_deg_t, shares_deg_2t) = double_share_batch
.iter()
.cloned()
.map(|d| (d.degree_t, d.degree_2t))
.unzip();
shares_deg_t_by_batch.push(shares_deg_t);
shares_deg_2t_by_batch.push(shares_deg_2t);
}
self.preprocess
.ran_dou_sha
.init_batch(
shares_deg_t_by_batch,
shares_deg_2t_by_batch,
*sessionid,
network.clone(),
)
.await?;
rds_sessions.push(*sessionid);
}
for sessionid in &rds_sessions {
let output = self
.preprocess
.ran_dou_sha
.wait_for_result(*sessionid, self.params.timeout)
.await?;
pair.extend(output);
assert!(self.preprocess.ran_dou_sha.clear_store(*sessionid).await);
}
Ok(pair)
}
async fn ensure_random_shares_small_field<G, N>(
&mut self,
network: Arc<N>,
rng: &mut G,
needed: usize,
) -> Result<(), HoneyBadgerError>
where
N: Network + Send + Sync + 'static,
G: Rng + Send,
{
if needed == 0 {
return Ok(());
}
let output_per_column = self.params.n_parties - 2 * self.params.threshold;
let columns_needed = (needed + output_per_column - 1) / output_per_column;
let max_columns_per_run = 2048usize;
let run = (columns_needed + max_columns_per_run - 1) / max_columns_per_run;
let mut round_id = 0u8;
let mut ran_sha_counter = self.counters.ran_sha_small_field_counter.get_next().await?;
for i in 0..run {
info!("Random share generation (small field) run {}", i);
let columns_remaining = columns_needed - i * max_columns_per_run;
let batch_size = columns_remaining.min(max_columns_per_run);
let sessionid = SessionId::new(
ProtocolType::RanShaSmallField,
SessionId::pack_slot(ran_sha_counter, 0, round_id),
self.params.instance_id,
);
self.preprocess
.small_field_preproc
.share_gen
.init_batch(sessionid, batch_size, rng, network.clone())
.await?;
let output = self
.preprocess
.small_field_preproc
.share_gen
.wait_for_result(sessionid, self.params.timeout)
.await?;
self.preprocessing_material.lock().await.add(
None,
None,
None,
Some(output),
None,
None,
);
assert!(
self.preprocess
.small_field_preproc
.share_gen
.clear_store(sessionid)
.await
);
if round_id == 255 {
ran_sha_counter = self
.counters
.ran_sha_small_field_counter
.get_next()
.await
.unwrap();
round_id = 0;
} else {
round_id += 1;
}
}
self.preprocess
.small_field_preproc
.share_gen
.rbc
.clear_store()
.await;
Ok(())
}
async fn ensure_beaver_triples_small_field<G, N>(
&mut self,
network: Arc<N>,
rng: &mut G,
needed: usize,
) -> Result<(), HoneyBadgerError>
where
N: Network + Send + Sync + 'static,
G: Rng + Send,
{
let current_triples = {
let guard = self.preprocessing_material.lock().await;
guard.length().beaver_triples_small_field
};
let missing_triples = needed.saturating_sub(current_triples);
if missing_triples == 0 {
return Ok(());
}
let group_size = 2 * self.params.threshold + 1;
let total_triples_to_generate =
((missing_triples + group_size - 1) / group_size) * group_size;
let random_shares_a = self
.preprocessing_material
.lock()
.await
.take_random_shares_small_field(total_triples_to_generate)?;
let random_shares_b = self
.preprocessing_material
.lock()
.await
.take_random_shares_small_field(total_triples_to_generate)?;
let ran_dou_sha_pair = self
.ensure_ran_dou_sha_pair_small_field(network.clone(), rng, total_triples_to_generate)
.await?;
let mut triple_counter = self.counters.triple_small_field_counter.get_next().await?;
let mut round_id = 0u8;
let mut group_index = 0;
let total_groups = total_triples_to_generate / group_size;
let max_batch_groups = triple_batch_groups_limit();
while group_index < total_groups {
let batch_groups = (total_groups - group_index).min(max_batch_groups);
let share_start = group_index * group_size;
let share_end = share_start + batch_groups * group_size;
let sessionid = SessionId::new(
ProtocolType::TripleSmallField,
SessionId::pack_slot(triple_counter, 0, round_id),
self.params.instance_id,
);
self.preprocess
.small_field_preproc
.triple_gen
.init_batch(
random_shares_a[share_start..share_end].to_vec(),
random_shares_b[share_start..share_end].to_vec(),
ran_dou_sha_pair[share_start..share_end].to_vec(),
sessionid,
network.clone(),
)
.await?;
let triples = self
.preprocess
.small_field_preproc
.triple_gen
.wait_for_result(sessionid, self.params.timeout)
.await?;
self.preprocessing_material.lock().await.add(
None,
Some(triples),
None,
None,
None,
None,
);
assert!(
self.preprocess
.small_field_preproc
.triple_gen
.clear_store(sessionid)
.await
);
if round_id == 255 {
triple_counter = self
.counters
.triple_small_field_counter
.get_next()
.await
.unwrap();
round_id = 0;
} else {
round_id += 1;
}
group_index += batch_groups;
}
Ok(())
}
async fn ensure_ran_dou_sha_pair_small_field<G, N>(
&mut self,
network: Arc<N>,
rng: &mut G,
needed: usize,
) -> Result<Vec<DoubleShamirShare<GoldilocksField>>, HoneyBadgerError>
where
N: Network + Send + Sync + 'static,
G: Rng + Send,
{
let mut pair = Vec::new();
let output_per_column = self.params.threshold + 1;
let columns_needed = (needed + output_per_column - 1) / output_per_column;
let max_columns_per_run = ran_dou_sha_batch_columns_limit();
let run = (columns_needed + max_columns_per_run - 1) / max_columns_per_run;
let mut round_id = 0u8;
let mut ran_dou_sha_counter = self
.counters
.ran_dou_sha_small_field_counter
.get_next()
.await?;
for i in 0..run {
let columns_remaining = columns_needed - i * max_columns_per_run;
let batch_size = columns_remaining.min(max_columns_per_run);
let sessionid = SessionId::new(
ProtocolType::RanDouShaSmallField,
SessionId::pack_slot(ran_dou_sha_counter, 0, round_id),
self.params.instance_id,
);
let double_shares = self
.ensure_double_shares_small_field(sessionid, batch_size, network.clone(), rng)
.await?;
let mut shares_deg_t_by_batch = Vec::with_capacity(batch_size);
let mut shares_deg_2t_by_batch = Vec::with_capacity(batch_size);
for double_share_batch in double_shares.chunks_exact(self.params.n_parties) {
let (shares_deg_t, shares_deg_2t) = double_share_batch
.iter()
.cloned()
.map(|d| (d.degree_t, d.degree_2t))
.unzip();
shares_deg_t_by_batch.push(shares_deg_t);
shares_deg_2t_by_batch.push(shares_deg_2t);
}
self.preprocess
.small_field_preproc
.ran_dou_sha
.init_batch(
shares_deg_t_by_batch,
shares_deg_2t_by_batch,
sessionid,
network.clone(),
)
.await?;
let output = self
.preprocess
.small_field_preproc
.ran_dou_sha
.wait_for_result(sessionid, self.params.timeout)
.await?;
pair.extend(output);
assert!(
self.preprocess
.small_field_preproc
.ran_dou_sha
.clear_store(sessionid)
.await
);
if round_id == 255 {
ran_dou_sha_counter = self
.counters
.ran_dou_sha_small_field_counter
.get_next()
.await
.unwrap();
round_id = 0;
} else {
round_id += 1;
}
}
self.preprocess
.small_field_preproc
.ran_dou_sha
.rbc
.clear_store()
.await;
Ok(pair)
}
async fn ensure_double_shares_small_field<G, N>(
&mut self,
sessionid: SessionId,
batch_size: usize,
network: Arc<N>,
rng: &mut G,
) -> Result<Vec<DoubleShamirShare<GoldilocksField>>, HoneyBadgerError>
where
N: Network + Send + Sync + 'static,
G: Rng + Send,
{
let dou_sha_session_id = SessionId::new(
ProtocolType::DouShaSmallField,
SessionId::pack_slot(
sessionid.exec_id(),
sessionid.sub_id(),
sessionid.round_id(),
),
self.params.instance_id,
);
self.preprocess
.small_field_preproc
.dou_sha
.init_batch(dou_sha_session_id, batch_size, rng, network.clone())
.await?;
let dou_sha = self
.preprocess
.small_field_preproc
.dou_sha
.wait_for_result(dou_sha_session_id, self.params.timeout)
.await?;
assert!(
self.preprocess
.small_field_preproc
.dou_sha
.clear_store(dou_sha_session_id)
.await
);
Ok(dou_sha)
}
async fn ensure_prandbit_shares<N, G>(
&mut self,
rng: &mut G,
network: Arc<N>,
) -> Result<(), HoneyBadgerError>
where
N: Network + Send + Sync + 'static,
G: Rng + Send,
{
let no_shares = {
let store = self.preprocessing_material.lock().await;
store.length().prandbit
};
if no_shares >= self.params.n_prandbit {
info!("There are enough PRandBit shares");
return Ok(());
}
let missing = self.params.n_prandbit.saturating_sub(no_shares);
let batch = self.params.threshold + 1;
let total_randbit_to_generate = ((missing + batch - 1) / batch) * batch;
let mut randbit_output: Vec<ShamirShare<GoldilocksField, 1, Robust>> = Vec::new();
let randbit_sessionid = SessionId::new(
ProtocolType::RandBit,
SessionId::pack_slot(self.counters.rand_bit_counter.get_next().await?, 0, 0),
self.params.instance_id,
);
let prandbit_sessionid = SessionId::new(
ProtocolType::PRandBit,
SessionId::pack_slot(self.counters.prand_bit_counter.get_next().await?, 0, 0),
self.params.instance_id,
);
let current_triples = {
let guard = self.preprocessing_material.lock().await;
guard.length().beaver_triples_small_field
};
let missing_triples = total_randbit_to_generate.saturating_sub(current_triples);
let group_size = 2 * self.params.threshold + 1;
let total_triples_to_generate =
((missing_triples + group_size - 1) / group_size) * group_size;
let random_shares_for_triples = 2 * total_triples_to_generate;
self.ensure_random_shares_small_field(
network.clone(),
rng,
total_randbit_to_generate + random_shares_for_triples,
)
.await?;
let random_shares_a = self
.preprocessing_material
.lock()
.await
.take_random_shares_small_field(total_randbit_to_generate)?;
self.ensure_beaver_triples_small_field(network.clone(), rng, total_randbit_to_generate)
.await?;
let beaver_triples = self
.preprocessing_material
.lock()
.await
.take_beaver_triples_small_field(total_randbit_to_generate)?;
self.preprocess
.small_field_preproc
.rand_bit
.init(
random_shares_a,
beaver_triples,
randbit_sessionid,
self.params.timeout,
network.clone(),
)
.await?;
let output = self
.preprocess
.small_field_preproc
.rand_bit
.wait_for_result(randbit_sessionid, self.params.timeout)
.await?;
randbit_output.extend(output);
self.preprocess
.small_field_preproc
.rand_bit
.clear_store(randbit_sessionid)
.await?;
info!(id = self.id, "PRandbit share generation");
self.preprocess
.prand_bit
.generate_riss(
prandbit_sessionid,
randbit_output,
self.params.l,
self.params.k,
total_randbit_to_generate,
network,
)
.await?;
let output = self
.preprocess
.prand_bit
.wait_for_bit_result(prandbit_sessionid, self.params.timeout)
.await?;
self.preprocessing_material
.lock()
.await
.add(None, None, None, None, Some(output), None);
self.preprocess
.prand_bit
.clear_store(prandbit_sessionid)
.await?;
Ok(())
}
async fn ensure_prandint_shares<N>(&mut self, network: Arc<N>) -> Result<(), HoneyBadgerError>
where
N: Network + Send + Sync + 'static,
{
let no_shares = {
let store = self.preprocessing_material.lock().await;
store.length().prandint
};
if no_shares >= self.params.n_prandint {
info!("There are enough prandbit shares");
return Ok(());
}
let missing = self.params.n_prandint.saturating_sub(no_shares);
info!("PRandInt share generation");
let max_prandint_batch = 64 * (self.params.threshold + 1);
let mut prandint_output = Vec::with_capacity(missing);
for batch_size in chunk_sizes(missing, max_prandint_batch) {
let sessionid = SessionId::new(
ProtocolType::PRandInt,
SessionId::pack_slot(self.counters.prand_int_counter.get_next().await?, 0, 0),
self.params.instance_id,
);
self.preprocess
.prand_bit
.generate_riss(
sessionid,
vec![],
self.params.l,
self.params.k,
batch_size,
network.clone(),
)
.await?;
let output = self
.preprocess
.prand_bit
.wait_for_int_result(sessionid, self.params.timeout)
.await?;
prandint_output.extend(output);
self.preprocess.prand_bit.clear_store(sessionid).await?;
}
self.preprocessing_material.lock().await.add(
None,
None,
None,
None,
None,
Some(prandint_output),
);
Ok(())
}
}
fn chunk_sizes(total: usize, max_chunk_size: usize) -> impl Iterator<Item = usize> {
let max_chunk_size = max_chunk_size.max(1);
(0..total)
.step_by(max_chunk_size)
.map(move |start| (total - start).min(max_chunk_size))
}
pub(crate) fn max_mul_pairs_per_session(threshold: usize) -> usize {
128 * threshold.saturating_add(1)
}
#[derive(Serialize, Deserialize, Debug)]
pub enum WrappedMessage {
RanDouSha(RanDouShaMessage),
Rbc(Msg<SessionId>),
BatchRecon(BatchReconMsg),
Input(InputMessage),
RanSha(RanShaMessage),
Dousha(DouShaMessage),
Output(OutputMessage),
PRandBitD(PRandBitDMessage),
}
impl WrappedMessage {
pub fn rbc_wrap(msg: Msg<SessionId>) -> Result<Vec<u8>, RbcError> {
let wrapped = WrappedMessage::Rbc(msg);
Ok(bincode::serialize(&wrapped)?)
}
}
#[repr(u8)]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub enum ProtocolType {
None = 0,
Randousha = 1,
Ransha = 2,
Input = 3,
Rbc = 4,
Triple = 5,
BatchRecon = 6,
Dousha = 7,
Mul = 8,
PRandInt = 9,
PRandBit = 10,
RandBit = 11,
FpMul = 12,
Trunc = 13,
FpDivConst = 14,
TripleSmallField = 15,
RanShaSmallField = 16,
RanDouShaSmallField = 17,
DouShaSmallField = 18,
}
impl ProtocolTag for ProtocolType {
#[inline]
fn to_u8(self) -> u8 {
self as u8
}
#[inline]
fn from_u8(v: u8) -> Option<Self> {
match v {
0 => Some(Self::None),
1 => Some(Self::Randousha),
2 => Some(Self::Ransha),
3 => Some(Self::Input),
4 => Some(Self::Rbc),
5 => Some(Self::Triple),
6 => Some(Self::BatchRecon),
7 => Some(Self::Dousha),
8 => Some(Self::Mul),
9 => Some(Self::PRandInt),
10 => Some(Self::PRandBit),
11 => Some(Self::RandBit),
12 => Some(Self::FpMul),
13 => Some(Self::Trunc),
14 => Some(Self::FpDivConst),
15 => Some(Self::TripleSmallField),
16 => Some(Self::RanShaSmallField),
17 => Some(Self::RanDouShaSmallField),
18 => Some(Self::DouShaSmallField),
_ => None,
}
}
}
#[derive(PartialOrd, Ord, Clone, Serialize, Deserialize, Copy, PartialEq, Eq, Hash)]
pub struct SessionId(u128);
impl fmt::Debug for SessionId {
fn fmt(&self, f: &mut fmt::Formatter) -> Result<(), fmt::Error> {
let caller = ((self.0 >> 112) & 0xFF) as u8;
let exec_id = self.exec_id();
let sub_id = self.sub_id();
let round_id = self.round_id();
let instance_id = self.instance_id();
write!(
f,
"[caller={},exec_id={},sub_id={},round_id={},instance_id={}]",
caller, exec_id, sub_id, round_id, instance_id
)
}
}
impl ProtocolSessionId for SessionId {
type Protocol = ProtocolType;
fn new(protocol: ProtocolType, slot: u128, instance_id: u32) -> Self {
let slot_mask: u128 = (1u128 << 80) - 1; let value = (((protocol as u128) & 0xFF) << 112)
| (((slot & slot_mask) as u128) << 32)
| (instance_id as u128);
SessionId(value)
}
fn calling_protocol(self) -> Option<ProtocolType> {
let val = ((self.0 >> 112) & 0xFF) as u8;
ProtocolType::from_u8(val)
}
fn slot(self) -> u128 {
(self.0 >> 32) & ((1u128 << 80) - 1)
}
fn instance_id(self) -> u32 {
self.0 as u32
}
fn as_u128(self) -> u128 {
self.0
}
unsafe fn from_u128(id: u128) -> Self {
SessionId(id)
}
}
impl SessionId {
pub fn exec_id(self) -> u64 {
(self.0 >> 48) as u64
}
pub fn sub_id(self) -> u8 {
((self.0 >> 40) & 0xFF) as u8
}
pub fn round_id(self) -> u8 {
((self.0 >> 32) & 0xFF) as u8
}
#[inline]
pub fn pack_slot(exec_id: u64, sub_id: u8, round_id: u8) -> u128 {
((exec_id as u128) << 16) | ((sub_id as u128) << 8) | (round_id as u128)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use tokio::sync::Mutex;
#[test]
fn test_session_id_debug_format() {
let caller = ProtocolType::from_u8(5u8).unwrap();
let exec_id = 42u64;
let sub_id = 7u8;
let round_id = 3u8;
let instance_id = 0xDEADBEEF;
let session_id = SessionId::new(
caller,
SessionId::pack_slot(exec_id, sub_id, round_id),
instance_id,
);
let debug_str = format!("{:?}", session_id);
assert_eq!(
debug_str,
"[caller=5,exec_id=42,sub_id=7,round_id=3,instance_id=3735928559]"
);
}
#[test]
fn test_session_id() {
let caller = ProtocolType::Triple;
let exec_id = 42u64;
let sub_id = 7u8;
let round_id = 3u8;
let instance_id = 0xDEADBEEF;
let session_id = SessionId::new(
caller,
SessionId::pack_slot(exec_id, sub_id, round_id),
instance_id,
);
assert_eq!(session_id.calling_protocol().unwrap(), caller);
assert_eq!(session_id.exec_id(), exec_id);
assert_eq!(session_id.sub_id(), sub_id);
assert_eq!(session_id.round_id(), round_id);
assert_eq!(session_id.instance_id(), instance_id);
let session_id2 = SessionId::new(
session_id.calling_protocol().unwrap(),
SessionId::pack_slot(
session_id.exec_id(),
session_id.sub_id(),
session_id.round_id(),
),
session_id.instance_id(),
);
assert_eq!(session_id, session_id2);
}
#[tokio::test]
async fn test_subprotocol_counter_limit_error() {
let counter = SubProtocolCounter(Arc::new(Mutex::new(Some(u64::MAX))));
let val = counter.get_next().await;
assert_eq!(val.unwrap(), u64::MAX);
let err = counter.get_next().await;
assert!(matches!(err, Err(HoneyBadgerError::LimitError)));
}
#[test]
fn test_max_mul_pairs_per_session_tracks_child_session_space() {
assert_eq!(max_mul_pairs_per_session(0), 128);
assert_eq!(max_mul_pairs_per_session(1), 256);
assert_eq!(max_mul_pairs_per_session(2), 384);
assert_eq!(max_mul_pairs_per_session(3), 512);
}
}