use std::collections::BTreeMap;
use base64::Engine as _;
use base64::engine::general_purpose::URL_SAFE_NO_PAD as B64;
use chacha20poly1305::aead::{Aead, KeyInit, Payload};
use chacha20poly1305::{ChaCha20Poly1305, Key, Nonce};
use crate::error::RoomKeyError;
use crate::mls::STORAGE_KEY_LEN;
use crate::wire::EpochLink;
pub type StorageKey = [u8; STORAGE_KEY_LEN];
fn link_aad(room_id: &str, epoch: u32) -> Vec<u8> {
format!("{room_id}|epoch-link|{epoch}|{}", epoch.saturating_sub(1)).into_bytes()
}
pub fn seal_link(
room_id: &str,
epoch: u32,
current: &StorageKey,
predecessor: &StorageKey,
) -> Result<EpochLink, RoomKeyError> {
if epoch < 2 {
return Err(RoomKeyError::Seal(format!(
"epoch {epoch} has no predecessor to link to"
)));
}
let cipher = ChaCha20Poly1305::new(&Key::from(*current));
let mut nonce_bytes = [0u8; 12];
getrandom::fill(&mut nonce_bytes).expect("OS randomness unavailable");
let wrapped = cipher
.encrypt(
&Nonce::from(nonce_bytes),
Payload {
msg: predecessor.as_slice(),
aad: &link_aad(room_id, epoch),
},
)
.map_err(|e| RoomKeyError::Seal(format!("seal the epoch link: {e}")))?;
Ok(EpochLink {
epoch,
wrapped: B64.encode(wrapped),
nonce: B64.encode(nonce_bytes),
})
}
pub fn open_link(
room_id: &str,
link: &EpochLink,
current: &StorageKey,
) -> Result<StorageKey, RoomKeyError> {
let wrapped = B64
.decode(&link.wrapped)
.map_err(|e| RoomKeyError::Seal(format!("decode the epoch link: {e}")))?;
let nonce = B64
.decode(&link.nonce)
.map_err(|e| RoomKeyError::Seal(format!("decode the epoch link nonce: {e}")))?;
if nonce.len() != 12 {
return Err(RoomKeyError::Seal(format!(
"epoch link nonce is {} bytes, expected 12",
nonce.len()
)));
}
let cipher = ChaCha20Poly1305::new(&Key::from(*current));
let plain = cipher
.decrypt(
&Nonce::try_from(&nonce[..])
.map_err(|e| RoomKeyError::Seal(format!("epoch link nonce: {e}")))?,
Payload {
msg: &wrapped,
aad: &link_aad(room_id, link.epoch),
},
)
.map_err(|_| RoomKeyError::DidNotOpen)?;
let key: StorageKey = plain
.try_into()
.map_err(|_| RoomKeyError::Seal("an epoch link did not wrap a storage key".into()))?;
Ok(key)
}
#[derive(Debug, Clone)]
pub struct EpochKeyChain {
room_id: String,
anchor_epoch: u32,
links: BTreeMap<u32, EpochLink>,
resolved: BTreeMap<u32, StorageKey>,
}
impl EpochKeyChain {
pub fn new(room_id: impl Into<String>) -> Self {
Self {
room_id: room_id.into(),
anchor_epoch: 0,
links: BTreeMap::new(),
resolved: BTreeMap::new(),
}
}
pub fn reanchor(&mut self, epoch: u32, key: StorageKey) {
self.resolved.insert(epoch, key);
self.anchor_epoch = epoch;
}
pub fn add_links(&mut self, links: impl IntoIterator<Item = EpochLink>) {
for link in links {
self.links.insert(link.epoch, link);
}
}
pub fn anchor_epoch(&self) -> u32 {
self.anchor_epoch
}
pub fn links(&self) -> Vec<EpochLink> {
self.links.values().cloned().collect()
}
pub fn earliest_reachable(&mut self) -> u32 {
let mut epoch = self.anchor_epoch;
while epoch > 1 && self.key_for(epoch - 1).is_ok() {
epoch -= 1;
}
epoch
}
pub fn key_for(&mut self, epoch: u32) -> Result<StorageKey, RoomKeyError> {
if let Some(key) = self.resolved.get(&epoch) {
return Ok(*key);
}
if epoch > self.anchor_epoch {
return Err(RoomKeyError::EpochAhead {
sealed: epoch,
held: self.anchor_epoch,
});
}
let mut cursor = *self.resolved.range(epoch..).next().map(|(e, _)| e).ok_or(
RoomKeyError::EpochUnreachable {
sealed: epoch,
earliest: epoch,
},
)?;
while cursor > epoch {
let key = *self
.resolved
.get(&cursor)
.expect("the cursor only ever names a resolved epoch");
let link = self
.links
.get(&cursor)
.ok_or(RoomKeyError::EpochUnreachable {
sealed: epoch,
earliest: cursor,
})?;
let previous = open_link(&self.room_id, link, &key)?;
cursor -= 1;
self.resolved.insert(cursor, previous);
}
Ok(*self
.resolved
.get(&epoch)
.expect("the walk ends with the target resolved"))
}
}
#[cfg(test)]
mod tests {
use super::*;
const ROOM: &str = "did:webvh:zRoom";
fn key(seed: u8) -> StorageKey {
[seed; STORAGE_KEY_LEN]
}
#[test]
fn a_link_round_trips() {
let current = key(2);
let previous = key(1);
let link = seal_link(ROOM, 2, ¤t, &previous).expect("seal");
assert_eq!(open_link(ROOM, &link, ¤t).expect("open"), previous);
}
#[test]
fn a_link_lifted_to_another_rung_does_not_open() {
let current = key(2);
let link = seal_link(ROOM, 2, ¤t, &key(1)).expect("seal");
let mut moved = link.clone();
moved.epoch = 3;
assert!(
open_link(ROOM, &moved, ¤t).is_err(),
"relabelling a link's epoch must fail authentication"
);
}
#[test]
fn a_link_served_under_another_room_does_not_open() {
let current = key(2);
let link = seal_link(ROOM, 2, ¤t, &key(1)).expect("seal");
assert!(
open_link("did:webvh:zOtherRoom", &link, ¤t).is_err(),
"a rung must not open under a room it was not sealed for"
);
assert_eq!(
open_link(ROOM, &link, ¤t).expect("opens under its own room"),
key(1)
);
}
#[test]
fn the_wrong_key_does_not_open_a_link() {
let link = seal_link(ROOM, 2, &key(2), &key(1)).expect("seal");
assert!(open_link(ROOM, &link, &key(9)).is_err());
}
#[test]
fn the_first_epoch_has_no_predecessor() {
assert!(seal_link(ROOM, 1, &key(1), &key(0)).is_err());
assert!(seal_link(ROOM, 0, &key(1), &key(0)).is_err());
}
#[test]
fn a_chain_walks_back_to_the_first_epoch() {
let keys: Vec<StorageKey> = (1..=5).map(key).collect();
let links: Vec<EpochLink> = (2..=5)
.map(|e| {
seal_link(ROOM, e, &keys[(e - 1) as usize], &keys[(e - 2) as usize]).expect("seal")
})
.collect();
let mut chain = EpochKeyChain::new(ROOM);
chain.reanchor(5, keys[4]);
chain.add_links(links);
for epoch in 1..=5u32 {
assert_eq!(
chain.key_for(epoch).expect("every retained epoch resolves"),
keys[(epoch - 1) as usize],
"epoch {epoch} must resolve to the key it was sealed with"
);
}
assert_eq!(chain.earliest_reachable(), 1);
}
#[test]
fn a_severed_chain_reaches_no_further() {
let keys: Vec<StorageKey> = (1..=4).map(key).collect();
let links: Vec<EpochLink> = [3u32, 4]
.iter()
.map(|&e| {
seal_link(ROOM, e, &keys[(e - 1) as usize], &keys[(e - 2) as usize]).expect("seal")
})
.collect();
let mut chain = EpochKeyChain::new(ROOM);
chain.reanchor(4, keys[3]);
chain.add_links(links);
assert_eq!(chain.key_for(2).expect("epoch 2 is retained"), keys[1]);
assert!(
matches!(chain.key_for(1), Err(RoomKeyError::EpochUnreachable { .. })),
"a severed epoch must be unreachable, and say so as unreachable"
);
assert_eq!(chain.earliest_reachable(), 2);
}
#[test]
fn an_epoch_ahead_is_reported_as_ahead() {
let mut chain = EpochKeyChain::new(ROOM);
chain.reanchor(2, key(2));
assert!(matches!(
chain.key_for(5),
Err(RoomKeyError::EpochAhead { sealed: 5, held: 2 })
));
}
#[test]
fn re_anchoring_commit_by_commit_keeps_everything_below_reachable() {
let mut chain = EpochKeyChain::new(ROOM);
chain.reanchor(1, key(1));
for epoch in 2..=4u32 {
let link =
seal_link(ROOM, epoch, &key(epoch as u8), &key((epoch - 1) as u8)).expect("seal");
chain.add_links(Some(link));
chain.reanchor(epoch, key(epoch as u8));
}
assert_eq!(chain.anchor_epoch(), 4);
for epoch in 1..=4u32 {
assert_eq!(chain.key_for(epoch).expect("resolves"), key(epoch as u8));
}
assert_eq!(chain.earliest_reachable(), 1);
}
#[test]
fn re_anchoring_is_idempotent() {
let link = seal_link(ROOM, 2, &key(2), &key(1)).expect("seal");
let mut chain = EpochKeyChain::new(ROOM);
chain.add_links(Some(link));
chain.reanchor(2, key(2));
assert_eq!(chain.key_for(1).expect("walks back"), key(1));
chain.reanchor(2, key(2));
assert_eq!(chain.key_for(1).expect("still resolves"), key(1));
}
#[test]
fn advancing_without_a_link_hands_on_no_history() {
let mut chain = EpochKeyChain::new(ROOM);
chain.reanchor(1, key(1));
chain.reanchor(2, key(2));
assert_eq!(chain.key_for(2).expect("the anchor resolves"), key(2));
assert!(
chain.links().is_empty(),
"an unlinked advance must leave nothing to hand on"
);
let mut rebuilt = EpochKeyChain::new(ROOM);
rebuilt.reanchor(2, key(2));
rebuilt.add_links(chain.links());
assert!(matches!(
rebuilt.key_for(1),
Err(RoomKeyError::EpochUnreachable { .. })
));
assert_eq!(rebuilt.earliest_reachable(), 2);
}
}