feat(crypto): Wait for a key query to be done if we're claming one-time keys
This commit is contained in:
@@ -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" }
|
||||
|
||||
|
||||
@@ -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<event_listener::Event>,
|
||||
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<UserKeyQueryResult, Elapsed> {
|
||||
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<UserId>,
|
||||
device_id: Arc<DeviceId>,
|
||||
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<UserId>, device_id: Arc<DeviceId>, 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))
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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<DashMap<OwnedUserId, DashSet<OwnedDeviceId>>>,
|
||||
key_request_machine: GossipMachine,
|
||||
outgoing_to_device_requests: Arc<DashMap<OwnedTransactionId, OutgoingRequest>>,
|
||||
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<DashMap<OwnedUserId, DashSet<OwnedDeviceId>>>,
|
||||
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<HashMap<OwnedDeviceId, ReadOnlyDevice>> {
|
||||
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]
|
||||
|
||||
Reference in New Issue
Block a user