1#![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
31pub const CIPHERSUITE: Ciphersuite = Ciphersuite::MLS_128_DHKEMX25519_AES128GCM_SHA256_Ed25519;
33
34#[derive(Copy, Clone, Debug, PartialEq, Eq)]
36pub enum StreamLabel {
37 Control,
39 Audio,
41 Text,
43 Signal,
45}
46
47impl StreamLabel {
48 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
59pub 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
72pub enum ProcessedKind {
73 Commit,
75 Application,
78 Proposal,
80 External,
82}
83
84#[derive(Debug, thiserror::Error)]
86pub enum MlsError {
87 #[error("openmls: {0}")]
89 OpenMls(String),
90 #[error("aead: {0}")]
92 Aead(String),
93 #[error("transition in progress: pending staged commit exists")]
96 TransitionInProgress,
97}
98
99pub struct MlsContext {
105 pub provider: OpenMlsRustCrypto,
107 pub signer: SignatureKeyPair,
109 pub group: MlsGroup,
111 pub credential: CredentialWithKey,
113 pub identity: Vec<u8>,
115 pub pending_staged: Option<StagedCommit>,
122}
123
124fn 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 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 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 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 pub fn remove_members(&mut self, leaf_indices: &[u32]) -> Result<Vec<u8>, MlsError> {
265 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 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 let _ = self.group.merge_pending_commit(&self.provider);
309 Ok(())
310 }
311
312 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 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 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 pub fn epoch(&self) -> u64 {
388 self.group.epoch().as_u64()
389 }
390
391 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 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 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 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 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 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 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 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 assert_eq!(alice.epoch(), 1);
634
635 bob.accept_welcome(&welcome).unwrap();
636 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 assert_eq!(alice.epoch(), 0);
682
683 alice.finalize_pending_commit().unwrap();
684 assert_eq!(alice.epoch(), 1);
686
687 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 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 let blob = alice.export_state().unwrap();
737 let restored_alice = MlsContext::restore_state(&blob).unwrap();
738 assert_eq!(restored_alice.epoch(), 1);
739
740 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 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 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 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 let (mut alice, _akp) = alice();
789 let (bob, bob_kp) = bob();
790
791 let bob_blob = bob.export_state().unwrap();
793 let bob_kp_inner = bob_kp.key_package().clone();
794 drop(bob);
795
796 let welcome = alice.invite(&[bob_kp_inner]).unwrap();
798 assert_eq!(alice.epoch(), 1);
799
800 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 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}