diff --git a/bindings/matrix-sdk-ffi/src/client.rs b/bindings/matrix-sdk-ffi/src/client.rs index 6150a9ff0..2813ced45 100644 --- a/bindings/matrix-sdk-ffi/src/client.rs +++ b/bindings/matrix-sdk-ffi/src/client.rs @@ -236,7 +236,7 @@ impl Client { } } -#[uniffi::export] +#[uniffi::export(async_runtime = "tokio")] impl Client { /// Login using a username and password. pub fn login( @@ -259,7 +259,7 @@ impl Client { }) } - pub fn get_media_file( + pub async fn get_media_file( &self, media_source: Arc, body: Option, @@ -267,24 +267,22 @@ impl Client { use_cache: bool, temp_dir: Option, ) -> Result, ClientError> { - let client = self.inner.clone(); let source = (*media_source).clone(); let mime_type: mime::Mime = mime_type.parse()?; - RUNTIME.block_on(async move { - let handle = client - .media() - .get_media_file( - &MediaRequest { source, format: MediaFormat::File }, - body, - &mime_type, - use_cache, - temp_dir, - ) - .await?; + let handle = self + .inner + .media() + .get_media_file( + &MediaRequest { source, format: MediaFormat::File }, + body, + &mime_type, + use_cache, + temp_dir, + ) + .await?; - Ok(Arc::new(MediaFileHandle::new(handle))) - }) + Ok(Arc::new(MediaFileHandle::new(handle))) } /// Restores the client from a `Session`. @@ -350,7 +348,7 @@ impl Client { } } -#[uniffi::export] +#[uniffi::export(async_runtime = "tokio")] impl Client { pub fn set_delegate( self: Arc, @@ -488,68 +486,65 @@ impl Client { }) } - pub fn upload_media( + pub async fn upload_media( &self, mime_type: String, data: Vec, progress_watcher: Option>, ) -> Result { - let l = self.inner.clone(); + let mime_type: mime::Mime = mime_type.parse().context("Parsing mime type")?; + let request = self.inner.media().upload(&mime_type, data); - RUNTIME.block_on(async move { - let mime_type: mime::Mime = mime_type.parse().context("Parsing mime type")?; - let request = l.media().upload(&mime_type, data); - if let Some(progress_watcher) = progress_watcher { - let mut subscriber = request.subscribe_to_send_progress(); - RUNTIME.spawn(async move { - while let Some(progress) = subscriber.next().await { - progress_watcher.transmission_progress(progress.into()); - } - }); - } - let response = request.await?; - Ok(String::from(response.content_uri)) - }) + if let Some(progress_watcher) = progress_watcher { + let mut subscriber = request.subscribe_to_send_progress(); + RUNTIME.spawn(async move { + while let Some(progress) = subscriber.next().await { + progress_watcher.transmission_progress(progress.into()); + } + }); + } + + let response = request.await?; + + Ok(String::from(response.content_uri)) } - pub fn get_media_content( + pub async fn get_media_content( &self, media_source: Arc, ) -> Result, ClientError> { - let l = self.inner.clone(); let source = (*media_source).clone(); - RUNTIME.block_on(async move { - Ok(l.media() - .get_media_content(&MediaRequest { source, format: MediaFormat::File }, true) - .await?) - }) + Ok(self + .inner + .media() + .get_media_content(&MediaRequest { source, format: MediaFormat::File }, true) + .await?) } - pub fn get_media_thumbnail( + pub async fn get_media_thumbnail( &self, media_source: Arc, width: u64, height: u64, ) -> Result, ClientError> { - let l = self.inner.clone(); let source = (*media_source).clone(); - RUNTIME.block_on(async move { - Ok(l.media() - .get_media_content( - &MediaRequest { - source, - format: MediaFormat::Thumbnail(MediaThumbnailSize { - method: Method::Scale, - width: UInt::new(width).unwrap(), - height: UInt::new(height).unwrap(), - }), - }, - true, - ) - .await?) - }) + Ok(self + .inner + .media() + .get_media_content( + &MediaRequest { + source, + format: MediaFormat::Thumbnail(MediaThumbnailSize { + method: Method::Scale, + width: UInt::new(width).unwrap(), + height: UInt::new(height).unwrap(), + }), + }, + true, + ) + .await?) } pub fn get_session_verification_controller( diff --git a/bindings/matrix-sdk-ffi/src/room.rs b/bindings/matrix-sdk-ffi/src/room.rs index 4816b6974..85c61d858 100644 --- a/bindings/matrix-sdk-ffi/src/room.rs +++ b/bindings/matrix-sdk-ffi/src/room.rs @@ -424,6 +424,24 @@ impl Room { Ok(self.inner.can_user_ban(&user_id).await?) } + pub async fn ban_user( + &self, + user_id: String, + reason: Option, + ) -> Result<(), ClientError> { + let user_id = UserId::parse(&user_id)?; + Ok(self.inner.ban_user(&user_id, reason.as_deref()).await?) + } + + pub async fn unban_user( + &self, + user_id: String, + reason: Option, + ) -> Result<(), ClientError> { + let user_id = UserId::parse(&user_id)?; + Ok(self.inner.unban_user(&user_id, reason.as_deref()).await?) + } + pub async fn can_user_invite(&self, user_id: String) -> Result { let user_id = UserId::parse(&user_id)?; Ok(self.inner.can_user_invite(&user_id).await?) @@ -434,6 +452,15 @@ impl Room { Ok(self.inner.can_user_kick(&user_id).await?) } + pub async fn kick_user( + &self, + user_id: String, + reason: Option, + ) -> Result<(), ClientError> { + let user_id = UserId::parse(&user_id)?; + Ok(self.inner.kick_user(&user_id, reason.as_deref()).await?) + } + pub async fn can_user_send_state( &self, user_id: String, diff --git a/bindings/matrix-sdk-ffi/src/session_verification.rs b/bindings/matrix-sdk-ffi/src/session_verification.rs index 0b160de78..8b77b6373 100644 --- a/bindings/matrix-sdk-ffi/src/session_verification.rs +++ b/bindings/matrix-sdk-ffi/src/session_verification.rs @@ -31,11 +31,17 @@ impl SessionVerificationEmoji { } } +#[derive(uniffi::Enum)] +pub enum SessionVerificationData { + Emojis { emojis: Vec>, indices: Vec }, + Decimals { values: Vec }, +} + #[uniffi::export(callback_interface)] pub trait SessionVerificationControllerDelegate: Sync + Send { fn did_accept_verification_request(&self); fn did_start_sas_verification(&self); - fn did_receive_verification_data(&self, data: Vec>); + fn did_receive_verification_data(&self, data: SessionVerificationData); fn did_fail(&self); fn did_cancel(&self); fn did_finish(&self); @@ -199,25 +205,31 @@ impl SessionVerificationController { while let Some(state) = stream.next().await { match state { - SasState::KeysExchanged { emojis, decimals: _ } => { - // TODO: If emojis is None, decimals should be used. - if let Some(emojis) = emojis { - if let Some(delegate) = &*delegate.read().unwrap() { - let emojis = emojis - .emojis - .iter() - .map(|e| { - Arc::new(SessionVerificationEmoji { - symbol: e.symbol.to_owned(), - description: e.description.to_owned(), - }) - }) - .collect::>(); - - delegate.did_receive_verification_data(emojis); + SasState::KeysExchanged { emojis, decimals } => { + if let Some(delegate) = &*delegate.read().unwrap() { + if let Some(emojis) = emojis { + delegate.did_receive_verification_data( + SessionVerificationData::Emojis { + emojis: emojis + .emojis + .into_iter() + .map(|emoji| { + Arc::new(SessionVerificationEmoji { + symbol: emoji.symbol.to_owned(), + description: emoji.description.to_owned(), + }) + }) + .collect(), + indices: emojis.indices.to_vec(), + }, + ); + } else { + delegate.did_receive_verification_data( + SessionVerificationData::Decimals { + values: vec![decimals.0, decimals.1, decimals.2], + }, + ) } - } else if let Some(delegate) = &*delegate.read().unwrap() { - delegate.did_fail() } } SasState::Done { .. } => { diff --git a/bindings/matrix-sdk-ffi/src/timeline/mod.rs b/bindings/matrix-sdk-ffi/src/timeline/mod.rs index b4bde590e..21787e02a 100644 --- a/bindings/matrix-sdk-ffi/src/timeline/mod.rs +++ b/bindings/matrix-sdk-ffi/src/timeline/mod.rs @@ -25,7 +25,6 @@ use matrix_sdk::attachment::{ use matrix_sdk_ui::timeline::{BackPaginationStatus, EventItemOrigin, Profile, TimelineDetails}; use mime::Mime; use ruma::{ - api::client::receipt::create_receipt::v3::ReceiptType, events::{ location::{AssetType as RumaAssetType, LocationContent, ZoomLevel}, poll::{ @@ -182,12 +181,16 @@ impl Timeline { RUNTIME.block_on(async { Ok(self.inner.paginate_backwards(opts.into()).await?) }) } - pub fn send_read_receipt(&self, event_id: String) -> Result<(), ClientError> { + pub fn send_read_receipt( + &self, + receipt_type: ReceiptType, + event_id: String, + ) -> Result<(), ClientError> { let event_id = EventId::parse(event_id)?; RUNTIME.block_on(async { self.inner - .send_single_receipt(ReceiptType::Read, ReceiptThread::Unthreaded, event_id) + .send_single_receipt(receipt_type.into(), ReceiptThread::Unthreaded, event_id) .await?; Ok(()) }) @@ -202,7 +205,7 @@ impl Timeline { pub fn send_image( self: Arc, url: String, - thumbnail_url: String, + thumbnail_url: Option, image_info: ImageInfo, progress_watcher: Option>, ) -> Arc { @@ -217,13 +220,13 @@ impl Timeline { let attachment_info = AttachmentInfo::Image(base_image_info); - let attachment_config = match image_info.thumbnail_info { - Some(thumbnail_image_info) => { + let attachment_config = match (thumbnail_url, image_info.thumbnail_info) { + (Some(thumbnail_url), Some(thumbnail_image_info)) => { let thumbnail = self.build_thumbnail_info(thumbnail_url, thumbnail_image_info)?; AttachmentConfig::with_thumbnail(thumbnail).info(attachment_info) } - None => AttachmentConfig::new().info(attachment_info), + _ => AttachmentConfig::new().info(attachment_info), }; self.send_attachment(url, mime_type, attachment_config, progress_watcher).await @@ -233,7 +236,7 @@ impl Timeline { pub fn send_video( self: Arc, url: String, - thumbnail_url: String, + thumbnail_url: Option, video_info: VideoInfo, progress_watcher: Option>, ) -> Arc { @@ -248,13 +251,13 @@ impl Timeline { let attachment_info = AttachmentInfo::Video(base_video_info); - let attachment_config = match video_info.thumbnail_info { - Some(thumbnail_image_info) => { + let attachment_config = match (thumbnail_url, video_info.thumbnail_info) { + (Some(thumbnail_url), Some(thumbnail_image_info)) => { let thumbnail = self.build_thumbnail_info(thumbnail_url, thumbnail_image_info)?; AttachmentConfig::with_thumbnail(thumbnail).info(attachment_info) } - None => AttachmentConfig::new().info(attachment_info), + _ => AttachmentConfig::new().info(attachment_info), }; self.send_attachment(url, mime_type, attachment_config, progress_watcher).await @@ -975,3 +978,21 @@ pub enum VirtualTimelineItem { /// The user's own read marker. ReadMarker, } + +/// A [`TimelineItem`](super::TimelineItem) that doesn't correspond to an event. +#[derive(uniffi::Enum)] +pub enum ReceiptType { + Read, + ReadPrivate, + FullyRead, +} + +impl From for ruma::api::client::receipt::create_receipt::v3::ReceiptType { + fn from(value: ReceiptType) -> Self { + match value { + ReceiptType::Read => Self::Read, + ReceiptType::ReadPrivate => Self::Read, + ReceiptType::FullyRead => Self::FullyRead, + } + } +} diff --git a/crates/matrix-sdk-base/src/media.rs b/crates/matrix-sdk-base/src/media.rs index 32356e989..4094a4da0 100644 --- a/crates/matrix-sdk-base/src/media.rs +++ b/crates/matrix-sdk-base/src/media.rs @@ -12,7 +12,7 @@ use ruma::{ }, sticker::StickerEventContent, }, - UInt, + MxcUri, UInt, }; const UNIQUE_SEPARATOR: &str = "_"; @@ -83,11 +83,22 @@ pub struct MediaRequest { pub format: MediaFormat, } +impl MediaRequest { + /// Get the [`MxcUri`] from `Self`. + pub fn uri(&self) -> &MxcUri { + match &self.source { + MediaSource::Plain(url) => url.as_ref(), + MediaSource::Encrypted(file) => file.url.as_ref(), + } + } +} + impl UniqueKey for MediaRequest { fn unique_key(&self) -> String { format!("{}{UNIQUE_SEPARATOR}{}", self.source.unique_key(), self.format.unique_key()) } } + /// Trait for media event content. pub trait MediaEventContent { /// Get the source of the file for `Self`. @@ -166,3 +177,47 @@ impl MediaEventContent for LocationMessageEventContent { self.info.as_ref()?.thumbnail_source.clone() } } + +#[cfg(test)] +mod tests { + use ruma::mxc_uri; + use serde_json::json; + + use super::*; + + #[test] + fn test_media_request_url() { + let mxc_uri = mxc_uri!("mxc://homeserver/media"); + + let plain = MediaRequest { + source: MediaSource::Plain(mxc_uri.to_owned()), + format: MediaFormat::File, + }; + + assert_eq!(plain.uri(), mxc_uri); + + let file = MediaRequest { + source: MediaSource::Encrypted(Box::new( + serde_json::from_value(json!({ + "url": mxc_uri, + "key": { + "kty": "oct", + "key_ops": ["encrypt", "decrypt"], + "alg": "A256CTR", + "k": "b50ACIv6LMn9AfMCFD1POJI_UAFWIclxAN1kWrEO2X8", + "ext": true, + }, + "iv": "AK1wyzigZtQAAAABAAAAKK", + "hashes": { + "sha256": "foobar", + }, + "v": "v2", + })) + .unwrap(), + )), + format: MediaFormat::File, + }; + + assert_eq!(file.uri(), mxc_uri); + } +} diff --git a/crates/matrix-sdk-base/src/store/integration_tests.rs b/crates/matrix-sdk-base/src/store/integration_tests.rs index 610d254d1..345c002d7 100644 --- a/crates/matrix-sdk-base/src/store/integration_tests.rs +++ b/crates/matrix-sdk-base/src/store/integration_tests.rs @@ -200,11 +200,8 @@ impl StateStoreIntegrationTests for DynStateStore { async fn test_media_content(&self) { let uri = mxc_uri!("mxc://localhost/media"); - let content: Vec = "somebinarydata".into(); - let request_file = MediaRequest { source: MediaSource::Plain(uri.to_owned()), format: MediaFormat::File }; - let request_thumbnail = MediaRequest { source: MediaSource::Plain(uri.to_owned()), format: MediaFormat::Thumbnail(MediaThumbnailSize { @@ -214,6 +211,17 @@ impl StateStoreIntegrationTests for DynStateStore { }), }; + let other_uri = mxc_uri!("mxc://localhost/media-other"); + let request_other_file = MediaRequest { + source: MediaSource::Plain(other_uri.to_owned()), + format: MediaFormat::File, + }; + + let content: Vec = "hello".into(); + let thumbnail_content: Vec = "world".into(); + let other_content: Vec = "foo".into(); + + // Media isn't present in the cache. assert!( self.get_media_content(&request_file).await.unwrap().is_none(), "unexpected media found" @@ -223,35 +231,63 @@ impl StateStoreIntegrationTests for DynStateStore { "media not found" ); + // Let's add the media. self.add_media_content(&request_file, content.clone()).await.expect("adding media failed"); - assert!( - self.get_media_content(&request_file).await.unwrap().is_some(), + + // Media is present in the cache. + assert_eq!( + self.get_media_content(&request_file).await.unwrap().as_ref(), + Some(&content), "media not found though added" ); + // Let's remove the media. self.remove_media_content(&request_file).await.expect("removing media failed"); + + // Media isn't present in the cache. assert!( self.get_media_content(&request_file).await.unwrap().is_none(), "media still there after removing" ); + // Let's add the media again. self.add_media_content(&request_file, content.clone()) .await .expect("adding media again failed"); - assert!( - self.get_media_content(&request_file).await.unwrap().is_some(), + + assert_eq!( + self.get_media_content(&request_file).await.unwrap().as_ref(), + Some(&content), "media not found after adding again" ); - self.add_media_content(&request_thumbnail, content.clone()) + // Let's add the thumbnail media. + self.add_media_content(&request_thumbnail, thumbnail_content.clone()) .await .expect("adding thumbnail failed"); - assert!( - self.get_media_content(&request_thumbnail).await.unwrap().is_some(), + + // Media's thumbnail is present. + assert_eq!( + self.get_media_content(&request_thumbnail).await.unwrap().as_ref(), + Some(&thumbnail_content), "thumbnail not found" ); + // Let's add another media with a different URI. + self.add_media_content(&request_other_file, other_content.clone()) + .await + .expect("adding other media failed"); + + // Other file is present. + assert_eq!( + self.get_media_content(&request_other_file).await.unwrap().as_ref(), + Some(&other_content), + "other file not found" + ); + + // Let's remove media based on URI. self.remove_media_content_for_uri(uri).await.expect("removing all media for uri failed"); + assert!( self.get_media_content(&request_file).await.unwrap().is_none(), "media wasn't removed" @@ -260,6 +296,10 @@ impl StateStoreIntegrationTests for DynStateStore { self.get_media_content(&request_thumbnail).await.unwrap().is_none(), "thumbnail wasn't removed" ); + assert!( + self.get_media_content(&request_other_file).await.unwrap().is_some(), + "other media was removed" + ); } async fn test_topic_redaction(&self) -> Result<()> { diff --git a/crates/matrix-sdk-base/src/store/memory_store.rs b/crates/matrix-sdk-base/src/store/memory_store.rs index 8ab4867e0..e00845b93 100644 --- a/crates/matrix-sdk-base/src/store/memory_store.rs +++ b/crates/matrix-sdk-base/src/store/memory_store.rs @@ -18,7 +18,7 @@ use std::{ }; use async_trait::async_trait; -use matrix_sdk_common::instant::Instant; +use matrix_sdk_common::{instant::Instant, ring_buffer::RingBuffer}; use ruma::{ canonical_json::{redact, RedactedBecause}, events::{ @@ -29,15 +29,16 @@ use ruma::{ AnySyncStateEvent, GlobalAccountDataEventType, RoomAccountDataEventType, StateEventType, }, serde::Raw, - CanonicalJsonObject, EventId, MxcUri, OwnedEventId, OwnedRoomId, OwnedUserId, RoomId, - RoomVersionId, UserId, + CanonicalJsonObject, EventId, MxcUri, OwnedEventId, OwnedMxcUri, OwnedRoomId, OwnedUserId, + RoomId, RoomVersionId, UserId, }; use tracing::{debug, warn}; use super::{Result, RoomInfo, StateChanges, StateStore, StoreError}; use crate::{ - deserialized_responses::RawAnySyncOrStrippedState, media::MediaRequest, MinimalRoomMemberEvent, - RoomMemberships, RoomState, StateStoreDataKey, StateStoreDataValue, + deserialized_responses::RawAnySyncOrStrippedState, + media::{MediaRequest, UniqueKey as _}, + MinimalRoomMemberEvent, RoomMemberships, RoomState, StateStoreDataKey, StateStoreDataValue, }; /// In-Memory, non-persistent implementation of the `StateStore` @@ -77,13 +78,14 @@ pub struct MemoryStore { HashMap<(String, Option), HashMap>>, >, >, + media: StdRwLock)>>, custom: StdRwLock, Vec>>, } impl MemoryStore { /// Create a new empty MemoryStore pub fn new() -> Self { - Default::default() + Self { media: StdRwLock::new(RingBuffer::new(20)), ..Default::default() } } fn get_user_room_receipt_event_impl( @@ -700,17 +702,55 @@ impl StateStore for MemoryStore { Ok(self.custom.write().unwrap().remove(key)) } - // The in-memory store doesn't cache media - async fn add_media_content(&self, _request: &MediaRequest, _data: Vec) -> Result<()> { + async fn add_media_content(&self, request: &MediaRequest, data: Vec) -> Result<()> { + // Avoid duplication. Let's try to remove it first. + self.remove_media_content(request).await?; + // Now, let's add it. + self.media.write().unwrap().push((request.uri().to_owned(), request.unique_key(), data)); + Ok(()) } - async fn get_media_content(&self, _request: &MediaRequest) -> Result>> { - Ok(None) + + async fn get_media_content(&self, request: &MediaRequest) -> Result>> { + let media = self.media.read().unwrap(); + let expected_key = request.unique_key(); + + Ok(media.iter().find_map(|(_media_uri, media_key, media_content)| { + (media_key == &expected_key).then(|| media_content.to_owned()) + })) } - async fn remove_media_content(&self, _request: &MediaRequest) -> Result<()> { + + async fn remove_media_content(&self, request: &MediaRequest) -> Result<()> { + let mut media = self.media.write().unwrap(); + let expected_key = request.unique_key(); + let Some(index) = media + .iter() + .position(|(_media_uri, media_key, _media_content)| media_key == &expected_key) + else { + return Ok(()); + }; + + media.remove(index); + Ok(()) } - async fn remove_media_content_for_uri(&self, _uri: &MxcUri) -> Result<()> { + + async fn remove_media_content_for_uri(&self, uri: &MxcUri) -> Result<()> { + let mut media = self.media.write().unwrap(); + let expected_key = uri.to_owned(); + let positions = media + .iter() + .enumerate() + .filter_map(|(position, (media_uri, _media_key, _media_content))| { + (media_uri == &expected_key).then_some(position) + }) + .collect::>(); + + // Iterate in reverse-order so that positions stay valid after first removals. + for position in positions.into_iter().rev() { + media.remove(position); + } + Ok(()) } @@ -738,5 +778,5 @@ mod tests { Ok(MemoryStore::new()) } - statestore_integration_tests!(); + statestore_integration_tests!(with_media_tests); } diff --git a/crates/matrix-sdk-common/Cargo.toml b/crates/matrix-sdk-common/Cargo.toml index dcab35295..cb11faa04 100644 --- a/crates/matrix-sdk-common/Cargo.toml +++ b/crates/matrix-sdk-common/Cargo.toml @@ -16,7 +16,7 @@ default-target = "x86_64-unknown-linux-gnu" targets = ["x86_64-unknown-linux-gnu", "wasm32-unknown-unknown"] [features] -js = ["instant/wasm-bindgen", "instant/inaccurate", "wasm-bindgen-futures"] +js = ["instant/wasm-bindgen", "wasm-bindgen-futures"] [dependencies] async-trait = { workspace = true } diff --git a/crates/matrix-sdk-common/src/ring_buffer.rs b/crates/matrix-sdk-common/src/ring_buffer.rs index 34cec4a6d..ba09471fd 100644 --- a/crates/matrix-sdk-common/src/ring_buffer.rs +++ b/crates/matrix-sdk-common/src/ring_buffer.rs @@ -75,12 +75,19 @@ impl RingBuffer { self.inner.pop_front() } + /// Removes and returns one specific element at `index` if it exists, + /// otherwise it returns `None`. + pub fn remove(&mut self, index: usize) -> Option { + self.inner.remove(index) + } + /// Returns an iterator that provides elements in front-to-back order, i.e. /// the same order you would get if you repeatedly called pop(). pub fn iter(&self) -> Iter<'_, T> { self.inner.iter() } + /// Returns an iterator that drains its items. pub fn drain(&mut self, range: R) -> Drain<'_, T> where R: RangeBounds, @@ -155,7 +162,7 @@ mod tests { } #[test] - pub fn test_push_and_pop_and_length() { + pub fn test_push_and_pop_and_remove_and_length() { let mut ring_buffer = RingBuffer::new(3); ring_buffer.push(1); @@ -167,23 +174,48 @@ mod tests { ring_buffer.push(3); assert_eq!(ring_buffer.len(), 3); - ring_buffer.pop(); + assert_eq!(ring_buffer.pop(), Some(1)); assert_eq!(ring_buffer.len(), 2); assert_eq!(ring_buffer.get(0), Some(&2)); assert_eq!(ring_buffer.get(1), Some(&3)); assert_eq!(ring_buffer.get(2), None); - ring_buffer.pop(); + assert_eq!(ring_buffer.pop(), Some(2)); assert_eq!(ring_buffer.len(), 1); assert_eq!(ring_buffer.get(0), Some(&3)); assert_eq!(ring_buffer.get(1), None); assert_eq!(ring_buffer.get(2), None); - ring_buffer.pop(); + assert_eq!(ring_buffer.pop(), Some(3)); assert_eq!(ring_buffer.len(), 0); assert_eq!(ring_buffer.get(0), None); assert_eq!(ring_buffer.get(1), None); assert_eq!(ring_buffer.get(2), None); + + assert_eq!(ring_buffer.pop(), None); + + ring_buffer.push(1); + ring_buffer.push(2); + ring_buffer.push(3); + assert_eq!(ring_buffer.len(), 3); + assert_eq!(ring_buffer.get(0), Some(&1)); + assert_eq!(ring_buffer.get(1), Some(&2)); + assert_eq!(ring_buffer.get(2), Some(&3)); + + assert_eq!(ring_buffer.remove(1), Some(2)); + assert_eq!(ring_buffer.len(), 2); + assert_eq!(ring_buffer.get(0), Some(&1)); + assert_eq!(ring_buffer.get(1), Some(&3)); + assert_eq!(ring_buffer.get(2), None); + + assert_eq!(ring_buffer.remove(0), Some(1)); + assert_eq!(ring_buffer.len(), 1); + assert_eq!(ring_buffer.get(0), Some(&3)); + assert_eq!(ring_buffer.get(1), None); + assert_eq!(ring_buffer.get(2), None); + + assert_eq!(ring_buffer.remove(1), None); + assert_eq!(ring_buffer.remove(10), None); } #[test] diff --git a/crates/matrix-sdk-crypto/src/identities/manager.rs b/crates/matrix-sdk-crypto/src/identities/manager.rs index 1790c498a..8de1d2a32 100644 --- a/crates/matrix-sdk-crypto/src/identities/manager.rs +++ b/crates/matrix-sdk-crypto/src/identities/manager.rs @@ -546,7 +546,7 @@ impl IdentityManager { } } - /// Try to deserialize the the master key and self-signing key of an + /// Try to deserialize the master key and self-signing key of an /// identity from a `/keys/query` response. /// /// Each user identity *must* at least contain a master and self-signing diff --git a/crates/matrix-sdk-crypto/src/session_manager/sessions.rs b/crates/matrix-sdk-crypto/src/session_manager/sessions.rs index fff9930cb..75abc7092 100644 --- a/crates/matrix-sdk-crypto/src/session_manager/sessions.rs +++ b/crates/matrix-sdk-crypto/src/session_manager/sessions.rs @@ -18,7 +18,6 @@ use std::{ time::Duration, }; -use itertools::Itertools; use matrix_sdk_common::failures_cache::FailuresCache; use ruma::{ api::client::keys::claim_keys::v3::{ @@ -295,18 +294,12 @@ impl SessionManager { if tracing::level_enabled!(tracing::Level::DEBUG) { // Reformat the map to skip the encryption algorithm, which isn't very useful. - // - // Note: we reify the debug string for `missing_session_devices_by_user` to work - // around a known footgun of `itertools` (it can be used only once). - let missing_session_devices_by_user = format!( - "{:?}", - missing_session_devices_by_user - .iter() - .map(|(user_id, devices)| (user_id, devices.keys().collect::>())) - .format(", ") - ); + let missing_session_devices_by_user = missing_session_devices_by_user + .iter() + .map(|(user_id, devices)| (user_id, devices.keys().collect::>())) + .collect::>(); debug!( - missing_session_devices_by_user, + ?missing_session_devices_by_user, ?timed_out_devices_by_user, "Collected user/device pairs that are missing an Olm session" ); diff --git a/crates/matrix-sdk-crypto/src/verification/mod.rs b/crates/matrix-sdk-crypto/src/verification/mod.rs index c62d33c3d..72fe1a31b 100644 --- a/crates/matrix-sdk-crypto/src/verification/mod.rs +++ b/crates/matrix-sdk-crypto/src/verification/mod.rs @@ -82,7 +82,7 @@ pub struct Emoji { pub description: &'static str, } -/// Format the the list of emojis as a two line string. +/// Format the list of emojis as a two line string. /// /// The first line will contain the emojis spread out so the second line can /// contain the descriptions centered bellow the emoji. diff --git a/crates/matrix-sdk-sqlite/src/state_store.rs b/crates/matrix-sdk-sqlite/src/state_store.rs index 906d5f0c9..58f476cfb 100644 --- a/crates/matrix-sdk-sqlite/src/state_store.rs +++ b/crates/matrix-sdk-sqlite/src/state_store.rs @@ -611,7 +611,7 @@ trait SqliteObjectStateStoreExt: SqliteObjectExt { async fn get_kv_blobs(&self, keys: Vec) -> Result>> { let keys_length = keys.len(); - chunk_large_query_over(keys, Some(keys_length), |keys| { + self.chunk_large_query_over(keys, Some(keys_length), |keys| { let sql_params = repeat_vars(keys.len()); let sql = format!("SELECT value FROM kv_blob WHERE key IN ({sql_params})"); @@ -639,7 +639,7 @@ trait SqliteObjectStateStoreExt: SqliteObjectExt { }) .await?) } else { - chunk_large_query_over(states, None, |states| { + self.chunk_large_query_over(states, None, |states| { let sql_params = repeat_vars(states.len()); let sql = format!("SELECT data FROM room_info WHERE state IN ({sql_params})"); @@ -659,7 +659,7 @@ trait SqliteObjectStateStoreExt: SqliteObjectExt { event_type: Key, state_keys: Vec, ) -> Result)>> { - chunk_large_query_over(state_keys, None, move |state_keys: Vec| { + self.chunk_large_query_over(state_keys, None, move |state_keys: Vec| { let sql_params = repeat_vars(state_keys.len()); let sql = format!( "SELECT stripped, data FROM state_event @@ -702,7 +702,7 @@ trait SqliteObjectStateStoreExt: SqliteObjectExt { ) -> Result, Vec)>> { let user_ids_length = user_ids.len(); - chunk_large_query_over(user_ids, Some(user_ids_length), move |user_ids| { + self.chunk_large_query_over(user_ids, Some(user_ids_length), move |user_ids| { let sql_params = repeat_vars(user_ids.len()); let sql = format!( "SELECT user_id, data FROM profile WHERE room_id = ? AND user_id IN ({sql_params})" @@ -724,7 +724,7 @@ trait SqliteObjectStateStoreExt: SqliteObjectExt { }) .await? } else { - chunk_large_query_over(memberships, None, move |memberships| { + self.chunk_large_query_over(memberships, None, move |memberships| { let sql_params = repeat_vars(memberships.len()); let sql = format!( "SELECT data FROM member WHERE room_id = ? AND membership IN ({sql_params})" @@ -776,7 +776,7 @@ trait SqliteObjectStateStoreExt: SqliteObjectExt { ) -> Result, Vec)>> { let names_length = names.len(); - chunk_large_query_over(names, Some(names_length), move |names| { + self.chunk_large_query_over(names, Some(names_length), move |names| { let sql_params = repeat_vars(names.len()); let sql = format!( "SELECT name, data FROM display_name WHERE room_id = ? AND name IN ({sql_params})" @@ -858,6 +858,55 @@ trait SqliteObjectStateStoreExt: SqliteObjectExt { self.execute("DELETE FROM media WHERE uri = ?", (uri,)).await?; Ok(()) } + + /// Chunk a large query over some keys. + /// + /// Imagine there is a _dynamic_ query that runs potentially large number of + /// parameters, so much that the maximum number of parameters can be hit. + /// Then, this helper is for you. It will execute the query on chunks of + /// parameters. + async fn chunk_large_query_over( + &self, + mut keys_to_chunk: Vec, + result_capacity: Option, + do_query: Query, + ) -> Result> + where + Query: Fn(Vec) -> Fut + Send + Sync, + Fut: Future, rusqlite::Error>> + Send, + Res: Send, + { + // Divide by 2 to allow space for more static parameters (not part of + // `keys_to_chunk`). + let maximum_chunk_size = self.limit(Limit::SQLITE_LIMIT_VARIABLE_NUMBER).await / 2; + let maximum_chunk_size: usize = maximum_chunk_size + .try_into() + .map_err(|_| Error::SqliteMaximumVariableNumber(maximum_chunk_size))?; + + if keys_to_chunk.len() < maximum_chunk_size { + // Chunking isn't necessary. + let chunk = keys_to_chunk; + + Ok(do_query(chunk).await?) + } else { + // Chunking _is_ necessary. + + // Define the accumulator. + let capacity = result_capacity.unwrap_or_default(); + let mut all_results = Vec::with_capacity(capacity); + + while !keys_to_chunk.is_empty() { + // Chunk and run the query. + let tail = keys_to_chunk.split_off(min(keys_to_chunk.len(), maximum_chunk_size)); + let chunk = keys_to_chunk; + keys_to_chunk = tail; + + all_results.extend(do_query(chunk).await?); + } + + Ok(all_results) + } + } } #[async_trait] @@ -1605,59 +1654,31 @@ struct ReceiptData { user_id: OwnedUserId, } -/// Chunk a large query over some keys. -/// -/// Imagine there is a _dynamic_ query that runs potentially large number of -/// parameters, so much that the maximum number of parameters can be hit. Then, -/// this helper is for you. It will execute the query on chunks of parameters. -async fn chunk_large_query_over( - mut keys_to_chunk: Vec, - result_capacity: Option, - do_query: Query, -) -> Result> -where - Query: Fn(Vec) -> Fut, - Fut: Future, rusqlite::Error>>, -{ - // `Limit` has a `repr(i32)`, it's safe to cast it to `i32`. Then divide by 2 to - // let space for more static parameters (not part of `keys_to_chunk`). - let maximum_chunk_size = Limit::SQLITE_LIMIT_VARIABLE_NUMBER as i32 / 2; - let maximum_chunk_size: usize = maximum_chunk_size - .try_into() - .map_err(|_| Error::SqliteMaximumVariableNumber(maximum_chunk_size))?; - - if keys_to_chunk.len() < maximum_chunk_size { - // Chunking isn't necessary. - let chunk = keys_to_chunk; - - Ok(do_query(chunk).await?) - } else { - // Chunking _is_ necessary. - - // Define the accumulator. - let capacity = result_capacity.unwrap_or_default(); - let mut all_results = Vec::with_capacity(capacity); - - while !keys_to_chunk.is_empty() { - // Chunk and run the query. - let tail = keys_to_chunk.split_off(min(keys_to_chunk.len(), maximum_chunk_size)); - let chunk = keys_to_chunk; - keys_to_chunk = tail; - - all_results.extend(do_query(chunk).await?); - } - - Ok(all_results) - } -} - /// Repeat `?` n times, where n is defined by `count`. `?` are comma-separated. fn repeat_vars(count: usize) -> impl fmt::Display { - assert_ne!(count, 0); + assert_ne!(count, 0, "Can't generate zero repeated vars"); iter::repeat("?").take(count).format(",") } +#[cfg(test)] +mod unit_tests { + use super::*; + + #[test] + fn can_generate_repeated_vars() { + assert_eq!(repeat_vars(1).to_string(), "?"); + assert_eq!(repeat_vars(2).to_string(), "?,?"); + assert_eq!(repeat_vars(5).to_string(), "?,?,?,?,?"); + } + + #[test] + #[should_panic(expected = "Can't generate zero repeated vars")] + fn generating_zero_vars_panics() { + repeat_vars(0); + } +} + #[cfg(test)] mod tests { use std::sync::atomic::{AtomicU32, Ordering::SeqCst}; diff --git a/crates/matrix-sdk-sqlite/src/utils.rs b/crates/matrix-sdk-sqlite/src/utils.rs index 06d06c211..9f14a5bbb 100644 --- a/crates/matrix-sdk-sqlite/src/utils.rs +++ b/crates/matrix-sdk-sqlite/src/utils.rs @@ -15,9 +15,9 @@ use std::{borrow::Borrow, ops::Deref}; use async_trait::async_trait; -use rusqlite::{OptionalExtension, Params, Row, Statement, Transaction}; +use rusqlite::{limits::Limit, OptionalExtension, Params, Row, Statement, Transaction}; -use crate::OpenStoreError; +use crate::{error::Result, OpenStoreError}; #[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord)] pub(crate) enum Key { @@ -86,7 +86,7 @@ pub(crate) trait SqliteObjectExt { E: From + Send + 'static, F: FnOnce(&Transaction<'_>) -> Result + Send + 'static; - async fn limit(&self, limit: rusqlite::limits::Limit) -> i32; + async fn limit(&self, limit: Limit) -> i32; } #[async_trait] @@ -148,7 +148,7 @@ impl SqliteObjectExt for deadpool_sqlite::Object { .unwrap() } - async fn limit(&self, limit: rusqlite::limits::Limit) -> i32 { + async fn limit(&self, limit: Limit) -> i32 { self.interact(move |conn| conn.limit(limit)).await.expect("Failed to fetch limit") } } diff --git a/crates/matrix-sdk/Cargo.toml b/crates/matrix-sdk/Cargo.toml index b9ce1c897..89bfc8b0c 100644 --- a/crates/matrix-sdk/Cargo.toml +++ b/crates/matrix-sdk/Cargo.toml @@ -159,3 +159,7 @@ wasm-bindgen-test = "0.3.33" [target.'cfg(not(target_arch = "wasm32"))'.dev-dependencies] tokio = { workspace = true, features = ["rt-multi-thread", "macros"] } wiremock = "0.5.13" + +[[test]] +name = "integration" +required-features = ["testing"] diff --git a/crates/matrix-sdk/src/error.rs b/crates/matrix-sdk/src/error.rs index 40ac14ca8..f74e30b80 100644 --- a/crates/matrix-sdk/src/error.rs +++ b/crates/matrix-sdk/src/error.rs @@ -452,7 +452,7 @@ pub enum NotificationSettingsError { #[error("Unable to update push rule")] UnableToUpdatePushRule, /// Rule not found - #[error("Rule not found")] + #[error("Rule `{0}` not found")] RuleNotFound(String), /// Unable to save the push rules #[error("Unable to save push rules")] diff --git a/crates/matrix-sdk/src/matrix_auth/mod.rs b/crates/matrix-sdk/src/matrix_auth/mod.rs index ceab02ea4..4da69ae6c 100644 --- a/crates/matrix-sdk/src/matrix_auth/mod.rs +++ b/crates/matrix-sdk/src/matrix_auth/mod.rs @@ -38,7 +38,7 @@ use ruma::{ serde::JsonObject, }; use serde::{Deserialize, Serialize}; -use tracing::{debug, info, instrument}; +use tracing::{debug, error, info, instrument}; use crate::{ authentication::AuthData, @@ -486,7 +486,9 @@ impl MatrixAuth { if let Some(save_session_callback) = self.client.inner.auth_ctx.save_session_callback.get() { - save_session_callback(self.client.clone()); + if let Err(err) = save_session_callback(self.client.clone()).await { + error!("when saving session after refresh: {err}"); + } } _ = self diff --git a/crates/matrix-sdk/src/media.rs b/crates/matrix-sdk/src/media.rs index c178adbfb..a5857947f 100644 --- a/crates/matrix-sdk/src/media.rs +++ b/crates/matrix-sdk/src/media.rs @@ -264,12 +264,12 @@ impl Media { request: &MediaRequest, use_cache: bool, ) -> Result> { - let content = - if use_cache { self.client.store().get_media_content(request).await? } else { None }; - - if let Some(content) = content { - return Ok(content); - } + // Read from the cache. + if use_cache { + if let Some(content) = self.client.store().get_media_content(request).await? { + return Ok(content); + } + }; let content: Vec = match &request.source { MediaSource::Encrypted(file) => { diff --git a/crates/matrix-sdk/src/notification_settings/command.rs b/crates/matrix-sdk/src/notification_settings/command.rs index 408c01d9e..ed87c9bb5 100644 --- a/crates/matrix-sdk/src/notification_settings/command.rs +++ b/crates/matrix-sdk/src/notification_settings/command.rs @@ -3,8 +3,8 @@ use std::fmt::Debug; use ruma::{ api::client::push::RuleScope, push::{ - Action, NewConditionalPushRule, NewPushRule, NewSimplePushRule, PushCondition, RuleKind, - Tweak, + Action, NewConditionalPushRule, NewPatternedPushRule, NewPushRule, NewSimplePushRule, + PushCondition, RuleKind, Tweak, }, OwnedRoomId, }; @@ -18,6 +18,8 @@ pub(crate) enum Command { SetRoomPushRule { scope: RuleScope, room_id: OwnedRoomId, notify: bool }, /// Set a new `Override` push rule matching a `RoomId` SetOverridePushRule { scope: RuleScope, rule_id: String, room_id: OwnedRoomId, notify: bool }, + /// Set a new push rule for a keyword. + SetKeywordPushRule { scope: RuleScope, keyword: String }, /// Set whether a push rule is enabled SetPushRuleEnabled { scope: RuleScope, kind: RuleKind, rule_id: String, enabled: bool }, /// Delete a push rule @@ -57,6 +59,16 @@ impl Command { Ok(NewPushRule::Override(new_rule)) } + Self::SetKeywordPushRule { scope: _, keyword } => { + // `Content` push rule matching this keyword + let new_rule = NewPatternedPushRule::new( + keyword.clone(), + keyword.clone(), + get_notify_actions(true), + ); + Ok(NewPushRule::Content(new_rule)) + } + Self::SetPushRuleEnabled { .. } | Self::DeletePushRule { .. } | Self::SetPushRuleActions { .. } => Err(NotificationSettingsError::InvalidParameter( diff --git a/crates/matrix-sdk/src/notification_settings/mod.rs b/crates/matrix-sdk/src/notification_settings/mod.rs index d885b23ed..1459c0d77 100644 --- a/crates/matrix-sdk/src/notification_settings/mod.rs +++ b/crates/matrix-sdk/src/notification_settings/mod.rs @@ -2,6 +2,7 @@ use std::sync::Arc; +use indexmap::IndexSet; use ruma::{ api::client::push::{ delete_pushrule, set_pushrule, set_pushrule_actions, set_pushrule_enabled, @@ -14,6 +15,7 @@ use tokio::sync::{ broadcast::{self, Receiver}, RwLock, }; +use tracing::{debug, error}; use self::{command::Command, rule_commands::RuleCommands, rules::Rules}; @@ -27,7 +29,7 @@ use crate::{ }; /// Enum representing the push notification modes for a room. -#[derive(Debug, Clone, PartialEq)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum RoomNotificationMode { /// Receive notifications for all messages. AllMessages, @@ -38,7 +40,7 @@ pub enum RoomNotificationMode { } /// Whether or not a room is encrypted -#[derive(Debug)] +#[derive(Debug, Clone, Copy)] pub enum IsEncrypted { /// The room is encrypted Yes, @@ -57,7 +59,7 @@ impl From for IsEncrypted { } /// Whether or not a room is a `one-to-one` -#[derive(Debug, Clone)] +#[derive(Debug, Clone, Copy)] pub enum IsOneToOne { /// A room is a `one-to-one` room if it has exactly two members. Yes, @@ -196,11 +198,6 @@ impl NotificationSettings { is_one_to_one: IsOneToOne, mode: RoomNotificationMode, ) -> Result<(), NotificationSettingsError> { - let rule_ids = vec![ - rules::get_predefined_underride_room_rule_id(is_encrypted, is_one_to_one.clone()), - is_one_to_one.into(), - ]; - let actions = match mode { RoomNotificationMode::AllMessages => { vec![Action::Notify, Action::SetTweak(Tweak::Sound("default".into()))] @@ -210,8 +207,21 @@ impl NotificationSettings { } }; - for rule_id in rule_ids { - self.set_underride_push_rule_actions(rule_id, actions.clone()).await? + let room_rule_id = + rules::get_predefined_underride_room_rule_id(is_encrypted, is_one_to_one); + self.set_underride_push_rule_actions(room_rule_id, actions.clone()).await?; + + let poll_start_rule_id = rules::get_predefined_underride_poll_start_rule_id(is_one_to_one); + if let Err(error) = + self.set_underride_push_rule_actions(poll_start_rule_id, actions.clone()).await + { + // The poll start event rules are currently unstable so they might not be found + // on every homeserver. Let's ignore this error for the moment. + if let NotificationSettingsError::RuleNotFound(rule_id) = &error { + debug!("Unable to update poll start push rule: rule `{rule_id}` not found"); + } else { + return Err(error); + } } Ok(()) @@ -259,7 +269,7 @@ impl NotificationSettings { let rules = self.rules.read().await.clone(); // Check that the current mode is not already the target mode. - if rules.get_user_defined_room_notification_mode(room_id) == Some(mode.clone()) { + if rules.get_user_defined_room_notification_mode(room_id) == Some(mode) { return Ok(()); } @@ -362,6 +372,72 @@ impl NotificationSettings { } } + /// Get the keywords which have enabled rules. + pub async fn enabled_keywords(&self) -> IndexSet { + self.rules.read().await.enabled_keywords() + } + + /// Add or enable a rule for the given keyword. + /// + /// # Arguments + /// + /// * `keyword` - The keyword to match. + pub async fn add_keyword(&self, keyword: String) -> Result<(), NotificationSettingsError> { + let rules = self.rules.read().await.clone(); + + let mut rule_commands = RuleCommands::new(rules.clone().ruleset); + + let existing_rules = rules.keyword_rules(&keyword); + + if existing_rules.is_empty() { + // Create a rule. + rule_commands.insert_keyword_rule(keyword)?; + } else { + if existing_rules.iter().any(|r| r.enabled) { + // Nothing to do. + return Ok(()); + } + + // Enable one of the rules. + rule_commands.set_rule_enabled(RuleKind::Content, &existing_rules[0].rule_id, true)?; + } + + self.run_server_commands(&rule_commands).await?; + + let rules = &mut *self.rules.write().await; + rules.apply(rule_commands); + + Ok(()) + } + + /// Remove the rules for the given keyword. + /// + /// # Arguments + /// + /// * `keyword` - The keyword to unmatch. + pub async fn remove_keyword(&self, keyword: &str) -> Result<(), NotificationSettingsError> { + let rules = self.rules.read().await.clone(); + + let mut rule_commands = RuleCommands::new(rules.clone().ruleset); + + let existing_rules = rules.keyword_rules(keyword); + + if existing_rules.is_empty() { + return Ok(()); + } + + for rule in existing_rules { + rule_commands.delete_rule(RuleKind::Content, rule.rule_id.clone())?; + } + + self.run_server_commands(&rule_commands).await?; + + let rules = &mut *self.rules.write().await; + rules.apply(rule_commands); + + Ok(()) + } + /// Convert commands into requests to the server, and run them. async fn run_server_commands( &self, @@ -376,20 +452,28 @@ impl NotificationSettings { kind.clone(), rule_id.clone(), ); - self.client - .send(request, request_config) - .await - .map_err(|_| NotificationSettingsError::UnableToRemovePushRule)?; + self.client.send(request, request_config).await.map_err(|error| { + error!("Unable to delete {kind} push rule `{rule_id}`: {error}"); + NotificationSettingsError::UnableToRemovePushRule + })?; } - Command::SetRoomPushRule { scope, room_id: _, notify: _ } => { + Command::SetRoomPushRule { scope, room_id, notify: _ } => { let push_rule = command.to_push_rule()?; let request = set_pushrule::v3::Request::new(scope.clone(), push_rule); - self.client - .send(request, request_config) - .await - .map_err(|_| NotificationSettingsError::UnableToAddPushRule)?; + self.client.send(request, request_config).await.map_err(|error| { + error!("Unable to set room push rule `{room_id}`: {error}"); + NotificationSettingsError::UnableToAddPushRule + })?; } - Command::SetOverridePushRule { scope, rule_id: _, room_id: _, notify: _ } => { + Command::SetOverridePushRule { scope, rule_id, room_id: _, notify: _ } => { + let push_rule = command.to_push_rule()?; + let request = set_pushrule::v3::Request::new(scope.clone(), push_rule); + self.client.send(request, request_config).await.map_err(|error| { + error!("Unable to set override push rule `{rule_id}`: {error}"); + NotificationSettingsError::UnableToAddPushRule + })?; + } + Command::SetKeywordPushRule { scope, keyword: _ } => { let push_rule = command.to_push_rule()?; let request = set_pushrule::v3::Request::new(scope.clone(), push_rule); self.client @@ -404,10 +488,10 @@ impl NotificationSettings { rule_id.clone(), *enabled, ); - self.client - .send(request, request_config) - .await - .map_err(|_| NotificationSettingsError::UnableToUpdatePushRule)?; + self.client.send(request, request_config).await.map_err(|error| { + error!("Unable to set {kind} push rule `{rule_id}` enabled: {error}"); + NotificationSettingsError::UnableToUpdatePushRule + })?; } Command::SetPushRuleActions { scope, kind, rule_id, actions } => { let request = set_pushrule_actions::v3::Request::new( @@ -416,10 +500,10 @@ impl NotificationSettings { rule_id.clone(), actions.clone(), ); - self.client - .send(request, request_config) - .await - .map_err(|_| NotificationSettingsError::UnableToUpdatePushRule)?; + self.client.send(request, request_config).await.map_err(|error| { + error!("Unable to set {kind} push rule `{rule_id}` actions: {error}"); + NotificationSettingsError::UnableToUpdatePushRule + })?; } } } @@ -778,16 +862,16 @@ mod tests { let mode = settings.get_user_defined_room_notification_mode(&room_id).await; assert!(mode.is_none()); - let new_modes = &[ + let new_modes = [ RoomNotificationMode::AllMessages, RoomNotificationMode::MentionsAndKeywordsOnly, RoomNotificationMode::Mute, ]; for new_mode in new_modes { - settings.set_room_notification_mode(&room_id, new_mode.clone()).await.unwrap(); + settings.set_room_notification_mode(&room_id, new_mode).await.unwrap(); assert_eq!( - new_mode.clone(), + new_mode, settings.get_user_defined_room_notification_mode(&room_id).await.unwrap() ); } @@ -1190,4 +1274,294 @@ mod tests { RoomNotificationMode::AllMessages ); } + + #[async_test] + async fn list_keywords() { + let server = MockServer::start().await; + let client = logged_in_client(Some(server.uri())).await; + + // Initial state: No keywords + let ruleset = get_server_default_ruleset(); + let settings = NotificationSettings::new(client.clone(), ruleset); + + let keywords = settings.enabled_keywords().await; + + assert!(keywords.is_empty()); + + // Initial state: 3 rules, 2 keywords + let mut ruleset = get_server_default_ruleset(); + ruleset + .insert( + NewPushRule::Content(NewPatternedPushRule::new( + "a".to_owned(), + "a".to_owned(), + vec![], + )), + None, + None, + ) + .unwrap(); + // Test deduplication. + ruleset + .insert( + NewPushRule::Content(NewPatternedPushRule::new( + "a_bis".to_owned(), + "a".to_owned(), + vec![], + )), + None, + None, + ) + .unwrap(); + ruleset + .insert( + NewPushRule::Content(NewPatternedPushRule::new( + "b".to_owned(), + "b".to_owned(), + vec![], + )), + None, + None, + ) + .unwrap(); + + let settings = NotificationSettings::new(client, ruleset); + + let keywords = settings.enabled_keywords().await; + assert_eq!(keywords.len(), 2); + assert!(keywords.get("a").is_some()); + assert!(keywords.get("b").is_some()) + } + + #[async_test] + async fn add_keyword_missing() { + let server = MockServer::start().await; + let client = logged_in_client(Some(server.uri())).await; + let settings = client.notification_settings().await; + + Mock::given(method("PUT")) + .and(path("/_matrix/client/r0/pushrules/global/content/banana")) + .respond_with(ResponseTemplate::new(200)) + .expect(1) + .mount(&server) + .await; + + settings.add_keyword("banana".to_owned()).await.unwrap(); + + // The ruleset must have been updated. + let keywords = settings.enabled_keywords().await; + assert_eq!(keywords.len(), 1); + assert!(keywords.get("banana").is_some()); + + // Rule exists. + let rule_enabled = + settings.is_push_rule_enabled(RuleKind::Content, "banana").await.unwrap(); + assert!(rule_enabled) + } + + #[async_test] + async fn add_keyword_disabled() { + let server = MockServer::start().await; + let client = logged_in_client(Some(server.uri())).await; + + let mut ruleset = get_server_default_ruleset(); + ruleset + .insert( + NewPushRule::Content(NewPatternedPushRule::new( + "banana_two".to_owned(), + "banana".to_owned(), + vec![], + )), + None, + None, + ) + .unwrap(); + ruleset.set_enabled(RuleKind::Content, "banana_two", false).unwrap(); + ruleset + .insert( + NewPushRule::Content(NewPatternedPushRule::new( + "banana_one".to_owned(), + "banana".to_owned(), + vec![], + )), + None, + None, + ) + .unwrap(); + ruleset.set_enabled(RuleKind::Content, "banana_one", false).unwrap(); + + let settings = NotificationSettings::new(client, ruleset); + Mock::given(method("PUT")) + .and(path("/_matrix/client/r0/pushrules/global/content/banana_one/enabled")) + .respond_with(ResponseTemplate::new(200)) + .expect(1) + .mount(&server) + .await; + + settings.add_keyword("banana".to_owned()).await.unwrap(); + + // The ruleset must have been updated. + let keywords = settings.enabled_keywords().await; + + assert_eq!(keywords.len(), 1); + assert!(keywords.get("banana").is_some()); + + // The first rule was enabled. + let first_rule_enabled = + settings.is_push_rule_enabled(RuleKind::Content, "banana_one").await.unwrap(); + assert!(first_rule_enabled); + let second_rule_enabled = + settings.is_push_rule_enabled(RuleKind::Content, "banana_two").await.unwrap(); + assert!(!second_rule_enabled); + } + + #[async_test] + async fn add_keyword_noop() { + let server = MockServer::start().await; + let client = logged_in_client(Some(server.uri())).await; + + let mut ruleset = get_server_default_ruleset(); + ruleset + .insert( + NewPushRule::Content(NewPatternedPushRule::new( + "banana_two".to_owned(), + "banana".to_owned(), + vec![], + )), + None, + None, + ) + .unwrap(); + ruleset + .insert( + NewPushRule::Content(NewPatternedPushRule::new( + "banana_one".to_owned(), + "banana".to_owned(), + vec![], + )), + None, + None, + ) + .unwrap(); + ruleset.set_enabled(RuleKind::Content, "banana_one", false).unwrap(); + + let settings = NotificationSettings::new(client, ruleset); + settings.add_keyword("banana".to_owned()).await.unwrap(); + + // Nothing changed. + let keywords = settings.enabled_keywords().await; + + assert_eq!(keywords.len(), 1); + assert!(keywords.get("banana").is_some()); + + let first_rule_enabled = + settings.is_push_rule_enabled(RuleKind::Content, "banana_one").await.unwrap(); + assert!(!first_rule_enabled); + let second_rule_enabled = + settings.is_push_rule_enabled(RuleKind::Content, "banana_two").await.unwrap(); + assert!(second_rule_enabled); + } + + #[async_test] + async fn remove_keyword_all() { + let server = MockServer::start().await; + let client = logged_in_client(Some(server.uri())).await; + + let mut ruleset = get_server_default_ruleset(); + ruleset + .insert( + NewPushRule::Content(NewPatternedPushRule::new( + "banana_two".to_owned(), + "banana".to_owned(), + vec![], + )), + None, + None, + ) + .unwrap(); + ruleset + .insert( + NewPushRule::Content(NewPatternedPushRule::new( + "banana_one".to_owned(), + "banana".to_owned(), + vec![], + )), + None, + None, + ) + .unwrap(); + ruleset.set_enabled(RuleKind::Content, "banana_one", false).unwrap(); + + let settings = NotificationSettings::new(client, ruleset); + + Mock::given(method("DELETE")) + .and(path("/_matrix/client/r0/pushrules/global/content/banana_one")) + .respond_with(ResponseTemplate::new(200)) + .expect(1) + .mount(&server) + .await; + Mock::given(method("DELETE")) + .and(path("/_matrix/client/r0/pushrules/global/content/banana_two")) + .respond_with(ResponseTemplate::new(200)) + .expect(1) + .mount(&server) + .await; + + settings.remove_keyword("banana").await.unwrap(); + + // The ruleset must have been updated. + let keywords = settings.enabled_keywords().await; + assert!(keywords.is_empty()); + + // Rules we removed. + let first_rule_error = + settings.is_push_rule_enabled(RuleKind::Content, "banana_one").await.unwrap_err(); + assert_matches!(first_rule_error, NotificationSettingsError::RuleNotFound(_)); + let second_rule_error = + settings.is_push_rule_enabled(RuleKind::Content, "banana_two").await.unwrap_err(); + assert_matches!(second_rule_error, NotificationSettingsError::RuleNotFound(_)); + } + + #[async_test] + async fn remove_keyword_noop() { + let server = MockServer::start().await; + let client = logged_in_client(Some(server.uri())).await; + let settings = client.notification_settings().await; + + settings.remove_keyword("banana").await.unwrap(); + } + + #[async_test] + async fn test_set_default_room_notification_mode_missing_poll_start() { + let server = MockServer::start().await; + Mock::given(method("PUT")).respond_with(ResponseTemplate::new(200)).mount(&server).await; + let client = logged_in_client(Some(server.uri())).await; + + // If the initial mode is `AllMessages` + let mut ruleset = get_server_default_ruleset(); + ruleset.underride.remove(PredefinedUnderrideRuleId::PollStart.as_str()); + + let settings = NotificationSettings::new(client, ruleset); + assert_eq!( + settings.get_default_room_notification_mode(IsEncrypted::No, IsOneToOne::No).await, + RoomNotificationMode::AllMessages + ); + + // After setting the default mode to `MentionsAndKeywordsOnly` + settings + .set_default_room_notification_mode( + IsEncrypted::No, + IsOneToOne::No, + RoomNotificationMode::MentionsAndKeywordsOnly, + ) + .await + .unwrap(); + + // the new mode returned by `get_default_room_notification_mode()` should + // reflect the change. + assert_matches!( + settings.get_default_room_notification_mode(IsEncrypted::No, IsOneToOne::No).await, + RoomNotificationMode::MentionsAndKeywordsOnly + ); + } } diff --git a/crates/matrix-sdk/src/notification_settings/rule_commands.rs b/crates/matrix-sdk/src/notification_settings/rule_commands.rs index 42f132547..3ac2cda73 100644 --- a/crates/matrix-sdk/src/notification_settings/rule_commands.rs +++ b/crates/matrix-sdk/src/notification_settings/rule_commands.rs @@ -55,6 +55,19 @@ impl RuleCommands { Ok(()) } + /// Insert a new rule for a keyword. + pub(crate) fn insert_keyword_rule( + &mut self, + keyword: String, + ) -> Result<(), NotificationSettingsError> { + let command = Command::SetKeywordPushRule { scope: RuleScope::Global, keyword }; + + self.rules.insert(command.to_push_rule()?, None, None)?; + self.commands.push(command); + + Ok(()) + } + /// Delete a rule pub(crate) fn delete_rule( &mut self, diff --git a/crates/matrix-sdk/src/notification_settings/rules.rs b/crates/matrix-sdk/src/notification_settings/rules.rs index 5dde925c8..941485258 100644 --- a/crates/matrix-sdk/src/notification_settings/rules.rs +++ b/crates/matrix-sdk/src/notification_settings/rules.rs @@ -1,9 +1,10 @@ //! Ruleset utility struct use imbl::HashSet; +use indexmap::IndexSet; use ruma::{ push::{ - AnyPushRuleRef, PredefinedContentRuleId, PredefinedOverrideRuleId, + AnyPushRuleRef, PatternedPushRule, PredefinedContentRuleId, PredefinedOverrideRuleId, PredefinedUnderrideRuleId, PushCondition, RuleKind, Ruleset, }, RoomId, @@ -42,8 +43,8 @@ impl Rules { } // add any `Room` rules matching this `room_id` - if let Some(rule) = self.ruleset.room.iter().find(|x| x.rule_id == room_id) { - custom_rules.push((RuleKind::Room, rule.rule_id.to_string())); + if let Some(rule) = self.ruleset.get(RuleKind::Room, room_id) { + custom_rules.push((RuleKind::Room, rule.rule_id().to_owned())); } // add any `Underride` rules matching this `room_id` @@ -83,9 +84,9 @@ impl Rules { } // Search for an enabled `Room` rule where `rule_id` is the `room_id` - if let Some(rule) = self.ruleset.room.iter().find(|x| x.enabled && x.rule_id == room_id) { + if let Some(rule) = self.ruleset.get(RuleKind::Room, room_id) { // if this rule contains a `Notify` action - if rule.actions.iter().any(|x| x.should_notify()) { + if rule.triggers_notification() { return Some(RoomNotificationMode::AllMessages); } return Some(RoomNotificationMode::MentionsAndKeywordsOnly); @@ -113,9 +114,11 @@ impl Rules { // If there is an `Underride` rule that should trigger a notification, the mode // is `AllMessages` - if self.ruleset.underride.iter().any(|r| { - r.enabled && r.rule_id == rule_id && r.actions.iter().any(|a| a.should_notify()) - }) { + if self + .ruleset + .get(RuleKind::Underride, rule_id) + .is_some_and(|r| r.enabled() && r.triggers_notification()) + { RoomNotificationMode::AllMessages } else { // Otherwise, the mode is `MentionsAndKeywordsOnly` @@ -171,7 +174,7 @@ impl Rules { if let Some(rule) = self.ruleset.get(RuleKind::Override, PredefinedOverrideRuleId::ContainsDisplayName) { - if rule.enabled() && rule.actions().iter().any(|a| a.should_notify()) { + if rule.enabled() && rule.triggers_notification() { return true; } } @@ -180,7 +183,7 @@ impl Rules { if let Some(rule) = self.ruleset.get(RuleKind::Content, PredefinedContentRuleId::ContainsUserName) { - if rule.enabled() && rule.actions().iter().any(|a| a.should_notify()) { + if rule.enabled() && rule.triggers_notification() { return true; } } @@ -200,12 +203,9 @@ impl Rules { // Fallback to deprecated rule for compatibility #[allow(deprecated)] - let room_notif_rule_id = PredefinedOverrideRuleId::RoomNotif.as_str(); - self.ruleset.override_.iter().any(|r| { - r.enabled - && r.rule_id == room_notif_rule_id - && r.actions.iter().any(|a| a.should_notify()) - }) + self.ruleset + .get(RuleKind::Override, PredefinedOverrideRuleId::RoomNotif) + .is_some_and(|r| r.enabled() && r.triggers_notification()) } /// Get whether the given ruleset contains some enabled keywords rules. @@ -214,6 +214,21 @@ impl Rules { self.ruleset.content.iter().any(|r| !r.default && r.enabled) } + /// The keywords which have enabled rules. + pub(crate) fn enabled_keywords(&self) -> IndexSet { + self.ruleset + .content + .iter() + .filter(|r| !r.default && r.enabled) + .map(|r| r.pattern.clone()) + .collect() + } + + /// The rules for a keyword, if any. + pub(crate) fn keyword_rules(&self, keyword: &str) -> Vec<&PatternedPushRule> { + self.ruleset.content.iter().filter(|r| !r.default && r.pattern == keyword).collect() + } + /// Get whether a rule is enabled. pub(crate) fn is_enabled( &self, @@ -241,7 +256,9 @@ impl Rules { Command::DeletePushRule { scope: _, kind, rule_id } => { _ = self.ruleset.remove(kind, rule_id); } - Command::SetRoomPushRule { .. } | Command::SetOverridePushRule { .. } => { + Command::SetRoomPushRule { .. } + | Command::SetOverridePushRule { .. } + | Command::SetKeywordPushRule { .. } => { if let Ok(push_rule) = command.to_push_rule() { _ = self.ruleset.insert(push_rule, None, None); } @@ -257,7 +274,7 @@ impl Rules { } } -/// Gets the `PredefinedUnderrideRuleId` corresponding to the given +/// Gets the `PredefinedUnderrideRuleId` for rooms corresponding to the given /// criteria. /// /// # Arguments @@ -276,12 +293,18 @@ pub(crate) fn get_predefined_underride_room_rule_id( } } -impl From for PredefinedUnderrideRuleId { - fn from(is_one_to_one: IsOneToOne) -> Self { - match is_one_to_one { - IsOneToOne::Yes => Self::PollStartOneToOne, - IsOneToOne::No => Self::PollStart, - } +/// Gets the `PredefinedUnderrideRuleId` for poll start events corresponding to +/// the given criteria. +/// +/// # Arguments +/// +/// * `is_one_to_one` - `Yes` if the room is a direct chat involving two people +pub(crate) fn get_predefined_underride_poll_start_rule_id( + is_one_to_one: IsOneToOne, +) -> PredefinedUnderrideRuleId { + match is_one_to_one { + IsOneToOne::Yes => PredefinedUnderrideRuleId::PollStartOneToOne, + IsOneToOne::No => PredefinedUnderrideRuleId::PollStart, } } @@ -418,6 +441,18 @@ pub(crate) mod tests { ); } + #[async_test] + async fn test_get_predefined_underride_poll_start_rule_id() { + assert_eq!( + rules::get_predefined_underride_poll_start_rule_id(IsOneToOne::No), + PredefinedUnderrideRuleId::PollStart + ); + assert_eq!( + rules::get_predefined_underride_poll_start_rule_id(IsOneToOne::Yes), + PredefinedUnderrideRuleId::PollStartOneToOne + ); + } + #[async_test] async fn test_get_default_room_notification_mode_mentions_and_keywords() { let mut ruleset = get_server_default_ruleset(); diff --git a/crates/matrix-sdk/src/room/mod.rs b/crates/matrix-sdk/src/room/mod.rs index b63f3d3e1..58b19bbd0 100644 --- a/crates/matrix-sdk/src/room/mod.rs +++ b/crates/matrix-sdk/src/room/mod.rs @@ -29,7 +29,7 @@ use ruma::{ membership::{ ban_user, forget_room, get_member_events, invite_user::{self, v3::InvitationRecipient}, - join_room_by_id, kick_user, leave_room, Invite3pid, + join_room_by_id, kick_user, leave_room, unban_user, Invite3pid, }, message::send_message_event, read_marker::set_read_marker, @@ -1041,6 +1041,23 @@ impl Room { Ok(()) } + /// Unban the user with `UserId` from this room. + /// + /// # Arguments + /// + /// * `user_id` - The user to unban with `UserId`. + /// + /// * `reason` - The reason for unbanning this user. + #[instrument(skip_all)] + pub async fn unban_user(&self, user_id: &UserId, reason: Option<&str>) -> Result<()> { + let request = assign!( + unban_user::v3::Request::new(self.room_id().to_owned(), user_id.to_owned()), + { reason: reason.map(ToOwned::to_owned) } + ); + self.client.send(request, None).await?; + Ok(()) + } + /// Kick a user out of this room. /// /// # Arguments diff --git a/crates/matrix-sdk/src/sliding_sync/README.md b/crates/matrix-sdk/src/sliding_sync/README.md index 747f7d40f..4687dc2e3 100644 --- a/crates/matrix-sdk/src/sliding_sync/README.md +++ b/crates/matrix-sdk/src/sliding_sync/README.md @@ -388,7 +388,7 @@ _Note_: This is not yet exposed via the API. See [#1475](https://github.com/matr Sliding Sync is modeled for faster and more efficient user-facing client applications, but offers significant speed ups even for bot cases through its filtering mechanism. The sort-order and specific subsets, however, are -usually not of interest for bots. For that use case the the +usually not of interest for bots. For that use case the [`v4::SyncRequestList`][] offers the [`slow_get_all_rooms`](`v4::SyncRequestList::slow_get_all_rooms`) flag. diff --git a/crates/matrix-sdk/tests/integration/client.rs b/crates/matrix-sdk/tests/integration/client.rs index 6b527c381..71ad6cce0 100644 --- a/crates/matrix-sdk/tests/integration/client.rs +++ b/crates/matrix-sdk/tests/integration/client.rs @@ -279,18 +279,68 @@ async fn left_rooms() { async fn get_media_content() { let (client, server) = logged_in_client().await; + let media = client.media(); + let request = MediaRequest { source: MediaSource::Plain(mxc_uri!("mxc://localhost/textfile").to_owned()), format: MediaFormat::File, }; - Mock::given(method("GET")) - .and(path("/_matrix/media/r0/download/localhost/textfile")) - .respond_with(ResponseTemplate::new(200).set_body_string("Some very interesting text.")) - .mount(&server) - .await; + // First time, without the cache. + { + let expected_content = "Hello, World!"; + let _mock_guard = Mock::given(method("GET")) + .and(path("/_matrix/media/r0/download/localhost/textfile")) + .respond_with(ResponseTemplate::new(200).set_body_string(expected_content)) + .mount_as_scoped(&server) + .await; - client.media().get_media_content(&request, false).await.unwrap(); + assert_eq!( + media.get_media_content(&request, false).await.unwrap(), + expected_content.as_bytes() + ); + } + + // Second time, without the cache, error from the HTTP server. + { + let _mock_guard = Mock::given(method("GET")) + .and(path("/_matrix/media/r0/download/localhost/textfile")) + .respond_with(ResponseTemplate::new(500)) + .mount_as_scoped(&server) + .await; + + assert!(media.get_media_content(&request, false).await.is_err()); + } + + let expected_content = "Hello, World (2)!"; + + // Third time, with the cache. + { + let _mock_guard = Mock::given(method("GET")) + .and(path("/_matrix/media/r0/download/localhost/textfile")) + .respond_with(ResponseTemplate::new(200).set_body_string(expected_content)) + .mount_as_scoped(&server) + .await; + + assert_eq!( + media.get_media_content(&request, true).await.unwrap(), + expected_content.as_bytes() + ); + } + + // Third time, with the cache, the HTTP server isn't reached. + { + let _mock_guard = Mock::given(method("GET")) + .and(path("/_matrix/media/r0/download/localhost/textfile")) + .respond_with(ResponseTemplate::new(500)) + .mount_as_scoped(&server) + .await; + + assert_eq!( + client.media().get_media_content(&request, true).await.unwrap(), + expected_content.as_bytes() + ); + } } #[async_test] diff --git a/crates/matrix-sdk/tests/integration/refresh_token.rs b/crates/matrix-sdk/tests/integration/refresh_token.rs index ec4e9e491..55b50ade3 100644 --- a/crates/matrix-sdk/tests/integration/refresh_token.rs +++ b/crates/matrix-sdk/tests/integration/refresh_token.rs @@ -1,5 +1,4 @@ use std::{ - future::ready, sync::{Arc, Mutex}, time::Duration, }; @@ -178,8 +177,11 @@ async fn test_refresh_token() { .set_session_callbacks(Box::new(|_| panic!("reload session never called")), { let num_save_session_callback_calls = num_save_session_callback_calls.clone(); Box::new(move |_client| { - *num_save_session_callback_calls.lock().unwrap() += 1; - Box::pin(ready(Ok(()))) + let num_save_session_callback_calls = num_save_session_callback_calls.clone(); + Box::pin(async move { + *num_save_session_callback_calls.lock().unwrap() += 1; + Ok(()) + }) }) }) .unwrap(); diff --git a/crates/matrix-sdk/tests/integration/room/joined.rs b/crates/matrix-sdk/tests/integration/room/joined.rs index 285a838db..af03d0170 100644 --- a/crates/matrix-sdk/tests/integration/room/joined.rs +++ b/crates/matrix-sdk/tests/integration/room/joined.rs @@ -129,6 +129,29 @@ async fn ban_user() { room.ban_user(user, None).await.unwrap(); } +#[async_test] +async fn unban_user() { + let (client, server) = logged_in_client().await; + + Mock::given(method("POST")) + .and(path_regex(r"^/_matrix/client/r0/rooms/.*/unban$")) + .and(header("authorization", "Bearer 1234")) + .respond_with(ResponseTemplate::new(200).set_body_json(&*test_json::EMPTY)) + .mount(&server) + .await; + + mock_sync(&server, &*test_json::SYNC, None).await; + + let sync_settings = SyncSettings::new().timeout(Duration::from_millis(3000)); + + let _response = client.sync_once(sync_settings).await.unwrap(); + + let user = user_id!("@example:localhost"); + let room = client.get_room(&DEFAULT_TEST_ROOM_ID).unwrap(); + + room.unban_user(user, None).await.unwrap(); +} + #[async_test] async fn kick_user() { let (client, server) = logged_in_client().await; diff --git a/xtask/src/ci.rs b/xtask/src/ci.rs index 51f7d067c..0a8fc0876 100644 --- a/xtask/src/ci.rs +++ b/xtask/src/ci.rs @@ -1,4 +1,7 @@ -use std::collections::BTreeMap; +use std::{ + collections::BTreeMap, + env::consts::{DLL_PREFIX, DLL_SUFFIX}, +}; use clap::{Args, Subcommand}; use xshell::{cmd, pushd}; @@ -130,22 +133,22 @@ fn check_bindings() -> Result<()> { cmd!( " rustup run stable cargo run -p uniffi-bindgen -- generate + --library --language kotlin --language swift - --lib-file target/debug/libmatrix_sdk_ffi.a --out-dir target/generated-bindings - bindings/matrix-sdk-ffi/src/api.udl + target/debug/{DLL_PREFIX}matrix_sdk_ffi{DLL_SUFFIX} " ) .run()?; cmd!( " rustup run stable cargo run -p uniffi-bindgen -- generate + --library --language kotlin --language swift - --lib-file target/debug/libmatrix_sdk_crypto_ffi.a --out-dir target/generated-bindings - bindings/matrix-sdk-crypto-ffi/src/olm.udl + target/debug/{DLL_PREFIX}matrix_sdk_crypto_ffi{DLL_SUFFIX} " ) .run()?; diff --git a/xtask/src/swift.rs b/xtask/src/swift.rs index dc4f8a863..184b2df10 100644 --- a/xtask/src/swift.rs +++ b/xtask/src/swift.rs @@ -2,7 +2,7 @@ use std::fs::{copy, create_dir_all, remove_dir_all, remove_file, rename}; use camino::{Utf8Path, Utf8PathBuf}; use clap::{Args, Subcommand}; -use uniffi_bindgen::bindings::TargetLanguage; +use uniffi_bindgen::{bindings::TargetLanguage, library_mode::generate_bindings}; use xshell::{cmd, pushd}; use crate::{workspace, Result}; @@ -54,58 +54,38 @@ impl SwiftArgs { } } +const FFI_LIBRARY_NAME: &str = "libmatrix_sdk_ffi.a"; + fn build_library() -> Result<()> { println!("Running debug library build."); - let release_type = "debug"; - let static_lib_filename = "libmatrix_sdk_ffi.a"; - let root_directory = workspace::root_path()?; let target_directory = workspace::target_path()?; let ffi_directory = root_directory.join("bindings/apple/generated/matrix_sdk_ffi"); - let library_file = ffi_directory.join(static_lib_filename); + let lib_output_dir = target_directory.join("debug"); create_dir_all(ffi_directory.as_path())?; cmd!("rustup run stable cargo build -p matrix-sdk-ffi").run()?; - rename( - target_directory.join(release_type).join(static_lib_filename), - ffi_directory.join(static_lib_filename), - )?; + rename(lib_output_dir.join(FFI_LIBRARY_NAME), ffi_directory.join(FFI_LIBRARY_NAME))?; let swift_directory = root_directory.join("bindings/apple/generated/swift"); create_dir_all(swift_directory.as_path())?; - generate_uniffi(&library_file, &ffi_directory)?; + generate_uniffi(&ffi_directory.join(FFI_LIBRARY_NAME), &ffi_directory)?; let module_map_file = ffi_directory.join("module.modulemap"); if module_map_file.exists() { remove_file(module_map_file.as_path())?; } - // TODO: Find the modulemap in the ffi directory. rename(ffi_directory.join("matrix_sdk_ffiFFI.modulemap"), module_map_file)?; - // TODO: Move all swift files. - rename( - ffi_directory.join("matrix_sdk_ffi.swift"), - swift_directory.join("matrix_sdk_ffi.swift"), - )?; + move_swift_files(&ffi_directory, &swift_directory)?; Ok(()) } -fn generate_uniffi(library_file: &Utf8Path, ffi_directory: &Utf8Path) -> Result<()> { - let root_directory = workspace::root_path()?; - let udl_file = root_directory.join("bindings/matrix-sdk-ffi/src/api.udl"); - - uniffi_bindgen::generate_bindings( - udl_file.as_path(), - None, - vec![TargetLanguage::Swift], - Some(ffi_directory), - Some(library_file), - None, - false, - )?; +fn generate_uniffi(library_path: &Utf8Path, ffi_directory: &Utf8Path) -> Result<()> { + generate_bindings(library_path, None, &[TargetLanguage::Swift], None, ffi_directory, false)?; Ok(()) } @@ -113,7 +93,7 @@ fn build_path_for_target(target: &str, profile: &str) -> Result { // The builtin dev profile has its files stored under target/debug, all // other targets have matching directory names let profile_dir_name = if profile == "dev" { "debug" } else { profile }; - Ok(workspace::target_path()?.join(target).join(profile_dir_name).join("libmatrix_sdk_ffi.a")) + Ok(workspace::target_path()?.join(target).join(profile_dir_name).join(FFI_LIBRARY_NAME)) } fn build_xcframework( @@ -134,7 +114,7 @@ fn build_xcframework( create_dir_all(headers_dir.clone())?; create_dir_all(swift_dir.clone())?; - let (libs, uniff_lib_path) = if let Some(target) = only_target { + let (libs, uniffi_lib_path) = if let Some(target) = only_target { println!("-- Building for {target} 1/1"); cmd!( @@ -150,11 +130,11 @@ fn build_xcframework( cmd!( "rustup run stable cargo build -p matrix-sdk-ffi - --target aarch64-apple-ios - --target aarch64-apple-darwin - --target x86_64-apple-darwin - --target aarch64-apple-ios-sim - --target x86_64-apple-ios + --target aarch64-apple-ios + --target aarch64-apple-darwin + --target x86_64-apple-darwin + --target aarch64-apple-ios-sim + --target x86_64-apple-ios --profile {profile}" ) .run()?; @@ -186,7 +166,7 @@ fn build_xcframework( }; println!("-- Generating uniffi files"); - generate_uniffi(&uniff_lib_path, &generated_dir)?; + generate_uniffi(&uniffi_lib_path, &generated_dir)?; rename(generated_dir.join("matrix_sdk_ffiFFI.h"), headers_dir.join("matrix_sdk_ffiFFI.h"))?; @@ -197,7 +177,7 @@ fn build_xcframework( headers_dir.join("module.modulemap"), )?; - rename(generated_dir.join("matrix_sdk_ffi.swift"), swift_dir.join("matrix_sdk_ffi.swift"))?; + move_swift_files(&generated_dir, &swift_dir)?; println!("-- Generating MatrixSDKFFI.xcframework framework"); let xcframework_path = generated_dir.join("MatrixSDKFFI.xcframework"); @@ -248,3 +228,18 @@ fn build_xcframework( Ok(()) } + +fn move_swift_files(source: &Utf8PathBuf, destination: &Utf8PathBuf) -> Result<()> { + for entry in source.read_dir_utf8()? { + let entry = entry?; + + if entry.file_type()?.is_file() { + let path = entry.path(); + if path.extension() == Some("swift") { + let file_name = path.file_name().expect("Failed to get file name"); + rename(path, destination.join(file_name)).expect("Failed to move swift file"); + } + } + } + Ok(()) +}