use std::collections::{BTreeMap, HashMap};
use std::sync::{Arc, Mutex};
use serde_json::Value;
use zeroize::Zeroizing;
use crate::crypto::{
derive_key, encryption_audience, CryptoError, MultiEpochOrgCipher, PayloadCipher,
SyncKeyProvider,
};
use crate::oplog::Scope;
struct DenyCipher {
reason: String,
}
impl PayloadCipher for DenyCipher {
fn encrypt(&self, _plaintext: &Value) -> Result<Value, CryptoError> {
Err(CryptoError::Key(self.reason.clone()))
}
fn decrypt(&self, _envelope: &Value) -> Result<Value, CryptoError> {
Err(CryptoError::Key(self.reason.clone()))
}
}
pub struct OrgAwareKeyProvider {
personal: Arc<dyn SyncKeyProvider>,
org_roots: HashMap<String, BTreeMap<u64, Zeroizing<[u8; 32]>>>,
ciphers: Mutex<HashMap<String, Arc<dyn PayloadCipher>>>,
}
impl std::fmt::Debug for OrgAwareKeyProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OrgAwareKeyProvider")
.field("orgs", &self.org_roots.keys().collect::<Vec<_>>())
.finish_non_exhaustive()
}
}
impl OrgAwareKeyProvider {
pub fn new(
personal: Arc<dyn SyncKeyProvider>,
org_roots: HashMap<String, BTreeMap<u64, Zeroizing<[u8; 32]>>>,
) -> Self {
Self {
personal,
org_roots,
ciphers: Mutex::new(HashMap::new()),
}
}
fn org_cipher(&self, org: &str, audience: &str) -> Arc<dyn PayloadCipher> {
let mut cache = self.ciphers.lock().expect("org cipher cache poisoned");
if let Some(c) = cache.get(audience) {
return c.clone();
}
let cipher: Arc<dyn PayloadCipher> = match self.org_roots.get(org) {
Some(epochs) if !epochs.is_empty() => {
let keys: BTreeMap<u64, Zeroizing<[u8; 32]>> = epochs
.iter()
.map(|(&epoch, k_org)| (epoch, Zeroizing::new(derive_key(&k_org[..], audience))))
.collect();
Arc::new(MultiEpochOrgCipher::new(audience, keys))
}
_ => Arc::new(DenyCipher {
reason: format!(
"no org key for {audience}: member not granted K_org, or org-scope not activated"
),
}),
};
cache.insert(audience.to_string(), cipher.clone());
cipher
}
}
impl SyncKeyProvider for OrgAwareKeyProvider {
fn cipher_for(&self, scope: &Scope) -> Arc<dyn PayloadCipher> {
match scope {
Scope::Personal => self.personal.cipher_for(scope),
Scope::Shared { org } => self.org_cipher(org, &encryption_audience(scope)),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::crypto::DerivedKeyProvider;
fn personal(user: &str) -> Arc<dyn SyncKeyProvider> {
Arc::new(DerivedKeyProvider::from_login_secret(b"pass", user))
}
fn provider(user: &str, roots: &[(&str, [u8; 32])]) -> OrgAwareKeyProvider {
OrgAwareKeyProvider::new(
personal(user),
roots
.iter()
.map(|(o, k)| {
let m = BTreeMap::from([(1u64, Zeroizing::new(*k))]);
(o.to_string(), m)
})
.collect(),
)
}
fn provider_multi(user: &str, roots: &[(&str, &[(u64, [u8; 32])])]) -> OrgAwareKeyProvider {
OrgAwareKeyProvider::new(
personal(user),
roots
.iter()
.map(|(o, epochs)| {
let m: BTreeMap<u64, Zeroizing<[u8; 32]>> = epochs
.iter()
.map(|(e, k)| (*e, Zeroizing::new(*k)))
.collect();
(o.to_string(), m)
})
.collect(),
)
}
#[test]
fn personal_scope_round_trips_via_the_inner_provider() {
let p = provider("alice", &[]);
let msg = serde_json::json!({"note": "personal"});
let ct = p.cipher_for(&Scope::Personal).encrypt(&msg).unwrap();
let pt = p.cipher_for(&Scope::Personal).decrypt(&ct).unwrap();
assert_eq!(pt, msg);
}
#[test]
fn shared_org_ops_are_mutually_readable_across_members() {
let k_org = [9u8; 32];
let alice = provider("alice", &[("acme", k_org)]);
let bob = provider("bob", &[("acme", k_org)]);
let scope = Scope::Shared { org: "acme".into() };
let msg = serde_json::json!({"shared": "brain"});
let ct = alice.cipher_for(&scope).encrypt(&msg).unwrap();
assert_eq!(bob.cipher_for(&scope).decrypt(&ct).unwrap(), msg);
}
#[test]
fn shared_org_without_key_fails_closed_both_ways() {
let p = provider("alice", &[]);
let scope = Scope::Shared { org: "acme".into() };
assert!(matches!(
p.cipher_for(&scope).encrypt(&serde_json::json!({"x": 1})),
Err(CryptoError::Key(_))
));
assert!(matches!(
p.cipher_for(&scope)
.decrypt(&serde_json::json!({"car_enc": "x"})),
Err(CryptoError::Key(_))
));
}
#[test]
fn distinct_orgs_derive_independent_keys() {
let acme = provider("alice", &[("acme", [1u8; 32])]);
let globex = provider("alice", &[("globex", [2u8; 32])]);
let msg = serde_json::json!({"secret": "acme-only"});
let ct = acme
.cipher_for(&Scope::Shared { org: "acme".into() })
.encrypt(&msg)
.unwrap();
assert!(matches!(
globex
.cipher_for(&Scope::Shared {
org: "globex".into()
})
.decrypt(&ct),
Err(CryptoError::Decrypt)
));
}
#[test]
fn personal_and_org_audiences_are_independent() {
let p = provider("alice", &[("acme", [7u8; 32])]);
let msg = serde_json::json!({"m": 1});
let org_ct = p
.cipher_for(&Scope::Shared { org: "acme".into() })
.encrypt(&msg)
.unwrap();
assert!(matches!(
p.cipher_for(&Scope::Personal).decrypt(&org_ct),
Err(CryptoError::Decrypt)
));
}
#[test]
fn org_cipher_is_cached_per_audience() {
let p = provider("alice", &[("acme", [3u8; 32])]);
let scope = Scope::Shared { org: "acme".into() };
let a = p.cipher_for(&scope);
let b = p.cipher_for(&scope);
assert!(
Arc::ptr_eq(&a, &b),
"HKDF should run once per org, not per op"
);
}
#[test]
fn debug_never_prints_key_material() {
let p = provider("alice", &[("acme", [0xabu8; 32])]);
let dbg = format!("{p:?}");
assert!(dbg.contains("acme"), "shows which orgs are resolved");
assert!(!dbg.contains("abab"), "must not leak key bytes");
}
#[test]
fn encrypt_uses_newest_epoch_and_older_ops_still_decrypt() {
let scope = Scope::Shared { org: "acme".into() };
let old_only = provider_multi("alice", &[("acme", &[(1, [1u8; 32])])]);
let both = provider_multi("alice", &[("acme", &[(1, [1u8; 32]), (2, [2u8; 32])])]);
let old_msg = serde_json::json!({"gen": 1});
let old_ct = old_only.cipher_for(&scope).encrypt(&old_msg).unwrap();
assert_eq!(both.cipher_for(&scope).decrypt(&old_ct).unwrap(), old_msg);
let new_ct = both
.cipher_for(&scope)
.encrypt(&serde_json::json!({"gen": 2}))
.unwrap();
assert_eq!(new_ct.get("kid").and_then(|v| v.as_u64()), Some(2));
}
#[test]
fn dropped_epoch_can_no_longer_decrypt_its_ops() {
let scope = Scope::Shared { org: "acme".into() };
let e1 = provider_multi("alice", &[("acme", &[(1, [1u8; 32])])]);
let old_ct = e1
.cipher_for(&scope)
.encrypt(&serde_json::json!({"gen": 1}))
.unwrap();
let e2_only = provider_multi("alice", &[("acme", &[(2, [2u8; 32])])]);
assert!(matches!(
e2_only.cipher_for(&scope).decrypt(&old_ct),
Err(CryptoError::Key(_))
));
}
#[test]
fn org_envelope_missing_kid_fails_closed() {
let scope = Scope::Shared { org: "acme".into() };
let p = provider("alice", &[("acme", [4u8; 32])]);
let no_kid = serde_json::json!({
"car_enc": "chacha20poly1305",
"nonce": "000000000000000000000000",
"ct": "00",
});
assert!(matches!(
p.cipher_for(&scope).decrypt(&no_kid),
Err(CryptoError::BadEnvelope(_))
));
}
#[test]
fn org_envelopes_carry_a_kid_personal_ones_do_not() {
let p = provider("alice", &[("acme", [5u8; 32])]);
let org_ct = p
.cipher_for(&Scope::Shared { org: "acme".into() })
.encrypt(&serde_json::json!({"x": 1}))
.unwrap();
assert!(org_ct.get("kid").is_some(), "org envelope stamps epoch");
let personal_ct = p
.cipher_for(&Scope::Personal)
.encrypt(&serde_json::json!({"x": 1}))
.unwrap();
assert!(
personal_ct.get("kid").is_none(),
"personal envelope omits kid (wire-compatible)"
);
}
}