use std::collections::{BTreeMap, HashMap, HashSet};
use base64::engine::general_purpose::STANDARD_NO_PAD;
use base64::Engine;
use ed25519_dalek::{Signature, VerifyingKey};
use rusqlite::Connection;
use crate::error::MatrixError;
use crate::keys::CrossSigningUsage;
use crate::store::{Membership, TxnDedupEntry};
pub const SEND_TO_DEVICE_MAX_TARGETS: usize = 1000;
pub const SEND_TO_DEVICE_MAX_CONTENT_BYTES: usize = 64 * 1024;
pub const BACKUP_ALGORITHM: &str = "m.megolm_backup.v1.aes-hmac-sha2";
pub fn peers_sharing_a_room_with(conn: &Connection, user_id: i64) -> rusqlite::Result<HashSet<i64>> {
let mut ids = HashSet::new();
ids.insert(user_id);
for room_id in crate::store::rooms_for_user(conn, user_id, Some(Membership::Join))? {
for member in crate::store::room_members(conn, &room_id, None)? {
if matches!(member.membership, Membership::Join | Membership::Invite) {
ids.insert(member.user_id);
}
}
}
Ok(ids)
}
pub struct DeviceListDelta {
pub changed: Vec<String>,
pub left: Vec<String>,
}
pub fn shared_room_transitions(conn: &Connection, caller_user_id: i64, from_exclusive: i64, to_inclusive: i64) -> rusqlite::Result<(Vec<i64>, Vec<i64>)> {
let mut entered = Vec::new();
let mut departed = Vec::new();
let Some(caller_mxid) = crate::store::mxid_of(conn, caller_user_id)? else {
return Ok((entered, departed));
};
let in_room = |membership: Membership| matches!(membership, Membership::Join | Membership::Invite);
let caller_rooms = crate::store::rooms_for_user(conn, caller_user_id, None)?;
for room_id in crate::store::rooms_with_member_events_in_window(conn, &caller_rooms, from_exclusive, to_inclusive)? {
let Some(room) = crate::store::get_room(conn, &room_id)? else { continue };
if !room.is_encrypted {
continue;
}
let Some(own) = crate::store::room_member(conn, &room_id, caller_user_id)? else { continue };
let was_joined = crate::store::membership_at(conn, &room_id, &caller_mxid, from_exclusive)? == Some(Membership::Join);
match own.membership {
Membership::Join if was_joined => {
for mxid in crate::store::member_state_keys_in_window(conn, &room_id, from_exclusive, to_inclusive)? {
let Some(user_id) = crate::store::user_id_of(conn, &mxid)? else { continue };
if user_id == caller_user_id {
continue;
}
let now_in = crate::store::room_member(conn, &room_id, user_id)?.is_some_and(|m| in_room(m.membership));
let before_in = crate::store::membership_at(conn, &room_id, &mxid, from_exclusive)?.is_some_and(in_room);
if now_in && !before_in {
entered.push(user_id);
}
}
}
Membership::Join => {
for member in crate::store::room_members(conn, &room_id, None)? {
if member.user_id != caller_user_id && in_room(member.membership) {
entered.push(member.user_id);
}
}
}
Membership::Leave | Membership::Ban if was_joined => {
for member in crate::store::room_members(conn, &room_id, None)? {
if member.user_id != caller_user_id && in_room(member.membership) {
departed.push(member.user_id);
}
}
}
Membership::Leave | Membership::Ban | Membership::Invite => {}
}
}
Ok((entered, departed))
}
pub fn sorted_mxids(conn: &Connection, user_ids: &HashSet<i64>) -> rusqlite::Result<Vec<String>> {
let mut mxids = Vec::with_capacity(user_ids.len());
for &user_id in user_ids {
if let Some(mxid) = crate::store::mxid_of(conn, user_id)? {
mxids.push(mxid);
}
}
mxids.sort();
mxids.dedup();
Ok(mxids)
}
pub fn device_list_delta(conn: &Connection, caller_user_id: i64, from_exclusive: i64, to_inclusive: i64) -> rusqlite::Result<DeviceListDelta> {
let visible = peers_sharing_a_room_with(conn, caller_user_id)?;
let (entered, departed) = shared_room_transitions(conn, caller_user_id, from_exclusive, to_inclusive)?;
let mut changed_ids: HashSet<i64> = HashSet::new();
for user_id in crate::keys::device_list_changes_between(conn, from_exclusive, to_inclusive)? {
if visible.contains(&user_id) {
changed_ids.insert(user_id);
}
}
changed_ids.extend(entered.into_iter().filter(|user_id| visible.contains(user_id)));
let caller_rooms = crate::store::rooms_for_user(conn, caller_user_id, None)?;
let mut left_ids: HashSet<i64> = HashSet::new();
for user_id in crate::store::user_ids_with_leave_transition_in_rooms(conn, &caller_rooms, from_exclusive, to_inclusive)? {
if !visible.contains(&user_id) {
left_ids.insert(user_id);
}
}
left_ids.extend(departed.into_iter().filter(|user_id| !visible.contains(user_id)));
Ok(DeviceListDelta { changed: sorted_mxids(conn, &changed_ids)?, left: sorted_mxids(conn, &left_ids)? })
}
#[derive(serde::Deserialize, Default)]
pub struct KeysUploadRequest {
#[serde(default)]
pub device_keys: Option<serde_json::Value>,
#[serde(default)]
pub one_time_keys: Option<BTreeMap<String, serde_json::Value>>,
#[serde(default)]
pub fallback_keys: Option<BTreeMap<String, serde_json::Value>>,
}
pub fn check_device_keys_ownership(device_keys: &serde_json::Value, caller_mxid: &str, caller_device_id: &str) -> Result<(), MatrixError> {
let user_id = device_keys.get("user_id").and_then(|v| v.as_str());
let device_id = device_keys.get("device_id").and_then(|v| v.as_str());
if user_id != Some(caller_mxid) || device_id != Some(caller_device_id) {
return Err(MatrixError::invalid_param("device_keys.user_id/device_id must name the authenticated caller"));
}
Ok(())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DeviceKeysUploadAction {
Insert,
Noop,
Reset,
}
pub fn decide_device_keys_upload(existing_keys_json: Option<&str>, new_keys_value: &serde_json::Value) -> Result<DeviceKeysUploadAction, MatrixError> {
let Some(existing_keys_json) = existing_keys_json else {
return Ok(DeviceKeysUploadAction::Insert);
};
let existing_keys: serde_json::Value = serde_json::from_str(existing_keys_json)?;
if &existing_keys == new_keys_value {
Ok(DeviceKeysUploadAction::Noop)
} else {
Ok(DeviceKeysUploadAction::Reset)
}
}
pub fn split_algorithm_key_id(full_id: &str) -> Result<(&str, &str), MatrixError> {
let (algorithm, key_id) = full_id
.split_once(':')
.ok_or_else(|| MatrixError::invalid_param(format!("malformed one-time key id: {full_id}")))?;
if algorithm.is_empty() || key_id.is_empty() {
return Err(MatrixError::invalid_param(format!("malformed one-time key id: {full_id}")));
}
Ok((algorithm, key_id))
}
pub struct KeysUploadOutcome {
pub otk_counts: HashMap<String, i64>,
pub device_keys_changed: bool,
}
pub fn apply_keys_upload(
conn: &mut Connection,
caller_user_id: i64,
caller_mxid: &str,
caller_device_id: &str,
request: &KeysUploadRequest,
now: &str,
) -> Result<KeysUploadOutcome, MatrixError> {
let mut device_keys_changed = false;
if let Some(device_keys) = &request.device_keys {
check_device_keys_ownership(device_keys, caller_mxid, caller_device_id)?;
let new_keys_value = device_keys.get("keys").cloned().unwrap_or(serde_json::Value::Null);
let existing = crate::keys::device_keys_for(conn, &[caller_user_id])?
.into_iter()
.find(|d| d.device_id == caller_device_id);
let action = decide_device_keys_upload(existing.as_ref().map(|d| d.keys.as_str()), &new_keys_value)?;
if matches!(action, DeviceKeysUploadAction::Insert | DeviceKeysUploadAction::Reset) {
if action == DeviceKeysUploadAction::Reset {
crate::keys::clear_device_one_time_material(conn, caller_user_id, caller_device_id)?;
}
let algorithms_json = device_keys.get("algorithms").cloned().unwrap_or(serde_json::json!([])).to_string();
let keys_json = new_keys_value.to_string();
let signatures_json = device_keys.get("signatures").cloned().unwrap_or(serde_json::json!({})).to_string();
crate::keys::upsert_device_keys(conn, caller_user_id, caller_device_id, &algorithms_json, &keys_json, &signatures_json, now)?;
device_keys_changed = true;
}
}
if let Some(one_time_keys) = &request.one_time_keys {
let mut batch = Vec::with_capacity(one_time_keys.len());
for (full_id, value) in one_time_keys {
let (algorithm, _) = split_algorithm_key_id(full_id)?;
batch.push((full_id.clone(), algorithm.to_string(), value.to_string()));
}
crate::keys::add_one_time_keys(conn, caller_user_id, caller_device_id, &batch)?;
}
if let Some(fallback_keys) = &request.fallback_keys {
for (full_id, value) in fallback_keys {
let (algorithm, _) = split_algorithm_key_id(full_id)?;
crate::keys::upsert_fallback_key(conn, caller_user_id, caller_device_id, algorithm, full_id, &value.to_string(), now)?;
}
}
let otk_counts = crate::keys::count_one_time_keys(conn, caller_user_id, caller_device_id)?;
Ok(KeysUploadOutcome { otk_counts, device_keys_changed })
}
#[derive(serde::Deserialize, Default)]
pub struct KeysQueryRequest {
#[serde(default)]
pub device_keys: BTreeMap<String, Vec<String>>,
}
pub fn build_keys_query_response(
conn: &Connection,
caller_user_id: i64,
caller_mxid: &str,
requested: &BTreeMap<String, Vec<String>>,
) -> Result<serde_json::Value, MatrixError> {
let visible = peers_sharing_a_room_with(conn, caller_user_id)?;
build_keys_query_visible(conn, &visible, Some((caller_user_id, caller_mxid)), requested)
}
pub fn build_keys_query_visible(
conn: &Connection,
visible: &HashSet<i64>,
caller: Option<(i64, &str)>,
requested: &BTreeMap<String, Vec<String>>,
) -> Result<serde_json::Value, MatrixError> {
let mut device_keys_out = serde_json::Map::new();
let mut master_keys_out = serde_json::Map::new();
let mut self_signing_keys_out = serde_json::Map::new();
for (mxid, requested_device_ids) in requested {
let allowed_user_id = crate::store::user_id_of(conn, mxid)?.filter(|uid| visible.contains(uid));
let Some(target_user_id) = allowed_user_id else {
device_keys_out.insert(mxid.clone(), serde_json::json!({}));
continue;
};
let mut devices_out = serde_json::Map::new();
for device_keys in crate::keys::device_keys_for(conn, &[target_user_id])? {
if !requested_device_ids.is_empty() && !requested_device_ids.contains(&device_keys.device_id) {
continue;
}
let algorithms: serde_json::Value = serde_json::from_str(&device_keys.algorithms)?;
let keys: serde_json::Value = serde_json::from_str(&device_keys.keys)?;
let signatures: serde_json::Value = serde_json::from_str(&device_keys.signatures)?;
devices_out.insert(
device_keys.device_id.clone(),
serde_json::json!({
"user_id": mxid,
"device_id": device_keys.device_id,
"algorithms": algorithms,
"keys": keys,
"signatures": signatures,
}),
);
}
device_keys_out.insert(mxid.clone(), serde_json::Value::Object(devices_out));
for cross_signing_key in crate::keys::cross_signing_keys_for(conn, &[target_user_id])? {
let value: serde_json::Value = serde_json::from_str(&cross_signing_key.key_json)?;
match cross_signing_key.usage {
CrossSigningUsage::Master => {
master_keys_out.insert(mxid.clone(), value);
}
CrossSigningUsage::SelfSigning => {
self_signing_keys_out.insert(mxid.clone(), value);
}
CrossSigningUsage::UserSigning => {} }
}
}
let mut user_signing_keys_out = serde_json::Map::new();
if let Some((caller_user_id, caller_mxid)) = caller {
if let Some(row) = crate::keys::cross_signing_key_for(conn, caller_user_id, CrossSigningUsage::UserSigning)? {
user_signing_keys_out.insert(caller_mxid.to_string(), serde_json::from_str(&row.key_json)?);
}
}
Ok(serde_json::json!({
"device_keys": device_keys_out,
"master_keys": master_keys_out,
"self_signing_keys": self_signing_keys_out,
"user_signing_keys": user_signing_keys_out,
"failures": {},
}))
}
#[derive(serde::Deserialize, Default)]
pub struct KeysClaimRequest {
#[serde(default)]
pub one_time_keys: BTreeMap<String, BTreeMap<String, String>>,
}
pub fn count_claim_targets(requested: &BTreeMap<String, BTreeMap<String, String>>) -> usize {
requested.values().map(BTreeMap::len).sum()
}
pub fn build_keys_claim_response(conn: &mut Connection, caller_user_id: i64, requested: &BTreeMap<String, BTreeMap<String, String>>) -> Result<serde_json::Value, MatrixError> {
let visible = peers_sharing_a_room_with(conn, caller_user_id)?;
build_keys_claim_visible(conn, &visible, requested)
}
pub fn build_keys_claim_visible(conn: &mut Connection, visible: &HashSet<i64>, requested: &BTreeMap<String, BTreeMap<String, String>>) -> Result<serde_json::Value, MatrixError> {
let mut out = serde_json::Map::new();
for (mxid, per_device) in requested {
let Some(target_user_id) = crate::store::user_id_of(conn, mxid)?.filter(|uid| visible.contains(uid)) else {
continue;
};
let mut devices_out = serde_json::Map::new();
for (device_id, algorithm) in per_device {
let Some((key_id, key_json)) = crate::keys::claim_one_time_key(conn, target_user_id, device_id, algorithm)? else {
continue;
};
let value: serde_json::Value = serde_json::from_str(&key_json)?;
devices_out.insert(device_id.clone(), serde_json::json!({ key_id: value }));
}
if !devices_out.is_empty() {
out.insert(mxid.clone(), serde_json::Value::Object(devices_out));
}
}
Ok(serde_json::json!({ "one_time_keys": out, "failures": {} }))
}
#[derive(serde::Deserialize, Default)]
pub struct KeysChangesQuery {
pub from: Option<String>,
pub to: Option<String>,
}
pub fn build_keys_changes_response(conn: &Connection, caller_user_id: i64, from_exclusive: i64, to_inclusive: i64) -> Result<serde_json::Value, MatrixError> {
let delta = device_list_delta(conn, caller_user_id, from_exclusive, to_inclusive)?;
Ok(serde_json::json!({ "changed": delta.changed, "left": delta.left }))
}
#[derive(serde::Deserialize, Default)]
pub struct DeviceSigningUploadRequest {
#[serde(default)]
pub master_key: Option<serde_json::Value>,
#[serde(default)]
pub self_signing_key: Option<serde_json::Value>,
#[serde(default)]
pub user_signing_key: Option<serde_json::Value>,
}
pub fn master_verifying_key(master_key_object: &serde_json::Value) -> Result<(String, VerifyingKey), MatrixError> {
let keys = master_key_object
.get("keys")
.and_then(|v| v.as_object())
.ok_or_else(|| MatrixError::invalid_param("master_key.keys missing"))?;
let (full_key_id, value) = keys
.iter()
.find(|(k, _)| k.starts_with("ed25519:"))
.ok_or_else(|| MatrixError::invalid_param("master_key has no ed25519 key"))?;
let key_id = full_key_id.trim_start_matches("ed25519:").to_string();
let b64 = value.as_str().ok_or_else(|| MatrixError::invalid_param("master_key.keys value must be a string"))?;
let bytes = STANDARD_NO_PAD.decode(b64).map_err(|_| MatrixError::invalid_param("master_key public key is not valid base64"))?;
let array: [u8; 32] = bytes.try_into().map_err(|_| MatrixError::invalid_param("master_key public key must be 32 bytes"))?;
let verifying_key = VerifyingKey::from_bytes(&array).map_err(|_| MatrixError::invalid_param("master_key is not a valid ed25519 key"))?;
Ok((key_id, verifying_key))
}
pub fn canonical_json_without_signatures(value: &serde_json::Value) -> String {
let mut stripped = value.clone();
if let Some(obj) = stripped.as_object_mut() {
obj.remove("signatures");
obj.remove("unsigned");
}
stripped.to_string()
}
pub fn verify_signed_by_master(target_key_object: &serde_json::Value, caller_mxid: &str, master_key_object: &serde_json::Value) -> Result<(), MatrixError> {
let (master_key_id, verifying_key) = master_verifying_key(master_key_object)?;
let signature_b64 = target_key_object
.get("signatures")
.and_then(|s| s.get(caller_mxid))
.and_then(|by_user| by_user.get(format!("ed25519:{master_key_id}").as_str()))
.and_then(|v| v.as_str())
.ok_or_else(|| MatrixError::invalid_param("missing signature by the master key"))?;
let sig_bytes = STANDARD_NO_PAD.decode(signature_b64).map_err(|_| MatrixError::invalid_param("signature is not valid base64"))?;
let sig_array: [u8; 64] = sig_bytes.try_into().map_err(|_| MatrixError::invalid_param("signature must be 64 bytes"))?;
let signature = Signature::from_bytes(&sig_array);
let message = canonical_json_without_signatures(target_key_object);
verifying_key
.verify_strict(message.as_bytes(), &signature)
.map_err(|_| MatrixError::invalid_param("signature does not verify against the master key"))
}
pub fn validate_cross_signing_key_object(value: &serde_json::Value, caller_mxid: &str, expected_usage: &str) -> Result<(), MatrixError> {
let obj = value.as_object().ok_or_else(|| MatrixError::invalid_param("cross-signing key must be an object"))?;
if obj.get("user_id").and_then(|v| v.as_str()) != Some(caller_mxid) {
return Err(MatrixError::invalid_param("cross-signing key user_id must be the caller"));
}
let usage = obj.get("usage").and_then(|v| v.as_array()).ok_or_else(|| MatrixError::invalid_param("cross-signing key missing usage"))?;
if usage.len() != 1 || usage[0].as_str() != Some(expected_usage) {
return Err(MatrixError::invalid_param(format!("cross-signing key usage must be exactly [\"{expected_usage}\"]")));
}
if obj.get("keys").and_then(|v| v.as_object()).is_none_or(|k| k.is_empty()) {
return Err(MatrixError::invalid_param("cross-signing key missing keys"));
}
Ok(())
}
pub fn apply_device_signing_upload(conn: &mut Connection, caller_user_id: i64, caller_mxid: &str, request: &DeviceSigningUploadRequest, now: &str) -> Result<bool, MatrixError> {
if let Some(master) = &request.master_key {
validate_cross_signing_key_object(master, caller_mxid, "master")?;
}
let master_for_verification: Option<serde_json::Value> = match &request.master_key {
Some(master) => Some(master.clone()),
None => crate::keys::cross_signing_key_for(conn, caller_user_id, CrossSigningUsage::Master)?
.map(|row| serde_json::from_str(&row.key_json))
.transpose()?,
};
if let Some(self_signing) = &request.self_signing_key {
validate_cross_signing_key_object(self_signing, caller_mxid, "self_signing")?;
let master = master_for_verification
.as_ref()
.ok_or_else(|| MatrixError::invalid_param("no master key on file to verify self_signing_key against"))?;
verify_signed_by_master(self_signing, caller_mxid, master)?;
}
if let Some(user_signing) = &request.user_signing_key {
validate_cross_signing_key_object(user_signing, caller_mxid, "user_signing")?;
let master = master_for_verification
.as_ref()
.ok_or_else(|| MatrixError::invalid_param("no master key on file to verify user_signing_key against"))?;
verify_signed_by_master(user_signing, caller_mxid, master)?;
}
let mut wrote = false;
if let Some(master) = &request.master_key {
crate::keys::upsert_cross_signing_key(conn, caller_user_id, CrossSigningUsage::Master, &master.to_string(), now)?;
wrote = true;
}
if let Some(self_signing) = &request.self_signing_key {
crate::keys::upsert_cross_signing_key(conn, caller_user_id, CrossSigningUsage::SelfSigning, &self_signing.to_string(), now)?;
wrote = true;
}
if let Some(user_signing) = &request.user_signing_key {
crate::keys::upsert_cross_signing_key(conn, caller_user_id, CrossSigningUsage::UserSigning, &user_signing.to_string(), now)?;
wrote = true;
}
Ok(wrote)
}
pub fn master_key_names_local_id(conn: &Connection, user_id: i64, local_id: &str) -> Result<bool, MatrixError> {
let Some(row) = crate::keys::cross_signing_key_for(conn, user_id, CrossSigningUsage::Master)? else {
return Ok(false);
};
let value: serde_json::Value = serde_json::from_str(&row.key_json)?;
let Some(keys) = value.get("keys").and_then(|k| k.as_object()) else {
return Ok(false);
};
Ok(keys.keys().any(|k| k.trim_start_matches("ed25519:") == local_id))
}
fn cross_signing_key_names_local_id(conn: &Connection, user_id: i64, local_id: &str) -> Result<bool, MatrixError> {
for usage in [CrossSigningUsage::Master, CrossSigningUsage::SelfSigning, CrossSigningUsage::UserSigning] {
let Some(row) = crate::keys::cross_signing_key_for(conn, user_id, usage)? else {
continue;
};
let value: serde_json::Value = serde_json::from_str(&row.key_json)?;
if let Some(keys) = value.get("keys").and_then(|k| k.as_object()) {
if keys.keys().any(|k| k.trim_start_matches("ed25519:") == local_id) {
return Ok(true);
}
}
}
Ok(false)
}
pub fn signature_target_is_authorized(conn: &Connection, caller_mxid: &str, target_mxid: &str, target_user_id: i64, target_key_id: &str) -> Result<bool, MatrixError> {
if target_mxid == caller_mxid {
Ok(crate::keys::get_device(conn, target_user_id, target_key_id)?.is_some() || cross_signing_key_names_local_id(conn, target_user_id, target_key_id)?)
} else {
master_key_names_local_id(conn, target_user_id, target_key_id)
}
}
pub fn apply_signatures_upload(conn: &mut Connection, caller_user_id: i64, caller_mxid: &str, body: &serde_json::Map<String, serde_json::Value>, now: &str) -> Result<serde_json::Value, MatrixError> {
let mut failures = serde_json::Map::new();
for (target_mxid, per_key) in body {
let Some(per_key_obj) = per_key.as_object() else { continue };
let mut user_failures = serde_json::Map::new();
let target_user_id = crate::store::user_id_of(conn, target_mxid)?;
for (target_key_id, signed_value) in per_key_obj {
let authorized = match target_user_id {
Some(target_user_id) => signature_target_is_authorized(conn, caller_mxid, target_mxid, target_user_id, target_key_id)?,
None => false,
};
if !authorized {
user_failures.insert(
target_key_id.clone(),
serde_json::json!({ "errcode": "M_INVALID_PARAM", "error": "signature target not permitted" }),
);
continue;
}
let target_user_id = target_user_id.ok_or_else(MatrixError::internal)?;
crate::keys::add_signatures(conn, &[(caller_user_id, target_user_id, target_key_id.clone(), signed_value.to_string(), now.to_string())])?;
}
if !user_failures.is_empty() {
failures.insert(target_mxid.clone(), serde_json::Value::Object(user_failures));
}
}
Ok(serde_json::Value::Object(failures))
}
#[derive(serde::Deserialize)]
pub struct SendToDeviceRequest {
pub messages: BTreeMap<String, BTreeMap<String, serde_json::Value>>,
}
pub fn expand_send_to_device_targets(conn: &Connection, sender_user_id: i64, messages: &BTreeMap<String, BTreeMap<String, serde_json::Value>>) -> Result<Vec<(i64, String, String)>, MatrixError> {
let visible = peers_sharing_a_room_with(conn, sender_user_id)?;
let mut targets = Vec::new();
for (mxid, per_device) in messages {
let Some(recipient_user_id) = crate::store::user_id_of(conn, mxid)?.filter(|uid| visible.contains(uid)) else {
continue;
};
let mut per_recipient: HashMap<String, String> = HashMap::new();
if let Some(wildcard_content) = per_device.get("*") {
let content_str = wildcard_content.to_string();
if content_str.len() > SEND_TO_DEVICE_MAX_CONTENT_BYTES {
return Err(MatrixError::bad_json("to-device content too large"));
}
for device in crate::keys::list_devices(conn, recipient_user_id)? {
per_recipient.insert(device.device_id, content_str.clone());
}
}
for (device_selector, content) in per_device {
if device_selector == "*" {
continue;
}
let content_str = content.to_string();
if content_str.len() > SEND_TO_DEVICE_MAX_CONTENT_BYTES {
return Err(MatrixError::bad_json("to-device content too large"));
}
per_recipient.insert(device_selector.clone(), content_str);
}
for (device_id, content_str) in per_recipient {
targets.push((recipient_user_id, device_id, content_str));
}
}
if targets.len() > SEND_TO_DEVICE_MAX_TARGETS {
return Err(MatrixError::invalid_param("too many target devices in one call"));
}
Ok(targets)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum SendToDeviceOutcome {
New(HashSet<i64>),
AlreadySent,
}
pub fn apply_send_to_device(
conn: &mut Connection,
sender_user_id: i64,
sender_device_id: &str,
event_type: &str,
txn_id: &str,
messages: &BTreeMap<String, BTreeMap<String, serde_json::Value>>,
now: &str,
) -> Result<SendToDeviceOutcome, MatrixError> {
if let TxnDedupEntry::Seen(_) = crate::store::txn_dedup_lookup(conn, sender_user_id, sender_device_id, txn_id)? {
return Ok(SendToDeviceOutcome::AlreadySent);
}
let targets = expand_send_to_device_targets(conn, sender_user_id, messages)?;
let wake_ids: HashSet<i64> = targets.iter().map(|(uid, ..)| *uid).collect();
let rows: Vec<(i64, String, String, String)> = targets.into_iter().map(|(uid, dev, content)| (uid, dev, event_type.to_string(), content)).collect();
match crate::keys::enqueue_to_device_deduped(conn, sender_user_id, sender_device_id, txn_id, &rows, now)? {
crate::keys::ToDeviceDedupOutcome::New => Ok(SendToDeviceOutcome::New(wake_ids)),
crate::keys::ToDeviceDedupOutcome::AlreadySent => Ok(SendToDeviceOutcome::AlreadySent),
}
}
#[derive(serde::Deserialize, Default)]
pub struct PutDeviceRequest {
#[serde(default)]
pub display_name: Option<String>,
}
#[derive(serde::Deserialize)]
pub struct DeleteDevicesRequest {
pub devices: Vec<String>,
}
pub fn device_to_json(device: &crate::keys::Device) -> serde_json::Value {
let last_seen_ts = chrono::DateTime::parse_from_rfc3339(&device.last_seen_at).ok().map(|dt| dt.timestamp_millis());
serde_json::json!({
"device_id": device.device_id,
"display_name": device.display_name,
"last_seen_ts": last_seen_ts,
})
}
#[derive(serde::Deserialize)]
pub struct BackupVersionCreateRequest {
pub algorithm: String,
pub auth_data: serde_json::Value,
}
#[derive(serde::Deserialize)]
pub struct BackupVersionUpdateRequest {
#[serde(default)]
pub algorithm: Option<String>,
pub auth_data: serde_json::Value,
}
pub fn validate_backup_algorithm(algorithm: &str) -> Result<(), MatrixError> {
if algorithm == BACKUP_ALGORITHM {
Ok(())
} else {
Err(MatrixError::invalid_param(format!("unsupported key-backup algorithm: {algorithm}")))
}
}
pub fn backup_version_to_response(row: &crate::keys::KeyBackupVersion, count: i64) -> Result<serde_json::Value, MatrixError> {
Ok(serde_json::json!({
"version": row.version.to_string(),
"algorithm": row.algorithm,
"auth_data": serde_json::from_str::<serde_json::Value>(&row.auth_data)?,
"etag": row.etag.to_string(),
"count": count,
}))
}
#[derive(serde::Deserialize, Default)]
pub struct VersionQuery {
pub version: Option<String>,
}
pub fn parse_required_version(raw: Option<&str>) -> Result<i64, MatrixError> {
let raw = raw.ok_or_else(|| MatrixError::invalid_param("version is required"))?;
raw.parse::<i64>().map_err(|_| MatrixError::invalid_param("version must be an integer"))
}
pub fn require_current_backup_version(conn: &Connection, user_id: i64, requested_version: i64) -> Result<(), MatrixError> {
let current = crate::keys::current_backup_version(conn, user_id)?;
match current {
Some(row) if row.version == requested_version => Ok(()),
Some(row) => Err(MatrixError::wrong_room_keys_version(Some(row.version))),
None => Err(MatrixError::wrong_room_keys_version(None)),
}
}
pub fn normalize_put_backup_body(room_id: Option<&str>, session_id: Option<&str>, body: &serde_json::Value) -> Result<Vec<(String, String, String)>, MatrixError> {
match (room_id, session_id) {
(Some(room_id), Some(session_id)) => Ok(vec![(room_id.to_string(), session_id.to_string(), body.to_string())]),
(Some(room_id), None) => {
let sessions = body
.get("sessions")
.and_then(|v| v.as_object())
.ok_or_else(|| MatrixError::invalid_param("body must have a sessions object"))?;
Ok(sessions.iter().map(|(sid, data)| (room_id.to_string(), sid.clone(), data.to_string())).collect())
}
(None, _) => {
let rooms = body
.get("rooms")
.and_then(|v| v.as_object())
.ok_or_else(|| MatrixError::invalid_param("body must have a rooms object"))?;
let mut out = Vec::new();
for (rid, room_value) in rooms {
let sessions = room_value
.get("sessions")
.and_then(|v| v.as_object())
.ok_or_else(|| MatrixError::invalid_param("each room must have a sessions object"))?;
for (sid, data) in sessions {
out.push((rid.clone(), sid.clone(), data.to_string()));
}
}
Ok(out)
}
}
}
pub fn backup_sessions_to_response(room_id: Option<&str>, session_id: Option<&str>, sessions: Vec<crate::keys::KeyBackupSession>) -> Result<serde_json::Value, MatrixError> {
match (room_id, session_id) {
(Some(_), Some(_)) => {
let one = sessions.into_iter().next().ok_or_else(|| MatrixError::not_found("no such backup session"))?;
Ok(serde_json::from_str(&one.session_data)?)
}
(Some(_), None) => {
let mut sessions_out = serde_json::Map::new();
for s in sessions {
sessions_out.insert(s.session_id.clone(), serde_json::from_str(&s.session_data)?);
}
Ok(serde_json::json!({ "sessions": sessions_out }))
}
(None, _) => {
let mut rooms_out = serde_json::Map::new();
for s in sessions {
let room_entry = rooms_out.entry(s.room_id.clone()).or_insert_with(|| serde_json::json!({ "sessions": {} }));
room_entry["sessions"][s.session_id.as_str()] = serde_json::from_str(&s.session_data)?;
}
Ok(serde_json::json!({ "rooms": rooms_out }))
}
}
}
#[cfg(test)]
mod device_keys_upload_tests {
use super::*;
use rusqlite::Connection;
const T0: &str = "2026-10-06T00:00:00+00:00";
fn test_conn() -> Connection {
let conn = Connection::open_in_memory().expect("memory");
crate::store::create_matrix_schema(&conn).expect("schema");
crate::keys::create_matrix_keys_schema(&conn).expect("keys");
conn
}
#[test]
fn decide_insert_noop_and_reset() {
let keys_a = serde_json::json!({"curve25519:D":"aaa","ed25519:D":"bbb"});
let keys_b = serde_json::json!({"curve25519:D":"ccc","ed25519:D":"ddd"});
assert_eq!(decide_device_keys_upload(None, &keys_a).unwrap(), DeviceKeysUploadAction::Insert);
assert_eq!(
decide_device_keys_upload(Some(&keys_a.to_string()), &keys_a).unwrap(),
DeviceKeysUploadAction::Noop
);
assert_eq!(
decide_device_keys_upload(Some(&keys_a.to_string()), &keys_b).unwrap(),
DeviceKeysUploadAction::Reset
);
}
#[test]
fn apply_keys_upload_reset_replaces_identity_clears_otks_and_logs_change() {
let mut conn = test_conn();
crate::store::ensure_matrix_user(&conn, 1, "alice00000000000000000000000001", T0).expect("user");
let device_id = crate::keys::create_device(&conn, 1, crate::keys::CredentialKind::Bearer, "tok", T0).expect("device");
let mxid = crate::store::mxid_of(&conn, 1).expect("mxid").expect("row");
let first = KeysUploadRequest {
device_keys: Some(serde_json::json!({
"user_id": mxid,
"device_id": device_id,
"algorithms": ["m.olm.v1.curve25519-aes-sha2"],
"keys": {"curve25519:D":"oldcurve","ed25519:D":"olded"},
"signatures": {}
})),
one_time_keys: Some(std::collections::BTreeMap::from([(
"signed_curve25519:AAAAAQ".to_string(),
serde_json::json!({"key":"otk1"}),
)])),
fallback_keys: Some(std::collections::BTreeMap::from([(
"signed_curve25519:FALLBACK".to_string(),
serde_json::json!({"key":"fb1"}),
)])),
};
let first_out = apply_keys_upload(&mut conn, 1, &mxid, &device_id, &first, T0).expect("first");
assert!(first_out.device_keys_changed);
assert_eq!(crate::keys::count_one_time_keys(&conn, 1, &device_id).unwrap().get("signed_curve25519"), Some(&1));
let before_changes: i64 = conn
.query_row("SELECT COUNT(*) FROM device_list_changes WHERE user_id = 1", [], |r| r.get(0))
.unwrap();
let reset = KeysUploadRequest {
device_keys: Some(serde_json::json!({
"user_id": mxid,
"device_id": device_id,
"algorithms": ["m.olm.v1.curve25519-aes-sha2"],
"keys": {"curve25519:D":"newcurve","ed25519:D":"newed"},
"signatures": {}
})),
one_time_keys: Some(std::collections::BTreeMap::from([(
"signed_curve25519:BBBBBQ".to_string(),
serde_json::json!({"key":"otk2"}),
)])),
fallback_keys: None,
};
let reset_out = apply_keys_upload(&mut conn, 1, &mxid, &device_id, &reset, T0).expect("reset");
assert!(reset_out.device_keys_changed);
let stored = crate::keys::device_keys_for(&conn, &[1]).unwrap();
assert_eq!(stored.len(), 1);
assert!(stored[0].keys.contains("newcurve"), "new identity stored: {}", stored[0].keys);
assert!(!stored[0].keys.contains("oldcurve"));
let otk = crate::keys::count_one_time_keys(&conn, 1, &device_id).unwrap();
assert_eq!(otk.get("signed_curve25519"), Some(&1), "old OTKs wiped; only the new batch remains");
let fallback = crate::keys::unused_fallback_key_types(&conn, 1, &device_id).unwrap();
assert!(fallback.is_empty(), "old fallback wiped on reset");
let after_changes: i64 = conn
.query_row("SELECT COUNT(*) FROM device_list_changes WHERE user_id = 1", [], |r| r.get(0))
.unwrap();
assert_eq!(after_changes, before_changes + 1, "reset logs device_lists.changed");
}
}