use base64::Engine as _;
use base64::engine::general_purpose::URL_SAFE_NO_PAD as B64;
use openmls::prelude::*;
use openmls_basic_credential::SignatureKeyPair;
use openmls_rust_crypto::OpenMlsRustCrypto;
use openmls_traits::OpenMlsProvider;
use tls_codec::{Deserialize as _, Serialize as _};
use crate::error::RoomKeyError;
pub const ROOM_CIPHERSUITE: Ciphersuite = Ciphersuite::MLS_128_DHKEMX25519_AES128GCM_SHA256_Ed25519;
const STORAGE_KEY_LABEL: &str = "openvtc/room/storage/v1";
pub const STORAGE_KEY_LEN: usize = 32;
pub struct RoomIdentity {
signer: SignatureKeyPair,
credential: CredentialWithKey,
member_did: String,
}
impl RoomIdentity {
pub fn new(member_did: &str, provider: &impl OpenMlsProvider) -> Result<Self, RoomKeyError> {
let signer = SignatureKeyPair::new(ROOM_CIPHERSUITE.signature_algorithm())
.map_err(|e| RoomKeyError::Group(format!("generate MLS signature key: {e:?}")))?;
signer
.store(provider.storage())
.map_err(|e| RoomKeyError::Group(format!("store MLS signature key: {e:?}")))?;
let credential = Credential::new(CredentialType::Basic, member_did.as_bytes().to_vec());
Ok(Self {
credential: CredentialWithKey {
credential,
signature_key: signer.public().into(),
},
signer,
member_did: member_did.to_string(),
})
}
pub fn key_package(&self, provider: &impl OpenMlsProvider) -> Result<KeyPackage, RoomKeyError> {
KeyPackage::builder()
.build(
ROOM_CIPHERSUITE,
provider,
&self.signer,
self.credential.clone(),
)
.map(|b| b.key_package().clone())
.map_err(|e| RoomKeyError::Group(format!("build MLS key package: {e:?}")))
}
}
pub struct RoomGroup {
group: MlsGroup,
identity: RoomIdentity,
provider: OpenMlsRustCrypto,
}
pub struct MembershipChange {
pub commit: Vec<u8>,
pub welcome: Option<Vec<u8>>,
pub epoch: u64,
}
impl RoomGroup {
pub fn create(member_did: &str) -> Result<Self, RoomKeyError> {
let provider = OpenMlsRustCrypto::default();
let identity = RoomIdentity::new(member_did, &provider)?;
let config = MlsGroupCreateConfig::builder()
.ciphersuite(ROOM_CIPHERSUITE)
.use_ratchet_tree_extension(true)
.build();
let group = MlsGroup::new(
&provider,
&identity.signer,
&config,
identity.credential.clone(),
)
.map_err(|e| RoomKeyError::Group(format!("create MLS group: {e:?}")))?;
Ok(Self {
group,
identity,
provider,
})
}
pub fn join(member_did: &str, welcome: &[u8]) -> Result<Self, RoomKeyError> {
let provider = OpenMlsRustCrypto::default();
let identity = RoomIdentity::new(member_did, &provider)?;
Self::join_with(identity, provider, welcome)
}
pub fn join_with(
identity: RoomIdentity,
provider: OpenMlsRustCrypto,
welcome: &[u8],
) -> Result<Self, RoomKeyError> {
let msg = MlsMessageIn::tls_deserialize_exact(welcome)
.map_err(|e| RoomKeyError::Group(format!("parse welcome: {e:?}")))?;
let welcome = match msg.extract() {
MlsMessageBodyIn::Welcome(w) => w,
_ => {
return Err(RoomKeyError::Group(
"expected a Welcome message, got another MLS body".into(),
));
}
};
let config = MlsGroupJoinConfig::builder()
.use_ratchet_tree_extension(true)
.build();
let staged = StagedWelcome::new_from_welcome(&provider, &config, welcome, None)
.map_err(|e| RoomKeyError::Group(format!("stage welcome: {e:?}")))?;
let group = staged
.into_group(&provider)
.map_err(|e| RoomKeyError::Group(format!("join group from welcome: {e:?}")))?;
Ok(Self {
group,
identity,
provider,
})
}
pub fn add_member(
&mut self,
key_package: KeyPackage,
) -> Result<MembershipChange, RoomKeyError> {
let (commit, welcome, _) = self
.group
.add_members(&self.provider, &self.identity.signer, &[key_package])
.map_err(|e| RoomKeyError::Group(format!("add member: {e:?}")))?;
self.group
.merge_pending_commit(&self.provider)
.map_err(|e| RoomKeyError::Group(format!("merge add commit: {e:?}")))?;
Ok(MembershipChange {
commit: commit
.tls_serialize_detached()
.map_err(|e| RoomKeyError::Group(format!("serialise commit: {e:?}")))?,
welcome: Some(
welcome
.tls_serialize_detached()
.map_err(|e| RoomKeyError::Group(format!("serialise welcome: {e:?}")))?,
),
epoch: self.group.epoch().as_u64(),
})
}
pub fn remove_member(
&mut self,
index: LeafNodeIndex,
) -> Result<MembershipChange, RoomKeyError> {
let (commit, _, _) = self
.group
.remove_members(&self.provider, &self.identity.signer, &[index])
.map_err(|e| RoomKeyError::Group(format!("remove member: {e:?}")))?;
self.group
.merge_pending_commit(&self.provider)
.map_err(|e| RoomKeyError::Group(format!("merge remove commit: {e:?}")))?;
Ok(MembershipChange {
commit: commit
.tls_serialize_detached()
.map_err(|e| RoomKeyError::Group(format!("serialise commit: {e:?}")))?,
welcome: None,
epoch: self.group.epoch().as_u64(),
})
}
pub fn apply_commit(&mut self, commit: &[u8]) -> Result<u64, RoomKeyError> {
let msg = MlsMessageIn::tls_deserialize_exact(commit)
.map_err(|e| RoomKeyError::Group(format!("parse commit: {e:?}")))?;
let protocol_message: ProtocolMessage = msg
.try_into_protocol_message()
.map_err(|e| RoomKeyError::Group(format!("not a protocol message: {e:?}")))?;
let processed = self
.group
.process_message(&self.provider, protocol_message)
.map_err(|e| RoomKeyError::Group(format!("process commit: {e:?}")))?;
match processed.into_content() {
ProcessedMessageContent::StagedCommitMessage(staged) => {
self.group
.merge_staged_commit(&self.provider, *staged)
.map_err(|e| RoomKeyError::Group(format!("merge staged commit: {e:?}")))?;
Ok(self.group.epoch().as_u64())
}
_ => Err(RoomKeyError::Group(
"expected a commit, got another message type".into(),
)),
}
}
pub fn epoch(&self) -> u64 {
self.group.epoch().as_u64()
}
pub fn storage_key(&self) -> Result<[u8; STORAGE_KEY_LEN], RoomKeyError> {
let secret = self
.group
.export_secret(
self.provider.crypto(),
STORAGE_KEY_LABEL,
&[],
STORAGE_KEY_LEN,
)
.map_err(|e| RoomKeyError::Group(format!("export storage key: {e:?}")))?;
let mut key = [0u8; STORAGE_KEY_LEN];
key.copy_from_slice(&secret);
Ok(key)
}
pub fn epoch_authenticator(&self) -> Vec<u8> {
self.group.epoch_authenticator().as_slice().to_vec()
}
pub fn member_count(&self) -> usize {
self.group.members().count()
}
pub fn leaf_of(&self, member_did: &str) -> Option<LeafNodeIndex> {
self.group.members().find_map(|m| {
(m.credential.serialized_content() == member_did.as_bytes()).then_some(m.index)
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_creator_forms_a_group_of_one() {
let room = RoomGroup::create("did:key:zAlice").expect("create");
assert_eq!(room.member_count(), 1);
assert_eq!(room.epoch(), 0, "a fresh group starts at epoch 0");
}
#[test]
fn a_storage_key_is_derived_and_is_stable_within_an_epoch() {
let room = RoomGroup::create("did:key:zAlice").expect("create");
let a = room.storage_key().expect("export");
let b = room.storage_key().expect("export again");
assert_eq!(a, b, "the same epoch must derive the same key");
assert_ne!(a, [0u8; STORAGE_KEY_LEN], "and it must not be zeroes");
}
#[test]
fn an_added_member_derives_the_same_storage_key() {
let mut alice = RoomGroup::create("did:key:zAlice").expect("alice");
let bob_provider = OpenMlsRustCrypto::default();
let bob_identity = RoomIdentity::new("did:key:zBob", &bob_provider).expect("bob identity");
let bob_kp = bob_identity.key_package(&bob_provider).expect("bob kp");
let change = alice.add_member(bob_kp).expect("add bob");
let welcome = change.welcome.expect("an add produces a welcome");
let bob = RoomGroup::join_with(bob_identity, bob_provider, &welcome).expect("bob joins");
assert_eq!(alice.member_count(), 2);
assert_eq!(
alice.storage_key().unwrap(),
bob.storage_key().unwrap(),
"both members must derive the same storage key, without the host seeing it"
);
assert_eq!(alice.epoch(), bob.epoch());
}
#[test]
fn removing_a_member_changes_the_storage_key() {
let mut alice = RoomGroup::create("did:key:zAlice").expect("alice");
let bob_provider = OpenMlsRustCrypto::default();
let bob_identity = RoomIdentity::new("did:key:zBob", &bob_provider).expect("bob identity");
let bob_kp = bob_identity.key_package(&bob_provider).expect("bob kp");
let change = alice.add_member(bob_kp).expect("add bob");
let bob = RoomGroup::join_with(
bob_identity,
bob_provider,
&change.welcome.expect("welcome"),
)
.expect("bob joins");
let shared = alice.storage_key().unwrap();
assert_eq!(shared, bob.storage_key().unwrap());
let bob_leaf = alice.leaf_of("did:key:zBob").expect("bob is a member");
alice.remove_member(bob_leaf).expect("remove bob");
let after = alice.storage_key().unwrap();
assert_ne!(
shared, after,
"after removal the key must differ, or removal removes nothing"
);
assert_ne!(
bob.storage_key().unwrap(),
after,
"and the removed member must not be able to derive the new one"
);
}
#[test]
fn members_in_the_same_epoch_share_an_epoch_authenticator() {
let mut alice = RoomGroup::create("did:key:zAlice").expect("alice");
let bob_provider = OpenMlsRustCrypto::default();
let bob_identity = RoomIdentity::new("did:key:zBob", &bob_provider).expect("bob identity");
let bob_kp = bob_identity.key_package(&bob_provider).expect("bob kp");
let change = alice.add_member(bob_kp).expect("add bob");
let bob = RoomGroup::join_with(
bob_identity,
bob_provider,
&change.welcome.expect("welcome"),
)
.expect("bob joins");
assert_eq!(
alice.epoch_authenticator(),
bob.epoch_authenticator(),
"a member whose authenticator differs from the anchored one has been forked"
);
assert!(!alice.epoch_authenticator().is_empty());
}
#[test]
fn an_epoch_advances_on_every_membership_change() {
let mut alice = RoomGroup::create("did:key:zAlice").expect("alice");
let start = alice.epoch();
let p = OpenMlsRustCrypto::default();
let id = RoomIdentity::new("did:key:zBob", &p).expect("identity");
alice
.add_member(id.key_package(&p).expect("kp"))
.expect("add");
assert!(
alice.epoch() > start,
"a membership change must move the epoch, or the host serves stale ciphertext"
);
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct IdentitySnapshot {
entries: Vec<(String, String)>,
member_did: String,
signature_public: String,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct GroupSnapshot {
entries: Vec<(String, String)>,
group_id: String,
member_did: String,
signature_public: String,
}
impl RoomGroup {
pub fn snapshot(&self) -> Result<GroupSnapshot, RoomKeyError> {
let store = self.provider.storage();
let values = store
.values
.read()
.map_err(|_| RoomKeyError::Group("group store lock poisoned".into()))?;
Ok(GroupSnapshot {
entries: values
.iter()
.map(|(k, v)| (B64.encode(k), B64.encode(v)))
.collect(),
group_id: B64.encode(self.group.group_id().as_slice()),
member_did: self.identity.member_did.clone(),
signature_public: B64.encode(self.identity.signer.public()),
})
}
pub fn restore(snapshot: &GroupSnapshot) -> Result<Self, RoomKeyError> {
let provider = OpenMlsRustCrypto::default();
{
let mut values = provider
.storage()
.values
.write()
.map_err(|_| RoomKeyError::Group("group store lock poisoned".into()))?;
for (k, v) in &snapshot.entries {
let key = B64
.decode(k)
.map_err(|e| RoomKeyError::Group(format!("decode a store key: {e}")))?;
let value = B64
.decode(v)
.map_err(|e| RoomKeyError::Group(format!("decode a store value: {e}")))?;
values.insert(key, value);
}
}
let group_id = GroupId::from_slice(
&B64.decode(&snapshot.group_id)
.map_err(|e| RoomKeyError::Group(format!("decode the group id: {e}")))?,
);
let group = MlsGroup::load(provider.storage(), &group_id)
.map_err(|e| RoomKeyError::Group(format!("load the group: {e:?}")))?
.ok_or_else(|| {
RoomKeyError::Group("the snapshot's store holds no group at that id".into())
})?;
let public = B64
.decode(&snapshot.signature_public)
.map_err(|e| RoomKeyError::Group(format!("decode the signature key: {e}")))?;
let signer = SignatureKeyPair::read(
provider.storage(),
&public,
ROOM_CIPHERSUITE.signature_algorithm(),
)
.ok_or_else(|| {
RoomKeyError::Group("the snapshot's store holds no signature keypair".into())
})?;
let credential = Credential::new(
CredentialType::Basic,
snapshot.member_did.as_bytes().to_vec(),
);
Ok(Self {
group,
identity: RoomIdentity {
member_did: snapshot.member_did.clone(),
credential: CredentialWithKey {
credential,
signature_key: signer.public().into(),
},
signer,
},
provider,
})
}
}
#[cfg(test)]
mod custody_tests {
use super::*;
use crate::sealed::SealedRoom;
#[test]
fn a_restored_group_opens_what_the_original_sealed() {
let room_id = "did:key:zRoom";
let group = RoomGroup::create("did:key:zAlice").expect("create");
let sealed_room = SealedRoom::new(room_id, group);
let ciphertext = sealed_room
.seal_record("k1", 1, b"survives a restart")
.expect("seal");
let snapshot = sealed_room.group().snapshot().expect("snapshot");
drop(sealed_room);
let mut restored =
SealedRoom::new(room_id, RoomGroup::restore(&snapshot).expect("restore"));
let opened = restored.open_record("k1", 1, &ciphertext).expect("open");
assert_eq!(opened, b"survives a restart");
}
#[test]
fn a_snapshot_round_trips_through_json() {
let group = RoomGroup::create("did:key:zAlice").expect("create");
let snapshot = group.snapshot().expect("snapshot");
let wire = serde_json::to_string(&snapshot).expect("serialise");
let back: GroupSnapshot = serde_json::from_str(&wire).expect("deserialise");
let restored = RoomGroup::restore(&back).expect("restore");
assert_eq!(restored.epoch(), group.epoch());
assert_eq!(restored.member_count(), group.member_count());
}
#[test]
fn a_restored_group_is_at_the_epoch_it_was_snapshotted_at() {
let mut alice = RoomGroup::create("did:key:zAlice").expect("create");
let bob_provider = OpenMlsRustCrypto::default();
let bob = RoomIdentity::new("did:key:zBob", &bob_provider).expect("bob");
let package = bob.key_package(&bob_provider).expect("key package");
alice.add_member(package).expect("add");
assert_eq!(alice.epoch(), 1, "a membership change advanced the epoch");
let restored = RoomGroup::restore(&alice.snapshot().expect("snapshot")).expect("restore");
assert_eq!(restored.epoch(), 1);
assert_eq!(restored.member_count(), 2);
}
#[test]
fn a_minted_identity_joins_and_can_read_the_room() {
let room_id = "did:key:zRoom";
let (pending, package_bytes) =
IdentitySnapshot::mint("did:key:zBob").expect("mint an identity");
let mut owner = RoomGroup::create("did:key:zAlice").expect("create");
let change = owner
.add_member_from_bytes(&package_bytes)
.expect("add from the bytes that travelled");
let welcome = change.welcome.expect("adding produces a welcome");
let joined = RoomGroup::join_from_identity(&pending, &welcome).expect("join");
assert_eq!(joined.member_count(), 2);
let sealed = SealedRoom::new(room_id, owner)
.seal_record("k1", 1, b"for the new member")
.expect("seal");
let opened = SealedRoom::new(room_id, joined)
.open_record("k1", 1, &sealed)
.expect("the joiner reads what the owner sealed");
assert_eq!(opened, b"for the new member");
}
#[test]
fn a_snapshot_with_no_group_is_refused() {
let group = RoomGroup::create("did:key:zAlice").expect("create");
let mut snapshot = group.snapshot().expect("snapshot");
snapshot.entries.clear();
let Err(err) = RoomGroup::restore(&snapshot) else {
panic!("a snapshot with no group must not restore");
};
assert!(format!("{err}").contains("no group"), "{err}");
}
}
fn provider_from(entries: &[(String, String)]) -> Result<OpenMlsRustCrypto, RoomKeyError> {
let provider = OpenMlsRustCrypto::default();
{
let mut values = provider
.storage()
.values
.write()
.map_err(|_| RoomKeyError::Group("group store lock poisoned".into()))?;
for (k, v) in entries {
let key = B64
.decode(k)
.map_err(|e| RoomKeyError::Group(format!("decode a store key: {e}")))?;
let value = B64
.decode(v)
.map_err(|e| RoomKeyError::Group(format!("decode a store value: {e}")))?;
values.insert(key, value);
}
}
Ok(provider)
}
fn entries_of(provider: &OpenMlsRustCrypto) -> Result<Vec<(String, String)>, RoomKeyError> {
let values = provider
.storage()
.values
.read()
.map_err(|_| RoomKeyError::Group("group store lock poisoned".into()))?;
Ok(values
.iter()
.map(|(k, v)| (B64.encode(k), B64.encode(v)))
.collect())
}
impl IdentitySnapshot {
pub fn mint(member_did: &str) -> Result<(Self, Vec<u8>), RoomKeyError> {
let provider = OpenMlsRustCrypto::default();
let identity = RoomIdentity::new(member_did, &provider)?;
let package = identity.key_package(&provider)?;
let bytes = package
.tls_serialize_detached()
.map_err(|e| RoomKeyError::Group(format!("serialise the key package: {e:?}")))?;
Ok((
Self {
entries: entries_of(&provider)?,
member_did: member_did.to_string(),
signature_public: B64.encode(identity.signer.public()),
},
bytes,
))
}
pub fn member_did(&self) -> &str {
&self.member_did
}
}
impl RoomGroup {
pub fn add_member_from_bytes(
&mut self,
key_package: &[u8],
) -> Result<MembershipChange, RoomKeyError> {
let incoming = KeyPackageIn::tls_deserialize_exact(key_package)
.map_err(|e| RoomKeyError::Group(format!("parse the key package: {e:?}")))?;
let validated = incoming
.validate(self.provider.crypto(), ProtocolVersion::Mls10)
.map_err(|e| RoomKeyError::Group(format!("validate the key package: {e:?}")))?;
self.add_member(validated)
}
pub fn join_from_identity(
snapshot: &IdentitySnapshot,
welcome: &[u8],
) -> Result<Self, RoomKeyError> {
let provider = provider_from(&snapshot.entries)?;
let public = B64
.decode(&snapshot.signature_public)
.map_err(|e| RoomKeyError::Group(format!("decode the signature key: {e}")))?;
let signer = SignatureKeyPair::read(
provider.storage(),
&public,
ROOM_CIPHERSUITE.signature_algorithm(),
)
.ok_or_else(|| {
RoomKeyError::Group("the snapshot's store holds no signature keypair".into())
})?;
let credential = Credential::new(
CredentialType::Basic,
snapshot.member_did.as_bytes().to_vec(),
);
let identity = RoomIdentity {
member_did: snapshot.member_did.clone(),
credential: CredentialWithKey {
credential,
signature_key: signer.public().into(),
},
signer,
};
Self::join_with(identity, provider, welcome)
}
}