Skip to main content

core_crypto/mls/session/
epoch_observer.rs

1use std::sync::Arc;
2
3use async_trait::async_trait;
4
5use super::{Error, Result, Session};
6use crate::ConversationId;
7
8/// An `EpochObserver` is notified whenever a conversation's epoch changes.
9#[cfg_attr(target_os = "unknown", async_trait(?Send))]
10#[cfg_attr(not(target_os = "unknown"), async_trait)]
11pub trait EpochObserver: Send + Sync {
12    /// This function will be called every time a conversation's epoch changes.
13    ///
14    /// The `epoch` parameter is the new epoch.
15    ///
16    /// <div class="warning">
17    /// This function must not block! Foreign implementors of this inteface can
18    /// spawn a task indirecting the notification, or (unblocking) send the notification
19    /// on some kind of channel, or anything else, as long as the operation completes
20    /// quickly.
21    /// </div>
22    async fn epoch_changed(&self, conversation_id: ConversationId, epoch: u64);
23}
24
25impl Session {
26    /// Add an epoch observer to this session.
27    /// (see [EpochObserver]).
28    ///
29    /// This function should be called 0 or 1 times in a session's lifetime. If called
30    /// when an epoch observer already exists, this will return an error.
31    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    /// Notify the observer that the epoch has changed, if one is present.
41    pub(crate) async fn notify_epoch_changed(&self, conversation_id: ConversationId, epoch: u64) {
42        // Clone the handle out and release the lock before awaiting: the callback is foreign code,
43        // and `async_lock`'s `RwLock` is write-preferring, so holding a read guard across it lets a
44        // concurrent `register_epoch_observer` block every subsequent notification behind us.
45        // This also ensures we don't deadlock if a client's handler for the epoch event attempts to
46        // replace the epoch observer.
47        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            // trigger an epoch
76            let id = test_conv.advance_epoch().await.id;
77            session_context.transaction.finish().await.unwrap();
78
79            // ensure we have observed the epoch change
80            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            //  bob has the observer
101            let observer = TestEpochObserver::new();
102            bob.session()
103                .await
104                .register_epoch_observer(observer.clone())
105                .await
106                .unwrap();
107
108            // alice triggers an epoch
109            let id = test_conv.advance_epoch().await.id;
110            bob.transaction.finish().await.unwrap();
111
112            // ensure we have observed the epoch change
113            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}