diff --git a/crates/matrix-sdk-crypto/src/store/integration_tests.rs b/crates/matrix-sdk-crypto/src/store/integration_tests.rs index 932a31598..a5b9063c9 100644 --- a/crates/matrix-sdk-crypto/src/store/integration_tests.rs +++ b/crates/matrix-sdk-crypto/src/store/integration_tests.rs @@ -567,6 +567,62 @@ macro_rules! cryptostore_integration_tests { assert_eq!(store.inbound_group_session_counts(None).await.unwrap().total, 1); } + #[async_test] + async fn test_get_inbound_group_sessions_by_room_id_empty() { + let dir = "get_inbound_group_session_by_room_id_empty"; + let (_, store) = get_loaded_store(dir).await; + assert_eq!(store.get_inbound_group_sessions().await.unwrap().len(), 0); + + let room_id = &room_id!("!testing:localhost"); + assert_eq!(store.get_inbound_group_sessions_by_room_id(room_id).await.unwrap().len(), 0); + } + + #[async_test] + async fn test_get_inbound_group_sessions_by_room_id() { + let dir = "get_inbound_group_session_by_room_id"; + let (account, store) = get_loaded_store(dir).await; + assert_eq!(store.get_inbound_group_sessions().await.unwrap().len(), 0); + + let room_id = &room_id!("!testing:localhost"); + let (_, session_1) = account.create_group_session_pair_with_defaults(room_id).await; + let (_, session_2) = account.create_group_session_pair_with_defaults(room_id).await; + + let second_room_id = &room_id!("!other_room_testing:localhost"); + let (_, session_3) = account.create_group_session_pair_with_defaults(second_room_id).await; + + let mut sessions = vec![ + session_1, + session_2, + session_3 + ]; + + let changes = Changes { + inbound_group_sessions: sessions.clone(), + ..Default::default() + }; + store.save_changes(changes).await.expect("Can't save group session"); + + drop(store); + + // The last session is in a different room, so should not be returned by + // get_inbound_group_sessions_by_room_id. Remove it from the list. + sessions.pop(); + + let store = get_store(dir, None, false).await; + // Make sure all the sessions are in the store + assert_eq!(store.get_inbound_group_sessions().await.unwrap().len(), 3); + + store.load_account().await.unwrap(); + + let loaded_sessions = store + .get_inbound_group_sessions_by_room_id(room_id) + .await + .unwrap(); + + assert_eq!(loaded_sessions.len(), 2); + assert_session_lists_eq(sessions, loaded_sessions, "room by id sessions"); + } + #[async_test] async fn test_fetch_inbound_group_sessions_for_device() { // Given a store exists, containing inbound group sessions from different devices diff --git a/crates/matrix-sdk-crypto/src/store/memorystore.rs b/crates/matrix-sdk-crypto/src/store/memorystore.rs index ef032520a..05c5f9282 100644 --- a/crates/matrix-sdk-crypto/src/store/memorystore.rs +++ b/crates/matrix-sdk-crypto/src/store/memorystore.rs @@ -465,6 +465,25 @@ impl CryptoStore for MemoryStore { Ok(RoomKeyCounts { total, backed_up }) } + async fn get_inbound_group_sessions_by_room_id( + &self, + room_id: &RoomId, + ) -> Result> { + let inbounds = match self.inbound_group_sessions.read().get(room_id) { + None => Vec::new(), + Some(v) => v + .values() + .map(|ser| { + let pickle: PickledInboundGroupSession = + serde_json::from_str(ser).expect("Pickle deserialization should work"); + InboundGroupSession::from_pickle(pickle) + .expect("Expect from pickle to always work") + }) + .collect(), + }; + Ok(inbounds) + } + async fn get_inbound_group_sessions_for_device_batch( &self, sender_key: Curve25519PublicKey, @@ -1370,6 +1389,13 @@ mod integration_tests { self.0.inbound_group_session_counts(backup_version).await } + async fn get_inbound_group_sessions_by_room_id( + &self, + room_id: &RoomId, + ) -> Result, Self::Error> { + self.0.get_inbound_group_sessions_by_room_id(room_id).await + } + async fn get_inbound_group_sessions_for_device_batch( &self, sender_key: Curve25519PublicKey, diff --git a/crates/matrix-sdk-crypto/src/store/traits.rs b/crates/matrix-sdk-crypto/src/store/traits.rs index 50b6f3962..133b8458b 100644 --- a/crates/matrix-sdk-crypto/src/store/traits.rs +++ b/crates/matrix-sdk-crypto/src/store/traits.rs @@ -128,6 +128,15 @@ pub trait CryptoStore: AsyncTraitDeps { backup_version: Option<&str>, ) -> Result; + /// Get all the inbound group sessions for a given room. + /// + /// # Arguments + /// * `room_id` - The ID of the room to return sessions for. + async fn get_inbound_group_sessions_by_room_id( + &self, + room_id: &RoomId, + ) -> Result, Self::Error>; + /// Get a batch of inbound group sessions for the device with the supplied /// curve key, whose sender data is of the supplied type. /// @@ -434,6 +443,13 @@ impl CryptoStore for EraseCryptoStoreError { self.0.get_inbound_group_sessions().await.map_err(Into::into) } + async fn get_inbound_group_sessions_by_room_id( + &self, + room_id: &RoomId, + ) -> Result> { + self.0.get_inbound_group_sessions_by_room_id(room_id).await.map_err(Into::into) + } + async fn get_inbound_group_sessions_for_device_batch( &self, curve_key: Curve25519PublicKey, diff --git a/crates/matrix-sdk-indexeddb/src/crypto_store/mod.rs b/crates/matrix-sdk-indexeddb/src/crypto_store/mod.rs index 3d6ff0efb..2eeecd557 100644 --- a/crates/matrix-sdk-indexeddb/src/crypto_store/mod.rs +++ b/crates/matrix-sdk-indexeddb/src/crypto_store/mod.rs @@ -986,6 +986,31 @@ impl_crypto_store! { ).await } + async fn get_inbound_group_sessions_by_room_id( + &self, + room_id: &RoomId, + ) -> Result> { + let range = self.serializer.encode_to_range(keys::INBOUND_GROUP_SESSIONS_V3, room_id)?; + Ok(self + .inner + .transaction_on_one_with_mode( + keys::INBOUND_GROUP_SESSIONS_V3, + IdbTransactionMode::Readonly, + )? + .object_store(keys::INBOUND_GROUP_SESSIONS_V3)? + .get_all_with_key(&range)? + .await? + .into_iter() + .filter_map(|v| match self.deserialize_inbound_group_session(v) { + Ok(session) => Some(session), + Err(e) => { + warn!("Failed to deserialize inbound group session: {e}"); + None + } + }) + .collect::>()) + } + async fn get_inbound_group_sessions_for_device_batch( &self, sender_key: Curve25519PublicKey, diff --git a/crates/matrix-sdk-sqlite/src/crypto_store.rs b/crates/matrix-sdk-sqlite/src/crypto_store.rs index 5e6917988..fd3c42759 100644 --- a/crates/matrix-sdk-sqlite/src/crypto_store.rs +++ b/crates/matrix-sdk-sqlite/src/crypto_store.rs @@ -531,6 +531,24 @@ trait SqliteObjectCryptoStoreExt: SqliteAsyncConnExt { Ok(RoomKeyCounts { total, backed_up }) } + async fn get_inbound_group_sessions_by_room_id( + &self, + room_id: Key, + ) -> Result, bool)>> { + Ok(self + .prepare( + "SELECT data, backed_up FROM inbound_group_session WHERE room_id = :room_id", + move |mut stmt| { + stmt.query(named_params! { + ":room_id": room_id, + })? + .mapped(|row| Ok((row.get(0)?, row.get(1)?))) + .collect() + }, + ) + .await?) + } + async fn get_inbound_group_sessions_for_device_batch( &self, sender_key: Key, @@ -1057,6 +1075,22 @@ impl CryptoStore for SqliteCryptoStore { .collect() } + async fn get_inbound_group_sessions_by_room_id( + &self, + room_id: &RoomId, + ) -> Result> { + let room_id = self.encode_key("inbound_group_session", room_id.as_bytes()); + self.acquire() + .await? + .get_inbound_group_sessions_by_room_id(room_id) + .await? + .into_iter() + .map(|(value, backed_up)| { + self.deserialize_and_unpickle_inbound_group_session(value, backed_up) + }) + .collect() + } + async fn get_inbound_group_sessions_for_device_batch( &self, sender_key: Curve25519PublicKey,