Skip to main content

gbp_mls/
lib.rs

1//! MLS (RFC 9420) integration for the Group Protocol Stack.
2//!
3//! This crate provides:
4//!
5//! * [`MlsContext`] — a member-side wrapper around an `openmls 0.8` group
6//!   (signing key, credential, provider, current group).
7//! * [`StreamLabel`] — labelled exporter constants used to derive AEAD keys
8//!   from the MLS exporter (`gbp/control`, `gbp/audio`, `gbp/text`,
9//!   `gbp/signal`).
10//! * `seal` / `open` — ChaCha20-Poly1305 AEAD with the labelled-exporter key.
11//!
12//! On every epoch change the old key material is invalidated automatically:
13//! the AEAD key is derived on the fly from `MlsGroup::export_secret`, never
14//! cached, and the previous epoch's secret becomes unreachable as soon as the
15//! group ratchets forward.
16
17#![deny(missing_docs)]
18
19use chacha20poly1305::{
20    ChaCha20Poly1305, Key, Nonce,
21    aead::{Aead, KeyInit},
22};
23use gbp_core::StreamType;
24use openmls::prelude::tls_codec::DeserializeBytes as _;
25use openmls::prelude::tls_codec::Serialize as _;
26use openmls::prelude::*;
27use openmls_basic_credential::SignatureKeyPair;
28use openmls_rust_crypto::{MemoryStorage, OpenMlsRustCrypto};
29use std::collections::HashMap;
30
31/// MLS ciphersuite used by the stack: X25519-AES128GCM-SHA256-Ed25519.
32pub const CIPHERSUITE: Ciphersuite = Ciphersuite::MLS_128_DHKEMX25519_AES128GCM_SHA256_Ed25519;
33
34/// Exporter label that binds the AEAD key to a stream class.
35#[derive(Copy, Clone, Debug, PartialEq, Eq)]
36pub enum StreamLabel {
37    /// `gbp/control` — control plane key.
38    Control,
39    /// `gbp/audio` — GAP key.
40    Audio,
41    /// `gbp/text` — GTP key.
42    Text,
43    /// `gbp/signal` — GSP key.
44    Signal,
45}
46
47impl StreamLabel {
48    /// Returns the stable string used as the `MlsGroup::export_secret` label.
49    pub fn as_str(self) -> &'static str {
50        match self {
51            Self::Control => "gbp/control",
52            Self::Audio => "gbp/audio",
53            Self::Text => "gbp/text",
54            Self::Signal => "gbp/signal",
55        }
56    }
57}
58
59/// Maps a [`StreamType`] to the corresponding [`StreamLabel`].
60pub fn label_for(st: StreamType) -> StreamLabel {
61    match st {
62        StreamType::Control => StreamLabel::Control,
63        StreamType::Audio => StreamLabel::Audio,
64        StreamType::Text => StreamLabel::Text,
65        StreamType::Signal => StreamLabel::Signal,
66    }
67}
68
69/// Categorises an MLS message processed via
70/// [`MlsContext::process_message`].
71#[derive(Debug, Clone, Copy, PartialEq, Eq)]
72pub enum ProcessedKind {
73    /// A Commit message was applied to the group; epoch advanced.
74    Commit,
75    /// An Application message was decrypted (not used by this stack — GBP
76    /// carries application data outside MLS application messages).
77    Application,
78    /// A Proposal-only message was staged.
79    Proposal,
80    /// An external message that did not advance the group.
81    External,
82}
83
84/// Errors raised by the MLS / AEAD layer.
85#[derive(Debug, thiserror::Error)]
86pub enum MlsError {
87    /// Any error returned by `openmls`, serialised as a string.
88    #[error("openmls: {0}")]
89    OpenMls(String),
90    /// AEAD seal or open failure.
91    #[error("aead: {0}")]
92    Aead(String),
93    /// A pending staged commit already exists — the previous transition must
94    /// be finalised or cleared before processing another commit.
95    #[error("transition in progress: pending staged commit exists")]
96    TransitionInProgress,
97}
98
99/// MLS context for a single group member.
100///
101/// Owns the OpenMLS provider, the signing key, the credential and the
102/// current `MlsGroup`. Ratcheting forward is performed by [`MlsContext::invite`]
103/// and [`MlsContext::accept_welcome`].
104pub struct MlsContext {
105    /// OpenMLS crypto provider.
106    pub provider: OpenMlsRustCrypto,
107    /// Signing key pair for this member.
108    pub signer: SignatureKeyPair,
109    /// Current MLS group.
110    pub group: MlsGroup,
111    /// Credential with the public signing key.
112    pub credential: CredentialWithKey,
113    /// Member identity (opaque application-defined bytes).
114    pub identity: Vec<u8>,
115    /// Staged commit produced by [`MlsContext::process_message`] but not
116    /// yet merged. Held until [`MlsContext::finalize_pending_commit`] (on
117    /// EXECUTE_TRANSITION) so that the local epoch only advances together
118    /// with the rest of the group, never earlier — otherwise this side's
119    /// READY frame would be sealed under an epoch the coordinator can't
120    /// open.
121    pub pending_staged: Option<StagedCommit>,
122}
123
124// ── Storage (de)serialisation for export_state / restore_state ────────────────
125// MemoryStorage exposes its key-value map as a public field; its built-in
126// serialize/deserialize are behind the `test-utils` feature, so we (de)serialise
127// the map ourselves with a simple length-prefixed (u32-LE) record format.
128
129fn serialize_storage(s: &MemoryStorage) -> Result<Vec<u8>, MlsError> {
130    let map = s
131        .values
132        .read()
133        .map_err(|_| MlsError::OpenMls("storage lock poisoned".into()))?;
134    let mut out = Vec::new();
135    out.extend_from_slice(&(map.len() as u32).to_le_bytes());
136    for (k, v) in map.iter() {
137        out.extend_from_slice(&(k.len() as u32).to_le_bytes());
138        out.extend_from_slice(k);
139        out.extend_from_slice(&(v.len() as u32).to_le_bytes());
140        out.extend_from_slice(v);
141    }
142    Ok(out)
143}
144
145fn deserialize_storage(bytes: &[u8]) -> Result<HashMap<Vec<u8>, Vec<u8>>, MlsError> {
146    let mut cur = bytes;
147    fn rd_u32(cur: &mut &[u8]) -> Result<usize, MlsError> {
148        if cur.len() < 4 {
149            return Err(MlsError::OpenMls("truncated storage blob".into()));
150        }
151        let n = u32::from_le_bytes([cur[0], cur[1], cur[2], cur[3]]) as usize;
152        *cur = &cur[4..];
153        Ok(n)
154    }
155    fn rd_bytes<'a>(cur: &mut &'a [u8], len: usize) -> Result<&'a [u8], MlsError> {
156        if cur.len() < len {
157            return Err(MlsError::OpenMls("truncated storage blob".into()));
158        }
159        let (head, tail) = cur.split_at(len);
160        *cur = tail;
161        Ok(head)
162    }
163    let count = rd_u32(&mut cur)?;
164    let mut map = HashMap::with_capacity(count);
165    for _ in 0..count {
166        let klen = rd_u32(&mut cur)?;
167        let k = rd_bytes(&mut cur, klen)?.to_vec();
168        let vlen = rd_u32(&mut cur)?;
169        let v = rd_bytes(&mut cur, vlen)?.to_vec();
170        map.insert(k, v);
171    }
172    Ok(map)
173}
174
175impl MlsContext {
176    /// Creates a new context with a single-member group, returning the
177    /// context together with a [`KeyPackageBundle`] that other members can
178    /// use to invite this one.
179    pub fn new_member(identity: &[u8]) -> Result<(Self, KeyPackageBundle), MlsError> {
180        let provider = OpenMlsRustCrypto::default();
181        let signer = SignatureKeyPair::new(CIPHERSUITE.signature_algorithm())
182            .map_err(|e| MlsError::OpenMls(format!("signer: {e:?}")))?;
183        signer
184            .store(provider.storage())
185            .map_err(|e| MlsError::OpenMls(format!("store signer: {e:?}")))?;
186
187        let credential = BasicCredential::new(identity.to_vec());
188        let credential_with_key = CredentialWithKey {
189            credential: credential.into(),
190            signature_key: signer.public().into(),
191        };
192
193        let kp_bundle = KeyPackage::builder()
194            .build(CIPHERSUITE, &provider, &signer, credential_with_key.clone())
195            .map_err(|e| MlsError::OpenMls(format!("kp: {e:?}")))?;
196
197        let cfg = MlsGroupCreateConfig::builder()
198            .ciphersuite(CIPHERSUITE)
199            .use_ratchet_tree_extension(true)
200            .build();
201        let group = MlsGroup::new(&provider, &signer, &cfg, credential_with_key.clone())
202            .map_err(|e| MlsError::OpenMls(format!("group: {e:?}")))?;
203
204        Ok((
205            Self {
206                provider,
207                signer,
208                group,
209                credential: credential_with_key,
210                identity: identity.to_vec(),
211                pending_staged: None,
212            },
213            kp_bundle,
214        ))
215    }
216
217    /// Result of [`MlsContext::invite_full`]: the Commit message that
218    /// existing members must apply via [`MlsContext::process_message`],
219    /// plus the Welcome that the new joiner must apply via
220    /// [`MlsContext::accept_welcome`].
221    ///
222    /// RFC 9420 §11/§12.4 — Welcome is for the joiner only; existing members
223    /// MUST receive the Commit to advance their epoch.
224    ///
225    /// IMPORTANT: this call **does not** merge the pending commit. The
226    /// caller MUST call [`MlsContext::finalize_pending_commit`] only after
227    /// they are confident the Commit/Welcome have been distributed (e.g.
228    /// the GBP coordinator has observed READY quorum). If the distribution
229    /// fails, call [`MlsContext::clear_pending_commit`] to roll back.
230    pub fn invite_full(
231        &mut self,
232        key_packages: &[KeyPackage],
233    ) -> Result<(Vec<u8>, Vec<u8>), MlsError> {
234        let (commit, welcome, _gi) = self
235            .group
236            .add_members(&self.provider, &self.signer, key_packages)
237            .map_err(|e| MlsError::OpenMls(format!("add_members: {e:?}")))?;
238        let commit_bytes = commit
239            .tls_serialize_detached()
240            .map_err(|e| MlsError::OpenMls(format!("commit serialize: {e:?}")))?;
241        let welcome_bytes = welcome
242            .tls_serialize_detached()
243            .map_err(|e| MlsError::OpenMls(format!("welcome serialize: {e:?}")))?;
244        Ok((commit_bytes, welcome_bytes))
245    }
246
247    /// Backwards-compatible wrapper. Builds the Commit, eagerly merges, and
248    /// returns only the Welcome bytes. Kept for callers that distribute the
249    /// Commit out-of-band and don't need atomic abort semantics.
250    pub fn invite(&mut self, key_packages: &[KeyPackage]) -> Result<Vec<u8>, MlsError> {
251        let (_commit, welcome) = self.invite_full(key_packages)?;
252        self.finalize_pending_commit()?;
253        Ok(welcome)
254    }
255
256    /// Removes members identified by their MLS LeafIndex via a Remove commit
257    /// and returns the TLS-serialised Commit message that remaining members
258    /// must apply via [`MlsContext::process_message`].
259    ///
260    /// Like [`MlsContext::invite_full`], the caller is responsible for
261    /// calling [`MlsContext::finalize_pending_commit`] after successful
262    /// distribution, or [`MlsContext::clear_pending_commit`] on failure.
263    /// RFC 9420 §12.3.
264    pub fn remove_members(&mut self, leaf_indices: &[u32]) -> Result<Vec<u8>, MlsError> {
265        // Validate indices against the current group size up front so the
266        // caller gets a clear error rather than an opaque openmls failure.
267        let group_size = self.group.members().count() as u32;
268        for &idx in leaf_indices {
269            if idx >= group_size {
270                return Err(MlsError::OpenMls(format!(
271                    "leaf_index {idx} out of range (group size {group_size})"
272                )));
273            }
274        }
275        let leaves: Vec<LeafNodeIndex> = leaf_indices
276            .iter()
277            .copied()
278            .map(LeafNodeIndex::new)
279            .collect();
280        let (commit, _welcome_opt, _gi) = self
281            .group
282            .remove_members(&self.provider, &self.signer, &leaves)
283            .map_err(|e| MlsError::OpenMls(format!("remove_members: {e:?}")))?;
284        commit
285            .tls_serialize_detached()
286            .map_err(|e| MlsError::OpenMls(format!("commit serialize: {e:?}")))
287    }
288
289    /// Merges any pending commit. Handles both:
290    /// * a self-issued commit produced by [`MlsContext::invite_full`] /
291    ///   [`MlsContext::remove_members`] (merged via `merge_pending_commit`);
292    /// * a staged commit deposited by [`MlsContext::process_message`]
293    ///   (merged via `merge_staged_commit`, consumed from
294    ///   [`MlsContext::pending_staged`]).
295    ///
296    /// Idempotent: if there is nothing to merge, returns Ok. Called from
297    /// the GBP control plane in response to `EXECUTE_TRANSITION`.
298    pub fn finalize_pending_commit(&mut self) -> Result<(), MlsError> {
299        if let Some(staged) = self.pending_staged.take() {
300            self.group
301                .merge_staged_commit(&self.provider, staged)
302                .map_err(|e| MlsError::OpenMls(format!("merge_staged: {e:?}")))?;
303        }
304        // merge_pending_commit errors if there's nothing to merge — for
305        // members that only received a commit (no self-issued one) that's
306        // expected, so swallow the error. Self-issued commits are merged
307        // via this path on the coordinator side.
308        let _ = self.group.merge_pending_commit(&self.provider);
309        Ok(())
310    }
311
312    /// Discards any pending commit (self-issued and/or staged) without
313    /// applying it. Used on `ABORT_TRANSITION`.
314    pub fn clear_pending_commit(&mut self) -> Result<(), MlsError> {
315        self.pending_staged = None;
316        self.group
317            .clear_pending_commit(self.provider.storage())
318            .map_err(|e| MlsError::OpenMls(format!("clear: {e:?}")))?;
319        Ok(())
320    }
321
322    /// Applies a Commit (or staged Proposal) message to the group. Existing
323    /// members invoke this after receiving the Commit broadcast embedded in
324    /// `PREPARE_TRANSITION` args.
325    ///
326    /// IMPORTANT: a Commit is staged but **not** merged here. It must be
327    /// merged via [`MlsContext::finalize_pending_commit`] in response to the
328    /// matching `EXECUTE_TRANSITION`, so that this side's MLS epoch
329    /// advances together with the rest of the group — never earlier.
330    /// Calling this twice without an intervening finalize/clear discards
331    /// the previously staged commit (the second call wins).
332    pub fn process_message(&mut self, msg_bytes: &[u8]) -> Result<ProcessedKind, MlsError> {
333        let msg_in = MlsMessageIn::tls_deserialize_exact_bytes(msg_bytes)
334            .map_err(|e| MlsError::OpenMls(format!("msg parse: {e:?}")))?;
335        let protocol_msg = match msg_in.extract() {
336            MlsMessageBodyIn::PublicMessage(m) => ProtocolMessage::from(m),
337            MlsMessageBodyIn::PrivateMessage(m) => ProtocolMessage::from(m),
338            other => {
339                return Err(MlsError::OpenMls(format!(
340                    "expected protocol message, got {other:?}"
341                )));
342            }
343        };
344        let processed = self
345            .group
346            .process_message(&self.provider, protocol_msg)
347            .map_err(|e| MlsError::OpenMls(format!("process: {e:?}")))?;
348        match processed.into_content() {
349            ProcessedMessageContent::StagedCommitMessage(staged) => {
350                if self.pending_staged.is_some() {
351                    return Err(MlsError::TransitionInProgress);
352                }
353                self.pending_staged = Some(*staged);
354                Ok(ProcessedKind::Commit)
355            }
356            ProcessedMessageContent::ApplicationMessage(_) => Ok(ProcessedKind::Application),
357            ProcessedMessageContent::ProposalMessage(_) => Ok(ProcessedKind::Proposal),
358            ProcessedMessageContent::ExternalJoinProposalMessage(_) => Ok(ProcessedKind::External),
359        }
360    }
361
362    /// Replaces the local group with the one described by the given
363    /// `Welcome` message.
364    pub fn accept_welcome(&mut self, welcome_bytes: &[u8]) -> Result<(), MlsError> {
365        let msg_in = MlsMessageIn::tls_deserialize_exact_bytes(welcome_bytes)
366            .map_err(|e| MlsError::OpenMls(format!("welcome parse: {e:?}")))?;
367        let welcome = match msg_in.extract() {
368            MlsMessageBodyIn::Welcome(w) => w,
369            other => {
370                return Err(MlsError::OpenMls(format!(
371                    "expected welcome, got {other:?}"
372                )));
373            }
374        };
375        let join_cfg = MlsGroupJoinConfig::builder()
376            .use_ratchet_tree_extension(true)
377            .build();
378        let staged = StagedWelcome::new_from_welcome(&self.provider, &join_cfg, welcome, None)
379            .map_err(|e| MlsError::OpenMls(format!("staged: {e:?}")))?;
380        self.group = staged
381            .into_group(&self.provider)
382            .map_err(|e| MlsError::OpenMls(format!("into_group: {e:?}")))?;
383        Ok(())
384    }
385
386    /// Returns the current group epoch.
387    pub fn epoch(&self) -> u64 {
388        self.group.epoch().as_u64()
389    }
390
391    /// Returns the 16-byte group identifier (truncated or zero-padded if the
392    /// underlying MLS group_id has a different length).
393    pub fn group_id_16(&self) -> [u8; 16] {
394        let raw = self.group.group_id().as_slice();
395        let mut out = [0u8; 16];
396        let n = raw.len().min(16);
397        out[..n].copy_from_slice(&raw[..n]);
398        out
399    }
400
401    /// Serialises the full local MLS state into an opaque blob that
402    /// [`MlsContext::restore_state`] can reconstruct verbatim. Lets a client
403    /// persist the context (disk / IndexedDB) so a chat survives a restart
404    /// without re-establishing the group — the basis for deterministic,
405    /// reload-surviving secret chats.
406    ///
407    /// The blob bundles four length-prefixed (u32-LE) sections:
408    /// `[provider storage | signer | identity | group_id]`. It contains
409    /// **private key material** — callers MUST store it encrypted at rest.
410    pub fn export_state(&self) -> Result<Vec<u8>, MlsError> {
411        let storage_buf = serialize_storage(self.provider.storage())?;
412        let signer_buf = self
413            .signer
414            .tls_serialize_detached()
415            .map_err(|e| MlsError::OpenMls(format!("signer serialize: {e:?}")))?;
416        let gid = self.group.group_id().as_slice().to_vec();
417
418        let mut out = Vec::with_capacity(
419            16 + storage_buf.len() + signer_buf.len() + self.identity.len() + gid.len(),
420        );
421        for part in [
422            storage_buf.as_slice(),
423            signer_buf.as_slice(),
424            self.identity.as_slice(),
425            gid.as_slice(),
426        ] {
427            out.extend_from_slice(&(part.len() as u32).to_le_bytes());
428            out.extend_from_slice(part);
429        }
430        Ok(out)
431    }
432
433    /// Reconstructs a context from a blob produced by
434    /// [`MlsContext::export_state`]. The restored context is at the same epoch
435    /// with the same group state, signer and identity, and can immediately
436    /// send / receive again.
437    pub fn restore_state(blob: &[u8]) -> Result<Self, MlsError> {
438        let mut cur = blob;
439        let mut take = || -> Result<&[u8], MlsError> {
440            if cur.len() < 4 {
441                return Err(MlsError::OpenMls("truncated state blob (length)".into()));
442            }
443            let len = u32::from_le_bytes([cur[0], cur[1], cur[2], cur[3]]) as usize;
444            cur = &cur[4..];
445            if cur.len() < len {
446                return Err(MlsError::OpenMls("truncated state blob (body)".into()));
447            }
448            let (head, tail) = cur.split_at(len);
449            cur = tail;
450            Ok(head)
451        };
452        let storage_bytes = take()?.to_vec();
453        let signer_bytes = take()?.to_vec();
454        let identity = take()?.to_vec();
455        let gid_bytes = take()?.to_vec();
456
457        // Rehydrate a fresh provider's (public) key-value map from the blob.
458        let provider = OpenMlsRustCrypto::default();
459        let map = deserialize_storage(&storage_bytes)?;
460        *provider
461            .storage()
462            .values
463            .write()
464            .map_err(|_| MlsError::OpenMls("storage lock poisoned".into()))? = map;
465
466        let signer = SignatureKeyPair::tls_deserialize_exact_bytes(&signer_bytes)
467            .map_err(|e| MlsError::OpenMls(format!("signer parse: {e:?}")))?;
468        let credential = CredentialWithKey {
469            credential: BasicCredential::new(identity.clone()).into(),
470            signature_key: signer.public().into(),
471        };
472        let group_id = GroupId::from_slice(&gid_bytes);
473        let group = MlsGroup::load(provider.storage(), &group_id)
474            .map_err(|e| MlsError::OpenMls(format!("group load: {e:?}")))?
475            .ok_or_else(|| MlsError::OpenMls("no group in restored state".into()))?;
476
477        Ok(Self {
478            provider,
479            signer,
480            group,
481            credential,
482            identity,
483            pending_staged: None,
484        })
485    }
486
487    /// Exports a 32-byte secret under the given stream label.
488    pub fn export_stream_key(&self, label: StreamLabel) -> Result<[u8; 32], MlsError> {
489        let secret = self
490            .group
491            .export_secret(self.provider.crypto(), label.as_str(), &[], 32)
492            .map_err(|e| MlsError::OpenMls(format!("export: {e:?}")))?;
493        let mut out = [0u8; 32];
494        out.copy_from_slice(&secret);
495        Ok(out)
496    }
497
498    /// Exports `len` bytes under an arbitrary `label` and `context`.
499    ///
500    /// Used by external crates (e.g. `hush-sframe`) that need custom KDF
501    /// labels without depending on OpenMLS directly.
502    pub fn export_raw(&self, label: &str, context: &[u8], len: usize) -> Result<Vec<u8>, MlsError> {
503        let secret = self
504            .group
505            .export_secret(self.provider.crypto(), label, context, len)
506            .map_err(|e| MlsError::OpenMls(format!("export_raw: {e:?}")))?;
507        Ok(secret.to_vec())
508    }
509
510    /// Encrypts `plaintext` with ChaCha20-Poly1305 using the stream-labelled
511    /// AEAD key and a nonce derived from the per-stream `seq`.
512    pub fn seal(
513        &self,
514        label: StreamLabel,
515        seq: u32,
516        plaintext: &[u8],
517    ) -> Result<Vec<u8>, MlsError> {
518        let key = self.export_stream_key(label)?;
519        let cipher = ChaCha20Poly1305::new(&Key::from(key));
520        let mut nonce = [0u8; 12];
521        nonce[..4].copy_from_slice(&seq.to_be_bytes());
522        cipher
523            .encrypt(&Nonce::from(nonce), plaintext)
524            .map_err(|e| MlsError::Aead(e.to_string()))
525    }
526
527    /// Decrypts `ciphertext` with the same parameters as [`MlsContext::seal`].
528    pub fn open(
529        &self,
530        label: StreamLabel,
531        seq: u32,
532        ciphertext: &[u8],
533    ) -> Result<Vec<u8>, MlsError> {
534        let key = self.export_stream_key(label)?;
535        let cipher = ChaCha20Poly1305::new(&Key::from(key));
536        let mut nonce = [0u8; 12];
537        nonce[..4].copy_from_slice(&seq.to_be_bytes());
538        cipher
539            .decrypt(&Nonce::from(nonce), ciphertext)
540            .map_err(|e| MlsError::Aead(e.to_string()))
541    }
542}
543
544#[cfg(test)]
545mod tests {
546    use super::*;
547
548    fn alice() -> (MlsContext, openmls::prelude::KeyPackageBundle) {
549        MlsContext::new_member(b"alice").unwrap()
550    }
551
552    fn bob() -> (MlsContext, openmls::prelude::KeyPackageBundle) {
553        MlsContext::new_member(b"bob").unwrap()
554    }
555
556    #[test]
557    fn stream_label_strings_are_correct() {
558        assert_eq!(StreamLabel::Control.as_str(), "gbp/control");
559        assert_eq!(StreamLabel::Audio.as_str(), "gbp/audio");
560        assert_eq!(StreamLabel::Text.as_str(), "gbp/text");
561        assert_eq!(StreamLabel::Signal.as_str(), "gbp/signal");
562    }
563
564    #[test]
565    fn label_for_maps_every_stream_type() {
566        assert_eq!(label_for(StreamType::Control), StreamLabel::Control);
567        assert_eq!(label_for(StreamType::Audio), StreamLabel::Audio);
568        assert_eq!(label_for(StreamType::Text), StreamLabel::Text);
569        assert_eq!(label_for(StreamType::Signal), StreamLabel::Signal);
570    }
571
572    #[test]
573    fn new_member_starts_at_epoch_zero() {
574        let (ctx, _kp) = alice();
575        assert_eq!(ctx.epoch(), 0);
576    }
577
578    #[test]
579    fn group_id_16_is_16_bytes() {
580        let (ctx, _kp) = alice();
581        let id = ctx.group_id_16();
582        assert_eq!(id.len(), 16);
583    }
584
585    #[test]
586    fn export_stream_key_is_32_bytes_and_stable() {
587        let (ctx, _kp) = alice();
588        let k1 = ctx.export_stream_key(StreamLabel::Text).unwrap();
589        let k2 = ctx.export_stream_key(StreamLabel::Text).unwrap();
590        assert_eq!(k1.len(), 32);
591        assert_eq!(k1, k2);
592    }
593
594    #[test]
595    fn different_labels_produce_different_keys() {
596        let (ctx, _kp) = alice();
597        let k_ctrl = ctx.export_stream_key(StreamLabel::Control).unwrap();
598        let k_text = ctx.export_stream_key(StreamLabel::Text).unwrap();
599        assert_ne!(k_ctrl, k_text);
600    }
601
602    #[test]
603    fn seal_open_single_member_round_trip() {
604        let (ctx, _kp) = alice();
605        let plaintext = b"hello world";
606        let ciphertext = ctx.seal(StreamLabel::Text, 1, plaintext).unwrap();
607        assert_ne!(ciphertext, plaintext);
608        let recovered = ctx.open(StreamLabel::Text, 1, &ciphertext).unwrap();
609        assert_eq!(recovered, plaintext);
610    }
611
612    #[test]
613    fn seal_wrong_seq_fails_to_open() {
614        let (ctx, _kp) = alice();
615        let ciphertext = ctx.seal(StreamLabel::Text, 1, b"secret").unwrap();
616        assert!(ctx.open(StreamLabel::Text, 2, &ciphertext).is_err());
617    }
618
619    #[test]
620    fn seal_wrong_label_fails_to_open() {
621        let (ctx, _kp) = alice();
622        let ciphertext = ctx.seal(StreamLabel::Text, 0, b"secret").unwrap();
623        assert!(ctx.open(StreamLabel::Audio, 0, &ciphertext).is_err());
624    }
625
626    #[test]
627    fn two_member_invite_and_welcome() {
628        let (mut alice, _akp) = alice();
629        let (mut bob, bob_kp) = bob();
630
631        let welcome = alice.invite(&[bob_kp.key_package().clone()]).unwrap();
632        // Alice's epoch advances after invite.
633        assert_eq!(alice.epoch(), 1);
634
635        bob.accept_welcome(&welcome).unwrap();
636        // Bob joins at epoch 1.
637        assert_eq!(bob.epoch(), 1);
638    }
639
640    #[test]
641    fn two_member_seal_open_cross_member() {
642        let (mut alice, _akp) = alice();
643        let (mut bob, bob_kp) = bob();
644
645        let welcome = alice.invite(&[bob_kp.key_package().clone()]).unwrap();
646        bob.accept_welcome(&welcome).unwrap();
647
648        let plaintext = b"cross-member secret";
649        let ct = alice.seal(StreamLabel::Control, 0, plaintext).unwrap();
650        let recovered = bob.open(StreamLabel::Control, 0, &ct).unwrap();
651        assert_eq!(recovered, plaintext);
652    }
653
654    #[test]
655    fn export_raw_returns_requested_length() {
656        let (ctx, _kp) = alice();
657        let raw = ctx.export_raw("test/label", b"ctx", 48).unwrap();
658        assert_eq!(raw.len(), 48);
659    }
660
661    #[test]
662    fn clear_pending_commit_is_idempotent() {
663        let (mut ctx, _kp) = alice();
664        ctx.clear_pending_commit().unwrap();
665        ctx.clear_pending_commit().unwrap();
666    }
667
668    #[test]
669    fn finalize_pending_commit_on_fresh_group_is_ok() {
670        let (mut ctx, _kp) = alice();
671        ctx.finalize_pending_commit().unwrap();
672    }
673
674    #[test]
675    fn invite_full_does_not_advance_epoch_until_finalize() {
676        let (mut alice, _akp) = alice();
677        let (_bob, bob_kp) = bob();
678
679        let (_commit, _welcome) = alice.invite_full(&[bob_kp.key_package().clone()]).unwrap();
680        // invite_full does NOT merge → epoch still 0
681        assert_eq!(alice.epoch(), 0);
682
683        alice.finalize_pending_commit().unwrap();
684        // after finalize → epoch 1
685        assert_eq!(alice.epoch(), 1);
686
687        // New members join via welcome, not via commit.
688        let (mut alice2, _akp2) = MlsContext::new_member(b"alice2").unwrap();
689        let (mut bob2, bob2_kp) = MlsContext::new_member(b"bob2").unwrap();
690        let (_commit_bytes, welcome_bytes) = alice2
691            .invite_full(&[bob2_kp.key_package().clone()])
692            .unwrap();
693        alice2.finalize_pending_commit().unwrap();
694        bob2.accept_welcome(&welcome_bytes).unwrap();
695        assert_eq!(alice2.epoch(), 1);
696        assert_eq!(bob2.epoch(), 1);
697    }
698
699    #[test]
700    fn export_restore_round_trip_preserves_state() {
701        let (ctx, _kp) = alice();
702        let blob = ctx.export_state().unwrap();
703        let restored = MlsContext::restore_state(&blob).unwrap();
704        assert_eq!(restored.epoch(), ctx.epoch());
705        assert_eq!(restored.group_id_16(), ctx.group_id_16());
706        // Identical exporter secret ⇒ the full group state was restored.
707        assert_eq!(
708            restored.export_stream_key(StreamLabel::Text).unwrap(),
709            ctx.export_stream_key(StreamLabel::Text).unwrap()
710        );
711    }
712
713    #[test]
714    fn restored_context_can_seal_and_open() {
715        let (ctx, _kp) = alice();
716        let blob = ctx.export_state().unwrap();
717        let restored = MlsContext::restore_state(&blob).unwrap();
718        let ct = restored
719            .seal(StreamLabel::Text, 7, b"after restore")
720            .unwrap();
721        assert_eq!(
722            restored.open(StreamLabel::Text, 7, &ct).unwrap(),
723            b"after restore"
724        );
725    }
726
727    #[test]
728    fn export_restore_preserves_multi_member_group() {
729        let (mut alice, _akp) = alice();
730        let (mut bob, bob_kp) = bob();
731        let welcome = alice.invite(&[bob_kp.key_package().clone()]).unwrap();
732        bob.accept_welcome(&welcome).unwrap();
733        assert_eq!(alice.epoch(), 1);
734
735        // Persist Alice at epoch 1, then restore from the blob.
736        let blob = alice.export_state().unwrap();
737        let restored_alice = MlsContext::restore_state(&blob).unwrap();
738        assert_eq!(restored_alice.epoch(), 1);
739
740        // Restored Alice still shares the group key with Bob.
741        let ct = restored_alice
742            .seal(StreamLabel::Control, 3, b"still in group")
743            .unwrap();
744        assert_eq!(
745            bob.open(StreamLabel::Control, 3, &ct).unwrap(),
746            b"still in group"
747        );
748    }
749
750    #[test]
751    fn multi_member_invite_one_welcome_serves_all_joiners() {
752        // A single Add commit for several KeyPackages yields ONE Welcome that
753        // every new member accepts with their own KeyPackage (RFC 9420 §12.4).
754        // This is what Hush secret groups rely on: claim N KeyPackages, one
755        // invite, broadcast one Welcome. (Existing tests only ever added one
756        // joiner at a time — this covers the multi-element slice.)
757        let (mut alice, _a) = alice();
758        let (mut bob, bob_kp) = bob();
759        let (mut carol, carol_kp) = MlsContext::new_member(b"carol").unwrap();
760
761        let welcome = alice
762            .invite(&[bob_kp.key_package().clone(), carol_kp.key_package().clone()])
763            .unwrap();
764        assert_eq!(alice.epoch(), 1, "one Add commit advances the epoch once");
765
766        // Both joiners accept the SAME Welcome and land at the same epoch.
767        bob.accept_welcome(&welcome).unwrap();
768        carol.accept_welcome(&welcome).unwrap();
769        assert_eq!(bob.epoch(), 1);
770        assert_eq!(carol.epoch(), 1);
771
772        // All three share the group key → mutual decryption.
773        let ct = alice.seal(StreamLabel::Text, 1, b"hello group").unwrap();
774        assert_eq!(bob.open(StreamLabel::Text, 1, &ct).unwrap(), b"hello group");
775        assert_eq!(
776            carol.open(StreamLabel::Text, 1, &ct).unwrap(),
777            b"hello group"
778        );
779    }
780
781    #[test]
782    fn restored_prekey_accepts_welcome() {
783        // A published KeyPackage's owner persists its context (export_state),
784        // then a fresh process restores it (restore_state) — the restored
785        // context MUST still accept a Welcome targeting that KeyPackage. This is
786        // the secret-DM reload path (Hush ADR-0023): the joiner's pre-key
787        // survives a reload (e.g. browser IndexedDB) and can still join.
788        let (mut alice, _akp) = alice();
789        let (bob, bob_kp) = bob();
790
791        // Persist bob's pre-key context, then drop the live one (simulate reload).
792        let bob_blob = bob.export_state().unwrap();
793        let bob_kp_inner = bob_kp.key_package().clone();
794        drop(bob);
795
796        // Alice invites bob's published KeyPackage.
797        let welcome = alice.invite(&[bob_kp_inner]).unwrap();
798        assert_eq!(alice.epoch(), 1);
799
800        // Bob restored from the blob accepts the Welcome — i.e. the private
801        // KeyPackage keys (init/encryption) survived export/restore.
802        let mut bob_restored = MlsContext::restore_state(&bob_blob).unwrap();
803        bob_restored.accept_welcome(&welcome).unwrap();
804        assert_eq!(bob_restored.epoch(), 1);
805
806        // Mutual decryption confirms the shared group.
807        let ct = alice.seal(StreamLabel::Text, 1, b"after reload").unwrap();
808        assert_eq!(
809            bob_restored.open(StreamLabel::Text, 1, &ct).unwrap(),
810            b"after reload"
811        );
812    }
813
814    #[test]
815    fn restore_state_rejects_truncated_blob() {
816        let (ctx, _kp) = alice();
817        let blob = ctx.export_state().unwrap();
818        assert!(MlsContext::restore_state(&blob[..blob.len() / 2]).is_err());
819        assert!(MlsContext::restore_state(&[]).is_err());
820    }
821}