use crate::error::{OptimError, Result};
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use std::collections::BTreeMap;
use x25519_dalek::{PublicKey, StaticSecret};
use scirs2_core::random::thread_rng;
const SEED_DOMAIN: &[u8] = b"OPTIRS-FED-SECAGG-PAIRWISE-SEED-v1";
const PRG_DOMAIN: &[u8] = b"OPTIRS-FED-SECAGG-MASK-PRG-v1";
const WORDS_PER_BLOCK: usize = 4;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
pub struct ClientPublicKey([u8; 32]);
impl ClientPublicKey {
pub fn from_bytes(bytes: [u8; 32]) -> Self {
Self(bytes)
}
pub fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
}
pub struct ClientKeyPair {
secret: StaticSecret,
public: ClientPublicKey,
}
impl std::fmt::Debug for ClientKeyPair {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ClientKeyPair")
.field("public", &self.public)
.field("secret", &"<redacted>")
.finish()
}
}
impl ClientKeyPair {
pub fn generate() -> Self {
let mut bytes = [0_u8; 32];
thread_rng().fill(&mut bytes[..]);
Self::from_secret_bytes(bytes)
}
pub fn from_secret_bytes(bytes: [u8; 32]) -> Self {
let secret = StaticSecret::from(bytes);
let public = ClientPublicKey(PublicKey::from(&secret).to_bytes());
Self { secret, public }
}
pub fn public_key(&self) -> ClientPublicKey {
self.public
}
pub fn shared_seed_with(&self, peer: &ClientPublicKey, round_seed: u64) -> Result<[u8; 32]> {
if *peer == self.public {
return Err(OptimError::InvalidParameter(
"a client cannot derive a pairwise mask with its own public key".to_string(),
));
}
let shared = self.secret.diffie_hellman(&PublicKey::from(peer.0));
let shared_bytes = shared.as_bytes();
if shared_bytes.iter().all(|&byte| byte == 0) {
return Err(OptimError::InvalidParameter(
"X25519 key agreement produced the all-zero shared secret; the peer supplied a \
low-order public key and the resulting mask would be predictable"
.to_string(),
));
}
let (low, high) = if self.public.0 <= peer.0 {
(&self.public.0, &peer.0)
} else {
(&peer.0, &self.public.0)
};
let mut hasher = Sha256::new();
hasher.update(SEED_DOMAIN);
hasher.update(round_seed.to_le_bytes());
hasher.update(low);
hasher.update(high);
hasher.update(shared_bytes);
let digest = hasher.finalize();
let mut seed = [0_u8; 32];
seed.copy_from_slice(&digest);
Ok(seed)
}
}
pub fn expand_mask(seed: &[u8; 32], dim: usize, modulus: i64) -> Result<Vec<i64>> {
if modulus <= 1 {
return Err(OptimError::InvalidParameter(format!(
"mask modulus must be greater than 1, got {modulus}"
)));
}
if dim == 0 {
return Ok(Vec::new());
}
let modulus_u = modulus as u64;
let acceptance_bound = (u64::MAX / modulus_u) * modulus_u;
let mut mask = Vec::with_capacity(dim);
let mut counter = 0_u64;
while mask.len() < dim {
let mut hasher = Sha256::new();
hasher.update(PRG_DOMAIN);
hasher.update(seed);
hasher.update(counter.to_le_bytes());
let block = hasher.finalize();
counter = counter.wrapping_add(1);
for word_index in 0..WORDS_PER_BLOCK {
if mask.len() == dim {
break;
}
let start = word_index * 8;
let mut word_bytes = [0_u8; 8];
word_bytes.copy_from_slice(&block[start..start + 8]);
let word = u64::from_le_bytes(word_bytes);
if word < acceptance_bound {
mask.push((word % modulus_u) as i64);
}
}
}
Ok(mask)
}
pub fn signed_pairwise_mask(
own_id: &str,
own_keys: &ClientKeyPair,
peer_id: &str,
peer_key: &ClientPublicKey,
round_seed: u64,
dim: usize,
modulus: i64,
) -> Result<Vec<i64>> {
if own_id == peer_id {
return Err(OptimError::InvalidParameter(format!(
"client {own_id} cannot hold a pairwise mask with itself"
)));
}
let seed = own_keys.shared_seed_with(peer_key, round_seed)?;
let mask = expand_mask(&seed, dim, modulus)?;
let positive = own_id < peer_id;
Ok(mask
.into_iter()
.map(|value| {
if positive {
value.rem_euclid(modulus)
} else {
(-value).rem_euclid(modulus)
}
})
.collect())
}
pub fn compute_client_mask(
own_id: &str,
own_keys: &ClientKeyPair,
peers: &BTreeMap<String, ClientPublicKey>,
round_seed: u64,
dim: usize,
modulus: i64,
) -> Result<Vec<i64>> {
if modulus <= 1 {
return Err(OptimError::InvalidParameter(format!(
"mask modulus must be greater than 1, got {modulus}"
)));
}
match peers.get(own_id) {
None => {
return Err(OptimError::InvalidParameter(format!(
"client {own_id} is not in the round's public-key directory"
)));
}
Some(published) if *published != own_keys.public_key() => {
return Err(OptimError::InvalidParameter(format!(
"the public key published for client {own_id} does not match the supplied key \
pair; the derived masks would not cancel"
)));
}
Some(_) => {}
}
if peers.len() < 2 {
return Err(OptimError::InvalidConfig(format!(
"client {own_id} has no peers in this round; a single-client cohort cannot be \
masked, so the upload would be the raw update"
)));
}
let mut total = vec![0_i64; dim];
for (peer_id, peer_key) in peers.iter() {
if peer_id == own_id {
continue;
}
let signed = signed_pairwise_mask(
own_id, own_keys, peer_id, peer_key, round_seed, dim, modulus,
)?;
for (accumulator, value) in total.iter_mut().zip(signed.iter()) {
*accumulator = (*accumulator + *value).rem_euclid(modulus);
}
}
Ok(total)
}
pub fn fresh_round_seed() -> u64 {
thread_rng().random::<u64>()
}
#[cfg(test)]
mod tests {
use super::*;
fn keys(tag: u8) -> ClientKeyPair {
let mut bytes = [0_u8; 32];
bytes[0] = tag;
bytes[31] = tag.wrapping_mul(7).wrapping_add(1);
ClientKeyPair::from_secret_bytes(bytes)
}
const MODULUS: i64 = 1 << 31;
#[test]
fn key_agreement_is_symmetric() {
let alice = keys(1);
let bob = keys(2);
let seed_ab = alice
.shared_seed_with(&bob.public_key(), 7)
.expect("alice -> bob");
let seed_ba = bob
.shared_seed_with(&alice.public_key(), 7)
.expect("bob -> alice");
assert_eq!(seed_ab, seed_ba);
}
#[test]
fn the_shared_seed_requires_a_secret_key_so_the_server_cannot_derive_it() {
let alice = keys(1);
let bob = keys(2);
let server = keys(3);
let truth = alice
.shared_seed_with(&bob.public_key(), 7)
.expect("alice -> bob");
let server_attempt = server
.shared_seed_with(&bob.public_key(), 7)
.expect("server -> bob");
assert_ne!(
truth, server_attempt,
"the pairwise seed must depend on a client secret, not only on public data"
);
let real = expand_mask(&truth, 64, MODULUS).expect("real mask");
let forged = expand_mask(&server_attempt, 64, MODULUS).expect("forged mask");
let matches = real
.iter()
.zip(forged.iter())
.filter(|(a, b)| a == b)
.count();
assert!(
matches < 4,
"a mask derived without the secret should not coincide with the real one \
({matches}/64 coordinates matched)"
);
}
#[test]
fn seeds_differ_across_rounds_and_across_pairs() {
let alice = keys(1);
let bob = keys(2);
let carol = keys(3);
let round_one = alice
.shared_seed_with(&bob.public_key(), 1)
.expect("round 1");
let round_two = alice
.shared_seed_with(&bob.public_key(), 2)
.expect("round 2");
assert_ne!(round_one, round_two);
let with_carol = alice
.shared_seed_with(&carol.public_key(), 1)
.expect("alice -> carol");
assert_ne!(round_one, with_carol);
}
#[test]
fn a_client_cannot_pair_with_itself() {
let alice = keys(1);
let err = alice
.shared_seed_with(&alice.public_key(), 1)
.expect_err("self pairing must fail");
assert!(format!("{err}").contains("own public key"));
}
#[test]
fn low_order_public_keys_are_rejected() {
let alice = keys(1);
let malicious = ClientPublicKey::from_bytes([0_u8; 32]);
let err = alice
.shared_seed_with(&malicious, 1)
.expect_err("low-order key must be rejected");
assert!(format!("{err}").contains("all-zero shared secret"));
}
#[test]
fn expand_mask_is_deterministic_and_in_range() {
let seed = [0x5A_u8; 32];
let first = expand_mask(&seed, 1000, MODULUS).expect("mask");
let second = expand_mask(&seed, 1000, MODULUS).expect("mask");
assert_eq!(first, second);
assert_eq!(first.len(), 1000);
assert!(first.iter().all(|&value| (0..MODULUS).contains(&value)));
assert!(expand_mask(&seed, 0, MODULUS).expect("empty").is_empty());
assert!(expand_mask(&seed, 4, 1).is_err());
}
#[test]
fn expand_mask_output_covers_the_whole_group() {
let seed = [0x11_u8; 32];
let mask = expand_mask(&seed, 4096, MODULUS).expect("mask");
let minimum = mask.iter().copied().min().expect("non-empty");
let maximum = mask.iter().copied().max().expect("non-empty");
assert!(
minimum < MODULUS / 100,
"minimum {minimum} is not near zero"
);
assert!(
maximum > MODULUS - MODULUS / 100,
"maximum {maximum} is not near the modulus"
);
let mut buckets = [0_usize; 8];
for &value in mask.iter() {
let bucket = ((value as i128 * 8) / MODULUS as i128) as usize;
buckets[bucket.min(7)] += 1;
}
for (index, &count) in buckets.iter().enumerate() {
assert!(
count > 4096 / 16 && count < 4096 / 4,
"bucket {index} holds {count} of 4096 samples, which is not roughly uniform"
);
}
}
#[test]
fn expand_mask_changes_with_the_seed() {
let a = expand_mask(&[1_u8; 32], 64, MODULUS).expect("mask");
let b = expand_mask(&[2_u8; 32], 64, MODULUS).expect("mask");
assert_ne!(a, b);
}
#[test]
fn signed_masks_of_a_pair_are_additive_inverses() {
let alice = keys(1);
let bob = keys(2);
let from_alice =
signed_pairwise_mask("alice", &alice, "bob", &bob.public_key(), 42, 32, MODULUS)
.expect("alice mask");
let from_bob =
signed_pairwise_mask("bob", &bob, "alice", &alice.public_key(), 42, 32, MODULUS)
.expect("bob mask");
assert_eq!(from_alice.len(), 32);
for (a, b) in from_alice.iter().zip(from_bob.iter()) {
assert_eq!(
(a + b).rem_euclid(MODULUS),
0,
"pairwise masks must cancel: {a} + {b} != 0 mod {MODULUS}"
);
}
}
#[test]
fn every_clients_mask_sums_to_zero_over_the_cohort() {
let dim = 128;
let round_seed = 0xDEAD_BEEF;
let pairs: Vec<(String, ClientKeyPair)> = (1..=6_u8)
.map(|tag| (format!("client{tag:02}"), keys(tag)))
.collect();
let directory: BTreeMap<String, ClientPublicKey> = pairs
.iter()
.map(|(id, keys)| (id.clone(), keys.public_key()))
.collect();
let mut total = vec![0_i64; dim];
for (id, key_pair) in pairs.iter() {
let mask = compute_client_mask(id, key_pair, &directory, round_seed, dim, MODULUS)
.expect("client mask");
assert_eq!(mask.len(), dim);
for (accumulator, value) in total.iter_mut().zip(mask.iter()) {
*accumulator = (*accumulator + *value).rem_euclid(MODULUS);
}
}
assert!(
total.iter().all(|&value| value == 0),
"cohort masks did not telescope to zero"
);
}
#[test]
fn an_individual_mask_is_not_trivial() {
let dim = 256;
let pairs: Vec<(String, ClientKeyPair)> = (1..=4_u8)
.map(|tag| (format!("client{tag}"), keys(tag)))
.collect();
let directory: BTreeMap<String, ClientPublicKey> = pairs
.iter()
.map(|(id, keys)| (id.clone(), keys.public_key()))
.collect();
let (id, key_pair) = &pairs[0];
let mask = compute_client_mask(id, key_pair, &directory, 1, dim, MODULUS).expect("mask");
let zeros = mask.iter().filter(|&&value| value == 0).count();
assert!(zeros < 4, "{zeros} of {dim} mask coordinates were zero");
}
#[test]
fn compute_client_mask_validates_the_directory() {
let alice = keys(1);
let bob = keys(2);
let mut directory = BTreeMap::new();
directory.insert("bob".to_string(), bob.public_key());
let err = compute_client_mask("alice", &alice, &directory, 1, 8, MODULUS)
.expect_err("missing from directory");
assert!(format!("{err}").contains("not in the round's public-key directory"));
directory.insert("alice".to_string(), keys(9).public_key());
let err = compute_client_mask("alice", &alice, &directory, 1, 8, MODULUS)
.expect_err("key mismatch");
assert!(format!("{err}").contains("does not match the supplied key pair"));
let mut solo = BTreeMap::new();
solo.insert("alice".to_string(), alice.public_key());
let err =
compute_client_mask("alice", &alice, &solo, 1, 8, MODULUS).expect_err("solo cohort");
assert!(format!("{err}").contains("no peers"));
}
#[test]
fn generated_key_pairs_are_distinct() {
let first = ClientKeyPair::generate();
let second = ClientKeyPair::generate();
assert_ne!(first.public_key(), second.public_key());
assert!(format!("{first:?}").contains("<redacted>"));
}
#[test]
fn fresh_round_seeds_are_not_a_counter() {
let seeds: Vec<u64> = (0..8).map(|_| fresh_round_seed()).collect();
let distinct: std::collections::HashSet<u64> = seeds.iter().copied().collect();
assert_eq!(distinct.len(), seeds.len());
assert!(seeds.iter().any(|&seed| seed > u64::MAX / 1024));
}
}