From ba7ccb40cc3111c6593bf642762b4c4c4439bc26 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Damir=20Jeli=C4=87?= Date: Thu, 2 Jun 2022 11:58:09 +0200 Subject: [PATCH] feat(crypto): Wait for a key query to be done if we're claming one-time keys --- crates/matrix-sdk-crypto/Cargo.toml | 2 + .../src/identities/manager.rs | 88 ++++++++++++++++++- .../matrix-sdk-crypto/src/identities/mod.rs | 2 +- crates/matrix-sdk-crypto/src/machine.rs | 7 +- .../src/session_manager/sessions.rs | 43 ++++++++- 5 files changed, 134 insertions(+), 8 deletions(-) diff --git a/crates/matrix-sdk-crypto/Cargo.toml b/crates/matrix-sdk-crypto/Cargo.toml index b2bf504c6..d4881aa09 100644 --- a/crates/matrix-sdk-crypto/Cargo.toml +++ b/crates/matrix-sdk-crypto/Cargo.toml @@ -33,6 +33,7 @@ bs58 = { version = "0.4.0", optional = true } byteorder = "1.4.3" ctr = "0.9.1" dashmap = "5.2.0" +event-listener = "2.5.2" futures-util = { version = "0.3.21", default-features = false, features = ["alloc"] } hmac = "0.12.1" http = { version = "0.2.6", optional = true } # feature = testing only @@ -49,6 +50,7 @@ tracing = "0.1.34" zeroize = { version = "1.3.0", features = ["zeroize_derive"] } [target.'cfg(not(target_arch = "wasm32"))'.dependencies] +tokio = { version = "1.18", default-features = false, features = ["time"] } ruma = { version = "0.6.2", features = ["client-api-c", "rand", "signatures", "unstable-msc2676", "unstable-msc2677"] } vodozemac = { git = "https://github.com/matrix-org/vodozemac/", rev = "d0e744287a14319c2a9148fef3747548c740fc36" } diff --git a/crates/matrix-sdk-crypto/src/identities/manager.rs b/crates/matrix-sdk-crypto/src/identities/manager.rs index 417e41d08..3d6986811 100644 --- a/crates/matrix-sdk-crypto/src/identities/manager.rs +++ b/crates/matrix-sdk-crypto/src/identities/manager.rs @@ -17,6 +17,7 @@ use std::{ convert::TryFrom, ops::Deref, sync::Arc, + time::Duration, }; use futures_util::future::join_all; @@ -46,10 +47,87 @@ enum DeviceChange { None, } +/// A listener that can notify if a `/keys/query` response has been received. +#[derive(Clone, Debug)] +pub(crate) struct KeysQueryListener { + inner: Arc, + store: Store, +} + +/// Error type notifying that a timeout has elapsed. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) struct Elapsed(()); + +/// Result type telling us if a `/keys/query` response was expected for a given +/// user. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum UserKeyQueryResult { + WasPending, + WasNotPending, +} + +impl KeysQueryListener { + pub(crate) fn new(store: Store) -> Self { + Self { inner: event_listener::Event::new().into(), store } + } + + /// Notify our listeners that we received a `/keys/query` response. + fn notify(&self) { + self.inner.notify(usize::MAX); + } + + /// Wait for a `/keys/query` response to be received if one is expected for + /// the given user. + /// + /// If the given timeout has elapsed the method will stop waiting and return + /// an error. + pub async fn wait_if_user_pending( + &self, + timeout: Duration, + user: &UserId, + ) -> Result { + let users_for_key_query = self.store.users_for_key_query(); + + if users_for_key_query.contains(user) { + if let Err(e) = self.wait(timeout).await { + warn!( + user_id =% user, + "The user has a pending `/key/query` request which did \ + not finish yet, some devices might be missing." + ); + + Err(e) + } else { + Ok(UserKeyQueryResult::WasPending) + } + } else { + Ok(UserKeyQueryResult::WasNotPending) + } + } + + /// Wait for a `/keys/query` response to be received. + /// + /// If the given timeout has elapsed the method will stop waiting and return + /// an error. + pub async fn wait(&self, timeout: Duration) -> Result<(), Elapsed> { + let listener = self.inner.listen(); + + #[cfg(not(target_arch = "wasm32"))] + tokio::time::timeout(timeout, async { listener.await }).await.map_err(|_| Elapsed(()))?; + + // TODO we should ensure that this is async on wasm as well. + #[cfg(target_arch = "wasm32")] + listener.wait_timeout(timeout); + + Ok(()) + } +} + #[derive(Debug, Clone)] pub(crate) struct IdentityManager { user_id: Arc, device_id: Arc, + keys_query_listener: KeysQueryListener, store: Store, } @@ -57,13 +135,19 @@ impl IdentityManager { const MAX_KEY_QUERY_USERS: usize = 250; pub fn new(user_id: Arc, device_id: Arc, store: Store) -> Self { - IdentityManager { user_id, device_id, store } + let keys_query_listener = KeysQueryListener::new(store.clone()); + + IdentityManager { user_id, device_id, store, keys_query_listener } } fn user_id(&self) -> &UserId { &self.user_id } + pub fn listen_for_received_queries(&self) -> KeysQueryListener { + self.keys_query_listener.clone() + } + /// Receive a successful keys query response. /// /// Returns a list of devices newly discovered devices and devices that @@ -129,6 +213,8 @@ impl IdentityManager { "Finished handling of the keys/query response" ); + self.keys_query_listener.notify(); + Ok((devices, identities)) } diff --git a/crates/matrix-sdk-crypto/src/identities/mod.rs b/crates/matrix-sdk-crypto/src/identities/mod.rs index d90206b36..1fad14145 100644 --- a/crates/matrix-sdk-crypto/src/identities/mod.rs +++ b/crates/matrix-sdk-crypto/src/identities/mod.rs @@ -50,7 +50,7 @@ use std::sync::{ }; pub use device::{Device, LocalTrust, ReadOnlyDevice, UserDevices}; -pub(crate) use manager::IdentityManager; +pub(crate) use manager::{IdentityManager, KeysQueryListener, UserKeyQueryResult}; use serde::{Deserialize, Deserializer, Serializer}; pub use user::{ MasterPubkey, OwnUserIdentity, ReadOnlyOwnUserIdentity, ReadOnlyUserIdentities, diff --git a/crates/matrix-sdk-crypto/src/machine.rs b/crates/matrix-sdk-crypto/src/machine.rs index a3bb2fd87..38fc02e91 100644 --- a/crates/matrix-sdk-crypto/src/machine.rs +++ b/crates/matrix-sdk-crypto/src/machine.rs @@ -166,15 +166,18 @@ impl OlmMachine { group_session_manager.session_cache(), users_for_key_claim.clone(), ); + let identity_manager = + IdentityManager::new(user_id.clone(), device_id.clone(), store.clone()); + + let event = identity_manager.listen_for_received_queries(); let session_manager = SessionManager::new( account.clone(), users_for_key_claim, key_request_machine.clone(), store.clone(), + event, ); - let identity_manager = - IdentityManager::new(user_id.clone(), device_id.clone(), store.clone()); #[cfg(feature = "backups_v1")] let backup_machine = BackupMachine::new(account.clone(), store.clone(), None); diff --git a/crates/matrix-sdk-crypto/src/session_manager/sessions.rs b/crates/matrix-sdk-crypto/src/session_manager/sessions.rs index 58404e27a..5c275ac4d 100644 --- a/crates/matrix-sdk-crypto/src/session_manager/sessions.rs +++ b/crates/matrix-sdk-crypto/src/session_manager/sessions.rs @@ -13,7 +13,7 @@ // limitations under the License. use std::{ - collections::{BTreeMap, BTreeSet}, + collections::{BTreeMap, BTreeSet, HashMap}, sync::Arc, time::Duration, }; @@ -33,6 +33,7 @@ use tracing::{debug, error, info, warn}; use crate::{ error::OlmResult, gossiping::GossipMachine, + identities::{KeysQueryListener, UserKeyQueryResult}, olm::Account, requests::{OutgoingRequest, ToDeviceRequest}, store::{Changes, Result as StoreResult, Store}, @@ -51,17 +52,20 @@ pub(crate) struct SessionManager { wedged_devices: Arc>>, key_request_machine: GossipMachine, outgoing_to_device_requests: Arc>, + keys_query_listener: KeysQueryListener, } impl SessionManager { const KEY_CLAIM_TIMEOUT: Duration = Duration::from_secs(10); const UNWEDGING_INTERVAL: Duration = Duration::from_secs(60 * 60); + const KEYS_QUERY_WAIT_TIME: Duration = Duration::from_secs(5); pub fn new( account: Account, users_for_key_claim: Arc>>, key_request_machine: GossipMachine, store: Store, + keys_query_listener: KeysQueryListener, ) -> Self { Self { account, @@ -70,6 +74,7 @@ impl SessionManager { users_for_key_claim, wedged_devices: Default::default(), outgoing_to_device_requests: Default::default(), + keys_query_listener, } } @@ -155,6 +160,30 @@ impl SessionManager { Ok(()) } + async fn get_user_devices( + &self, + user_id: &UserId, + ) -> StoreResult> { + use UserKeyQueryResult::*; + + let user_devices = self.store.get_readonly_devices_filtered(user_id).await?; + + let user_devices = if user_devices.is_empty() { + match self + .keys_query_listener + .wait_if_user_pending(Self::KEYS_QUERY_WAIT_TIME, user_id) + .await + { + Ok(WasPending) => self.store.get_readonly_devices_filtered(user_id).await?, + _ => user_devices, + } + } else { + user_devices + }; + + Ok(user_devices) + } + /// Get the a key claiming request for the user/device pairs that we are /// missing Olm sessions for. /// @@ -191,7 +220,7 @@ impl SessionManager { // Add the list of devices that the user wishes to establish sessions // right now. for user_id in users { - let user_devices = self.store.get_readonly_devices_filtered(user_id).await?; + let user_devices = self.get_user_devices(user_id).await?; for (device_id, device) in user_devices { if !device.algorithms().contains(&EventEncryptionAlgorithm::OlmV1Curve25519AesSha2) @@ -354,7 +383,7 @@ mod tests { use super::SessionManager; use crate::{ gossiping::GossipMachine, - identities::ReadOnlyDevice, + identities::{KeysQueryListener, ReadOnlyDevice}, olm::{Account, PrivateCrossSigningIdentity, ReadOnlyAccount}, session_manager::GroupSessionCache, store::{CryptoStore, MemoryStore, Store}, @@ -402,7 +431,13 @@ mod tests { users_for_key_claim.clone(), ); - SessionManager::new(account, users_for_key_claim, key_request, store) + SessionManager::new( + account, + users_for_key_claim, + key_request, + store.clone(), + KeysQueryListener::new(store), + ) } #[async_test]