Merge branch 'main' into andybalaam/mark_sessions_as_backed_up

This commit is contained in:
Andy Balaam
2023-12-14 15:25:43 +00:00
29 changed files with 1074 additions and 308 deletions
+52 -57
View File
@@ -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<MediaSource>,
body: Option<String>,
@@ -267,24 +267,22 @@ impl Client {
use_cache: bool,
temp_dir: Option<String>,
) -> Result<Arc<MediaFileHandle>, 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<Self>,
@@ -488,68 +486,65 @@ impl Client {
})
}
pub fn upload_media(
pub async fn upload_media(
&self,
mime_type: String,
data: Vec<u8>,
progress_watcher: Option<Box<dyn ProgressWatcher>>,
) -> Result<String, ClientError> {
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<MediaSource>,
) -> Result<Vec<u8>, 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<MediaSource>,
width: u64,
height: u64,
) -> Result<Vec<u8>, 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(
+27
View File
@@ -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<String>,
) -> 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<String>,
) -> 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<bool, ClientError> {
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<String>,
) -> 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,
@@ -31,11 +31,17 @@ impl SessionVerificationEmoji {
}
}
#[derive(uniffi::Enum)]
pub enum SessionVerificationData {
Emojis { emojis: Vec<Arc<SessionVerificationEmoji>>, indices: Vec<u8> },
Decimals { values: Vec<u16> },
}
#[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<Arc<SessionVerificationEmoji>>);
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::<Vec<_>>();
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 { .. } => {
+32 -11
View File
@@ -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<Self>,
url: String,
thumbnail_url: String,
thumbnail_url: Option<String>,
image_info: ImageInfo,
progress_watcher: Option<Box<dyn ProgressWatcher>>,
) -> Arc<SendAttachmentJoinHandle> {
@@ -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<Self>,
url: String,
thumbnail_url: String,
thumbnail_url: Option<String>,
video_info: VideoInfo,
progress_watcher: Option<Box<dyn ProgressWatcher>>,
) -> Arc<SendAttachmentJoinHandle> {
@@ -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<ReceiptType> 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,
}
}
}
+56 -1
View File
@@ -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);
}
}
@@ -200,11 +200,8 @@ impl StateStoreIntegrationTests for DynStateStore {
async fn test_media_content(&self) {
let uri = mxc_uri!("mxc://localhost/media");
let content: Vec<u8> = "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<u8> = "hello".into();
let thumbnail_content: Vec<u8> = "world".into();
let other_content: Vec<u8> = "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<()> {
@@ -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<String>), HashMap<OwnedEventId, HashMap<OwnedUserId, Receipt>>>,
>,
>,
media: StdRwLock<RingBuffer<(OwnedMxcUri, String /* unique key */, Vec<u8>)>>,
custom: StdRwLock<HashMap<Vec<u8>, Vec<u8>>>,
}
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<u8>) -> Result<()> {
async fn add_media_content(&self, request: &MediaRequest, data: Vec<u8>) -> 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<Option<Vec<u8>>> {
Ok(None)
async fn get_media_content(&self, request: &MediaRequest) -> Result<Option<Vec<u8>>> {
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::<Vec<_>>();
// 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);
}
+1 -1
View File
@@ -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 }
+36 -4
View File
@@ -75,12 +75,19 @@ impl<T> RingBuffer<T> {
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<T> {
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<R>(&mut self, range: R) -> Drain<'_, T>
where
R: RangeBounds<usize>,
@@ -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]
@@ -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
@@ -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::<Vec<_>>()))
.format(", ")
);
let missing_session_devices_by_user = missing_session_devices_by_user
.iter()
.map(|(user_id, devices)| (user_id, devices.keys().collect::<BTreeSet<_>>()))
.collect::<BTreeMap<_, _>>();
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"
);
@@ -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.
+74 -53
View File
@@ -611,7 +611,7 @@ trait SqliteObjectStateStoreExt: SqliteObjectExt {
async fn get_kv_blobs(&self, keys: Vec<Key>) -> Result<Vec<Vec<u8>>> {
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<Key>,
) -> Result<Vec<(bool, Vec<u8>)>> {
chunk_large_query_over(state_keys, None, move |state_keys: Vec<Key>| {
self.chunk_large_query_over(state_keys, None, move |state_keys: Vec<Key>| {
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<(Vec<u8>, Vec<u8>)>> {
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<(Vec<u8>, Vec<u8>)>> {
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<Query, Fut, Res>(
&self,
mut keys_to_chunk: Vec<Key>,
result_capacity: Option<usize>,
do_query: Query,
) -> Result<Vec<Res>>
where
Query: Fn(Vec<Key>) -> Fut + Send + Sync,
Fut: Future<Output = Result<Vec<Res>, 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<Query, Fut, Res>(
mut keys_to_chunk: Vec<Key>,
result_capacity: Option<usize>,
do_query: Query,
) -> Result<Vec<Res>>
where
Query: Fn(Vec<Key>) -> Fut,
Fut: Future<Output = Result<Vec<Res>, 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};
+4 -4
View File
@@ -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<rusqlite::Error> + Send + 'static,
F: FnOnce(&Transaction<'_>) -> Result<T, E> + 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")
}
}
+4
View File
@@ -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"]
+1 -1
View File
@@ -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")]
+4 -2
View File
@@ -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
+6 -6
View File
@@ -264,12 +264,12 @@ impl Media {
request: &MediaRequest,
use_cache: bool,
) -> Result<Vec<u8>> {
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<u8> = match &request.source {
MediaSource::Encrypted(file) => {
@@ -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(
@@ -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<bool> 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<String> {
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
);
}
}
@@ -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,
@@ -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<String> {
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<IsOneToOne> 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();
+18 -1
View File
@@ -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
+1 -1
View File
@@ -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.
+56 -6
View File
@@ -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]
@@ -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();
@@ -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;
+8 -5
View File
@@ -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()?;
+33 -38
View File
@@ -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<Utf8PathBuf> {
// 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(())
}