use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use base64::{engine::general_purpose::STANDARD as B64, Engine};
use zeroize::{Zeroize, Zeroizing};
use crate::error::{CryptoError, Result};
use crate::primitives;
use crate::storage::KeyStore;
use crate::types::{
E2EEnvelope, ECDHKeyPair, KeyData, KeyExchangeMessage, PublicKeyAnnouncement, TofuRecord,
};
pub type OnKeyChange = Arc<dyn Fn(&str, &str, &str) -> bool + Send + Sync>;
const MIN_ROTATION_INTERVAL: Duration = Duration::from_secs(60);
const DEFAULT_MAX_ANNOUNCEMENT_FUTURE: Duration = Duration::from_secs(30);
pub struct E2ESessionConfig {
pub identity_id: String,
pub base_path: String,
pub store: Arc<dyn KeyStore>,
pub on_key_change: Option<OnKeyChange>,
pub password_hash: Option<String>,
pub rotation_interval: Option<Duration>,
pub on_rotation: Option<Arc<dyn Fn() + Send + Sync>>,
pub max_announcement_age: Option<Duration>,
}
pub struct E2ESession {
config: E2ESessionConfig,
group_key: Option<Vec<u8>>,
ecdh_key_pair: Option<ECDHKeyPair>,
peer_public_keys: HashMap<String, Vec<u8>>,
started: bool,
destroyed: bool,
last_rotation: Option<u64>,
rotation_count: u64,
seen_nonces: HashSet<String>,
}
impl E2ESession {
pub fn new(config: E2ESessionConfig) -> Self {
Self {
config,
group_key: None,
ecdh_key_pair: None,
peer_public_keys: HashMap::new(),
started: false,
destroyed: false,
last_rotation: None,
rotation_count: 0,
seen_nonces: HashSet::new(),
}
}
pub fn encrypted(&self) -> bool {
self.group_key.is_some()
}
pub fn base_path(&self) -> &str {
&self.config.base_path
}
pub async fn start(&mut self) -> Result<()> {
if self.destroyed {
return Err(CryptoError::SessionDestroyed);
}
if self.started {
return Ok(());
}
self.started = true;
let session_id = self.session_id();
if let Some(data) = self.config.store.load_group_key(&session_id).await? {
match primitives::jwk_to_group_key(&data.key) {
Ok(key) => {
self.group_key = Some(key);
self.last_rotation = Some(data.stored_at);
}
Err(_) => {
self.config.store.delete_group_key(&session_id).await?;
}
}
}
Ok(())
}
pub async fn enable_encryption(&mut self) -> Result<PublicKeyAnnouncement> {
if self.destroyed {
return Err(CryptoError::SessionDestroyed);
}
let mut key = Zeroizing::new(primitives::generate_group_key());
self.group_key = Some(key.to_vec());
let creation_time = now_ms();
self.last_rotation = Some(creation_time);
let jwk = primitives::group_key_to_jwk(&key)?;
key.zeroize();
self.config
.store
.save_group_key(
&self.session_id(),
KeyData {
key: jwk,
stored_at: creation_time,
},
)
.await?;
self.make_public_key_announcement()
}
pub fn request_group_key(&mut self) -> Result<Option<PublicKeyAnnouncement>> {
if self.destroyed {
return Err(CryptoError::SessionDestroyed);
}
if self.group_key.is_some() {
return Ok(None);
}
self.make_public_key_announcement().map(Some)
}
pub fn encrypt(&self, value: &str) -> Result<E2EEnvelope> {
if self.destroyed {
return Err(CryptoError::SessionDestroyed);
}
let key = self.group_key.as_ref().ok_or(CryptoError::NoGroupKey)?;
let plaintext = value.as_bytes();
let (ciphertext, iv) = primitives::encrypt(key, plaintext)?;
Ok(E2EEnvelope {
_e2e: 1,
ct: B64.encode(&ciphertext),
iv: B64.encode(&iv),
v: 1,
})
}
pub async fn decrypt(&mut self, envelope: &E2EEnvelope) -> Result<String> {
if self.destroyed {
return Err(CryptoError::SessionDestroyed);
}
if envelope._e2e != 1 {
return Err(CryptoError::DecryptionFailed("invalid E2E marker".into()));
}
if envelope.v != 1 {
return Err(CryptoError::DecryptionFailed(
"unsupported envelope version".into(),
));
}
let key = Zeroizing::new(match &self.group_key {
Some(k) => k.clone(),
None => {
let session_id = self.session_id();
match self.config.store.load_group_key(&session_id).await? {
Some(data) => {
let k = primitives::jwk_to_group_key(&data.key)?;
self.group_key = Some(k.clone());
k
}
None => return Err(CryptoError::NoGroupKey),
}
}
});
let ciphertext = B64
.decode(&envelope.ct)
.map_err(|e| CryptoError::DecryptionFailed(format!("invalid base64 ct: {e}")))?;
let iv = B64
.decode(&envelope.iv)
.map_err(|e| CryptoError::DecryptionFailed(format!("invalid base64 iv: {e}")))?;
let plaintext = primitives::decrypt(&key, &ciphertext, &iv)?;
String::from_utf8(plaintext)
.map_err(|e| CryptoError::DecryptionFailed(format!("invalid UTF-8: {e}")))
}
pub async fn handle_peer_pubkey(
&mut self,
peer_id: &str,
announcement: &PublicKeyAnnouncement,
) -> Result<Option<KeyExchangeMessage>> {
if self.destroyed {
return Err(CryptoError::SessionDestroyed);
}
if peer_id == self.config.identity_id {
return Ok(None);
}
if let Some(max_age) = self.config.max_announcement_age {
let now = now_ms();
let ts = announcement.timestamp;
let max_age_ms = max_age.as_millis() as u64;
let max_future_ms = DEFAULT_MAX_ANNOUNCEMENT_FUTURE.as_millis() as u64;
if ts + max_age_ms < now {
return Err(CryptoError::Other(format!(
"announcement from {peer_id} is too old ({} ms)",
now.saturating_sub(ts)
)));
}
if ts > now + max_future_ms {
return Err(CryptoError::Other(format!(
"announcement from {peer_id} is too far in the future ({} ms)",
ts.saturating_sub(now)
)));
}
}
let peer_pub_bytes = primitives::jwk_to_public_key(&announcement.public_key)?;
self.verify_peer_key(peer_id, &announcement.public_key)
.await?;
self.peer_public_keys
.insert(peer_id.to_string(), peer_pub_bytes.clone());
let group_key = Zeroizing::new(match &self.group_key {
Some(k) => k.clone(),
None => return Ok(None),
});
self.ensure_ecdh_key_pair();
let kp = self.ecdh_key_pair();
let shared = Zeroizing::new(primitives::derive_shared_key(
&kp.private_key,
&peer_pub_bytes,
None,
)?);
let group_key_jwk = primitives::group_key_to_jwk(&group_key)?;
let mut group_key_json = Zeroizing::new(
serde_json::to_string(&group_key_jwk)
.map_err(|e| CryptoError::Serialization(e.to_string()))?,
);
let (ct, iv) = primitives::encrypt(&shared, group_key_json.as_bytes())?;
group_key_json.zeroize();
let sender_pub_jwk = primitives::public_key_to_jwk(&kp.public_key)?;
Ok(Some(KeyExchangeMessage {
from_id: self.config.identity_id.clone(),
encrypted_key: B64.encode(&ct),
iv: B64.encode(&iv),
sender_public_key: sender_pub_jwk,
}))
}
pub async fn handle_key_exchange(&mut self, msg: &KeyExchangeMessage) -> Result<()> {
if self.destroyed {
return Err(CryptoError::SessionDestroyed);
}
if msg.from_id.is_empty() {
return Err(CryptoError::InvalidKey(
"key exchange message missing sender ID".into(),
));
}
let nonce_key = format!("{}:{}", msg.from_id, msg.iv);
if !self.seen_nonces.insert(nonce_key) {
return Err(CryptoError::Other(format!(
"replayed key exchange message from {}",
msg.from_id
)));
}
if self.seen_nonces.len() > 10_000 {
self.seen_nonces.clear();
}
let sender_pub = primitives::jwk_to_public_key(&msg.sender_public_key)?;
self.verify_peer_key(&msg.from_id, &msg.sender_public_key)
.await?;
self.peer_public_keys
.insert(msg.from_id.clone(), sender_pub.clone());
self.ensure_ecdh_key_pair();
let kp = self.ecdh_key_pair();
let shared = Zeroizing::new(primitives::derive_shared_key(
&kp.private_key,
&sender_pub,
None,
)?);
let ct = B64
.decode(&msg.encrypted_key)
.map_err(|e| CryptoError::DecryptionFailed(format!("invalid base64: {e}")))?;
let iv = B64
.decode(&msg.iv)
.map_err(|e| CryptoError::DecryptionFailed(format!("invalid base64: {e}")))?;
let decrypted = primitives::decrypt(&shared, &ct, &iv)?;
let mut key_json = Zeroizing::new(
String::from_utf8(decrypted)
.map_err(|e| CryptoError::DecryptionFailed(format!("invalid UTF-8: {e}")))?,
);
let key_jwk: serde_json::Value = serde_json::from_str(&key_json)
.map_err(|e| CryptoError::DecryptionFailed(format!("invalid JWK JSON: {e}")))?;
key_json.zeroize();
let mut group_key = Zeroizing::new(primitives::jwk_to_group_key(&key_jwk)?);
self.group_key = Some(group_key.to_vec());
group_key.zeroize();
self.config
.store
.save_group_key(
&self.session_id(),
KeyData {
key: key_jwk,
stored_at: now_ms(),
},
)
.await?;
Ok(())
}
pub async fn rotate_key(&mut self) -> Result<Vec<(String, KeyExchangeMessage)>> {
if self.destroyed {
return Err(CryptoError::SessionDestroyed);
}
if self.group_key.is_none() {
return Ok(vec![]);
}
if let Some(ref mut old_key) = self.group_key {
old_key.zeroize();
}
let mut new_key = Zeroizing::new(primitives::generate_group_key());
self.group_key = Some(new_key.to_vec());
let rotation_time = now_ms();
self.last_rotation = Some(rotation_time);
self.rotation_count += 1;
let jwk = primitives::group_key_to_jwk(&new_key)?;
let mut group_key_json = Zeroizing::new(
serde_json::to_string(&jwk).map_err(|e| CryptoError::Serialization(e.to_string()))?,
);
new_key.zeroize();
self.config
.store
.save_group_key(
&self.session_id(),
KeyData {
key: jwk,
stored_at: rotation_time,
},
)
.await?;
self.ensure_ecdh_key_pair();
let kp = self.ecdh_key_pair();
let sender_pub_jwk = primitives::public_key_to_jwk(&kp.public_key)?;
let mut messages = Vec::new();
for (peer_id, peer_pub) in &self.peer_public_keys {
if *peer_id == self.config.identity_id {
continue;
}
if let Ok(shared) = primitives::derive_shared_key(&kp.private_key, peer_pub, None) {
let mut shared = Zeroizing::new(shared);
if let Ok((ct, iv)) = primitives::encrypt(&shared, group_key_json.as_bytes()) {
messages.push((
peer_id.clone(),
KeyExchangeMessage {
from_id: self.config.identity_id.clone(),
encrypted_key: B64.encode(&ct),
iv: B64.encode(&iv),
sender_public_key: sender_pub_jwk.clone(),
},
));
}
shared.zeroize();
}
}
group_key_json.zeroize();
Ok(messages)
}
pub fn remove_peer(&mut self, peer_id: &str) {
self.peer_public_keys.remove(peer_id);
}
pub fn should_rotate(&self) -> bool {
let interval = match self.config.rotation_interval {
Some(d) => d.max(MIN_ROTATION_INTERVAL),
None => return false,
};
if self.group_key.is_none() || self.destroyed {
return false;
}
let last = self.last_rotation.unwrap_or(0);
if last == 0 {
return false;
}
let elapsed_ms = now_ms().saturating_sub(last);
elapsed_ms >= interval.as_millis() as u64
}
pub async fn maybe_rotate(
&mut self,
) -> Result<Option<(Vec<(String, KeyExchangeMessage)>, PublicKeyAnnouncement)>> {
if !self.should_rotate() {
return Ok(None);
}
let messages = self.rotate_key().await?;
let announcement = self.make_public_key_announcement()?;
if let Some(ref cb) = self.config.on_rotation {
cb();
}
Ok(Some((messages, announcement)))
}
pub fn rotation_count(&self) -> u64 {
self.rotation_count
}
pub fn last_rotation(&self) -> Option<u64> {
self.last_rotation
}
pub fn destroy(&mut self) {
self.destroyed = true;
if let Some(ref mut key) = self.group_key {
key.zeroize();
}
self.group_key = None;
self.ecdh_key_pair = None;
self.peer_public_keys.clear();
self.seen_nonces.clear();
}
fn session_id(&self) -> String {
self.config.base_path.clone()
}
fn ensure_ecdh_key_pair(&mut self) {
if self.ecdh_key_pair.is_none() {
self.ecdh_key_pair = Some(primitives::generate_ecdh_key_pair());
}
}
fn ecdh_key_pair(&self) -> &ECDHKeyPair {
self.ecdh_key_pair.as_ref().unwrap()
}
fn make_public_key_announcement(&mut self) -> Result<PublicKeyAnnouncement> {
self.ensure_ecdh_key_pair();
let kp = self.ecdh_key_pair();
let jwk = primitives::public_key_to_jwk(&kp.public_key)?;
Ok(PublicKeyAnnouncement {
public_key: jwk,
timestamp: now_ms(),
})
}
async fn verify_peer_key(
&self,
peer_id: &str,
public_key_jwk: &serde_json::Value,
) -> Result<()> {
let fp = primitives::fingerprint_jwk(public_key_jwk);
let record_id = format!("{}:{}", self.config.base_path, peer_id);
let stored = self.config.store.load_tofu_record(&record_id).await?;
match stored {
None => {
self.config
.store
.save_tofu_record(
&record_id,
TofuRecord {
fingerprint: fp,
first_seen: now_ms(),
},
)
.await?;
}
Some(record) => {
if !primitives::constant_time_eq(record.fingerprint.as_bytes(), fp.as_bytes()) {
let accepted = self
.config
.on_key_change
.as_ref()
.map(|cb| cb(peer_id, &record.fingerprint, &fp))
.unwrap_or(false);
if !accepted {
return Err(CryptoError::TofuViolation(peer_id.to_string()));
}
self.config
.store
.save_tofu_record(
&record_id,
TofuRecord {
fingerprint: fp,
first_seen: record.first_seen,
},
)
.await?;
}
}
}
Ok(())
}
}
impl Drop for E2ESession {
fn drop(&mut self) {
if !self.destroyed {
self.destroy();
}
}
}
fn now_ms() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as u64
}
#[cfg(test)]
mod tests {
use super::*;
use crate::storage::MemoryKeyStore;
fn test_config(store: Arc<dyn KeyStore>) -> E2ESessionConfig {
E2ESessionConfig {
identity_id: "alice".to_string(),
base_path: "/test/room/1".to_string(),
store,
on_key_change: None,
password_hash: None,
rotation_interval: None,
on_rotation: None,
max_announcement_age: None,
}
}
#[tokio::test]
async fn session_starts_without_key() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(test_config(store));
session.start().await.unwrap();
assert!(!session.encrypted());
}
#[tokio::test]
async fn enable_encryption_creates_key() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(test_config(store));
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
assert!(session.encrypted());
}
#[tokio::test]
async fn encrypt_decrypt_round_trip() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(test_config(store));
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
let envelope = session.encrypt("Hello, world!").unwrap();
assert_eq!(envelope._e2e, 1);
assert_eq!(envelope.v, 1);
let decrypted = session.decrypt(&envelope).await.unwrap();
assert_eq!(decrypted, "Hello, world!");
}
#[tokio::test]
async fn persists_and_loads_key() {
let store = Arc::new(MemoryKeyStore::new());
let mut session1 = E2ESession::new(test_config(store.clone()));
session1.start().await.unwrap();
session1.enable_encryption().await.unwrap();
let envelope = session1.encrypt("hello").unwrap();
let mut session2 = E2ESession::new(test_config(store));
session2.start().await.unwrap();
assert!(session2.encrypted());
let decrypted = session2.decrypt(&envelope).await.unwrap();
assert_eq!(decrypted, "hello");
}
#[tokio::test]
async fn key_exchange_between_peers() {
let store_a = Arc::new(MemoryKeyStore::new());
let store_b = Arc::new(MemoryKeyStore::new());
let mut alice = E2ESession::new(E2ESessionConfig {
identity_id: "alice".to_string(),
base_path: "/test/room/1".to_string(),
store: store_a,
on_key_change: None,
password_hash: None,
rotation_interval: None,
on_rotation: None,
max_announcement_age: None,
});
alice.start().await.unwrap();
let _alice_announcement = alice.enable_encryption().await.unwrap();
let mut bob = E2ESession::new(E2ESessionConfig {
identity_id: "bob".to_string(),
base_path: "/test/room/1".to_string(),
store: store_b,
on_key_change: None,
password_hash: None,
rotation_interval: None,
on_rotation: None,
max_announcement_age: None,
});
bob.start().await.unwrap();
let bob_announcement = bob.request_group_key().unwrap().unwrap();
let keyex = alice
.handle_peer_pubkey("bob", &bob_announcement)
.await
.unwrap();
assert!(keyex.is_some());
bob.handle_key_exchange(&keyex.unwrap()).await.unwrap();
assert!(bob.encrypted());
let envelope = alice.encrypt("secret message").unwrap();
let decrypted = bob.decrypt(&envelope).await.unwrap();
assert_eq!(decrypted, "secret message");
}
#[tokio::test]
async fn rotate_key_invalidates_old_messages() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(test_config(store));
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
let old_envelope = session.encrypt("before rotation").unwrap();
session.rotate_key().await.unwrap();
let result = session.decrypt(&old_envelope).await;
assert!(result.is_err());
}
#[tokio::test]
async fn tofu_detects_key_change_and_accepts() {
use std::sync::atomic::{AtomicBool, Ordering};
let store = Arc::new(MemoryKeyStore::new());
let changed = Arc::new(AtomicBool::new(false));
let changed_clone = changed.clone();
let mut session = E2ESession::new(E2ESessionConfig {
identity_id: "alice".to_string(),
base_path: "/test/room/1".to_string(),
store: store.clone(),
on_key_change: Some(Arc::new(move |_peer, _old, _new| {
changed_clone.store(true, Ordering::SeqCst);
true })),
password_hash: None,
rotation_interval: None,
on_rotation: None,
max_announcement_age: None,
});
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
let bob_kp1 = primitives::generate_ecdh_key_pair();
let ann1 = PublicKeyAnnouncement {
public_key: primitives::public_key_to_jwk(&bob_kp1.public_key).unwrap(),
timestamp: now_ms(),
};
session.handle_peer_pubkey("bob", &ann1).await.unwrap();
assert!(!changed.load(Ordering::SeqCst));
let bob_kp2 = primitives::generate_ecdh_key_pair();
let ann2 = PublicKeyAnnouncement {
public_key: primitives::public_key_to_jwk(&bob_kp2.public_key).unwrap(),
timestamp: now_ms(),
};
session.handle_peer_pubkey("bob", &ann2).await.unwrap();
assert!(changed.load(Ordering::SeqCst));
}
#[tokio::test]
async fn tofu_rejects_key_change_without_callback() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(test_config(store.clone()));
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
let bob_kp1 = primitives::generate_ecdh_key_pair();
let ann1 = PublicKeyAnnouncement {
public_key: primitives::public_key_to_jwk(&bob_kp1.public_key).unwrap(),
timestamp: now_ms(),
};
session.handle_peer_pubkey("bob", &ann1).await.unwrap();
let bob_kp2 = primitives::generate_ecdh_key_pair();
let ann2 = PublicKeyAnnouncement {
public_key: primitives::public_key_to_jwk(&bob_kp2.public_key).unwrap(),
timestamp: now_ms(),
};
let result = session.handle_peer_pubkey("bob", &ann2).await;
assert!(matches!(result, Err(CryptoError::TofuViolation(_))));
}
#[tokio::test]
async fn tofu_rejects_key_change_when_callback_returns_false() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(E2ESessionConfig {
identity_id: "alice".to_string(),
base_path: "/test/room/1".to_string(),
store: store.clone(),
on_key_change: Some(Arc::new(|_peer, _old, _new| false)),
password_hash: None,
rotation_interval: None,
on_rotation: None,
max_announcement_age: None,
});
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
let bob_kp1 = primitives::generate_ecdh_key_pair();
let ann1 = PublicKeyAnnouncement {
public_key: primitives::public_key_to_jwk(&bob_kp1.public_key).unwrap(),
timestamp: now_ms(),
};
session.handle_peer_pubkey("bob", &ann1).await.unwrap();
let bob_kp2 = primitives::generate_ecdh_key_pair();
let ann2 = PublicKeyAnnouncement {
public_key: primitives::public_key_to_jwk(&bob_kp2.public_key).unwrap(),
timestamp: now_ms(),
};
let result = session.handle_peer_pubkey("bob", &ann2).await;
assert!(matches!(result, Err(CryptoError::TofuViolation(_))));
}
#[tokio::test]
async fn tofu_stores_records_without_callback() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(test_config(store.clone()));
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
let bob_kp = primitives::generate_ecdh_key_pair();
let ann = PublicKeyAnnouncement {
public_key: primitives::public_key_to_jwk(&bob_kp.public_key).unwrap(),
timestamp: now_ms(),
};
session.handle_peer_pubkey("bob", &ann).await.unwrap();
let record = store.load_tofu_record("/test/room/1:bob").await.unwrap();
assert!(record.is_some());
}
#[tokio::test]
async fn empty_from_id_rejected() {
let store_a = Arc::new(MemoryKeyStore::new());
let store_b = Arc::new(MemoryKeyStore::new());
let mut alice = E2ESession::new(E2ESessionConfig {
identity_id: "alice".to_string(),
base_path: "/test/room/1".to_string(),
store: store_a,
on_key_change: None,
password_hash: None,
rotation_interval: None,
on_rotation: None,
max_announcement_age: None,
});
alice.start().await.unwrap();
let mut bob = E2ESession::new(E2ESessionConfig {
identity_id: "bob".to_string(),
base_path: "/test/room/1".to_string(),
store: store_b,
on_key_change: None,
password_hash: None,
rotation_interval: None,
on_rotation: None,
max_announcement_age: None,
});
bob.start().await.unwrap();
bob.enable_encryption().await.unwrap();
let bob_announcement = bob.request_group_key().unwrap();
let msg = KeyExchangeMessage {
from_id: String::new(),
encrypted_key: "AAAA".to_string(),
iv: "BBBB".to_string(),
sender_public_key: serde_json::json!({}),
};
let result = alice.handle_key_exchange(&msg).await;
assert!(matches!(result, Err(CryptoError::InvalidKey(_))));
drop(bob_announcement);
}
#[tokio::test]
async fn encrypt_after_destroy_fails() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(test_config(store));
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
session.destroy();
let result = session.encrypt("test");
assert!(matches!(result, Err(CryptoError::SessionDestroyed)));
}
#[tokio::test]
async fn decrypt_after_destroy_fails() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(test_config(store));
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
let envelope = session.encrypt("test").unwrap();
session.destroy();
let result = session.decrypt(&envelope).await;
assert!(matches!(result, Err(CryptoError::SessionDestroyed)));
}
#[tokio::test]
async fn handle_peer_pubkey_after_destroy_fails() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(test_config(store));
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
session.destroy();
let bob_kp = primitives::generate_ecdh_key_pair();
let ann = PublicKeyAnnouncement {
public_key: primitives::public_key_to_jwk(&bob_kp.public_key).unwrap(),
timestamp: now_ms(),
};
let result = session.handle_peer_pubkey("bob", &ann).await;
assert!(matches!(result, Err(CryptoError::SessionDestroyed)));
}
#[tokio::test]
async fn handle_key_exchange_after_destroy_fails() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(test_config(store));
session.start().await.unwrap();
session.destroy();
let msg = KeyExchangeMessage {
from_id: "bob".to_string(),
encrypted_key: "AAAA".to_string(),
iv: "BBBB".to_string(),
sender_public_key: serde_json::json!({}),
};
let result = session.handle_key_exchange(&msg).await;
assert!(matches!(result, Err(CryptoError::SessionDestroyed)));
}
#[tokio::test]
async fn decrypt_rejects_unknown_envelope_version() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(test_config(store));
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
let envelope = E2EEnvelope {
_e2e: 1,
ct: "AAAA".to_string(),
iv: "BBBB".to_string(),
v: 2,
};
let result = session.decrypt(&envelope).await;
assert!(matches!(result, Err(CryptoError::DecryptionFailed(_))));
}
#[tokio::test]
async fn should_rotate_returns_false_without_interval() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(test_config(store));
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
assert!(!session.should_rotate());
}
#[tokio::test]
async fn rotation_tracks_count_and_timestamp() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(test_config(store));
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
assert_eq!(session.rotation_count(), 0);
assert!(session.last_rotation().is_some());
session.rotate_key().await.unwrap();
assert_eq!(session.rotation_count(), 1);
session.rotate_key().await.unwrap();
assert_eq!(session.rotation_count(), 2);
}
#[tokio::test]
async fn maybe_rotate_triggers_when_due() {
use std::sync::atomic::{AtomicU32, Ordering};
let store = Arc::new(MemoryKeyStore::new());
let rotation_cb_count = Arc::new(AtomicU32::new(0));
let cb_clone = rotation_cb_count.clone();
let mut session = E2ESession::new(E2ESessionConfig {
identity_id: "alice".to_string(),
base_path: "/test/room/1".to_string(),
store,
on_key_change: None,
password_hash: None,
rotation_interval: Some(Duration::from_secs(60)),
on_rotation: Some(Arc::new(move || {
cb_clone.fetch_add(1, Ordering::SeqCst);
})),
max_announcement_age: None,
});
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
let result = session.maybe_rotate().await.unwrap();
assert!(result.is_none());
session.last_rotation = Some(now_ms() - 120_000);
let result = session.maybe_rotate().await.unwrap();
assert!(result.is_some());
assert_eq!(rotation_cb_count.load(Ordering::SeqCst), 1);
assert_eq!(session.rotation_count(), 1);
}
#[tokio::test]
async fn timestamp_validation_rejects_old_announcement() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(E2ESessionConfig {
identity_id: "alice".to_string(),
base_path: "/test/room/1".to_string(),
store,
on_key_change: None,
password_hash: None,
rotation_interval: None,
on_rotation: None,
max_announcement_age: Some(Duration::from_secs(300)),
});
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
let bob_kp = primitives::generate_ecdh_key_pair();
let old_announcement = PublicKeyAnnouncement {
public_key: primitives::public_key_to_jwk(&bob_kp.public_key).unwrap(),
timestamp: now_ms() - 600_000, };
let result = session.handle_peer_pubkey("bob", &old_announcement).await;
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("too old"));
}
#[tokio::test]
async fn timestamp_validation_rejects_future_announcement() {
let store = Arc::new(MemoryKeyStore::new());
let mut session = E2ESession::new(E2ESessionConfig {
identity_id: "alice".to_string(),
base_path: "/test/room/1".to_string(),
store,
on_key_change: None,
password_hash: None,
rotation_interval: None,
on_rotation: None,
max_announcement_age: Some(Duration::from_secs(300)),
});
session.start().await.unwrap();
session.enable_encryption().await.unwrap();
let bob_kp = primitives::generate_ecdh_key_pair();
let future_announcement = PublicKeyAnnouncement {
public_key: primitives::public_key_to_jwk(&bob_kp.public_key).unwrap(),
timestamp: now_ms() + 60_000, };
let result = session
.handle_peer_pubkey("bob", &future_announcement)
.await;
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("future"));
}
#[tokio::test]
async fn replay_protection_rejects_duplicate_key_exchange() {
let store_a = Arc::new(MemoryKeyStore::new());
let store_b = Arc::new(MemoryKeyStore::new());
let mut alice = E2ESession::new(E2ESessionConfig {
identity_id: "alice".to_string(),
base_path: "/test/room/1".to_string(),
store: store_a,
on_key_change: None,
password_hash: None,
rotation_interval: None,
on_rotation: None,
max_announcement_age: None,
});
alice.start().await.unwrap();
alice.enable_encryption().await.unwrap();
let mut bob = E2ESession::new(E2ESessionConfig {
identity_id: "bob".to_string(),
base_path: "/test/room/1".to_string(),
store: store_b,
on_key_change: None,
password_hash: None,
rotation_interval: None,
on_rotation: None,
max_announcement_age: None,
});
bob.start().await.unwrap();
let bob_announcement = bob.request_group_key().unwrap().unwrap();
let keyex = alice
.handle_peer_pubkey("bob", &bob_announcement)
.await
.unwrap()
.unwrap();
bob.handle_key_exchange(&keyex).await.unwrap();
let result = bob.handle_key_exchange(&keyex).await;
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("replayed"));
}
#[tokio::test]
async fn persisted_key_restores_last_rotation() {
let store = Arc::new(MemoryKeyStore::new());
let mut session1 = E2ESession::new(test_config(store.clone()));
session1.start().await.unwrap();
session1.enable_encryption().await.unwrap();
let rotation_ts = session1.last_rotation().unwrap();
assert!(rotation_ts > 0);
let mut session2 = E2ESession::new(test_config(store));
session2.start().await.unwrap();
assert_eq!(session2.last_rotation(), Some(rotation_ts));
}
}