use std::{future::IntoFuture, time::Duration};
use futures_core::Stream;
use matrix_sdk_base::crypto::{
OlmError,
dehydrated_devices::{DehydrationError, RehydratedDevice},
store::types::DehydratedDeviceKey,
vodozemac::base64_decode,
};
use matrix_sdk_common::{
boxed_into_future, locks::Mutex as StdMutex, sleep::sleep, task_monitor::BackgroundTaskHandle,
};
use ruma::{
OwnedDeviceId,
api::{
client::dehydrated_device::{
DehydratedDeviceData, delete_dehydrated_device, get_dehydrated_device, get_events,
},
error::ErrorKind,
},
events::secret::request::SecretName,
serde::Raw,
};
use thiserror::Error;
use tokio::sync::broadcast;
use tokio_stream::wrappers::{BroadcastStream, errors::BroadcastStreamRecvError};
use tracing::{Instrument, Span, debug, info, instrument, trace, warn};
use zeroize::Zeroizing;
use crate::{
Client, HttpError,
client::WeakClient,
encryption::{CryptoStoreError, secret_storage::SecretStore},
};
const DEFAULT_DEVICE_DISPLAY_NAME: &str = "Dehydrated device";
const PICKLE_KEY_SECRET_NAME: &str = "org.matrix.msc3814";
const DEHYDRATION_INTERVAL: Duration = Duration::from_secs(7 * 24 * 60 * 60);
const LAST_UPLOADED_DEVICE_ID_KEY: &str = "matrix-sdk-dehydrated-devices.last-uploaded-device-id";
const MAX_TO_DEVICE_EVENTS: usize = 100_000;
#[derive(Debug, Error)]
pub enum DehydratedDeviceError {
#[error(transparent)]
Http(#[from] HttpError),
#[error(transparent)]
Crypto(#[from] DehydrationError),
#[error(transparent)]
Olm(#[from] OlmError),
#[error(
"the to-device drain stopped after {to_device_events} events with more still queued; the dehydrated device was kept so a retry can resume"
)]
DrainTruncated {
to_device_events: usize,
},
#[error(transparent)]
Store(#[from] CryptoStoreError),
#[error(transparent)]
SecretStorage(#[from] crate::encryption::secret_storage::SecretStorageError),
#[error("the dehydrated-device pickle key in Secret Storage is not valid base64: {0}")]
PickleKeyDecode(#[from] vodozemac::Base64DecodeError),
#[error("the client is not logged in")]
NotLoggedIn,
}
fn pickle_key_secret_name() -> SecretName {
SecretName::from(PICKLE_KEY_SECRET_NAME)
}
#[derive(Clone, Debug)]
pub enum DehydratedDeviceEvent {
Created {
device_id: OwnedDeviceId,
},
Uploaded {
device_id: OwnedDeviceId,
},
Deleted,
KeyCached,
RehydrationStarted {
device_id: OwnedDeviceId,
},
RehydrationProgress {
room_keys_imported: usize,
to_device_events: usize,
},
RehydrationCompleted {
device_id: OwnedDeviceId,
room_keys_imported: usize,
to_device_events: usize,
},
RehydrationError {
error: String,
},
RotationError {
error: String,
},
}
pub(crate) struct DehydratedDevicesState {
event_sender: broadcast::Sender<DehydratedDeviceEvent>,
rotation_task: StdMutex<Option<BackgroundTaskHandle>>,
}
impl Default for DehydratedDevicesState {
fn default() -> Self {
let (event_sender, _) = broadcast::channel(100);
Self { event_sender, rotation_task: StdMutex::new(None) }
}
}
#[derive(Debug, Clone)]
pub struct DehydratedDevices {
pub(super) client: Client,
}
struct DownloadedDevice {
device_id: OwnedDeviceId,
device_data: Raw<DehydratedDeviceData>,
}
struct DrainOutcome {
room_keys_imported: usize,
to_device_events: usize,
truncated: bool,
}
impl DehydratedDevices {
pub fn state_stream(
&self,
) -> impl Stream<Item = Result<DehydratedDeviceEvent, BroadcastStreamRecvError>> + use<> {
BroadcastStream::new(self.state().event_sender.subscribe())
}
fn state(&self) -> &DehydratedDevicesState {
&self.client.inner.e2ee.dehydrated_devices_state
}
fn emit(&self, event: DehydratedDeviceEvent) {
let _ = self.state().event_sender.send(event);
}
#[instrument(skip_all)]
pub async fn is_supported(&self) -> Result<bool, DehydratedDeviceError> {
let request = get_dehydrated_device::unstable::Request::new();
match self.client.send(request).await {
Ok(_) => Ok(true),
Err(e) => match e.client_api_error_kind() {
Some(ErrorKind::Unrecognized) => Ok(false),
Some(ErrorKind::NotFound) => Ok(true),
_ => Err(e.into()),
},
}
}
#[instrument(skip_all)]
pub async fn create(
&self,
display_name: Option<&str>,
pickle_key: &DehydratedDeviceKey,
) -> Result<OwnedDeviceId, DehydratedDeviceError> {
let olm = self.client.olm_machine().await;
let machine = olm.as_ref().ok_or(DehydratedDeviceError::NotLoggedIn)?;
debug!("Creating a new dehydrated device in the crypto store");
let dehydrated_device = machine.dehydrated_devices().create().await?;
let display_name = display_name.unwrap_or(DEFAULT_DEVICE_DISPLAY_NAME);
let request =
dehydrated_device.keys_for_upload(display_name.to_owned(), pickle_key).await?;
let device_id = request.device_id.clone();
self.emit(DehydratedDeviceEvent::Created { device_id: device_id.clone() });
debug!(?device_id, "Uploading dehydrated device to the homeserver");
self.client.send(request).await?;
info!(?device_id, "Successfully uploaded dehydrated device");
machine.store().set_value(LAST_UPLOADED_DEVICE_ID_KEY, &device_id).await?;
self.emit(DehydratedDeviceEvent::Uploaded { device_id: device_id.clone() });
Ok(device_id)
}
#[instrument(skip_all)]
pub async fn rehydrate(
&self,
pickle_key: &DehydratedDeviceKey,
) -> Result<bool, DehydratedDeviceError> {
let Some(downloaded) = self.download_device().await? else { return Ok(false) };
info!(device_id = ?downloaded.device_id, "Dehydrated device found");
if let Some(expected) = self.last_uploaded_device_id().await?
&& expected != downloaded.device_id
{
warn!(
?expected,
got = ?downloaded.device_id,
"Server returned a different dehydrated-device id than the one we last uploaded; continuing but the payload may be stale"
);
}
self.emit(DehydratedDeviceEvent::RehydrationStarted {
device_id: downloaded.device_id.clone(),
});
let rehydrated = self.rehydrate_device(&downloaded, pickle_key).await?;
let drained = self.absorb_events(&downloaded.device_id, &rehydrated).await?;
if drained.truncated {
return Err(DehydratedDeviceError::DrainTruncated {
to_device_events: drained.to_device_events,
});
}
self.emit(DehydratedDeviceEvent::RehydrationCompleted {
device_id: downloaded.device_id.clone(),
room_keys_imported: drained.room_keys_imported,
to_device_events: drained.to_device_events,
});
if let Err(e) = self.delete_device().await {
warn!(device_id = ?downloaded.device_id, error = %e, "Post-rehydration delete failed; the next rotation will replace the device");
}
Ok(true)
}
#[instrument(skip_all)]
pub(crate) async fn cache_key(
&self,
pickle_key: &DehydratedDeviceKey,
) -> Result<(), DehydratedDeviceError> {
let olm = self.client.olm_machine().await;
let machine = olm.as_ref().ok_or(DehydratedDeviceError::NotLoggedIn)?;
machine.dehydrated_devices().save_dehydrated_device_pickle_key(pickle_key).await?;
self.emit(DehydratedDeviceEvent::KeyCached);
Ok(())
}
#[instrument(skip_all)]
pub(crate) async fn cached_key(
&self,
) -> Result<Option<DehydratedDeviceKey>, DehydratedDeviceError> {
let olm = self.client.olm_machine().await;
let machine = olm.as_ref().ok_or(DehydratedDeviceError::NotLoggedIn)?;
Ok(machine.dehydrated_devices().get_dehydrated_device_pickle_key().await?)
}
async fn last_uploaded_device_id(
&self,
) -> Result<Option<OwnedDeviceId>, DehydratedDeviceError> {
let olm = self.client.olm_machine().await;
let machine = olm.as_ref().ok_or(DehydratedDeviceError::NotLoggedIn)?;
Ok(machine.store().get_value(LAST_UPLOADED_DEVICE_ID_KEY).await?)
}
#[instrument(skip_all)]
pub async fn is_key_stored(
&self,
secret_store: &SecretStore,
) -> Result<bool, DehydratedDeviceError> {
Ok(secret_store.get_secret(pickle_key_secret_name()).await?.is_some())
}
#[instrument(skip_all)]
pub async fn reset_key(
&self,
secret_store: &SecretStore,
) -> Result<DehydratedDeviceKey, DehydratedDeviceError> {
let key = DehydratedDeviceKey::new();
secret_store.put_secret(pickle_key_secret_name(), &key.to_base64()).await?;
self.cache_key(&key).await?;
Ok(key)
}
#[instrument(skip_all)]
pub(crate) async fn load_key(
&self,
secret_store: &SecretStore,
create_if_missing: bool,
) -> Result<Option<DehydratedDeviceKey>, DehydratedDeviceError> {
if let Some(cached) = self.cached_key().await? {
return Ok(Some(cached));
}
let Some(base64) = secret_store.get_secret(pickle_key_secret_name()).await? else {
return if create_if_missing {
Ok(Some(self.reset_key(secret_store).await?))
} else {
Ok(None)
};
};
let bytes = Zeroizing::new(base64_decode(&base64)?);
let key = DehydratedDeviceKey::from_slice(&bytes)?;
self.cache_key(&key).await?;
Ok(Some(key))
}
pub fn start<'a>(&'a self, secret_store: &'a SecretStore) -> StartDehydration<'a> {
StartDehydration::new(self, secret_store)
}
pub fn stop(&self) {
self.state().rotation_task.lock().take();
}
async fn schedule_dehydration(
&self,
secret_store: &SecretStore,
) -> Result<(), DehydratedDeviceError> {
let key = self
.load_key(secret_store, true)
.await?
.expect("load_key(create_if_missing=true) always yields a key");
self.create(None, &key).await?;
let weak_client = WeakClient::from_client(&self.client);
let handle = self
.client
.task_monitor()
.spawn_infinite_task("dehydrated_devices::rotation", async move {
loop {
sleep(DEHYDRATION_INTERVAL).await;
let Some(client) = weak_client.get() else {
continue;
};
client.encryption().dehydrated_devices().rotate_tick().await;
}
})
.abort_on_drop();
*self.state().rotation_task.lock() = Some(handle);
Ok(())
}
async fn rotate_tick(&self) {
let key = match self.cached_key().await {
Ok(Some(key)) => key,
Ok(None) => {
let msg = "no cached pickle key for dehydrated-device rotation".to_owned();
warn!("{msg}; skipping this rotation");
self.emit(DehydratedDeviceEvent::RotationError { error: msg });
return;
}
Err(e) => {
let msg = e.to_string();
warn!(error = %e, "Failed to load cached pickle key for rotation");
self.emit(DehydratedDeviceEvent::RotationError { error: msg });
return;
}
};
if let Err(e) = self.create(None, &key).await {
let msg = e.to_string();
warn!(error = msg, "Failed to rotate dehydrated device");
self.emit(DehydratedDeviceEvent::RotationError { error: msg });
}
}
#[instrument(skip_all)]
pub async fn delete(&self) -> Result<(), DehydratedDeviceError> {
self.stop();
self.delete_device().await
}
async fn delete_device(&self) -> Result<(), DehydratedDeviceError> {
let request = delete_dehydrated_device::unstable::Request::new();
match self.client.send(request).await {
Ok(_) => {
self.emit(DehydratedDeviceEvent::Deleted);
Ok(())
}
Err(e) => match e.client_api_error_kind() {
Some(ErrorKind::Unrecognized) | Some(ErrorKind::NotFound) => Ok(()),
_ => Err(e.into()),
},
}
}
async fn download_device(&self) -> Result<Option<DownloadedDevice>, DehydratedDeviceError> {
let request = get_dehydrated_device::unstable::Request::new();
match self.client.send(request).await {
Ok(response) => Ok(Some(DownloadedDevice {
device_id: response.device_id,
device_data: response.device_data,
})),
Err(e) => match e.client_api_error_kind() {
Some(ErrorKind::NotFound) | Some(ErrorKind::Unrecognized) => Ok(None),
_ => Err(e.into()),
},
}
}
async fn rehydrate_device(
&self,
downloaded: &DownloadedDevice,
pickle_key: &DehydratedDeviceKey,
) -> Result<RehydratedDevice, DehydratedDeviceError> {
let olm = self.client.olm_machine().await;
let machine = olm.as_ref().ok_or(DehydratedDeviceError::NotLoggedIn)?;
Ok(machine
.dehydrated_devices()
.rehydrate(pickle_key, &downloaded.device_id, downloaded.device_data.clone())
.await?)
}
async fn absorb_events(
&self,
device_id: &OwnedDeviceId,
rehydrated: &RehydratedDevice,
) -> Result<DrainOutcome, DehydratedDeviceError> {
let settings = self.client.decryption_settings();
let mut next_batch: Option<String> = None;
let mut to_device_count: usize = 0;
let mut room_key_count: usize = 0;
let truncated = loop {
let mut request = get_events::unstable::Request::new(device_id.clone());
request.next_batch.clone_from(&next_batch);
let response = self.client.send(request).await?;
if response.events.is_empty() {
break false;
}
to_device_count = to_device_count.saturating_add(response.events.len());
let imported = rehydrated.receive_events(response.events, settings).await?;
room_key_count = room_key_count.saturating_add(imported.len());
trace!(to_device_count, room_key_count, "Absorbed a batch of to-device events");
self.emit(DehydratedDeviceEvent::RehydrationProgress {
room_keys_imported: room_key_count,
to_device_events: to_device_count,
});
if to_device_count >= MAX_TO_DEVICE_EVENTS {
warn!(
to_device_count,
"Reached the dehydrated-device drain limit; stopping to bound the loop"
);
break true;
}
match &response.next_batch {
None => break false,
Some(token) if Some(token) == next_batch.as_ref() => {
warn!(
?next_batch,
"Server returned the same next_batch twice; stopping to avoid a loop"
);
break true;
}
Some(token) => next_batch = Some(token.clone()),
}
};
info!(
to_device_count,
room_key_count, truncated, "Finished draining the dehydrated device's to-device queue"
);
Ok(DrainOutcome {
room_keys_imported: room_key_count,
to_device_events: to_device_count,
truncated,
})
}
}
#[derive(Debug)]
pub struct StartDehydration<'a> {
devices: &'a DehydratedDevices,
secret_store: &'a SecretStore,
create_new_key: bool,
rehydrate: bool,
only_if_key_cached: bool,
tracing_span: Span,
}
impl<'a> StartDehydration<'a> {
fn new(devices: &'a DehydratedDevices, secret_store: &'a SecretStore) -> Self {
Self {
devices,
secret_store,
create_new_key: false,
rehydrate: true,
only_if_key_cached: false,
tracing_span: Span::current(),
}
}
pub fn create_new_key(mut self) -> Self {
self.create_new_key = true;
self
}
pub fn skip_rehydration(mut self) -> Self {
self.rehydrate = false;
self
}
pub fn only_if_key_cached(mut self) -> Self {
self.only_if_key_cached = true;
self
}
}
impl<'a> IntoFuture for StartDehydration<'a> {
type Output = Result<(), DehydratedDeviceError>;
boxed_into_future!(extra_bounds: 'a);
fn into_future(self) -> Self::IntoFuture {
let Self {
devices,
secret_store,
create_new_key,
rehydrate,
only_if_key_cached,
tracing_span,
} = self;
let future = async move {
if only_if_key_cached && devices.cached_key().await?.is_none() {
return Ok(());
}
devices.stop();
let mut rehydrate_failed = false;
if rehydrate
&& let Some(key) = devices.load_key(secret_store, false).await?
&& let Err(e) = devices.rehydrate(&key).await
{
let msg = e.to_string();
warn!(error = %e, "Rehydration failed during start; continuing");
devices.emit(DehydratedDeviceEvent::RehydrationError { error: msg });
rehydrate_failed = true;
}
if create_new_key {
if rehydrate_failed {
warn!(
"Skipping pickle-key reset after failed rehydration to preserve the chance of recovering the existing dehydrated device on another client"
);
} else {
devices.reset_key(secret_store).await?;
}
}
devices.schedule_dehydration(secret_store).await
};
Box::pin(future.instrument(tracing_span))
}
}