diff --git a/.editorconfig b/.editorconfig new file mode 100644 index 000000000..dddb23dbe --- /dev/null +++ b/.editorconfig @@ -0,0 +1,6 @@ +root = true + +# Unix-style newlines with a newline ending every file +[*] +end_of_line = lf +insert_final_newline = true diff --git a/.github/workflows/documentation.yml b/.github/workflows/documentation.yml index 9a1642dd8..b0da847cd 100644 --- a/.github/workflows/documentation.yml +++ b/.github/workflows/documentation.yml @@ -22,11 +22,16 @@ jobs: toolchain: nightly override: true + - name: Install Node.js + uses: actions/setup-node@v3 + with: + node-version: 18 + - name: Load cache uses: Swatinem/rust-cache@v1 # Keep in sync with xtask docs - - name: Build documentation + - name: Build rust documentation uses: actions-rs/cargo@v1 env: # Work around https://github.com/rust-lang/cargo/issues/10744 @@ -36,10 +41,31 @@ jobs: command: doc args: --no-deps --workspace --features docsrs + - name: Build `matrix-sdk-crypto-nodejs` doc + run: | + cd bindings/matrix-sdk-crypto-nodejs + npm install + npm run build && npm run doc + + - name: Build `matrix-sdk-crypto-js` doc + run: | + cd bindings/matrix-sdk-crypto-js + npm install + npm run build && npm run doc + + - name: Prepare the doc hierarchy + shell: bash + run: | + mkdir -p doc/bindings/matrix-sdk-crypto-nodejs/ + mkdir -p doc/bindings/matrix-sdk-crypto-js/ + mv target/doc/* doc/ + mv bindings/matrix-sdk-crypto-nodejs/docs/* doc/bindings/matrix-sdk-crypto-nodejs/ + mv bindings/matrix-sdk-crypto-js/docs/* doc/bindings/matrix-sdk-crypto-js/ + - name: Deploy documentation if: github.event_name == 'push' && github.ref == 'refs/heads/main' uses: peaceiris/actions-gh-pages@v3 with: github_token: ${{ secrets.GITHUB_TOKEN }} - publish_dir: ./target/doc/ + publish_dir: ./doc/ force_orphan: true diff --git a/crates/matrix-sdk-common/Cargo.toml b/crates/matrix-sdk-common/Cargo.toml index 8c8d49de0..d6b4d674b 100644 --- a/crates/matrix-sdk-common/Cargo.toml +++ b/crates/matrix-sdk-common/Cargo.toml @@ -16,6 +16,7 @@ default-target = "x86_64-unknown-linux-gnu" targets = ["x86_64-unknown-linux-gnu", "wasm32-unknown-unknown"] [dependencies] +futures-core = "0.3.21" ruma = { git = "https://github.com/ruma/ruma", rev = "ca8c66c885241a7ba3805399604eda4a38979f6b", features = ["client-api-c"] } serde = "1.0.136" @@ -24,7 +25,12 @@ async-lock = "2.5.0" instant = { version = "0.1.12", features = ["wasm-bindgen", "inaccurate"] } futures-util = { version = "0.3.21", default-features = false, features = ["channel"] } wasm-bindgen-futures = "0.4.30" +wasm-timer = "0.2.5" [target.'cfg(not(target_arch = "wasm32"))'.dependencies] -tokio = { version = "1.17.0", default-features = false, features = ["rt", "sync"] } +tokio = { version = "1.17.0", default-features = false, features = ["rt", "sync", "time"] } instant = { version = "0.1.12", features = ["now"] } + +[dev-dependencies] +matrix-sdk-test = { path = "../matrix-sdk-test/", version="0.5.0" } +wasm-bindgen-test = "0.3.30" diff --git a/crates/matrix-sdk-common/src/lib.rs b/crates/matrix-sdk-common/src/lib.rs index 2d17609eb..0f5213499 100644 --- a/crates/matrix-sdk-common/src/lib.rs +++ b/crates/matrix-sdk-common/src/lib.rs @@ -6,6 +6,7 @@ pub use instant; pub mod deserialized_responses; pub mod executor; pub mod locks; +pub mod timeout; /// Super trait that is used for our store traits, this trait will differ if /// it's used on WASM. WASM targets will not require `Send` and `Sync` to have diff --git a/crates/matrix-sdk-common/src/timeout.rs b/crates/matrix-sdk-common/src/timeout.rs new file mode 100644 index 000000000..e34d9d825 --- /dev/null +++ b/crates/matrix-sdk-common/src/timeout.rs @@ -0,0 +1,97 @@ +use std::{error::Error, fmt, time::Duration}; + +use futures_core::Future; +#[cfg(not(target_arch = "wasm32"))] +use tokio::time::timeout as tokio_timeout; +#[cfg(target_arch = "wasm32")] +use wasm_timer::ext::TryFutureExt; + +/// Error type notifying that a timeout has elapsed. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ElapsedError(()); + +impl fmt::Display for ElapsedError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "time waiting for future has elapsed!") + } +} + +impl Error for ElapsedError {} + +impl From for ElapsedError { + fn from(io_error: std::io::Error) -> Self { + io_error.into() + } +} + +/// Wait for `future` to be completed. `future` needs to return +/// a `Result`. +/// +/// If the given timeout has elapsed the method will stop waiting and return +/// an error. +pub async fn timeout(future: F, duration: Duration) -> Result +where + F: Future, +{ + #[cfg(not(target_arch = "wasm32"))] + return tokio_timeout(duration, future).await.map_err(|_| ElapsedError(())); + + #[cfg(target_arch = "wasm32")] + { + let try_future = async { + let output = future.await; + + // Contrary to clippy's note, the qualification of the Result is necessary + #[allow(unused_qualifications)] + Result::::Ok(output) + }; + + return try_future.timeout(duration).await; + } +} + +// TODO: Enable tests for wasm32 and debug why +// `with_timeout` test fails https://github.com/matrix-org/matrix-rust-sdk/issues/896 +#[cfg(all(test, not(target_arch = "wasm32")))] +pub(crate) mod tests { + use matrix_sdk_test::async_test; + #[cfg(not(target_arch = "wasm32"))] + use tokio::time::{sleep, Duration}; + #[cfg(target_arch = "wasm32")] + use {std::time::Duration, wasm_timer::Delay}; + + use super::timeout; + + #[cfg(target_arch = "wasm32")] + wasm_bindgen_test::wasm_bindgen_test_configure!(run_in_browser); + + #[async_test] + async fn without_timeout() { + let duration_future = Duration::from_millis(500); + let duration_timeout = Duration::from_millis(800); + + #[cfg(not(target_arch = "wasm32"))] + let fut = sleep(duration_future); + #[cfg(target_arch = "wasm32")] + let fut = Delay::new(duration_future); + + let res = timeout(fut, duration_timeout).await; + + res.expect("future should have completed without ElapsedError"); + } + + #[async_test] + async fn with_timeout() { + let duration_future = Duration::from_millis(900); + let duration_timeout = Duration::from_millis(800); + + #[cfg(not(target_arch = "wasm32"))] + let fut = sleep(duration_future); + #[cfg(target_arch = "wasm32")] + let fut = Delay::new(duration_future); + + let res = timeout(fut, duration_timeout).await; + + res.expect_err("future should throw an ElapsedError"); + } +} diff --git a/crates/matrix-sdk-crypto/src/gossiping/machine.rs b/crates/matrix-sdk-crypto/src/gossiping/machine.rs index 96053debc..204eeabbc 100644 --- a/crates/matrix-sdk-crypto/src/gossiping/machine.rs +++ b/crates/matrix-sdk-crypto/src/gossiping/machine.rs @@ -743,7 +743,7 @@ impl GossipMachine { /// Mark the given outgoing key info as done. /// /// This will queue up a request cancellation. - async fn mark_as_done(&self, key_info: GossipRequest) -> Result<(), CryptoStoreError> { + async fn mark_as_done(&self, key_info: &GossipRequest) -> Result<(), CryptoStoreError> { trace!( recipient = key_info.request_recipient.as_str(), request_type = key_info.request_type(), @@ -754,7 +754,7 @@ impl GossipMachine { self.outgoing_requests.remove(&key_info.request_id); // TODO return the key info instead of deleting it so the sync handler // can delete it in one transaction. - self.delete_key_info(&key_info).await?; + self.delete_key_info(key_info).await?; let request = key_info.to_cancellation(self.device_id()); self.outgoing_requests.insert(request.request_id.clone(), request); @@ -762,7 +762,90 @@ impl GossipMachine { Ok(()) } - pub async fn receive_secret( + async fn accept_secret( + &self, + event: &mut SecretSendEvent, + request: &GossipRequest, + secret_name: &SecretName, + ) -> Result<(), CryptoStoreError> { + if secret_name != &SecretName::RecoveryKey { + match self.store.import_secret(secret_name, &event.content.secret).await { + Ok(_) => self.mark_as_done(request).await?, + Err(e) => { + // If this is a store error propagate it up + // the call stack. + if let SecretImportError::Store(e) = e { + return Err(e); + } else { + // Otherwise warn that there was + // something wrong with the secret. + warn!( + secret_name = secret_name.as_ref(), + error = ?e, + "Error while importing a secret" + ) + } + } + } + } else { + // Skip importing the recovery key here since + // we'll want to check if the public key matches + // to the latest version on the server. The key + // will not be zeroized and + // instead leave the key in the event and let + // the user import it later. + } + + Ok(()) + } + + async fn receive_secret( + &self, + sender_key: &str, + event: &mut SecretSendEvent, + request: &GossipRequest, + secret_name: &SecretName, + ) -> Result<(), CryptoStoreError> { + // Set the secret name so other consumers of the event know + // what this event is about. + event.content.secret_name = Some(secret_name.to_owned()); + + debug!( + sender = event.sender.as_str(), + request_id = event.content.request_id.as_str(), + secret_name = secret_name.as_ref(), + "Received a m.secret.send event with a matching request" + ); + + if let Some(device) = + self.store.get_device_from_curve_key(&event.sender, sender_key).await? + { + // Only accept secrets from one of our own trusted devices. + if device.user_id() == self.user_id() && device.verified() { + self.accept_secret(event, request, secret_name).await?; + } else { + warn!( + sender = event.sender.as_str(), + request_id = event.content.request_id.as_str(), + secret_name = secret_name.as_ref(), + "Received a m.secret.send event from another user or from \ + unverified device" + ); + } + } else { + warn!( + sender = event.sender.as_str(), + request_id = event.content.request_id.as_str(), + secret_name = secret_name.as_ref(), + "Received a m.secret.send event from an unknown device" + ); + self.store.update_tracked_user(&event.sender, true).await?; + } + + Ok(()) + } + + pub async fn receive_secret_event( &self, sender_key: &str, event: &mut SecretSendEvent, @@ -785,69 +868,7 @@ impl GossipMachine { ); } SecretInfo::SecretRequest(secret_name) => { - // Set the secret name so other consumers of the event know - // what this event is about. - event.content.secret_name = Some(secret_name.to_owned()); - - debug!( - sender = event.sender.as_str(), - request_id = event.content.request_id.as_str(), - secret_name = secret_name.as_ref(), - "Received a m.secret.send event with a matching request" - ); - - if let Some(device) = - self.store.get_device_from_curve_key(&event.sender, sender_key).await? - { - if device.verified() { - if secret_name != &SecretName::RecoveryKey { - match self - .store - .import_secret(secret_name, &event.content.secret) - .await - { - Ok(_) => self.mark_as_done(request).await?, - Err(e) => { - // If this is a store error propagate it up - // the call stack. - if let SecretImportError::Store(e) = e { - return Err(e); - } else { - // Otherwise warn that there was - // something wrong with the secret. - warn!( - secret_name = secret_name.as_ref(), - error = ?e, - "Error while importing a secret" - ) - } - } - } - } else { - // Skip importing the recovery key here since - // we'll want to check if the public key matches - // to the latest version on the server. The key - // will not be zeroized and - // instead leave the key in the event and let - // the user import it later. - } - } else { - warn!( - sender = event.sender.as_str(), - request_id = event.content.request_id.as_str(), - secret_name = secret_name.as_ref(), - "Received a m.secret.send event from an unverified device" - ); - } - } else { - warn!( - sender = event.sender.as_str(), - request_id = event.content.request_id.as_str(), - secret_name = secret_name.as_ref(), - "Received a m.secret.send event from an unknown device" - ); - self.store.update_tracked_user(&event.sender, true).await?; - } + self.receive_secret(sender_key, event, &request, secret_name).await?; } } } @@ -884,14 +905,14 @@ impl GossipMachine { let first_index = session.first_known_index(); if first_old_index > first_index { - self.mark_as_done(info).await?; + self.mark_as_done(&info).await?; Some(session) } else { None } // If we didn't have a previous session, store it. } else { - self.mark_as_done(info).await?; + self.mark_as_done(&info).await?; Some(session) }; diff --git a/crates/matrix-sdk-crypto/src/identities/manager.rs b/crates/matrix-sdk-crypto/src/identities/manager.rs index e3d6fec86..2d6857adf 100644 --- a/crates/matrix-sdk-crypto/src/identities/manager.rs +++ b/crates/matrix-sdk-crypto/src/identities/manager.rs @@ -20,7 +20,10 @@ use std::{ }; use futures_util::future::join_all; -use matrix_sdk_common::executor::spawn; +use matrix_sdk_common::{ + executor::spawn, + timeout::{timeout, ElapsedError}, +}; use ruma::{ api::client::keys::get_keys::v3::Response as KeysQueryResponse, serde::Raw, DeviceId, OwnedDeviceId, OwnedUserId, UserId, @@ -53,10 +56,6 @@ pub(crate) struct KeysQueryListener { store: Store, } -/// Error type notifying that a timeout has elapsed. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub(crate) struct Elapsed(()); - /// Result type telling us if a `/keys/query` response was expected for a given /// user. #[derive(Clone, Copy, Debug, PartialEq, Eq)] @@ -84,7 +83,7 @@ impl KeysQueryListener { &self, timeout: Duration, user: &UserId, - ) -> Result { + ) -> Result { let users_for_key_query = self.store.users_for_key_query(); if users_for_key_query.contains(user) { @@ -108,17 +107,8 @@ impl KeysQueryListener { /// /// If the given timeout has elapsed the method will stop waiting and return /// an error. - pub async fn wait(&self, timeout: Duration) -> Result<(), Elapsed> { - let listener = self.inner.listen(); - - #[cfg(not(target_arch = "wasm32"))] - tokio::time::timeout(timeout, async { listener.await }).await.map_err(|_| Elapsed(()))?; - - // TODO we should ensure that this is async on wasm as well. - #[cfg(target_arch = "wasm32")] - listener.wait_timeout(timeout); - - Ok(()) + pub async fn wait(&self, duration: Duration) -> Result<(), ElapsedError> { + timeout(self.inner.listen(), duration).await } } diff --git a/crates/matrix-sdk-crypto/src/machine.rs b/crates/matrix-sdk-crypto/src/machine.rs index 5e5f1eced..2d6daf314 100644 --- a/crates/matrix-sdk-crypto/src/machine.rs +++ b/crates/matrix-sdk-crypto/src/machine.rs @@ -750,7 +750,9 @@ impl OlmMachine { decrypted.inbound_group_session = session; } ToDeviceEvents::SecretSend(mut e) => { - self.key_request_machine.receive_secret(&decrypted.sender_key, &mut e).await?; + self.key_request_machine + .receive_secret_event(&decrypted.sender_key, &mut e) + .await?; decrypted.event = Raw::from_json(to_raw_value(&e)?) } _ => { diff --git a/crates/matrix-sdk/src/client/mod.rs b/crates/matrix-sdk/src/client/mod.rs index fa8ccc07f..0fa06b1d7 100644 --- a/crates/matrix-sdk/src/client/mod.rs +++ b/crates/matrix-sdk/src/client/mod.rs @@ -15,11 +15,10 @@ // limitations under the License. use std::{ - collections::BTreeMap, + collections::{btree_map, BTreeMap}, fmt::{self, Debug}, future::Future, io::Read, - pin::Pin, sync::{ atomic::{AtomicU64, Ordering::SeqCst}, Arc, RwLock as StdRwLock, @@ -85,8 +84,8 @@ use crate::{ config::RequestConfig, error::{HttpError, HttpResult}, event_handler::{ - EventHandler, EventHandlerData, EventHandlerHandle, EventHandlerResult, - EventHandlerWrapper, EventKind, SyncEvent, + EventHandler, EventHandlerFn, EventHandlerFut, EventHandlerHandle, EventHandlerKey, + EventHandlerResult, EventHandlerWrapper, SyncEvent, }, http_client::HttpClient, room, Account, Error, Result, @@ -107,9 +106,7 @@ const DEFAULT_UPLOAD_SPEED: u64 = 125_000; /// 5 min minimal upload request timeout, used to clamp the request timeout. const MIN_UPLOAD_REQUEST_TIMEOUT: Duration = Duration::from_secs(60 * 5); -type EventHandlerFut = Pin + Send>>; -pub(crate) type EventHandlerFn = dyn Fn(EventHandlerData<'_>) -> EventHandlerFut + Send + Sync; -type EventHandlerMap = BTreeMap<(EventKind, &'static str), Vec>; +type EventHandlerMap = BTreeMap>; type NotificationHandlerFut = EventHandlerFut; type NotificationHandlerFn = @@ -458,8 +455,38 @@ impl Client { H: EventHandler, ::Output: EventHandlerResult, { - let key = (Ev::KIND, Ev::TYPE); + self.add_event_handler_impl(handler, None).await + } + /// Register a handler for a specific room, and event type. + /// + /// This method works the same way as + /// [`add_event_handler`][Self::add_event_handler], except that the handler + /// will only be called for events in the room with the specified ID. See + /// that method for more details on event handler functions. + pub async fn add_room_event_handler( + &self, + room_id: &RoomId, + handler: H, + ) -> EventHandlerHandle + where + Ev: SyncEvent + DeserializeOwned + Send + 'static, + H: EventHandler, + ::Output: EventHandlerResult, + { + self.add_event_handler_impl(handler, Some(room_id.to_owned())).await + } + + async fn add_event_handler_impl( + &self, + handler: H, + room_id: Option, + ) -> EventHandlerHandle + where + Ev: SyncEvent + DeserializeOwned + Send + 'static, + H: EventHandler, + ::Output: EventHandlerResult, + { let handler_fn: Box = Box::new(move |data| { let maybe_fut = serde_json::from_str(data.raw.get()) .map(|ev| handler.clone().handle_event(ev, data)); @@ -487,18 +514,17 @@ impl Client { }); let handler_id = self.inner.event_handler_counter.fetch_add(1, SeqCst); - - let handle = EventHandlerHandle { handler_id, ev_id: key }; + let key = EventHandlerKey::new(Ev::KIND, Ev::TYPE, room_id); self.inner .event_handlers .write() .await - .entry(key) + .entry(key.clone()) .or_default() - .push(EventHandlerWrapper { handler_fn, handle }); + .push(EventHandlerWrapper { handler_fn, handler_id }); - handle + EventHandlerHandle { key, handler_id } } #[allow(missing_docs)] @@ -565,11 +591,12 @@ impl Client { pub async fn remove_event_handler(&self, handle: EventHandlerHandle) { let mut event_handlers = self.inner.event_handlers.write().await; - if let Some(v) = event_handlers.get_mut(&handle.ev_id) { - v.retain(|e| e.handle.handler_id != handle.handler_id); + if let btree_map::Entry::Occupied(mut entry) = event_handlers.entry(handle.key) { + let v = entry.get_mut(); + v.retain(|e| e.handler_id != handle.handler_id); if v.is_empty() { - event_handlers.remove(&handle.ev_id); + entry.remove(); } } } diff --git a/crates/matrix-sdk/src/event_handler.rs b/crates/matrix-sdk/src/event_handler.rs index bb38e3231..e5f1789c4 100644 --- a/crates/matrix-sdk/src/event_handler.rs +++ b/crates/matrix-sdk/src/event_handler.rs @@ -33,15 +33,25 @@ #[cfg(any(feature = "anyhow", feature = "eyre"))] use std::any::TypeId; -use std::{borrow::Cow, fmt, future::Future, ops::Deref}; +use std::{ + borrow::{Borrow, Cow}, + fmt, + future::Future, + iter, + ops::Deref, + pin::Pin, +}; use matrix_sdk_base::deserialized_responses::{EncryptionInfo, SyncRoomEvent}; -use ruma::{events::AnySyncStateEvent, serde::Raw}; +use ruma::{events::AnySyncStateEvent, serde::Raw, OwnedRoomId}; use serde::Deserialize; use serde_json::value::RawValue as RawJsonValue; use tracing::error; -use crate::{client::EventHandlerFn, room, Client}; +use crate::{room, Client}; + +pub(crate) type EventHandlerFut = Pin + Send>>; +pub(crate) type EventHandlerFn = dyn Fn(EventHandlerData<'_>) -> EventHandlerFut + Send + Sync; #[doc(hidden)] #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)] @@ -87,22 +97,50 @@ pub trait SyncEvent { const TYPE: &'static str; } +#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)] +pub(crate) struct EventHandlerKey(EventHandlerKeyInner<'static>); + +impl EventHandlerKey { + pub(crate) fn new( + ev_kind: EventKind, + ev_type: &'static str, + room_id: Option, + ) -> Self { + Self(EventHandlerKeyInner { ev_kind, ev_type, room_id }) + } +} + +// This lifetime-generic impl is what makes it possible to obtain a +// &'static str event type from get_key_value in call_event_handlers. +impl<'a> Borrow> for EventHandlerKey { + fn borrow(&self) -> &EventHandlerKeyInner<'a> { + &self.0 + } +} + +#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)] +pub(crate) struct EventHandlerKeyInner<'a> { + pub ev_kind: EventKind, + pub ev_type: &'a str, + pub room_id: Option, +} + pub(crate) struct EventHandlerWrapper { pub handler_fn: Box, - pub handle: EventHandlerHandle, + pub handler_id: u64, } /// Handle to remove a registered event handler by passing it to /// [`Client::remove_event_handler`]. -#[derive(Clone, Copy, Debug)] +#[derive(Clone, Debug)] pub struct EventHandlerHandle { - pub(crate) ev_id: (EventKind, &'static str), + pub(crate) key: EventHandlerKey, pub(crate) handler_id: u64, } impl EventHandlerContext for EventHandlerHandle { fn from_data(data: &EventHandlerData<'_>) -> Option { - Some(data.handle) + Some(data.handle.clone()) } } @@ -289,13 +327,12 @@ impl Client { event_type: Cow<'a, str>, } - self.handle_sync_events_wrapped_with( - room, - events, - |ev| (ev, None), - |raw| Ok((kind, raw.deserialize_as::>()?.event_type)), - ) - .await + for raw_event in events { + let event_type = raw_event.deserialize_as::>()?.event_type; + self.call_event_handlers(room, raw_event.json(), kind, &event_type, None).await; + } + + Ok(()) } pub(crate) async fn handle_sync_state_events( @@ -314,17 +351,13 @@ impl Client { self.handle_sync_events(EventKind::State, room, state_events).await?; // Event handlers specifically for redacted OR unredacted state events - self.handle_sync_events_wrapped_with( - room, - state_events, - |ev| (ev, None), - |raw| { - let StateEventDetails { event_type, unsigned } = raw.deserialize_as()?; - let redacted = unsigned.and_then(|u| u.redacted_because).is_some(); - Ok((EventKind::state_redacted(redacted), event_type)) - }, - ) - .await?; + for raw_event in state_events { + let StateEventDetails { event_type, unsigned } = raw_event.deserialize_as()?; + let redacted = unsigned.and_then(|u| u.redacted_because).is_some(); + let event_kind = EventKind::state_redacted(redacted); + + self.call_event_handlers(room, raw_event.json(), event_kind, &event_type, None).await; + } Ok(()) } @@ -343,85 +376,90 @@ impl Client { } // Event handlers for possibly-redacted timeline events - self.handle_sync_events_wrapped_with( - room, - timeline_events, - |e| (&e.event, e.encryption_info.as_ref()), - |raw| { - let TimelineEventDetails { event_type, state_key, .. } = raw.deserialize_as()?; + for item in timeline_events { + let TimelineEventDetails { event_type, state_key, .. } = item.event.deserialize_as()?; - let kind = match state_key { - Some(_) => EventKind::State, - None => EventKind::MessageLike, - }; + let event_kind = match state_key { + Some(_) => EventKind::State, + None => EventKind::MessageLike, + }; - Ok((kind, event_type)) - }, - ) - .await?; + let raw_event = &item.event.json(); + let encryption_info = item.encryption_info.as_ref(); + + self.call_event_handlers(room, raw_event, event_kind, &event_type, encryption_info) + .await; + } // Event handlers specifically for redacted OR unredacted timeline events - self.handle_sync_events_wrapped_with( - room, - timeline_events, - |e| (&e.event, e.encryption_info.as_ref()), - |raw| { - let TimelineEventDetails { event_type, state_key, unsigned } = - raw.deserialize_as()?; + for item in timeline_events { + let TimelineEventDetails { event_type, state_key, unsigned } = + item.event.deserialize_as()?; - let redacted = unsigned.and_then(|u| u.redacted_because).is_some(); - let kind = match state_key { - Some(_) => EventKind::state_redacted(redacted), - None => EventKind::message_like_redacted(redacted), - }; + let redacted = unsigned.and_then(|u| u.redacted_because).is_some(); + let event_kind = match state_key { + Some(_) => EventKind::state_redacted(redacted), + None => EventKind::message_like_redacted(redacted), + }; - Ok((kind, event_type)) - }, - ) - .await?; + let raw_event = &item.event.json(); + let encryption_info = item.encryption_info.as_ref(); + + self.call_event_handlers(room, raw_event, event_kind, &event_type, encryption_info) + .await; + } Ok(()) } - async fn handle_sync_events_wrapped_with<'a, T: 'a, U: 'a>( + async fn call_event_handlers( &self, room: &Option, - list: &'a [U], - get_event_details: impl Fn(&'a U) -> (&'a Raw, Option<&'a EncryptionInfo>), - get_id: impl Fn(&Raw) -> serde_json::Result<(EventKind, Cow<'_, str>)>, - ) -> serde_json::Result<()> { - for x in list { - let (raw_event, encryption_info) = get_event_details(x); - let (ev_kind, ev_type) = get_id(raw_event)?; - let event_handler_id = (ev_kind, &*ev_type); + raw: &RawJsonValue, + ev_kind: EventKind, + ev_type: &str, + encryption_info: Option<&EncryptionInfo>, + ) { + // Construct event handler futures + let futures: Vec<_> = { + let non_room_handler_key = EventHandlerKeyInner { ev_kind, ev_type, room_id: None }; + let room_handler_key = room.as_ref().map(|r| { + let room_id = Some(r.room_id().to_owned()); + EventHandlerKeyInner { ev_kind, ev_type, room_id } + }); - // Construct event handler futures - let futures: Vec<_> = self - .event_handlers() - .await - .get(&event_handler_id) - .into_iter() - .flatten() - .map(|handler_wrapper| { - let data = EventHandlerData { - client: self.clone(), - room: room.clone(), - raw: raw_event.json(), - encryption_info, - handle: handler_wrapper.handle, - }; - (handler_wrapper.handler_fn)(data) + let handlers_lock = self.event_handlers().await; + + iter::once(non_room_handler_key) + .chain(room_handler_key) + .flat_map(|b_key| { + // Use get_key_value instead of just get to be able to access the event_type + // from the BTreeMap key as &'static str, required for EventHandlerHandle. + handlers_lock.get_key_value(&b_key).into_iter().flat_map(|(key, handlers)| { + handlers.iter().map(|wrap| { + let data = EventHandlerData { + client: self.clone(), + room: room.clone(), + raw, + encryption_info, + handle: EventHandlerHandle { + key: key.clone(), + handler_id: wrap.handler_id, + }, + }; + + (wrap.handler_fn)(data) + }) + }) }) - .collect(); + .collect() + }; - // Run the event handler futures with the `self.event_handlers` lock - // no longer being held, in order. - for fut in futures { - fut.await; - } + // Run the event handler futures with the `self.event_handlers` lock + // no longer being held, in order. + for fut in futures { + fut.await; } - - Ok(()) } } @@ -586,7 +624,9 @@ mod static_events { #[cfg(all(test, not(target_arch = "wasm32")))] mod tests { - use matrix_sdk_test::{async_test, InvitedRoomBuilder, JoinedRoomBuilder}; + use matrix_sdk_test::{ + async_test, test_json::DEFAULT_SYNC_ROOM_ID, InvitedRoomBuilder, JoinedRoomBuilder, + }; #[cfg(target_arch = "wasm32")] wasm_bindgen_test::wasm_bindgen_test_configure!(run_in_browser); use std::{future, sync::Arc}; @@ -598,6 +638,7 @@ mod tests { events::{ room::{ member::{OriginalSyncRoomMemberEvent, StrippedRoomMemberEvent}, + name::OriginalSyncRoomNameEvent, power_levels::OriginalSyncRoomPowerLevelsEvent, }, typing::SyncTypingEvent, @@ -712,6 +753,77 @@ mod tests { Ok(()) } + #[async_test] + async fn add_room_event_handler() -> crate::Result<()> { + use std::sync::atomic::{AtomicU8, Ordering::SeqCst}; + + let client = crate::client::tests::logged_in_client(None).await; + + let room_id_a = room_id!("!foo:example.org"); + let room_id_b = room_id!("!bar:matrix.org"); + + let member_count = Arc::new(AtomicU8::new(0)); + let power_levels_count = Arc::new(AtomicU8::new(0)); + + // Room event handlers for member events in both rooms + client + .add_room_event_handler(room_id_a, { + let member_count = member_count.clone(); + move |_ev: OriginalSyncRoomMemberEvent, _room: room::Room| { + member_count.fetch_add(1, SeqCst); + future::ready(()) + } + }) + .await; + client + .add_room_event_handler(room_id_b, { + let member_count = member_count.clone(); + move |_ev: OriginalSyncRoomMemberEvent, _room: room::Room| { + member_count.fetch_add(1, SeqCst); + future::ready(()) + } + }) + .await; + + // Power levels event handlers for member events in room A + client + .add_room_event_handler(room_id_a, { + let power_levels_count = power_levels_count.clone(); + move |_ev: OriginalSyncRoomPowerLevelsEvent, _client: Client, _room: room::Room| { + power_levels_count.fetch_add(1, SeqCst); + future::ready(()) + } + }) + .await; + + // Room name event handler for room name events in room B + client + .add_room_event_handler(room_id_b, move |_ev: OriginalSyncRoomNameEvent| async { + unreachable!("No room event in room B") + }) + .await; + + let response = EventBuilder::default() + .add_joined_room( + JoinedRoomBuilder::new(room_id_a) + .add_timeline_event(TimelineTestEvent::Member) + .add_state_event(StateTestEvent::PowerLevels) + .add_state_event(StateTestEvent::RoomName), + ) + .add_joined_room( + JoinedRoomBuilder::new(room_id_b) + .add_timeline_event(TimelineTestEvent::Member) + .add_state_event(StateTestEvent::PowerLevels), + ) + .build_sync_response(); + client.process_sync(response).await?; + + assert_eq!(member_count.load(SeqCst), 2); + assert_eq!(power_levels_count.load(SeqCst), 1); + + Ok(()) + } + #[async_test] async fn remove_event_handler() -> crate::Result<()> { use std::sync::atomic::{AtomicU8, Ordering::SeqCst}; @@ -730,15 +842,20 @@ mod tests { }) .await; - let handle = client - .add_event_handler({ - move |_ev: OriginalSyncRoomMemberEvent| { - panic!("handler should have been removed"); - #[allow(unreachable_code)] - future::ready(()) - } + let handle_a = client + .add_event_handler(move |_ev: OriginalSyncRoomMemberEvent| async { + panic!("handler should have been removed"); }) .await; + let handle_b = client + .add_room_event_handler( + #[allow(unknown_lints, clippy::explicit_auto_deref)] // lint is buggy + *DEFAULT_SYNC_ROOM_ID, + move |_ev: OriginalSyncRoomMemberEvent| async { + panic!("handler should have been removed"); + }, + ) + .await; client .add_event_handler({ @@ -756,7 +873,8 @@ mod tests { ) .build_sync_response(); - client.remove_event_handler(handle).await; + client.remove_event_handler(handle_a).await; + client.remove_event_handler(handle_b).await; client.process_sync(response).await?; diff --git a/testing/matrix-sdk-test/Cargo.toml b/testing/matrix-sdk-test/Cargo.toml index b94de0fe8..bfe738bb6 100644 --- a/testing/matrix-sdk-test/Cargo.toml +++ b/testing/matrix-sdk-test/Cargo.toml @@ -22,3 +22,9 @@ once_cell = "1.10.0" ruma = { git = "https://github.com/ruma/ruma", rev = "ca8c66c885241a7ba3805399604eda4a38979f6b", features = ["client-api-c"] } serde = "1.0.136" serde_json = "1.0.79" + +[target.'cfg(not(target_arch = "wasm32"))'.dependencies] +tokio = { version = "1.17.0", default-features = false, features = ["rt", "macros"] } + +[target.'cfg(target_arch = "wasm32")'.dependencies] +wasm-bindgen-test = "0.3.30"