use crate::messaging::system::{KeyedSig, SigShare};
use dashmap::DashMap;
use std::{
collections::BTreeMap,
sync::Arc,
time::{Duration, Instant},
};
use thiserror::Error;
use tiny_keccak::{Hasher, Sha3};
use tokio::sync::RwLock;
const DEFAULT_EXPIRATION: Duration = Duration::from_secs(120);
type Digest256 = [u8; 32];
#[derive(Debug, Clone)]
pub struct SignatureAggregator {
map: Arc<DashMap<Digest256, State>>,
expiration: Duration,
}
impl SignatureAggregator {
pub(crate) fn new() -> Self {
Self::with_expiration(DEFAULT_EXPIRATION)
}
pub(crate) fn with_expiration(expiration: Duration) -> Self {
Self {
map: Default::default(),
expiration,
}
}
pub(crate) async fn add(&self, payload: &[u8], sig_share: SigShare) -> Result<KeyedSig, Error> {
self.remove_expired().await;
if !sig_share.verify(payload) {
return Err(Error::InvalidShare);
}
let public_key = sig_share.public_key_set.public_key();
let mut hasher = Sha3::v256();
let mut hash = Digest256::default();
hasher.update(payload);
hasher.update(&public_key.to_bytes());
hasher.finalize(&mut hash);
let mut entry = self.map.entry(hash).or_insert_with(State::new);
entry.add(sig_share).await.map(|signature| KeyedSig {
public_key,
signature,
})
}
async fn remove_expired(&self) {
let expiration = self.expiration;
let mut to_remove = vec![];
for ref_multi in self.map.iter() {
let (digest, state) = ref_multi.pair();
if state.modified.read().await.elapsed() >= expiration {
to_remove.push(*digest);
}
}
self.map.retain(|digest, _| !to_remove.contains(digest))
}
}
impl Default for SignatureAggregator {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Error)]
pub enum Error {
#[error("not enough signature shares")]
NotEnoughShares,
#[error("signature share is invalid")]
InvalidShare,
#[error("failed to combine signature shares: {0}")]
Combine(#[from] bls::error::Error),
}
#[derive(Debug, Clone)]
struct State {
shares: Arc<DashMap<usize, bls::SignatureShare>>,
modified: Arc<RwLock<Instant>>,
}
impl State {
fn new() -> Self {
Self {
shares: Default::default(),
modified: Arc::new(RwLock::new(Instant::now())),
}
}
async fn add(&mut self, sig_share: SigShare) -> Result<bls::Signature, Error> {
if self
.shares
.insert(sig_share.index, sig_share.signature_share)
.is_none()
{
*self.modified.write().await = Instant::now();
} else {
return Err(Error::NotEnoughShares);
}
if self.shares.len() > sig_share.public_key_set.threshold() {
let mut shares_map = BTreeMap::default();
for ref_multi in self.shares.iter() {
let (index, share) = ref_multi.pair();
let _old_share = shares_map.insert(*index, share.clone());
}
let signature = sig_share
.public_key_set
.combine_signatures(&shares_map)
.map_err(Error::Combine)?;
self.shares.clear();
Ok(signature)
} else {
Err(Error::NotEnoughShares)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use rand::thread_rng;
use std::thread::sleep;
#[tokio::test(flavor = "multi_thread")]
async fn smoke() -> Result<(), Error> {
let mut rng = thread_rng();
let threshold = 3;
let sk_set = bls::SecretKeySet::random(threshold, &mut rng);
let aggregator = SignatureAggregator::default();
let payload = b"hello";
for index in 0..threshold {
let sig_share = create_sig_share(&sk_set, index, payload);
let result = aggregator.add(payload, sig_share).await;
match result {
Err(Error::NotEnoughShares) => (),
_ => panic!("unexpected result: {:?}", result),
}
}
let sig_share = create_sig_share(&sk_set, threshold, payload);
let sig = aggregator.add(payload, sig_share).await?;
assert!(sig.verify(payload));
let sig_share = create_sig_share(&sk_set, threshold + 1, payload);
let result = aggregator.add(payload, sig_share).await;
match result {
Err(Error::NotEnoughShares) => Ok(()),
_ => panic!("unexpected result: {:?}", result),
}
}
#[tokio::test(flavor = "multi_thread")]
async fn invalid_share() -> Result<(), Error> {
let mut rng = thread_rng();
let threshold = 3;
let sk_set = bls::SecretKeySet::random(threshold, &mut rng);
let aggregator = SignatureAggregator::new();
let payload = b"good";
for index in 0..threshold {
let sig_share = create_sig_share(&sk_set, index, payload);
let _keyed_sig = aggregator.add(payload, sig_share).await;
}
let invalid_sig_share = create_sig_share(&sk_set, threshold, b"bad");
let result = aggregator.add(payload, invalid_sig_share).await;
match result {
Err(Error::InvalidShare) => (),
_ => panic!("unexpected result: {:?}", result),
}
let sig_share = create_sig_share(&sk_set, threshold + 1, payload);
let sig = aggregator.add(payload, sig_share).await?;
assert!(sig.verify(payload));
Ok(())
}
#[tokio::test(flavor = "multi_thread")]
async fn expiration() {
let mut rng = thread_rng();
let threshold = 3;
let sk_set = bls::SecretKeySet::random(threshold, &mut rng);
let aggregator = SignatureAggregator::with_expiration(Duration::from_millis(500));
let payload = b"hello";
for index in 0..threshold {
let sig_share = create_sig_share(&sk_set, index, payload);
let _keyed_sig = aggregator.add(payload, sig_share).await;
}
sleep(Duration::from_secs(1));
let sig_share = create_sig_share(&sk_set, threshold, payload);
let result = aggregator.add(payload, sig_share).await;
match result {
Err(Error::NotEnoughShares) => (),
_ => panic!("unexpected result: {:?}", result),
}
}
#[tokio::test(flavor = "multi_thread")]
async fn repeated_voting() {
let mut rng = thread_rng();
let threshold = 3;
let sk_set = bls::SecretKeySet::random(threshold, &mut rng);
let aggregator = SignatureAggregator::new();
let payload = b"hello";
for index in 0..threshold {
let sig_share = create_sig_share(&sk_set, index, payload);
assert!(aggregator.add(payload, sig_share).await.is_err());
}
let sig_share = create_sig_share(&sk_set, threshold, payload);
assert!(aggregator.add(payload, sig_share).await.is_ok());
let offset = 2;
for index in offset..(threshold + offset) {
let sig_share = create_sig_share(&sk_set, index, payload);
assert!(aggregator.add(payload, sig_share).await.is_err());
}
let sig_share = create_sig_share(&sk_set, threshold + offset + 1, payload);
assert!(aggregator.add(payload, sig_share).await.is_ok());
}
fn create_sig_share(sk_set: &bls::SecretKeySet, index: usize, payload: &[u8]) -> SigShare {
let sk_share = sk_set.secret_key_share(index);
SigShare::new(sk_set.public_keys(), index, &sk_share, payload)
}
}