core_crypto/mls/session/
epoch_observer.rs1use std::sync::Arc;
2
3use async_trait::async_trait;
4
5use super::{Error, Result, Session};
6use crate::ConversationId;
7
8#[cfg_attr(target_os = "unknown", async_trait(?Send))]
10#[cfg_attr(not(target_os = "unknown"), async_trait)]
11pub trait EpochObserver: Send + Sync {
12 async fn epoch_changed(&self, conversation_id: ConversationId, epoch: u64);
23}
24
25impl Session {
26 pub async fn register_epoch_observer(&self, epoch_observer: Arc<dyn EpochObserver>) -> Result<()> {
32 let mut observer_guard = self.epoch_observer.write().await;
33 if observer_guard.is_some() {
34 return Err(Error::EpochObserverAlreadyExists);
35 }
36 observer_guard.replace(epoch_observer);
37 Ok(())
38 }
39
40 pub(crate) async fn notify_epoch_changed(&self, conversation_id: ConversationId, epoch: u64) {
42 let observer = self.epoch_observer.read().await.clone();
48 if let Some(observer) = observer {
49 observer.epoch_changed(conversation_id, epoch).await;
50 }
51 }
52}
53
54#[cfg(test)]
55mod tests {
56 use rstest::rstest;
57 use rstest_reuse::apply;
58
59 use crate::test_utils::{TestContext, TestEpochObserver, all_cred_cipher};
60
61 #[apply(all_cred_cipher)]
62 pub async fn observe_local_epoch_change(case: TestContext) {
63 let [session_context] = case.sessions().await;
64 Box::pin(async move {
65 let test_conv = case.create_conversation([&session_context]).await;
66
67 let observer = TestEpochObserver::new();
68 session_context
69 .session()
70 .await
71 .register_epoch_observer(observer.clone())
72 .await
73 .unwrap();
74
75 let id = test_conv.advance_epoch().await.id;
77 session_context.transaction.finish().await.unwrap();
78
79 let observed_epochs = observer.observed_epochs().await;
81 assert_eq!(
82 observed_epochs.len(),
83 1,
84 "we triggered exactly one epoch change and so should observe one epoch change"
85 );
86 assert_eq!(
87 observed_epochs[0].0, id,
88 "conversation id of observed epoch change must match"
89 );
90 })
91 .await
92 }
93
94 #[apply(all_cred_cipher)]
95 pub async fn observe_remote_epoch_change(case: TestContext) {
96 let [alice, bob] = case.sessions().await;
97 Box::pin(async move {
98 let test_conv = case.create_conversation([&alice, &bob]).await;
99
100 let observer = TestEpochObserver::new();
102 bob.session()
103 .await
104 .register_epoch_observer(observer.clone())
105 .await
106 .unwrap();
107
108 let id = test_conv.advance_epoch().await.id;
110 bob.transaction.finish().await.unwrap();
111
112 let observed_epochs = observer.observed_epochs().await;
114 assert_eq!(
115 observed_epochs.len(),
116 1,
117 "we triggered exactly one epoch change and so should observe one epoch change"
118 );
119 assert_eq!(
120 observed_epochs[0].0, id,
121 "conversation id of observed epoch change must match"
122 );
123 })
124 .await
125 }
126}