use getrandom::SysRng;
use p256::elliptic_curve::rand_core::Rng;
use p256::elliptic_curve::rand_core::UnwrapErr;
use std::collections::HashMap;
#[derive(Debug, Clone)]
pub struct AggregationParty {
pub id: u32,
pub value: i64,
pub masks_to_apply: HashMap<u32, i64>,
}
#[derive(Debug)]
pub struct SecureAggregation {
pub parties: Vec<AggregationParty>,
}
impl SecureAggregation {
pub fn new(party_count: u32) -> Self {
let parties = (0..party_count)
.map(|id| AggregationParty {
id,
value: 0,
masks_to_apply: HashMap::new(),
})
.collect();
Self { parties }
}
pub fn set_value(&mut self, party_id: u32, value: i64) {
if let Some(party) = self.parties.iter_mut().find(|p| p.id == party_id) {
party.value = value;
}
}
pub fn establish_masks(&mut self) {
let n = self.parties.len();
let mut rng = UnwrapErr(SysRng);
for i in 0..n {
for j in (i + 1)..n {
let mask = (rng.next_u32() as i64) % 1000;
let id_i = self.parties[i].id;
let id_j = self.parties[j].id;
self.parties[i].masks_to_apply.insert(id_j, mask);
self.parties[j].masks_to_apply.insert(id_i, -mask);
}
}
}
pub fn masked_values(&self) -> Vec<i64> {
self.parties
.iter()
.map(|p| {
let mask_sum: i64 = p.masks_to_apply.values().sum();
p.value + mask_sum
})
.collect()
}
pub fn aggregate(&self) -> i64 {
self.masked_values().iter().sum()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sum_preserved_with_masks() {
let mut agg = SecureAggregation::new(3);
agg.set_value(0, 10);
agg.set_value(1, 20);
agg.set_value(2, 30);
agg.establish_masks();
assert_eq!(agg.aggregate(), 60);
}
#[test]
fn single_party() {
let mut agg = SecureAggregation::new(1);
agg.set_value(0, 42);
agg.establish_masks();
assert_eq!(agg.aggregate(), 42);
}
#[test]
fn negative_values() {
let mut agg = SecureAggregation::new(3);
agg.set_value(0, -10);
agg.set_value(1, 20);
agg.set_value(2, -5);
agg.establish_masks();
assert_eq!(agg.aggregate(), 5);
}
#[test]
fn many_parties() {
let n = 10;
let mut agg = SecureAggregation::new(n);
let total: i64 = (1..=n as i64).sum();
for i in 0..n {
agg.set_value(i, (i + 1) as i64);
}
agg.establish_masks();
assert_eq!(agg.aggregate(), total);
}
#[test]
fn masked_values_hide_individuals() {
let mut agg = SecureAggregation::new(2);
agg.set_value(0, 100);
agg.set_value(1, 200);
agg.establish_masks();
let masked = agg.masked_values();
assert_ne!(masked[0], 100);
assert_ne!(masked[1], 200);
}
#[test]
fn zero_values() {
let mut agg = SecureAggregation::new(3);
agg.establish_masks();
assert_eq!(agg.aggregate(), 0);
}
#[test]
fn masks_symmetric() {
let mut agg = SecureAggregation::new(2);
agg.establish_masks();
let m01 = agg.parties[0].masks_to_apply.get(&1).copied().unwrap_or(0);
let m10 = agg.parties[1].masks_to_apply.get(&0).copied().unwrap_or(0);
assert_eq!(m01, -m10);
}
}