Skip to main content

matrix_sdk/event_handler/
mod.rs

1// Copyright 2021 Jonas Platte
2// Copyright 2022 Famedly GmbH
3//
4// Licensed under the Apache License, Version 2.0 (the "License");
5// you may not use this file except in compliance with the License.
6// You may obtain a copy of the License at
7//
8//     http://www.apache.org/licenses/LICENSE-2.0
9//
10// Unless required by applicable law or agreed to in writing, software
11// distributed under the License is distributed on an "AS IS" BASIS,
12// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13// See the License for the specific language governing permissions and
14// limitations under the License.
15
16//! Types and traits related for event handlers. For usage, see
17//! [`Client::add_event_handler`].
18//!
19//! ### How it works
20//!
21//! The `add_event_handler` method registers event handlers of different
22//! signatures by actually storing boxed closures that all have the same
23//! signature of `async (EventHandlerData) -> ()` where `EventHandlerData` is a
24//! private type that contains all of the data an event handler *might* need.
25//!
26//! The stored closure takes care of deserializing the event which the
27//! `EventHandlerData` contains as a (borrowed) [`serde_json::value::RawValue`],
28//! extracting the context arguments from other fields of `EventHandlerData` and
29//! calling / `.await`ing the event handler if the previous steps succeeded.
30//! It also logs any errors from the above chain of function calls.
31//!
32//! For more details, see the [`EventHandler`] trait.
33
34#[cfg(any(feature = "anyhow", feature = "eyre"))]
35use std::any::TypeId;
36use std::{
37    borrow::Cow,
38    fmt,
39    future::Future,
40    pin::Pin,
41    sync::{
42        Arc, RwLock, Weak,
43        atomic::{AtomicU64, Ordering::SeqCst},
44    },
45    task::{Context, Poll},
46};
47
48#[cfg(target_family = "wasm")]
49use anymap2::any::CloneAny;
50#[cfg(not(target_family = "wasm"))]
51use anymap2::any::CloneAnySendSync;
52use eyeball::{SharedObservable, Subscriber};
53use futures_core::Stream;
54use futures_util::stream::{FuturesUnordered, StreamExt};
55use matrix_sdk_base::{
56    SendOutsideWasm, SyncOutsideWasm,
57    deserialized_responses::{EncryptionInfo, TimelineEvent},
58    sync::State,
59};
60use matrix_sdk_common::deserialized_responses::ProcessedToDeviceEvent;
61use pin_project_lite::pin_project;
62use ruma::{OwnedRoomId, events::BooleanType, push::Action, serde::Raw};
63use serde::{Deserialize, de::DeserializeOwned};
64use serde_json::value::RawValue as RawJsonValue;
65use tracing::{debug, error, field::debug, instrument, warn};
66
67use self::maps::EventHandlerMaps;
68use crate::{Client, Room};
69
70mod context;
71mod maps;
72mod static_events;
73
74pub use self::context::{Ctx, EventHandlerContext, RawEvent};
75
76#[cfg(not(target_family = "wasm"))]
77type EventHandlerFut = Pin<Box<dyn Future<Output = ()> + Send>>;
78#[cfg(target_family = "wasm")]
79type EventHandlerFut = Pin<Box<dyn Future<Output = ()>>>;
80
81#[cfg(not(target_family = "wasm"))]
82type EventHandlerFn = dyn Fn(EventHandlerData<'_>) -> EventHandlerFut + Send + Sync;
83#[cfg(target_family = "wasm")]
84type EventHandlerFn = dyn Fn(EventHandlerData<'_>) -> EventHandlerFut;
85
86#[cfg(not(target_family = "wasm"))]
87type AnyMap = anymap2::Map<dyn CloneAnySendSync + Send + Sync>;
88#[cfg(target_family = "wasm")]
89type AnyMap = anymap2::Map<dyn CloneAny>;
90
91#[derive(Default)]
92pub(crate) struct EventHandlerStore {
93    handlers: RwLock<EventHandlerMaps>,
94    context: RwLock<AnyMap>,
95    counter: AtomicU64,
96}
97
98impl EventHandlerStore {
99    pub fn add_handler(&self, handle: EventHandlerHandle, handler_fn: Box<EventHandlerFn>) {
100        self.handlers.write().unwrap().add(handle, handler_fn);
101    }
102
103    pub fn add_context<T>(&self, ctx: T)
104    where
105        T: Clone + Send + Sync + 'static,
106    {
107        self.context.write().unwrap().insert(ctx);
108    }
109
110    pub fn remove(&self, handle: EventHandlerHandle) {
111        self.handlers.write().unwrap().remove(handle);
112    }
113
114    #[cfg(test)]
115    fn len(&self) -> usize {
116        self.handlers.read().unwrap().len()
117    }
118}
119
120#[doc(hidden)]
121#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
122pub enum HandlerKind {
123    GlobalAccountData,
124    RoomAccountData,
125    EphemeralRoomData,
126    Timeline,
127    MessageLike,
128    OriginalMessageLike,
129    RedactedMessageLike,
130    State,
131    OriginalState,
132    RedactedState,
133    StrippedState,
134    ToDevice,
135    Presence,
136}
137
138impl HandlerKind {
139    fn message_like_redacted(redacted: bool) -> Self {
140        if redacted { Self::RedactedMessageLike } else { Self::OriginalMessageLike }
141    }
142
143    fn state_redacted(redacted: bool) -> Self {
144        if redacted { Self::RedactedState } else { Self::OriginalState }
145    }
146}
147
148/// A statically-known event kind/type that can be retrieved from an event sync.
149pub trait SyncEvent {
150    #[doc(hidden)]
151    const KIND: HandlerKind;
152    #[doc(hidden)]
153    const TYPE: Option<&'static str>;
154    #[doc(hidden)]
155    type IsPrefix: BooleanType;
156}
157
158pub(crate) struct EventHandlerWrapper {
159    handler_fn: Box<EventHandlerFn>,
160    pub handler_id: u64,
161}
162
163/// Handle to remove a registered event handler by passing it to
164/// [`Client::remove_event_handler`].
165#[derive(Clone, Debug)]
166pub struct EventHandlerHandle {
167    pub(crate) ev_kind: HandlerKind,
168    pub(crate) ev_type: Option<StaticEventTypePart>,
169    pub(crate) room_id: Option<OwnedRoomId>,
170    pub(crate) handler_id: u64,
171}
172
173/// The static part of an event type.
174#[derive(Clone, Copy, Debug)]
175pub(crate) enum StaticEventTypePart {
176    /// The full event type is static.
177    Full(&'static str),
178    /// Only the prefix of the event type is static.
179    Prefix(&'static str),
180}
181
182/// Interface for event handlers.
183///
184/// This trait is an abstraction for a certain kind of functions / closures,
185/// specifically:
186///
187/// - They must have at least one argument, which is the event itself, a type
188///   that implements [`SyncEvent`]. Any additional arguments need to implement
189///   the [`EventHandlerContext`] trait.
190/// - Their return type has to be one of: `()`, `Result<(), impl Display + Debug
191///   - 'static>` (if you are using `anyhow::Result` or `eyre::Result` you can
192///   additionally enable the `anyhow` / `eyre` feature to get the verbose
193///   `Debug` output printed on error)
194///
195/// ### How it works
196///
197/// This trait is basically a very constrained version of `Fn`: It requires at
198/// least one argument, which is represented as its own generic parameter `Ev`
199/// with the remaining parameter types being represented by the second generic
200/// parameter `Ctx`; they have to be stuffed into one generic parameter as a
201/// tuple because Rust doesn't have variadic generics.
202///
203/// `Ev` and `Ctx` are generic parameters rather than associated types because
204/// the argument list is a generic parameter for the `Fn` traits too, so a
205/// single type could implement `Fn` multiple times with different argument
206/// lists¹. Luckily, when calling [`Client::add_event_handler`] with a closure
207/// argument the trait solver takes into account that only a single one of the
208/// implementations applies (even though this could theoretically change through
209/// a dependency upgrade) and uses that rather than raising an ambiguity error.
210/// This is the same trick used by web frameworks like actix-web and axum.
211///
212/// ¹ the only thing stopping such types from existing in stable Rust is that
213/// all manual implementations of the `Fn` traits require a Nightly feature
214pub trait EventHandler<Ev, Ctx>: Clone + SendOutsideWasm + SyncOutsideWasm + 'static {
215    /// The future returned by `handle_event`.
216    #[doc(hidden)]
217    type Future: EventHandlerFuture;
218
219    /// Create a future for handling the given event.
220    ///
221    /// `data` provides additional data about the event, for example the room it
222    /// appeared in.
223    ///
224    /// Returns `None` if one of the context extractors failed.
225    #[doc(hidden)]
226    fn handle_event(self, ev: Ev, data: EventHandlerData<'_>) -> Option<Self::Future>;
227}
228
229#[doc(hidden)]
230pub trait EventHandlerFuture:
231    Future<Output = <Self as EventHandlerFuture>::Output> + SendOutsideWasm + 'static
232{
233    type Output: EventHandlerResult;
234}
235
236impl<T> EventHandlerFuture for T
237where
238    T: Future + SendOutsideWasm + 'static,
239    <T as Future>::Output: EventHandlerResult,
240{
241    type Output = <T as Future>::Output;
242}
243
244#[doc(hidden)]
245#[derive(Debug)]
246pub struct EventHandlerData<'a> {
247    client: Client,
248    room: Option<Room>,
249    raw: &'a RawJsonValue,
250    encryption_info: Option<&'a EncryptionInfo>,
251    push_actions: &'a [Action],
252    handle: EventHandlerHandle,
253}
254
255/// Return types supported for event handlers implement this trait.
256///
257/// It is not meant to be implemented outside of matrix-sdk.
258pub trait EventHandlerResult: Sized {
259    #[doc(hidden)]
260    fn print_error(&self, event_type: Option<&str>);
261}
262
263impl EventHandlerResult for () {
264    fn print_error(&self, _event_type: Option<&str>) {}
265}
266
267impl<E: fmt::Debug + fmt::Display + 'static> EventHandlerResult for Result<(), E> {
268    fn print_error(&self, event_type: Option<&str>) {
269        let msg_fragment = match event_type {
270            Some(event_type) => format!(" for `{event_type}`"),
271            None => "".to_owned(),
272        };
273
274        match self {
275            #[cfg(feature = "anyhow")]
276            Err(e) if TypeId::of::<E>() == TypeId::of::<anyhow::Error>() => {
277                error!("Event handler{msg_fragment} failed: {e:?}");
278            }
279            #[cfg(feature = "eyre")]
280            Err(e) if TypeId::of::<E>() == TypeId::of::<eyre::Report>() => {
281                error!("Event handler{msg_fragment} failed: {e:?}");
282            }
283            Err(e) => {
284                error!("Event handler{msg_fragment} failed: {e}");
285            }
286            Ok(_) => {}
287        }
288    }
289}
290
291#[derive(Deserialize)]
292struct UnsignedDetails {
293    redacted_because: Option<serde::de::IgnoredAny>,
294}
295
296/// Event handling internals.
297impl Client {
298    pub(crate) fn add_event_handler_impl<Ev, Ctx, H>(
299        &self,
300        handler: H,
301        room_id: Option<OwnedRoomId>,
302    ) -> EventHandlerHandle
303    where
304        Ev: SyncEvent + DeserializeOwned + SendOutsideWasm + 'static,
305        H: EventHandler<Ev, Ctx>,
306    {
307        let handler_fn: Box<EventHandlerFn> = Box::new(move |data| {
308            let maybe_fut = serde_json::from_str(data.raw.get())
309                .map(|ev| handler.clone().handle_event(ev, data));
310
311            Box::pin(async move {
312                match maybe_fut {
313                    Ok(Some(fut)) => {
314                        fut.await.print_error(Ev::TYPE);
315                    }
316                    Ok(None) => {
317                        error!(
318                            event_type = Ev::TYPE, event_kind = ?Ev::KIND,
319                            "Event handler has an invalid context argument",
320                        );
321                    }
322                    Err(e) => {
323                        warn!(
324                            event_type = Ev::TYPE, event_kind = ?Ev::KIND,
325                            "Failed to deserialize event, skipping event handler.\n
326                             Deserialization error: {e}",
327                        );
328                    }
329                }
330            })
331        });
332
333        let handler_id = self.inner.event_handlers.counter.fetch_add(1, SeqCst);
334        let ev_type = Ev::TYPE.map(|ev_type| {
335            if Ev::IsPrefix::as_bool() {
336                StaticEventTypePart::Prefix(ev_type)
337            } else {
338                StaticEventTypePart::Full(ev_type)
339            }
340        });
341        let handle = EventHandlerHandle { ev_kind: Ev::KIND, ev_type, room_id, handler_id };
342
343        self.inner.event_handlers.add_handler(handle.clone(), handler_fn);
344
345        handle
346    }
347
348    pub(crate) async fn handle_sync_events<T>(
349        &self,
350        kind: HandlerKind,
351        room: Option<&Room>,
352        events: &[Raw<T>],
353    ) -> serde_json::Result<()> {
354        #[derive(Deserialize)]
355        struct ExtractType<'a> {
356            #[serde(borrow, rename = "type")]
357            event_type: Cow<'a, str>,
358        }
359
360        for raw_event in events {
361            let event_type = raw_event.deserialize_as_unchecked::<ExtractType<'_>>()?.event_type;
362            self.call_event_handlers(room, raw_event.json(), kind, &event_type, None, &[]).await;
363        }
364
365        Ok(())
366    }
367
368    pub(crate) async fn handle_sync_to_device_events(
369        &self,
370        events: &[ProcessedToDeviceEvent],
371    ) -> serde_json::Result<()> {
372        #[derive(Deserialize)]
373        struct ExtractType<'a> {
374            #[serde(borrow, rename = "type")]
375            event_type: Cow<'a, str>,
376        }
377
378        for processed_to_device in events {
379            let (raw_event, encryption_info) = match processed_to_device {
380                ProcessedToDeviceEvent::Decrypted { raw, encryption_info } => {
381                    (raw, Some(encryption_info))
382                }
383                other => (&other.to_raw(), None),
384            };
385            let event_type = raw_event.deserialize_as_unchecked::<ExtractType<'_>>()?.event_type;
386            self.call_event_handlers(
387                None,
388                raw_event.json(),
389                HandlerKind::ToDevice,
390                &event_type,
391                encryption_info,
392                &[],
393            )
394            .await;
395        }
396
397        Ok(())
398    }
399
400    pub(crate) async fn handle_sync_state_events(
401        &self,
402        room: Option<&Room>,
403        state: &State,
404    ) -> serde_json::Result<()> {
405        #[derive(Deserialize)]
406        struct StateEventDetails<'a> {
407            #[serde(borrow, rename = "type")]
408            event_type: Cow<'a, str>,
409            unsigned: Option<UnsignedDetails>,
410        }
411
412        let state_events = match state {
413            State::Before(events) => events,
414            State::After(events) => events,
415        };
416
417        // Event handlers for possibly-redacted state events
418        self.handle_sync_events(HandlerKind::State, room, state_events).await?;
419
420        // Event handlers specifically for redacted OR unredacted state events
421        for raw_event in state_events {
422            let StateEventDetails { event_type, unsigned } =
423                raw_event.deserialize_as_unchecked()?;
424            let redacted = unsigned.and_then(|u| u.redacted_because).is_some();
425            let handler_kind = HandlerKind::state_redacted(redacted);
426
427            self.call_event_handlers(room, raw_event.json(), handler_kind, &event_type, None, &[])
428                .await;
429        }
430
431        Ok(())
432    }
433
434    pub(crate) async fn handle_sync_timeline_events(
435        &self,
436        room: Option<&Room>,
437        timeline_events: &[TimelineEvent],
438    ) -> serde_json::Result<()> {
439        #[derive(Deserialize)]
440        struct TimelineEventDetails<'a> {
441            #[serde(borrow, rename = "type")]
442            event_type: Cow<'a, str>,
443            state_key: Option<serde::de::IgnoredAny>,
444            unsigned: Option<UnsignedDetails>,
445        }
446
447        for item in timeline_events {
448            let TimelineEventDetails { event_type, state_key, unsigned } =
449                item.raw().deserialize_as_unchecked()?;
450
451            let redacted = unsigned.and_then(|u| u.redacted_because).is_some();
452            let (handler_kind_g, handler_kind_r) = match state_key {
453                Some(_) => (HandlerKind::State, HandlerKind::state_redacted(redacted)),
454                None => (HandlerKind::MessageLike, HandlerKind::message_like_redacted(redacted)),
455            };
456
457            let raw_event = item.raw().json();
458            let encryption_info = item.encryption_info().map(|i| &**i);
459            let push_actions = item.push_actions().unwrap_or(&[]);
460
461            // Event handlers for possibly-redacted timeline events
462            self.call_event_handlers(
463                room,
464                raw_event,
465                handler_kind_g,
466                &event_type,
467                encryption_info,
468                push_actions,
469            )
470            .await;
471
472            // Event handlers specifically for redacted OR unredacted timeline
473            // events
474            self.call_event_handlers(
475                room,
476                raw_event,
477                handler_kind_r,
478                &event_type,
479                encryption_info,
480                push_actions,
481            )
482            .await;
483
484            // Event handlers for `AnySyncTimelineEvent`
485            let kind = HandlerKind::Timeline;
486            self.call_event_handlers(
487                room,
488                raw_event,
489                kind,
490                &event_type,
491                encryption_info,
492                push_actions,
493            )
494            .await;
495        }
496
497        Ok(())
498    }
499
500    #[instrument(skip_all, fields(?event_kind, ?event_type, room_id))]
501    async fn call_event_handlers(
502        &self,
503        room: Option<&Room>,
504        raw: &RawJsonValue,
505        event_kind: HandlerKind,
506        event_type: &str,
507        encryption_info: Option<&EncryptionInfo>,
508        push_actions: &[Action],
509    ) {
510        let room_id = room.map(|r| r.room_id());
511        if let Some(room_id) = room_id {
512            tracing::Span::current().record("room_id", debug(room_id));
513        }
514
515        // Construct event handler futures
516        let mut futures: FuturesUnordered<_> = self
517            .inner
518            .event_handlers
519            .handlers
520            .read()
521            .unwrap()
522            .get_handlers(event_kind, event_type, room_id)
523            .map(|(handle, handler_fn)| {
524                let data = EventHandlerData {
525                    client: self.clone(),
526                    room: room.cloned(),
527                    raw,
528                    encryption_info,
529                    push_actions,
530                    handle,
531                };
532
533                (handler_fn)(data)
534            })
535            .collect();
536
537        if !futures.is_empty() {
538            debug!(amount = futures.len(), "Calling event handlers");
539
540            // Run the event handler futures with the
541            // `self.event_handlers.handlers` lock no longer being held.
542            while let Some(()) = futures.next().await {}
543        }
544    }
545}
546
547/// A guard type that removes an event handler when it drops (goes out of
548/// scope).
549///
550/// Created with [`Client::event_handler_drop_guard`].
551#[derive(Debug)]
552pub struct EventHandlerDropGuard {
553    handle: EventHandlerHandle,
554    client: Client,
555}
556
557impl EventHandlerDropGuard {
558    pub(crate) fn new(handle: EventHandlerHandle, client: Client) -> Self {
559        Self { handle, client }
560    }
561}
562
563impl Drop for EventHandlerDropGuard {
564    fn drop(&mut self) {
565        self.client.remove_event_handler(self.handle.clone());
566    }
567}
568
569macro_rules! impl_event_handler {
570    ($($ty:ident),* $(,)?) => {
571        impl<Ev, Fun, Fut, $($ty),*> EventHandler<Ev, ($($ty,)*)> for Fun
572        where
573            Ev: SyncEvent,
574            Fun: FnOnce(Ev, $($ty),*) -> Fut + Clone + SendOutsideWasm + SyncOutsideWasm + 'static,
575            Fut: EventHandlerFuture,
576            $($ty: EventHandlerContext),*
577        {
578            type Future = Fut;
579
580            fn handle_event(self, ev: Ev, _d: EventHandlerData<'_>) -> Option<Self::Future> {
581                Some((self)(ev, $($ty::from_data(&_d)?),*))
582            }
583        }
584    };
585}
586
587impl_event_handler!();
588impl_event_handler!(A);
589impl_event_handler!(A, B);
590impl_event_handler!(A, B, C);
591impl_event_handler!(A, B, C, D);
592impl_event_handler!(A, B, C, D, E);
593impl_event_handler!(A, B, C, D, E, F);
594impl_event_handler!(A, B, C, D, E, F, G);
595impl_event_handler!(A, B, C, D, E, F, G, H);
596
597/// An observer of events (may be tailored to a room).
598///
599/// Only the most recent value can be observed. Subscribers are notified when a
600/// new value is sent, but there is no guarantee that they will see all values.
601///
602/// To create such observer, use [`Client::observe_events`] or
603/// [`Client::observe_room_events`].
604#[derive(Debug)]
605pub struct ObservableEventHandler<T> {
606    /// This type is actually nothing more than a thin glue layer between the
607    /// [`EventHandler`] mechanism and the reactive programming types from
608    /// [`eyeball`]. Here, we use a [`SharedObservable`] that is updated by the
609    /// [`EventHandler`].
610    shared_observable: SharedObservable<Option<T>>,
611
612    /// This type owns the [`EventHandlerDropGuard`]. As soon as this type goes
613    /// out of scope, the event handler is unregistered/removed.
614    ///
615    /// [`EventHandlerSubscriber`] holds a weak, non-owning reference, to this
616    /// guard. It is useful to detect when to close the [`Stream`]: as soon as
617    /// this type goes out of scope, the subscriber will close itself on poll.
618    event_handler_guard: Arc<EventHandlerDropGuard>,
619}
620
621impl<T> ObservableEventHandler<T> {
622    pub(crate) fn new(
623        shared_observable: SharedObservable<Option<T>>,
624        event_handler_guard: EventHandlerDropGuard,
625    ) -> Self {
626        Self { shared_observable, event_handler_guard: Arc::new(event_handler_guard) }
627    }
628
629    /// Subscribe to this observer.
630    ///
631    /// It returns an [`EventHandlerSubscriber`], which implements [`Stream`].
632    /// See its documentation to learn more.
633    pub fn subscribe(&self) -> EventHandlerSubscriber<T> {
634        EventHandlerSubscriber::new(
635            self.shared_observable.subscribe(),
636            // The subscriber holds a weak non-owning reference to the event
637            // handler guard, so that it can detect when this observer is
638            // dropped, and can close the subscriber's stream.
639            Arc::downgrade(&self.event_handler_guard),
640        )
641    }
642}
643
644pin_project! {
645    /// The subscriber of an [`ObservableEventHandler`].
646    ///
647    /// To create such subscriber, use [`ObservableEventHandler::subscribe`].
648    ///
649    /// This type implements [`Stream`], which means it is possible to poll the
650    /// next value asynchronously. In other terms, polling this type will return
651    /// the new event as soon as they are synced. See [`Client::observe_events`]
652    /// to learn more.
653    #[derive(Debug)]
654    pub struct EventHandlerSubscriber<T> {
655        // The `Subscriber` associated to the `SharedObservable` inside
656        // `ObservableEventHandle`.
657        //
658        // Keep in mind all this API is just a thin glue layer between
659        // `EventHandle` and `SharedObservable`, that's… maagiic!
660        #[pin]
661        subscriber: Subscriber<Option<T>>,
662
663        // A weak non-owning reference to the event handler guard from
664        // `ObservableEventHandler`. When this type is polled (via its `Stream`
665        // implementation), it is possible to detect whether the observable has
666        // been dropped by upgrading this weak reference, and close the `Stream`
667        // if it needs to.
668        event_handler_guard: Weak<EventHandlerDropGuard>,
669    }
670}
671
672impl<T> EventHandlerSubscriber<T> {
673    fn new(
674        subscriber: Subscriber<Option<T>>,
675        event_handler_handle: Weak<EventHandlerDropGuard>,
676    ) -> Self {
677        Self { subscriber, event_handler_guard: event_handler_handle }
678    }
679}
680
681impl<T> Stream for EventHandlerSubscriber<T>
682where
683    T: Clone,
684{
685    type Item = T;
686
687    fn poll_next(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Option<Self::Item>> {
688        let mut this = self.project();
689
690        let Some(_) = this.event_handler_guard.upgrade() else {
691            // The `EventHandlerHandle` has been dropped via
692            // `EventHandlerDropGuard`. It means the `ObservableEventHandler`
693            // has been dropped. It's time to close this stream.
694            return Poll::Ready(None);
695        };
696
697        // First off, the subscriber is of type `Subscriber<Option<T>>` because
698        // the `SharedObservable` starts with a `None` value to indicate it has
699        // no yet received any update. We want the `Stream` to return `T`, not
700        // `Option<T>`. We then filter out all `None` value.
701        //
702        // Second, when a `None` value is met, we want to poll again (hence the
703        // `loop`). At best, there is a new value to return. At worst, the
704        // subscriber will return `Poll::Pending` and will register the wakers
705        // accordingly.
706
707        loop {
708            match this.subscriber.as_mut().poll_next(context) {
709                // Stream has been closed somehow.
710                Poll::Ready(None) => return Poll::Ready(None),
711
712                // The initial value (of the `SharedObservable` behind
713                // `self.subscriber`) has been polled. We want to filter it out.
714                Poll::Ready(Some(None)) => {
715                    // Loop over.
716                    continue;
717                }
718
719                // We have a new value!
720                Poll::Ready(Some(Some(value))) => return Poll::Ready(Some(value)),
721
722                // Classical pending.
723                Poll::Pending => return Poll::Pending,
724            }
725        }
726    }
727}
728
729#[cfg(test)]
730mod tests {
731    use matrix_sdk_test::{
732        DEFAULT_TEST_ROOM_ID, InvitedRoomBuilder, JoinedRoomBuilder, async_test,
733        event_factory::{EventFactory, PreviousMembership},
734    };
735    use serde::Serialize;
736    use stream_assert::{assert_closed, assert_pending, assert_ready};
737    #[cfg(target_family = "wasm")]
738    wasm_bindgen_test::wasm_bindgen_test_configure!(run_in_browser);
739    use std::{
740        future,
741        sync::{
742            Arc, LazyLock,
743            atomic::{AtomicU8, Ordering::SeqCst},
744        },
745    };
746
747    use matrix_sdk_common::{deserialized_responses::EncryptionInfo, locks::Mutex};
748    use matrix_sdk_test::SyncResponseBuilder;
749    use ruma::{
750        event_id,
751        events::{
752            AnySyncStateEvent, AnySyncTimelineEvent, AnyToDeviceEvent,
753            macros::EventContent,
754            room::{
755                member::{MembershipState, OriginalSyncRoomMemberEvent, StrippedRoomMemberEvent},
756                name::OriginalSyncRoomNameEvent,
757                power_levels::OriginalSyncRoomPowerLevelsEvent,
758            },
759            secret_storage::key::SecretStorageKeyEvent,
760            typing::SyncTypingEvent,
761        },
762        mxc_uri,
763        room::JoinRule,
764        room_id,
765        serde::Raw,
766        user_id,
767    };
768    use serde_json::json;
769    use strass::assert_let;
770
771    use crate::{
772        Client, Room,
773        event_handler::Ctx,
774        test_utils::{logged_in_client, no_retry_test_client},
775    };
776
777    static MEMBER_EVENT: LazyLock<Raw<AnySyncTimelineEvent>> = LazyLock::new(|| {
778        EventFactory::new()
779            .member(user_id!("@example:localhost"))
780            .membership(MembershipState::Join)
781            .display_name("example")
782            .event_id(event_id!("$151800140517rfvjc:localhost"))
783            .previous(PreviousMembership::new(MembershipState::Invite).display_name("example"))
784            .into()
785    });
786
787    #[async_test]
788    async fn test_add_event_handler() -> crate::Result<()> {
789        let client = logged_in_client(None).await;
790
791        let member_count = Arc::new(AtomicU8::new(0));
792        let typing_count = Arc::new(AtomicU8::new(0));
793        let power_levels_count = Arc::new(AtomicU8::new(0));
794        let invited_member_count = Arc::new(AtomicU8::new(0));
795
796        client.add_event_handler({
797            let member_count = member_count.clone();
798            move |_ev: OriginalSyncRoomMemberEvent, _room: Room| async move {
799                member_count.fetch_add(1, SeqCst);
800            }
801        });
802        client.add_event_handler({
803            let typing_count = typing_count.clone();
804            move |_ev: SyncTypingEvent| async move {
805                typing_count.fetch_add(1, SeqCst);
806            }
807        });
808        client.add_event_handler({
809            let power_levels_count = power_levels_count.clone();
810            move |_ev: OriginalSyncRoomPowerLevelsEvent, _client: Client, _room: Room| async move {
811                power_levels_count.fetch_add(1, SeqCst);
812            }
813        });
814        client.add_event_handler({
815            let invited_member_count = invited_member_count.clone();
816            move |_ev: StrippedRoomMemberEvent| async move {
817                invited_member_count.fetch_add(1, SeqCst);
818            }
819        });
820
821        let f = EventFactory::new().sender(user_id!("@example:localhost"));
822        let response = SyncResponseBuilder::default()
823            .add_joined_room(
824                JoinedRoomBuilder::default()
825                    .add_timeline_event(MEMBER_EVENT.clone())
826                    .add_typing(
827                        f.typing(vec![user_id!("@alice:matrix.org"), user_id!("@bob:example.com")]),
828                    )
829                    .add_state_event(f.default_power_levels()),
830            )
831            .add_invited_room(
832                InvitedRoomBuilder::new(room_id!("!test_invited:example.org")).add_state_event({
833                    let bob = user_id!("@bob:example.org");
834                    EventFactory::new()
835                        .sender(user_id!("@example:example.org"))
836                        .member(user_id!("@alice:example.org"))
837                        .membership(MembershipState::Invite)
838                        .display_name("Alice")
839                        .avatar_url(mxc_uri!("mxc://example.org/SEsfnsuifSDFSSEF"))
840                        .age(1234_i32)
841                        .invite_room_state(vec![
842                            Raw::from(f.room_name("Example Room").sender(bob)),
843                            Raw::from(f.room_join_rules(JoinRule::Invite).sender(bob)),
844                        ])
845                }),
846            )
847            .build_sync_response();
848        client.process_sync(response).await?;
849
850        assert_eq!(member_count.load(SeqCst), 1);
851        assert_eq!(typing_count.load(SeqCst), 1);
852        assert_eq!(power_levels_count.load(SeqCst), 1);
853        assert_eq!(invited_member_count.load(SeqCst), 1);
854
855        Ok(())
856    }
857
858    #[async_test]
859    async fn test_add_to_device_event_handler() -> crate::Result<()> {
860        let client = logged_in_client(None).await;
861
862        let captured_event: Arc<Mutex<Option<AnyToDeviceEvent>>> = Arc::new(Mutex::new(None));
863        let captured_info: Arc<Mutex<Option<EncryptionInfo>>> = Arc::new(Mutex::new(None));
864
865        client.add_event_handler({
866            let captured = captured_event.clone();
867            let captured_info = captured_info.clone();
868            move |ev: AnyToDeviceEvent, encryption_info: Option<EncryptionInfo>| {
869                let mut captured_lock = captured.lock();
870                *captured_lock = Some(ev);
871                let mut captured_info_lock = captured_info.lock();
872                *captured_info_lock = encryption_info;
873                future::ready(())
874            }
875        });
876
877        let response = SyncResponseBuilder::default()
878            .add_to_device_event(json!({
879              "sender": "@alice:example.com",
880              "type": "m.custom.to.device.type",
881              "content": {
882                "a": "test",
883              }
884            }))
885            .build_sync_response();
886        client.process_sync(response).await?;
887
888        let captured = captured_event.lock().clone();
889        assert_let!(Some(received_event) = captured);
890        assert_eq!(received_event.event_type().to_string(), "m.custom.to.device.type");
891        let info = captured_info.lock().clone();
892        assert!(info.is_none());
893        Ok(())
894    }
895
896    #[async_test]
897    async fn test_add_room_event_handler() -> crate::Result<()> {
898        let client = logged_in_client(None).await;
899
900        let room_id_a = room_id!("!foo:example.org");
901        let room_id_b = room_id!("!bar:matrix.org");
902
903        let member_count = Arc::new(AtomicU8::new(0));
904        let power_levels_count = Arc::new(AtomicU8::new(0));
905
906        // Room event handlers for member events in both rooms
907        client.add_room_event_handler(room_id_a, {
908            let member_count = member_count.clone();
909            move |_ev: OriginalSyncRoomMemberEvent, _room: Room| {
910                member_count.fetch_add(1, SeqCst);
911                future::ready(())
912            }
913        });
914        client.add_room_event_handler(room_id_b, {
915            let member_count = member_count.clone();
916            move |_ev: OriginalSyncRoomMemberEvent, _room: Room| {
917                member_count.fetch_add(1, SeqCst);
918                future::ready(())
919            }
920        });
921
922        // Power levels event handlers for member events in room A
923        client.add_room_event_handler(room_id_a, {
924            let power_levels_count = power_levels_count.clone();
925            move |_ev: OriginalSyncRoomPowerLevelsEvent, _client: Client, _room: Room| {
926                power_levels_count.fetch_add(1, SeqCst);
927                future::ready(())
928            }
929        });
930
931        // Room name event handler for room name events in room B
932        client.add_room_event_handler(
933            room_id_b,
934            // lint is buggy: rustc wants the explicit conversion from ! to ()
935            // here, but clippy thinks it's useless.
936            #[allow(clippy::unused_unit)]
937            async move |_ev: OriginalSyncRoomNameEvent| -> () {
938                unreachable!("No room event in room B")
939            },
940        );
941
942        let f = EventFactory::new().sender(user_id!("@example:localhost"));
943        let response = SyncResponseBuilder::default()
944            .add_joined_room(
945                JoinedRoomBuilder::new(room_id_a)
946                    .add_timeline_event(MEMBER_EVENT.clone())
947                    .add_state_event(f.default_power_levels())
948                    .add_state_event(f.room_name("room name")),
949            )
950            .add_joined_room(
951                JoinedRoomBuilder::new(room_id_b)
952                    .add_timeline_event(MEMBER_EVENT.clone())
953                    .add_state_event(f.default_power_levels()),
954            )
955            .build_sync_response();
956        client.process_sync(response).await?;
957
958        assert_eq!(member_count.load(SeqCst), 2);
959        assert_eq!(power_levels_count.load(SeqCst), 1);
960
961        Ok(())
962    }
963
964    #[async_test]
965    async fn test_add_event_handler_with_tuples() -> crate::Result<()> {
966        let client = logged_in_client(None).await;
967
968        client.add_event_handler(
969            |_ev: OriginalSyncRoomMemberEvent, (_room, _client): (Room, Client)| future::ready(()),
970        );
971
972        // If it compiles, it works. No need to assert anything.
973
974        Ok(())
975    }
976
977    #[async_test]
978    async fn test_remove_event_handler() -> crate::Result<()> {
979        let client = logged_in_client(None).await;
980
981        let member_count = Arc::new(AtomicU8::new(0));
982
983        client.add_event_handler({
984            let member_count = member_count.clone();
985            move |_ev: OriginalSyncRoomMemberEvent| async move {
986                member_count.fetch_add(1, SeqCst);
987            }
988        });
989
990        let handle_a = client.add_event_handler(
991            // lint is buggy: rustc wants the explicit conversion from ! to ()
992            // here, but clippy thinks it's useless.
993            #[allow(clippy::unused_unit)]
994            async move |_ev: OriginalSyncRoomMemberEvent| -> () {
995                panic!("handler should have been removed");
996            },
997        );
998        let handle_b = client.add_room_event_handler(
999            #[allow(unknown_lints, clippy::explicit_auto_deref)] // lint is buggy
1000            *DEFAULT_TEST_ROOM_ID,
1001            // lint is buggy: rustc wants the explicit conversion from ! to ()
1002            // here, but clippy thinks it's useless.
1003            #[allow(clippy::unused_unit)]
1004            async move |_ev: OriginalSyncRoomMemberEvent| -> () {
1005                panic!("handler should have been removed");
1006            },
1007        );
1008
1009        client.add_event_handler({
1010            let member_count = member_count.clone();
1011            move |_ev: OriginalSyncRoomMemberEvent| async move {
1012                member_count.fetch_add(1, SeqCst);
1013            }
1014        });
1015
1016        let response = SyncResponseBuilder::default()
1017            .add_joined_room(JoinedRoomBuilder::default().add_timeline_event(MEMBER_EVENT.clone()))
1018            .build_sync_response();
1019
1020        client.remove_event_handler(handle_a);
1021        client.remove_event_handler(handle_b);
1022
1023        client.process_sync(response).await?;
1024
1025        assert_eq!(member_count.load(SeqCst), 2);
1026
1027        Ok(())
1028    }
1029
1030    #[async_test]
1031    async fn test_event_handler_drop_guard() {
1032        let client = no_retry_test_client(None).await;
1033
1034        let handle = client.add_event_handler(|_ev: OriginalSyncRoomMemberEvent| async {});
1035        assert_eq!(client.inner.event_handlers.len(), 1);
1036
1037        {
1038            let _guard = client.event_handler_drop_guard(handle);
1039            assert_eq!(client.inner.event_handlers.len(), 1);
1040            // guard dropped here
1041        }
1042
1043        assert_eq!(client.inner.event_handlers.len(), 0);
1044    }
1045
1046    #[async_test]
1047    async fn test_use_client_in_handler() {
1048        // This used to not work because we were requiring `Send` of event
1049        // handler futures even on WASM, where practically all futures that do
1050        // I/O aren't.
1051        let client = no_retry_test_client(None).await;
1052
1053        client.add_event_handler(|_ev: OriginalSyncRoomMemberEvent, client: Client| async move {
1054            // All of Client's async methods that do network requests (and
1055            // possibly some that don't) are `!Send` on wasm. We obviously want
1056            // to be able to use them in event handlers.
1057            client
1058                .homeserver_capabilities()
1059                .refresh()
1060                .await
1061                .map_err(|e| anyhow::anyhow!("{}", e))?;
1062            anyhow::Ok(())
1063        });
1064    }
1065
1066    #[async_test]
1067    async fn test_raw_event_handler() -> crate::Result<()> {
1068        let client = logged_in_client(None).await;
1069        let counter = Arc::new(AtomicU8::new(0));
1070        client.add_event_handler_context(counter.clone());
1071        client.add_event_handler(
1072            |_ev: Raw<OriginalSyncRoomMemberEvent>, counter: Ctx<Arc<AtomicU8>>| async move {
1073                counter.fetch_add(1, SeqCst);
1074            },
1075        );
1076
1077        let response = SyncResponseBuilder::default()
1078            .add_joined_room(JoinedRoomBuilder::default().add_timeline_event(MEMBER_EVENT.clone()))
1079            .build_sync_response();
1080        client.process_sync(response).await?;
1081
1082        assert_eq!(counter.load(SeqCst), 1);
1083        Ok(())
1084    }
1085
1086    #[async_test]
1087    async fn test_enum_event_handler() -> crate::Result<()> {
1088        let client = logged_in_client(None).await;
1089        let counter = Arc::new(AtomicU8::new(0));
1090        client.add_event_handler_context(counter.clone());
1091        client.add_event_handler(
1092            |_ev: AnySyncStateEvent, counter: Ctx<Arc<AtomicU8>>| async move {
1093                counter.fetch_add(1, SeqCst);
1094            },
1095        );
1096
1097        let response = SyncResponseBuilder::default()
1098            .add_joined_room(JoinedRoomBuilder::default().add_timeline_event(MEMBER_EVENT.clone()))
1099            .build_sync_response();
1100        client.process_sync(response).await?;
1101
1102        assert_eq!(counter.load(SeqCst), 1);
1103        Ok(())
1104    }
1105
1106    #[async_test]
1107    async fn test_observe_events() -> crate::Result<()> {
1108        let client = logged_in_client(None).await;
1109
1110        let room_id_0 = room_id!("!r0.matrix.org");
1111        let room_id_1 = room_id!("!r1.matrix.org");
1112
1113        let observable = client.observe_events::<OriginalSyncRoomNameEvent, Room>();
1114
1115        let mut subscriber = observable.subscribe();
1116
1117        assert_pending!(subscriber);
1118
1119        let f = EventFactory::new().sender(user_id!("@mnt_io:matrix.org"));
1120        let mut response_builder = SyncResponseBuilder::new();
1121        let response = response_builder
1122            .add_joined_room(
1123                JoinedRoomBuilder::new(room_id_0)
1124                    .add_state_event(f.room_name("Name 0").event_id(event_id!("$ev0"))),
1125            )
1126            .build_sync_response();
1127        client.process_sync(response).await?;
1128
1129        let (room_name, room) = assert_ready!(subscriber);
1130
1131        assert_eq!(room_name.event_id.as_str(), "$ev0");
1132        assert_eq!(room.room_id(), room_id_0);
1133        assert_eq!(room.name().unwrap(), "Name 0");
1134
1135        assert_pending!(subscriber);
1136
1137        let response = response_builder
1138            .add_joined_room(
1139                JoinedRoomBuilder::new(room_id_1)
1140                    .add_state_event(f.room_name("Name 1").event_id(event_id!("$ev1"))),
1141            )
1142            .build_sync_response();
1143        client.process_sync(response).await?;
1144
1145        let (room_name, room) = assert_ready!(subscriber);
1146
1147        assert_eq!(room_name.event_id.as_str(), "$ev1");
1148        assert_eq!(room.room_id(), room_id_1);
1149        assert_eq!(room.name().unwrap(), "Name 1");
1150
1151        assert_pending!(subscriber);
1152
1153        drop(observable);
1154        assert_closed!(subscriber);
1155
1156        Ok(())
1157    }
1158
1159    #[async_test]
1160    async fn test_observe_room_events() -> crate::Result<()> {
1161        let client = logged_in_client(None).await;
1162
1163        let room_id = room_id!("!r0.matrix.org");
1164
1165        let observable_for_room =
1166            client.observe_room_events::<OriginalSyncRoomNameEvent, (Room, Client)>(room_id);
1167
1168        let mut subscriber_for_room = observable_for_room.subscribe();
1169
1170        assert_pending!(subscriber_for_room);
1171
1172        let f = EventFactory::new().sender(user_id!("@mnt_io:matrix.org"));
1173        let mut response_builder = SyncResponseBuilder::new();
1174        let response = response_builder
1175            .add_joined_room(
1176                JoinedRoomBuilder::new(room_id)
1177                    .add_state_event(f.room_name("Name 0").event_id(event_id!("$ev0"))),
1178            )
1179            .build_sync_response();
1180        client.process_sync(response).await?;
1181
1182        let (room_name, (room, _client)) = assert_ready!(subscriber_for_room);
1183
1184        assert_eq!(room_name.event_id.as_str(), "$ev0");
1185        assert_eq!(room.name().unwrap(), "Name 0");
1186
1187        assert_pending!(subscriber_for_room);
1188
1189        let response = response_builder
1190            .add_joined_room(
1191                JoinedRoomBuilder::new(room_id)
1192                    .add_state_event(f.room_name("Name 1").event_id(event_id!("$ev1"))),
1193            )
1194            .build_sync_response();
1195        client.process_sync(response).await?;
1196
1197        let (room_name, (room, _client)) = assert_ready!(subscriber_for_room);
1198
1199        assert_eq!(room_name.event_id.as_str(), "$ev1");
1200        assert_eq!(room.name().unwrap(), "Name 1");
1201
1202        assert_pending!(subscriber_for_room);
1203
1204        drop(observable_for_room);
1205        assert_closed!(subscriber_for_room);
1206
1207        Ok(())
1208    }
1209
1210    #[async_test]
1211    async fn test_observe_several_room_events() -> crate::Result<()> {
1212        let client = logged_in_client(None).await;
1213
1214        let room_id = room_id!("!r0.matrix.org");
1215
1216        let observable_for_room =
1217            client.observe_room_events::<OriginalSyncRoomNameEvent, (Room, Client)>(room_id);
1218
1219        let mut subscriber_for_room = observable_for_room.subscribe();
1220
1221        assert_pending!(subscriber_for_room);
1222
1223        let f = EventFactory::new().sender(user_id!("@mnt_io:matrix.org"));
1224        let mut response_builder = SyncResponseBuilder::new();
1225        let response = response_builder
1226            .add_joined_room(
1227                JoinedRoomBuilder::new(room_id)
1228                    .add_state_event(f.room_name("Name 0").event_id(event_id!("$ev0")))
1229                    .add_state_event(f.room_name("Name 1").event_id(event_id!("$ev1")))
1230                    .add_state_event(f.room_name("Name 2").event_id(event_id!("$ev2"))),
1231            )
1232            .build_sync_response();
1233        client.process_sync(response).await?;
1234
1235        let (room_name, (room, _client)) = assert_ready!(subscriber_for_room);
1236
1237        // Check we only get notified about the latest received event
1238        assert_eq!(room_name.event_id.as_str(), "$ev2");
1239        assert_eq!(room.name().unwrap(), "Name 2");
1240
1241        assert_pending!(subscriber_for_room);
1242
1243        drop(observable_for_room);
1244        assert_closed!(subscriber_for_room);
1245
1246        Ok(())
1247    }
1248
1249    #[async_test]
1250    async fn test_observe_events_with_type_prefix() -> crate::Result<()> {
1251        let client = logged_in_client(None).await;
1252
1253        let observable = client.observe_events::<SecretStorageKeyEvent, ()>();
1254
1255        let mut subscriber = observable.subscribe();
1256
1257        assert_pending!(subscriber);
1258
1259        let mut response_builder = SyncResponseBuilder::new();
1260        let response = response_builder
1261            .add_custom_global_account_data(json!({
1262                "content": {
1263                    "algorithm": "m.secret_storage.v1.aes-hmac-sha2",
1264                    "iv": "gH2iNpiETFhApvW6/FFEJQ",
1265                    "mac": "9Lw12m5SKDipNghdQXKjgpfdj1/K7HFI2brO+UWAGoM",
1266                    "passphrase": {
1267                        "algorithm": "m.pbkdf2",
1268                        "salt": "IuLnH7S85YtZmkkBJKwNUKxWF42g9O1H",
1269                        "iterations": 10,
1270                    },
1271                },
1272                "type": "m.secret_storage.key.foobar",
1273            }))
1274            .build_sync_response();
1275        client.process_sync(response).await?;
1276
1277        let (secret_storage_key, ()) = assert_ready!(subscriber);
1278
1279        assert_eq!(secret_storage_key.content.key_id, "foobar");
1280
1281        assert_pending!(subscriber);
1282
1283        drop(observable);
1284        assert_closed!(subscriber);
1285
1286        Ok(())
1287    }
1288
1289    #[async_test]
1290    async fn test_observe_room_events_with_type_prefix() -> crate::Result<()> {
1291        // To create an event handler for a room account data event type with
1292        // prefix, we need to create a custom event type, none exist in the
1293        // Matrix specification yet.
1294        #[derive(Debug, Clone, EventContent, Serialize)]
1295        #[ruma_event(type = "fake.event.*", kind = RoomAccountData)]
1296        struct AccountDataWithPrefixEventContent {
1297            #[ruma_event(type_fragment)]
1298            #[serde(skip)]
1299            key_id: String,
1300        }
1301
1302        let room_id = room_id!("!r0.matrix.org");
1303        let client = logged_in_client(None).await;
1304
1305        let observable = client.observe_room_events::<AccountDataWithPrefixEvent, Room>(room_id);
1306
1307        let mut subscriber = observable.subscribe();
1308
1309        assert_pending!(subscriber);
1310
1311        let mut response_builder = SyncResponseBuilder::new();
1312        let response = response_builder
1313            .add_joined_room(
1314                JoinedRoomBuilder::new(room_id).add_account_data_bulk([Raw::new(&json!({
1315                    "content": {},
1316                    "type": "fake.event.foobar",
1317                }))
1318                .unwrap()
1319                .cast_unchecked()]),
1320            )
1321            .build_sync_response();
1322        client.process_sync(response).await?;
1323
1324        let (secret_storage_key, _room) = assert_ready!(subscriber);
1325
1326        assert_eq!(secret_storage_key.content.key_id, "foobar");
1327
1328        assert_pending!(subscriber);
1329
1330        drop(observable);
1331        assert_closed!(subscriber);
1332
1333        Ok(())
1334    }
1335}