1use rusqlite::{params, Connection, OptionalExtension, Transaction};
39use std::collections::HashMap;
40
41
42
43#[derive(Debug)]
50pub enum MatrixKeysStoreError {
51 Db(rusqlite::Error),
52 OneTimeKeyConflict(String),
58 WrongBackupVersion,
61}
62
63impl From<rusqlite::Error> for MatrixKeysStoreError {
64 fn from(e: rusqlite::Error) -> Self {
65 MatrixKeysStoreError::Db(e)
66 }
67}
68
69fn decode_enum<T>(idx: usize, column: &'static str, raw: &str, parse: fn(&str) -> Option<T>) -> rusqlite::Result<T> {
74 parse(raw).ok_or_else(|| rusqlite::Error::InvalidColumnType(idx, column.to_string(), rusqlite::types::Type::Text))
75}
76
77#[derive(Debug, Clone, Copy, PartialEq, Eq)]
83pub enum CredentialKind {
84 Bearer,
85 Web,
86}
87
88impl CredentialKind {
89 pub fn as_str(self) -> &'static str {
90 match self {
91 CredentialKind::Bearer => "bearer",
92 CredentialKind::Web => "web",
93 }
94 }
95
96 pub fn from_wire_name(s: &str) -> Option<Self> {
97 match s {
98 "bearer" => Some(CredentialKind::Bearer),
99 "web" => Some(CredentialKind::Web),
100 _ => None,
101 }
102 }
103}
104
105#[derive(Debug, Clone, Copy, PartialEq, Eq)]
107pub enum CrossSigningUsage {
108 Master,
109 SelfSigning,
110 UserSigning,
111}
112
113impl CrossSigningUsage {
114 pub fn as_str(self) -> &'static str {
115 match self {
116 CrossSigningUsage::Master => "master",
117 CrossSigningUsage::SelfSigning => "self_signing",
118 CrossSigningUsage::UserSigning => "user_signing",
119 }
120 }
121
122 pub fn from_wire_name(s: &str) -> Option<Self> {
123 match s {
124 "master" => Some(CrossSigningUsage::Master),
125 "self_signing" => Some(CrossSigningUsage::SelfSigning),
126 "user_signing" => Some(CrossSigningUsage::UserSigning),
127 _ => None,
128 }
129 }
130}
131
132fn random_device_id() -> String {
135 use base64::engine::general_purpose::URL_SAFE_NO_PAD;
136 use base64::Engine;
137 use rand::Rng;
138 let bytes: [u8; 8] = rand::thread_rng().gen();
139 URL_SAFE_NO_PAD.encode(bytes)
140}
141
142pub fn create_matrix_keys_schema(conn: &Connection) -> rusqlite::Result<()> {
162 conn.execute_batch(
163 r#"
164 CREATE TABLE IF NOT EXISTS devices (
165 user_id INTEGER NOT NULL,
166 device_id TEXT NOT NULL,
167 credential_kind TEXT NOT NULL,
168 credential_ref TEXT NOT NULL,
169 display_name TEXT,
170 created_at TEXT NOT NULL,
171 last_seen_at TEXT NOT NULL,
172 PRIMARY KEY (user_id, device_id)
173 );
174 CREATE UNIQUE INDEX IF NOT EXISTS idx_devices_credential ON devices(credential_kind, credential_ref);
175
176 CREATE TABLE IF NOT EXISTS device_keys (
177 user_id INTEGER NOT NULL,
178 device_id TEXT NOT NULL,
179 algorithms TEXT NOT NULL,
180 keys TEXT NOT NULL,
181 signatures TEXT NOT NULL,
182 uploaded_at TEXT NOT NULL,
183 PRIMARY KEY (user_id, device_id)
184 );
185
186 CREATE TABLE IF NOT EXISTS one_time_keys (
187 user_id INTEGER NOT NULL,
188 device_id TEXT NOT NULL,
189 key_id TEXT NOT NULL,
190 algorithm TEXT NOT NULL,
191 key_json TEXT NOT NULL,
192 PRIMARY KEY (user_id, device_id, key_id)
193 );
194 CREATE INDEX IF NOT EXISTS idx_otk_owner_algorithm ON one_time_keys(user_id, device_id, algorithm);
195
196 -- DEVIATION: `key_id` column added — see this function's doc comment.
197 CREATE TABLE IF NOT EXISTS fallback_keys (
198 user_id INTEGER NOT NULL,
199 device_id TEXT NOT NULL,
200 algorithm TEXT NOT NULL,
201 key_id TEXT NOT NULL,
202 key_json TEXT NOT NULL,
203 used INTEGER NOT NULL DEFAULT 0,
204 uploaded_at TEXT NOT NULL,
205 PRIMARY KEY (user_id, device_id, algorithm)
206 );
207
208 CREATE TABLE IF NOT EXISTS to_device_messages (
209 stream_id INTEGER PRIMARY KEY,
210 recipient_user_id INTEGER NOT NULL,
211 recipient_device_id TEXT NOT NULL,
212 sender_user_id INTEGER NOT NULL,
213 event_type TEXT NOT NULL,
214 content TEXT NOT NULL
215 );
216 CREATE INDEX IF NOT EXISTS idx_to_device_recipient ON to_device_messages(recipient_user_id, recipient_device_id, stream_id);
217
218 CREATE TABLE IF NOT EXISTS device_list_changes (
219 stream_id INTEGER PRIMARY KEY,
220 user_id INTEGER NOT NULL,
221 changed_at TEXT NOT NULL
222 );
223 CREATE INDEX IF NOT EXISTS idx_device_list_changes_user ON device_list_changes(user_id, stream_id);
224
225 CREATE TABLE IF NOT EXISTS cross_signing_keys (
226 user_id INTEGER NOT NULL,
227 usage TEXT NOT NULL,
228 key_json TEXT NOT NULL,
229 uploaded_at TEXT NOT NULL,
230 PRIMARY KEY (user_id, usage)
231 );
232
233 CREATE TABLE IF NOT EXISTS cross_signing_signatures (
234 id INTEGER PRIMARY KEY AUTOINCREMENT,
235 signer_user_id INTEGER NOT NULL,
236 target_user_id INTEGER NOT NULL,
237 target_key_id TEXT NOT NULL,
238 signature_json TEXT NOT NULL,
239 uploaded_at TEXT NOT NULL
240 );
241 CREATE INDEX IF NOT EXISTS idx_xsig_target ON cross_signing_signatures(target_user_id, target_key_id);
242
243 CREATE TABLE IF NOT EXISTS key_backup_versions (
244 version INTEGER PRIMARY KEY AUTOINCREMENT,
245 user_id INTEGER NOT NULL,
246 algorithm TEXT NOT NULL,
247 auth_data TEXT NOT NULL,
248 etag INTEGER NOT NULL DEFAULT 0,
249 is_deleted INTEGER NOT NULL DEFAULT 0,
250 created_at TEXT NOT NULL
251 );
252 CREATE INDEX IF NOT EXISTS idx_key_backup_versions_user ON key_backup_versions(user_id, version);
253
254 CREATE TABLE IF NOT EXISTS key_backup_sessions (
255 user_id INTEGER NOT NULL,
256 version INTEGER NOT NULL REFERENCES key_backup_versions(version),
257 room_id TEXT NOT NULL,
258 session_id TEXT NOT NULL,
259 session_data TEXT NOT NULL,
260 updated_at TEXT NOT NULL,
261 PRIMARY KEY (user_id, version, room_id, session_id)
262 );
263 "#,
264 )
265}
266
267#[derive(Debug, Clone, PartialEq)]
272pub struct Device {
273 pub user_id: i64,
274 pub device_id: String,
275 pub credential_kind: CredentialKind,
276 pub credential_ref: String,
277 pub display_name: Option<String>,
278 pub created_at: String,
279 pub last_seen_at: String,
280}
281
282const DEVICE_SELECT_COLUMNS: &str = "user_id, device_id, credential_kind, credential_ref, display_name, created_at, last_seen_at";
283
284fn device_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<Device> {
285 let credential_kind_raw: String = row.get(2)?;
286 Ok(Device {
287 user_id: row.get(0)?,
288 device_id: row.get(1)?,
289 credential_kind: decode_enum(2, "credential_kind", &credential_kind_raw, CredentialKind::from_wire_name)?,
290 credential_ref: row.get(3)?,
291 display_name: row.get(4)?,
292 created_at: row.get(5)?,
293 last_seen_at: row.get(6)?,
294 })
295}
296
297pub fn device_for_credential(conn: &Connection, credential_kind: CredentialKind, credential_ref: &str) -> rusqlite::Result<Option<Device>> {
300 conn.query_row(
301 &format!("SELECT {DEVICE_SELECT_COLUMNS} FROM devices WHERE credential_kind = ?1 AND credential_ref = ?2"),
302 params![credential_kind.as_str(), credential_ref],
303 device_from_row,
304 )
305 .optional()
306}
307
308pub fn create_device(conn: &Connection, user_id: i64, credential_kind: CredentialKind, credential_ref: &str, now: &str) -> rusqlite::Result<String> {
314 let device_id = random_device_id();
315 conn.execute(
316 "INSERT INTO devices (user_id, device_id, credential_kind, credential_ref, display_name, created_at, last_seen_at)
317 VALUES (?1, ?2, ?3, ?4, NULL, ?5, ?6)",
318 params![user_id, device_id, credential_kind.as_str(), credential_ref, now, now],
319 )?;
320 Ok(device_id)
321}
322
323pub fn touch_device(conn: &Connection, user_id: i64, device_id: &str, now: &str) -> rusqlite::Result<()> {
326 conn.execute(
327 "UPDATE devices SET last_seen_at = ?1 WHERE user_id = ?2 AND device_id = ?3",
328 params![now, user_id, device_id],
329 )?;
330 Ok(())
331}
332
333pub fn list_devices(conn: &Connection, user_id: i64) -> rusqlite::Result<Vec<Device>> {
335 let mut stmt = conn.prepare(&format!("SELECT {DEVICE_SELECT_COLUMNS} FROM devices WHERE user_id = ?1"))?;
336 let rows = stmt.query_map(params![user_id], device_from_row)?;
337 rows.collect()
338}
339
340pub fn get_device(conn: &Connection, user_id: i64, device_id: &str) -> rusqlite::Result<Option<Device>> {
342 conn.query_row(
343 &format!("SELECT {DEVICE_SELECT_COLUMNS} FROM devices WHERE user_id = ?1 AND device_id = ?2"),
344 params![user_id, device_id],
345 device_from_row,
346 )
347 .optional()
348}
349
350pub fn set_device_display_name(conn: &Connection, user_id: i64, device_id: &str, display_name: Option<&str>) -> rusqlite::Result<bool> {
353 let changed = conn.execute(
354 "UPDATE devices SET display_name = ?1 WHERE user_id = ?2 AND device_id = ?3",
355 params![display_name, user_id, device_id],
356 )?;
357 Ok(changed > 0)
358}
359
360fn delete_device_key_material(tx: &Transaction, user_id: i64, device_id: &str) -> rusqlite::Result<()> {
368 tx.execute("DELETE FROM device_keys WHERE user_id = ?1 AND device_id = ?2", params![user_id, device_id])?;
369 tx.execute("DELETE FROM one_time_keys WHERE user_id = ?1 AND device_id = ?2", params![user_id, device_id])?;
370 tx.execute("DELETE FROM fallback_keys WHERE user_id = ?1 AND device_id = ?2", params![user_id, device_id])?;
371 tx.execute(
372 "DELETE FROM to_device_messages WHERE recipient_user_id = ?1 AND recipient_device_id = ?2",
373 params![user_id, device_id],
374 )?;
375 Ok(())
376}
377
378pub fn delete_device(conn: &mut Connection, user_id: i64, device_id: &str, now: &str) -> rusqlite::Result<bool> {
383 let tx = conn.transaction()?;
384 let existed = tx.execute("DELETE FROM devices WHERE user_id = ?1 AND device_id = ?2", params![user_id, device_id])? > 0;
385 if existed {
386 delete_device_key_material(&tx, user_id, device_id)?;
387 log_device_list_change_tx(&tx, user_id, now)?;
388 }
389 tx.commit()?;
390 Ok(existed)
391}
392
393pub fn delete_device_by_credential(
401 conn: &mut Connection,
402 credential_kind: CredentialKind,
403 credential_ref: &str,
404 now: &str,
405) -> rusqlite::Result<Option<(i64, String)>> {
406 let tx = conn.transaction()?;
407 let found: Option<(i64, String)> = tx
408 .query_row(
409 "DELETE FROM devices WHERE credential_kind = ?1 AND credential_ref = ?2 RETURNING user_id, device_id",
410 params![credential_kind.as_str(), credential_ref],
411 |row| Ok((row.get(0)?, row.get(1)?)),
412 )
413 .optional()?;
414 if let Some((user_id, device_id)) = &found {
415 delete_device_key_material(&tx, *user_id, device_id)?;
416 log_device_list_change_tx(&tx, *user_id, now)?;
417 }
418 tx.commit()?;
419 Ok(found)
420}
421
422pub fn devices_for_reaper(conn: &Connection) -> rusqlite::Result<Vec<(i64, String, CredentialKind, String)>> {
427 let mut stmt = conn.prepare("SELECT user_id, device_id, credential_kind, credential_ref FROM devices")?;
428 let mut rows = stmt.query([])?;
429 let mut out = Vec::new();
430 while let Some(row) = rows.next()? {
431 let kind_raw: String = row.get(2)?;
432 let credential_kind = decode_enum(2, "credential_kind", &kind_raw, CredentialKind::from_wire_name)?;
433 out.push((row.get(0)?, row.get(1)?, credential_kind, row.get(3)?));
434 }
435 Ok(out)
436}
437
438#[derive(Debug, Clone, PartialEq)]
443pub struct DeviceKeys {
444 pub user_id: i64,
445 pub device_id: String,
446 pub algorithms: String,
447 pub keys: String,
448 pub signatures: String,
449 pub uploaded_at: String,
450}
451
452const DEVICE_KEYS_SELECT_COLUMNS: &str = "user_id, device_id, algorithms, keys, signatures, uploaded_at";
453
454fn device_keys_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<DeviceKeys> {
455 Ok(DeviceKeys {
456 user_id: row.get(0)?,
457 device_id: row.get(1)?,
458 algorithms: row.get(2)?,
459 keys: row.get(3)?,
460 signatures: row.get(4)?,
461 uploaded_at: row.get(5)?,
462 })
463}
464
465pub fn upsert_device_keys(
470 conn: &mut Connection,
471 user_id: i64,
472 device_id: &str,
473 algorithms_json: &str,
474 keys_json: &str,
475 signatures_json: &str,
476 now: &str,
477) -> rusqlite::Result<()> {
478 let tx = conn.transaction()?;
479 tx.execute(
480 "INSERT INTO device_keys (user_id, device_id, algorithms, keys, signatures, uploaded_at)
481 VALUES (?1, ?2, ?3, ?4, ?5, ?6)
482 ON CONFLICT(user_id, device_id) DO UPDATE SET
483 algorithms = excluded.algorithms, keys = excluded.keys, signatures = excluded.signatures, uploaded_at = excluded.uploaded_at",
484 params![user_id, device_id, algorithms_json, keys_json, signatures_json, now],
485 )?;
486 log_device_list_change_tx(&tx, user_id, now)?;
487 tx.commit()?;
488 Ok(())
489}
490
491pub fn clear_device_one_time_material(conn: &Connection, user_id: i64, device_id: &str) -> rusqlite::Result<()> {
496 conn.execute("DELETE FROM one_time_keys WHERE user_id = ?1 AND device_id = ?2", params![user_id, device_id])?;
497 conn.execute("DELETE FROM fallback_keys WHERE user_id = ?1 AND device_id = ?2", params![user_id, device_id])?;
498 Ok(())
499}
500
501pub fn device_keys_for(conn: &Connection, user_ids: &[i64]) -> rusqlite::Result<Vec<DeviceKeys>> {
505 if user_ids.is_empty() {
506 return Ok(Vec::new());
507 }
508 let placeholders = vec!["?"; user_ids.len()].join(",");
509 let sql = format!("SELECT {DEVICE_KEYS_SELECT_COLUMNS} FROM device_keys WHERE user_id IN ({placeholders})");
510 let mut stmt = conn.prepare(&sql)?;
511 let bound: Vec<&dyn rusqlite::ToSql> = user_ids.iter().map(|id| id as &dyn rusqlite::ToSql).collect();
512 let rows = stmt.query_map(bound.as_slice(), device_keys_from_row)?;
513 rows.collect()
514}
515
516pub fn add_one_time_keys(conn: &mut Connection, user_id: i64, device_id: &str, keys: &[(String, String, String)]) -> Result<(), MatrixKeysStoreError> {
527 let tx = conn.transaction()?;
528 for (key_id, algorithm, key_json) in keys {
529 let existing: Option<String> = tx
530 .query_row(
531 "SELECT key_json FROM one_time_keys WHERE user_id = ?1 AND device_id = ?2 AND key_id = ?3",
532 params![user_id, device_id, key_id],
533 |row| row.get(0),
534 )
535 .optional()?;
536 match existing {
537 None => {
538 tx.execute(
539 "INSERT INTO one_time_keys (user_id, device_id, key_id, algorithm, key_json) VALUES (?1, ?2, ?3, ?4, ?5)",
540 params![user_id, device_id, key_id, algorithm, key_json],
541 )?;
542 }
543 Some(ref existing_json) if existing_json == key_json => {
544 }
546 Some(_) => {
547 return Err(MatrixKeysStoreError::OneTimeKeyConflict(key_id.clone()));
548 }
549 }
550 }
551 tx.commit()?;
552 Ok(())
553}
554
555pub fn count_one_time_keys(conn: &Connection, user_id: i64, device_id: &str) -> rusqlite::Result<HashMap<String, i64>> {
558 let mut stmt = conn.prepare("SELECT algorithm, COUNT(*) FROM one_time_keys WHERE user_id = ?1 AND device_id = ?2 GROUP BY algorithm")?;
559 let rows = stmt.query_map(params![user_id, device_id], |row| Ok((row.get::<_, String>(0)?, row.get::<_, i64>(1)?)))?;
560 let mut out = HashMap::new();
561 for row in rows {
562 let (algorithm, count) = row?;
563 out.insert(algorithm, count);
564 }
565 Ok(out)
566}
567
568pub fn claim_one_time_key(conn: &mut Connection, user_id: i64, device_id: &str, algorithm: &str) -> rusqlite::Result<Option<(String, String)>> {
576 let tx = conn.transaction()?;
577 let claimed: Option<(String, String)> = tx
578 .query_row(
579 "DELETE FROM one_time_keys
580 WHERE user_id = ?1 AND device_id = ?2 AND algorithm = ?3
581 AND key_id = (
582 SELECT key_id FROM one_time_keys
583 WHERE user_id = ?1 AND device_id = ?2 AND algorithm = ?3
584 ORDER BY key_id ASC LIMIT 1
585 )
586 RETURNING key_id, key_json",
587 params![user_id, device_id, algorithm],
588 |row| Ok((row.get(0)?, row.get(1)?)),
589 )
590 .optional()?;
591 let result = match claimed {
592 Some(pair) => Some(pair),
593 None => {
594 let fallback: Option<(String, String)> = tx
595 .query_row(
596 "SELECT key_id, key_json FROM fallback_keys WHERE user_id = ?1 AND device_id = ?2 AND algorithm = ?3",
597 params![user_id, device_id, algorithm],
598 |row| Ok((row.get(0)?, row.get(1)?)),
599 )
600 .optional()?;
601 if fallback.is_some() {
602 tx.execute(
603 "UPDATE fallback_keys SET used = 1 WHERE user_id = ?1 AND device_id = ?2 AND algorithm = ?3",
604 params![user_id, device_id, algorithm],
605 )?;
606 }
607 fallback
608 }
609 };
610 tx.commit()?;
611 Ok(result)
612}
613
614pub fn upsert_fallback_key(
622 conn: &Connection,
623 user_id: i64,
624 device_id: &str,
625 algorithm: &str,
626 key_id: &str,
627 key_json: &str,
628 now: &str,
629) -> rusqlite::Result<()> {
630 conn.execute(
631 "INSERT INTO fallback_keys (user_id, device_id, algorithm, key_id, key_json, used, uploaded_at)
632 VALUES (?1, ?2, ?3, ?4, ?5, 0, ?6)
633 ON CONFLICT(user_id, device_id, algorithm) DO UPDATE SET
634 key_id = excluded.key_id, key_json = excluded.key_json, used = 0, uploaded_at = excluded.uploaded_at",
635 params![user_id, device_id, algorithm, key_id, key_json, now],
636 )?;
637 Ok(())
638}
639
640pub fn unused_fallback_key_types(conn: &Connection, user_id: i64, device_id: &str) -> rusqlite::Result<Vec<String>> {
644 let mut stmt = conn.prepare("SELECT algorithm FROM fallback_keys WHERE user_id = ?1 AND device_id = ?2 AND used = 0")?;
645 let rows = stmt.query_map(params![user_id, device_id], |row| row.get(0))?;
646 rows.collect()
647}
648
649#[derive(Debug, Clone, PartialEq)]
654pub struct ToDeviceMessage {
655 pub stream_id: i64,
656 pub recipient_user_id: i64,
657 pub recipient_device_id: String,
658 pub sender_user_id: i64,
659 pub event_type: String,
660 pub content: String,
661}
662
663const TO_DEVICE_SELECT_COLUMNS: &str = "stream_id, recipient_user_id, recipient_device_id, sender_user_id, event_type, content";
664
665fn to_device_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<ToDeviceMessage> {
666 Ok(ToDeviceMessage {
667 stream_id: row.get(0)?,
668 recipient_user_id: row.get(1)?,
669 recipient_device_id: row.get(2)?,
670 sender_user_id: row.get(3)?,
671 event_type: row.get(4)?,
672 content: row.get(5)?,
673 })
674}
675
676pub fn enqueue_to_device(conn: &mut Connection, sender_user_id: i64, messages: &[(i64, String, String, String)]) -> rusqlite::Result<i64> {
683 let tx = conn.transaction()?;
684 let mut last_stream_id = crate::store::max_stream_id(&tx)?;
685 for (recipient_user_id, recipient_device_id, event_type, content) in messages {
686 let stream_id = crate::store::next_stream_id(&tx)?;
687 tx.execute(
688 "INSERT INTO to_device_messages (stream_id, recipient_user_id, recipient_device_id, sender_user_id, event_type, content)
689 VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
690 params![stream_id, recipient_user_id, recipient_device_id, sender_user_id, event_type, content],
691 )?;
692 last_stream_id = stream_id;
693 }
694 tx.commit()?;
695 Ok(last_stream_id)
696}
697
698#[derive(Debug, Clone, Copy, PartialEq, Eq)]
700pub enum ToDeviceDedupOutcome {
701 New,
704 AlreadySent,
707}
708
709pub fn enqueue_to_device_deduped(
719 conn: &mut Connection,
720 sender_user_id: i64,
721 sender_device_id: &str,
722 txn_id: &str,
723 messages: &[(i64, String, String, String)],
724 now: &str,
725) -> rusqlite::Result<ToDeviceDedupOutcome> {
726 let tx = conn.transaction()?;
727 if let crate::store::TxnDedupEntry::Seen(_) = crate::store::txn_dedup_lookup(&tx, sender_user_id, sender_device_id, txn_id)? {
728 tx.commit()?;
729 return Ok(ToDeviceDedupOutcome::AlreadySent);
730 }
731 for (recipient_user_id, recipient_device_id, event_type, content) in messages {
732 let stream_id = crate::store::next_stream_id(&tx)?;
733 tx.execute(
734 "INSERT INTO to_device_messages (stream_id, recipient_user_id, recipient_device_id, sender_user_id, event_type, content)
735 VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
736 params![stream_id, recipient_user_id, recipient_device_id, sender_user_id, event_type, content],
737 )?;
738 }
739 crate::store::txn_dedup_record(&tx, sender_user_id, sender_device_id, txn_id, None, now)?;
740 tx.commit()?;
741 Ok(ToDeviceDedupOutcome::New)
742}
743
744pub fn to_device_for(conn: &Connection, user_id: i64, device_id: &str, after_stream: i64, limit: i64) -> rusqlite::Result<Vec<ToDeviceMessage>> {
747 let mut stmt = conn.prepare(&format!(
748 "SELECT {TO_DEVICE_SELECT_COLUMNS} FROM to_device_messages
749 WHERE recipient_user_id = ?1 AND recipient_device_id = ?2 AND stream_id > ?3
750 ORDER BY stream_id ASC LIMIT ?4"
751 ))?;
752 let rows = stmt.query_map(params![user_id, device_id, after_stream, limit], to_device_from_row)?;
753 rows.collect()
754}
755
756pub fn delete_to_device_up_to(conn: &Connection, user_id: i64, device_id: &str, stream_id: i64) -> rusqlite::Result<usize> {
762 conn.execute(
763 "DELETE FROM to_device_messages WHERE recipient_user_id = ?1 AND recipient_device_id = ?2 AND stream_id <= ?3",
764 params![user_id, device_id, stream_id],
765 )
766}
767
768fn log_device_list_change_tx(tx: &Transaction, user_id: i64, now: &str) -> rusqlite::Result<i64> {
773 let stream_id = crate::store::next_stream_id(tx)?;
774 tx.execute(
775 "INSERT INTO device_list_changes (stream_id, user_id, changed_at) VALUES (?1, ?2, ?3)",
776 params![stream_id, user_id, now],
777 )?;
778 crate::fed_edus::enqueue_device_list(tx, user_id, stream_id);
779 Ok(stream_id)
780}
781
782pub fn log_device_list_change(conn: &mut Connection, user_id: i64, now: &str) -> rusqlite::Result<i64> {
788 let tx = conn.transaction()?;
789 let stream_id = log_device_list_change_tx(&tx, user_id, now)?;
790 tx.commit()?;
791 Ok(stream_id)
792}
793
794pub fn device_list_changes_between(conn: &Connection, from_exclusive: i64, to_inclusive: i64) -> rusqlite::Result<Vec<i64>> {
799 let mut stmt = conn.prepare("SELECT DISTINCT user_id FROM device_list_changes WHERE stream_id > ?1 AND stream_id <= ?2")?;
800 let rows = stmt.query_map(params![from_exclusive, to_inclusive], |row| row.get(0))?;
801 rows.collect()
802}
803
804#[derive(Debug, Clone, PartialEq)]
809pub struct CrossSigningKey {
810 pub user_id: i64,
811 pub usage: CrossSigningUsage,
812 pub key_json: String,
813 pub uploaded_at: String,
814}
815
816const CROSS_SIGNING_KEY_SELECT_COLUMNS: &str = "user_id, usage, key_json, uploaded_at";
817
818fn cross_signing_key_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<CrossSigningKey> {
819 let usage_raw: String = row.get(1)?;
820 Ok(CrossSigningKey {
821 user_id: row.get(0)?,
822 usage: decode_enum(1, "usage", &usage_raw, CrossSigningUsage::from_wire_name)?,
823 key_json: row.get(2)?,
824 uploaded_at: row.get(3)?,
825 })
826}
827
828pub fn upsert_cross_signing_key(conn: &mut Connection, user_id: i64, usage: CrossSigningUsage, key_json: &str, now: &str) -> rusqlite::Result<()> {
832 let tx = conn.transaction()?;
833 tx.execute(
834 "INSERT INTO cross_signing_keys (user_id, usage, key_json, uploaded_at) VALUES (?1, ?2, ?3, ?4)
835 ON CONFLICT(user_id, usage) DO UPDATE SET key_json = excluded.key_json, uploaded_at = excluded.uploaded_at",
836 params![user_id, usage.as_str(), key_json, now],
837 )?;
838 log_device_list_change_tx(&tx, user_id, now)?;
839 tx.commit()?;
840 Ok(())
841}
842
843pub fn cross_signing_key_for(conn: &Connection, user_id: i64, usage: CrossSigningUsage) -> rusqlite::Result<Option<CrossSigningKey>> {
850 conn.query_row(
851 &format!("SELECT {CROSS_SIGNING_KEY_SELECT_COLUMNS} FROM cross_signing_keys WHERE user_id = ?1 AND usage = ?2"),
852 params![user_id, usage.as_str()],
853 cross_signing_key_from_row,
854 )
855 .optional()
856}
857
858pub fn cross_signing_keys_for(conn: &Connection, user_ids: &[i64]) -> rusqlite::Result<Vec<CrossSigningKey>> {
861 if user_ids.is_empty() {
862 return Ok(Vec::new());
863 }
864 let placeholders = vec!["?"; user_ids.len()].join(",");
865 let sql = format!("SELECT {CROSS_SIGNING_KEY_SELECT_COLUMNS} FROM cross_signing_keys WHERE user_id IN ({placeholders})");
866 let mut stmt = conn.prepare(&sql)?;
867 let bound: Vec<&dyn rusqlite::ToSql> = user_ids.iter().map(|id| id as &dyn rusqlite::ToSql).collect();
868 let rows = stmt.query_map(bound.as_slice(), cross_signing_key_from_row)?;
869 rows.collect()
870}
871
872#[derive(Debug, Clone, PartialEq)]
873pub struct CrossSigningSignature {
874 pub id: i64,
875 pub signer_user_id: i64,
876 pub target_user_id: i64,
877 pub target_key_id: String,
878 pub signature_json: String,
879 pub uploaded_at: String,
880}
881
882const CROSS_SIGNING_SIGNATURE_SELECT_COLUMNS: &str = "id, signer_user_id, target_user_id, target_key_id, signature_json, uploaded_at";
883
884fn cross_signing_signature_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<CrossSigningSignature> {
885 Ok(CrossSigningSignature {
886 id: row.get(0)?,
887 signer_user_id: row.get(1)?,
888 target_user_id: row.get(2)?,
889 target_key_id: row.get(3)?,
890 signature_json: row.get(4)?,
891 uploaded_at: row.get(5)?,
892 })
893}
894
895pub fn add_signatures(conn: &mut Connection, signatures: &[(i64, i64, String, String, String)]) -> rusqlite::Result<()> {
899 let tx = conn.transaction()?;
900 for (signer_user_id, target_user_id, target_key_id, signature_json, uploaded_at) in signatures {
901 tx.execute(
902 "INSERT INTO cross_signing_signatures (signer_user_id, target_user_id, target_key_id, signature_json, uploaded_at)
903 VALUES (?1, ?2, ?3, ?4, ?5)",
904 params![signer_user_id, target_user_id, target_key_id, signature_json, uploaded_at],
905 )?;
906 }
907 tx.commit()?;
908 Ok(())
909}
910
911pub fn signatures_for(conn: &Connection, target_user_id: i64, target_key_id: &str) -> rusqlite::Result<Vec<CrossSigningSignature>> {
913 let mut stmt = conn.prepare(&format!(
914 "SELECT {CROSS_SIGNING_SIGNATURE_SELECT_COLUMNS} FROM cross_signing_signatures WHERE target_user_id = ?1 AND target_key_id = ?2"
915 ))?;
916 let rows = stmt.query_map(params![target_user_id, target_key_id], cross_signing_signature_from_row)?;
917 rows.collect()
918}
919
920#[derive(Debug, Clone, PartialEq)]
925pub struct KeyBackupVersion {
926 pub version: i64,
927 pub user_id: i64,
928 pub algorithm: String,
929 pub auth_data: String,
930 pub etag: i64,
931 pub is_deleted: bool,
932 pub created_at: String,
933}
934
935const KEY_BACKUP_VERSION_SELECT_COLUMNS: &str = "version, user_id, algorithm, auth_data, etag, is_deleted, created_at";
936
937fn key_backup_version_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<KeyBackupVersion> {
938 Ok(KeyBackupVersion {
939 version: row.get(0)?,
940 user_id: row.get(1)?,
941 algorithm: row.get(2)?,
942 auth_data: row.get(3)?,
943 etag: row.get(4)?,
944 is_deleted: row.get(5)?,
945 created_at: row.get(6)?,
946 })
947}
948
949pub fn create_backup_version(conn: &Connection, user_id: i64, algorithm: &str, auth_data: &str, now: &str) -> rusqlite::Result<i64> {
953 conn.execute(
954 "INSERT INTO key_backup_versions (user_id, algorithm, auth_data, etag, is_deleted, created_at)
955 VALUES (?1, ?2, ?3, 0, 0, ?4)",
956 params![user_id, algorithm, auth_data, now],
957 )?;
958 Ok(conn.last_insert_rowid())
959}
960
961pub fn current_backup_version(conn: &Connection, user_id: i64) -> rusqlite::Result<Option<KeyBackupVersion>> {
964 conn.query_row(
965 &format!(
966 "SELECT {KEY_BACKUP_VERSION_SELECT_COLUMNS} FROM key_backup_versions
967 WHERE user_id = ?1 AND is_deleted = 0 ORDER BY version DESC LIMIT 1"
968 ),
969 params![user_id],
970 key_backup_version_from_row,
971 )
972 .optional()
973}
974
975pub fn get_backup_version(conn: &Connection, user_id: i64, version: i64) -> rusqlite::Result<Option<KeyBackupVersion>> {
980 conn.query_row(
981 &format!("SELECT {KEY_BACKUP_VERSION_SELECT_COLUMNS} FROM key_backup_versions WHERE user_id = ?1 AND version = ?2"),
982 params![user_id, version],
983 key_backup_version_from_row,
984 )
985 .optional()
986}
987
988pub fn update_backup_version_auth_data(conn: &Connection, user_id: i64, version: i64, auth_data: &str) -> rusqlite::Result<bool> {
991 let changed = conn.execute(
992 "UPDATE key_backup_versions SET auth_data = ?1, etag = etag + 1 WHERE user_id = ?2 AND version = ?3 AND is_deleted = 0",
993 params![auth_data, user_id, version],
994 )?;
995 Ok(changed > 0)
996}
997
998pub fn delete_backup_version(conn: &Connection, user_id: i64, version: i64) -> rusqlite::Result<bool> {
1002 let changed = conn.execute(
1003 "UPDATE key_backup_versions SET is_deleted = 1, etag = etag + 1 WHERE user_id = ?1 AND version = ?2 AND is_deleted = 0",
1004 params![user_id, version],
1005 )?;
1006 Ok(changed > 0)
1007}
1008
1009#[derive(Debug, Clone, PartialEq)]
1010pub struct KeyBackupSession {
1011 pub user_id: i64,
1012 pub version: i64,
1013 pub room_id: String,
1014 pub session_id: String,
1015 pub session_data: String,
1016 pub updated_at: String,
1017}
1018
1019const KEY_BACKUP_SESSION_SELECT_COLUMNS: &str = "user_id, version, room_id, session_id, session_data, updated_at";
1020
1021fn key_backup_session_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<KeyBackupSession> {
1022 Ok(KeyBackupSession {
1023 user_id: row.get(0)?,
1024 version: row.get(1)?,
1025 room_id: row.get(2)?,
1026 session_id: row.get(3)?,
1027 session_data: row.get(4)?,
1028 updated_at: row.get(5)?,
1029 })
1030}
1031
1032pub fn put_backup_sessions(
1037 conn: &mut Connection,
1038 user_id: i64,
1039 version: i64,
1040 sessions: &[(String, String, String)],
1041 now: &str,
1042) -> Result<(), MatrixKeysStoreError> {
1043 let tx = conn.transaction()?;
1044 let current_version: Option<i64> = tx.query_row(
1045 "SELECT MAX(version) FROM key_backup_versions WHERE user_id = ?1 AND is_deleted = 0",
1046 params![user_id],
1047 |row| row.get(0),
1048 )?;
1049 if current_version != Some(version) {
1050 return Err(MatrixKeysStoreError::WrongBackupVersion);
1051 }
1052 for (room_id, session_id, session_data) in sessions {
1053 tx.execute(
1054 "INSERT INTO key_backup_sessions (user_id, version, room_id, session_id, session_data, updated_at)
1055 VALUES (?1, ?2, ?3, ?4, ?5, ?6)
1056 ON CONFLICT(user_id, version, room_id, session_id) DO UPDATE SET
1057 session_data = excluded.session_data, updated_at = excluded.updated_at",
1058 params![user_id, version, room_id, session_id, session_data, now],
1059 )?;
1060 }
1061 tx.execute(
1062 "UPDATE key_backup_versions SET etag = etag + 1 WHERE user_id = ?1 AND version = ?2",
1063 params![user_id, version],
1064 )?;
1065 tx.commit()?;
1066 Ok(())
1067}
1068
1069pub fn get_backup_sessions(
1074 conn: &Connection,
1075 user_id: i64,
1076 version: i64,
1077 room_id: Option<&str>,
1078 session_id: Option<&str>,
1079) -> rusqlite::Result<Vec<KeyBackupSession>> {
1080 match (room_id, session_id) {
1081 (Some(room_id), Some(session_id)) => {
1082 let mut stmt = conn.prepare(&format!(
1083 "SELECT {KEY_BACKUP_SESSION_SELECT_COLUMNS} FROM key_backup_sessions
1084 WHERE user_id = ?1 AND version = ?2 AND room_id = ?3 AND session_id = ?4"
1085 ))?;
1086 let rows = stmt.query_map(params![user_id, version, room_id, session_id], key_backup_session_from_row)?;
1087 rows.collect()
1088 }
1089 (Some(room_id), None) => {
1090 let mut stmt = conn.prepare(&format!(
1091 "SELECT {KEY_BACKUP_SESSION_SELECT_COLUMNS} FROM key_backup_sessions
1092 WHERE user_id = ?1 AND version = ?2 AND room_id = ?3"
1093 ))?;
1094 let rows = stmt.query_map(params![user_id, version, room_id], key_backup_session_from_row)?;
1095 rows.collect()
1096 }
1097 (None, _) => {
1098 let mut stmt = conn.prepare(&format!(
1099 "SELECT {KEY_BACKUP_SESSION_SELECT_COLUMNS} FROM key_backup_sessions WHERE user_id = ?1 AND version = ?2"
1100 ))?;
1101 let rows = stmt.query_map(params![user_id, version], key_backup_session_from_row)?;
1102 rows.collect()
1103 }
1104 }
1105}
1106
1107pub fn delete_backup_sessions(
1111 conn: &mut Connection,
1112 user_id: i64,
1113 version: i64,
1114 room_id: Option<&str>,
1115 session_id: Option<&str>,
1116) -> rusqlite::Result<usize> {
1117 let tx = conn.transaction()?;
1118 let deleted = match (room_id, session_id) {
1119 (Some(room_id), Some(session_id)) => tx.execute(
1120 "DELETE FROM key_backup_sessions WHERE user_id = ?1 AND version = ?2 AND room_id = ?3 AND session_id = ?4",
1121 params![user_id, version, room_id, session_id],
1122 )?,
1123 (Some(room_id), None) => tx.execute(
1124 "DELETE FROM key_backup_sessions WHERE user_id = ?1 AND version = ?2 AND room_id = ?3",
1125 params![user_id, version, room_id],
1126 )?,
1127 (None, _) => tx.execute(
1128 "DELETE FROM key_backup_sessions WHERE user_id = ?1 AND version = ?2",
1129 params![user_id, version],
1130 )?,
1131 };
1132 if deleted > 0 {
1133 tx.execute(
1134 "UPDATE key_backup_versions SET etag = etag + 1 WHERE user_id = ?1 AND version = ?2",
1135 params![user_id, version],
1136 )?;
1137 }
1138 tx.commit()?;
1139 Ok(deleted)
1140}
1141
1142pub fn backup_count_and_etag(conn: &Connection, user_id: i64, version: i64) -> rusqlite::Result<(i64, i64)> {
1145 let count: i64 = conn.query_row(
1146 "SELECT COUNT(*) FROM key_backup_sessions WHERE user_id = ?1 AND version = ?2",
1147 params![user_id, version],
1148 |row| row.get(0),
1149 )?;
1150 let etag: i64 = conn.query_row(
1151 "SELECT etag FROM key_backup_versions WHERE user_id = ?1 AND version = ?2",
1152 params![user_id, version],
1153 |row| row.get(0),
1154 )?;
1155 Ok((count, etag))
1156}
1157
1158#[cfg(test)]
1159mod tests {
1160 use super::*;
1161
1162 const T0: &str = "2026-09-24T00:00:00+00:00";
1163
1164 fn test_conn() -> Connection {
1165 let conn = Connection::open_in_memory().expect("in-memory sqlite");
1166 crate::store::create_matrix_schema(&conn).expect("matrix schema (stream_counter lives there)");
1167 create_matrix_keys_schema(&conn).expect("matrix keys schema");
1168 conn
1169 }
1170
1171 fn count_device_list_changes(conn: &Connection, user_id: i64) -> i64 {
1172 conn.query_row("SELECT COUNT(*) FROM device_list_changes WHERE user_id = ?1", params![user_id], |row| row.get(0))
1173 .expect("count device_list_changes")
1174 }
1175
1176 #[test]
1179 fn claim_one_time_key_deletes_it_so_a_second_claim_gets_a_different_key_or_none() {
1180 let mut conn = test_conn();
1181 add_one_time_keys(
1182 &mut conn,
1183 1,
1184 "DEV1",
1185 &[
1186 ("signed_curve25519:AAAAAQ".to_string(), "signed_curve25519".to_string(), r#"{"key":"k1"}"#.to_string()),
1187 ("signed_curve25519:AAAAAg".to_string(), "signed_curve25519".to_string(), r#"{"key":"k2"}"#.to_string()),
1188 ],
1189 )
1190 .expect("add otks");
1191
1192 let first = claim_one_time_key(&mut conn, 1, "DEV1", "signed_curve25519").expect("claim 1").expect("has a key");
1193 let second = claim_one_time_key(&mut conn, 1, "DEV1", "signed_curve25519").expect("claim 2").expect("has a different key");
1194 assert_ne!(first.0, second.0, "the two claims must return different key ids");
1195
1196 let third = claim_one_time_key(&mut conn, 1, "DEV1", "signed_curve25519").expect("claim 3");
1197 assert_eq!(third, None, "no one-time keys or fallback keys remain");
1198 }
1199
1200 #[test]
1203 fn claim_falls_back_to_a_fallback_key_without_deleting_it_and_marks_it_used() {
1204 let mut conn = test_conn();
1205 upsert_fallback_key(&conn, 1, "DEV1", "signed_curve25519", "signed_curve25519:FALLBACK", r#"{"key":"fb"}"#, T0).expect("upsert fallback");
1206
1207 let claimed = claim_one_time_key(&mut conn, 1, "DEV1", "signed_curve25519").expect("claim").expect("fallback returned");
1208 assert_eq!(claimed.0, "signed_curve25519:FALLBACK");
1209 assert_eq!(claimed.1, r#"{"key":"fb"}"#);
1210
1211 let claimed_again = claim_one_time_key(&mut conn, 1, "DEV1", "signed_curve25519").expect("claim again").expect("fallback still there");
1212 assert_eq!(claimed_again.0, "signed_curve25519:FALLBACK", "a fallback key is never deleted on claim");
1213
1214 let unused = unused_fallback_key_types(&conn, 1, "DEV1").expect("unused types");
1215 assert!(unused.is_empty(), "the fallback key must be marked used after its first claim");
1216 }
1217
1218 #[test]
1221 fn device_unused_fallback_key_types_excludes_a_used_one() {
1222 let mut conn = test_conn();
1223 upsert_fallback_key(&conn, 1, "DEV1", "signed_curve25519", "signed_curve25519:FB1", r#"{"key":"fb1"}"#, T0).expect("fallback 1");
1224 upsert_fallback_key(&conn, 1, "DEV1", "olm_curve25519", "olm_curve25519:FB2", r#"{"key":"fb2"}"#, T0).expect("fallback 2");
1225
1226 let before = unused_fallback_key_types(&conn, 1, "DEV1").expect("before claim");
1227 assert_eq!(before.len(), 2);
1228
1229 claim_one_time_key(&mut conn, 1, "DEV1", "signed_curve25519").expect("claim marks it used");
1230
1231 let after = unused_fallback_key_types(&conn, 1, "DEV1").expect("after claim");
1232 assert_eq!(after, vec!["olm_curve25519".to_string()]);
1233 }
1234
1235 #[test]
1239 fn delete_device_by_credential_deletes_the_device_and_its_key_rows_and_logs_a_device_list_change() {
1240 let mut conn = test_conn();
1241 let device_id = create_device(&conn, 1, CredentialKind::Bearer, "tok-hash-1", T0).expect("create device");
1242 upsert_device_keys(&mut conn, 1, &device_id, "[]", "{}", "{}", T0).expect("upload device keys");
1243 add_one_time_keys(
1244 &mut conn,
1245 1,
1246 &device_id,
1247 &[("signed_curve25519:AAAAAQ".to_string(), "signed_curve25519".to_string(), "{}".to_string())],
1248 )
1249 .expect("otk");
1250 upsert_fallback_key(&conn, 1, &device_id, "signed_curve25519", "signed_curve25519:FB", "{}", T0).expect("fallback");
1251 enqueue_to_device(&mut conn, 2, &[(1, device_id.clone(), "m.text".to_string(), "{}".to_string())]).expect("to-device");
1252
1253 let before = count_device_list_changes(&conn, 1);
1254 let deleted = delete_device_by_credential(&mut conn, CredentialKind::Bearer, "tok-hash-1", T0).expect("revoke");
1255 assert_eq!(deleted, Some((1, device_id.clone())));
1256
1257 assert_eq!(get_device(&conn, 1, &device_id).expect("get"), None);
1258 let keys_count: i64 = conn
1259 .query_row("SELECT COUNT(*) FROM device_keys WHERE user_id = 1 AND device_id = ?1", params![device_id], |row| row.get(0))
1260 .expect("keys count");
1261 assert_eq!(keys_count, 0);
1262 let otk_count: i64 = conn
1263 .query_row("SELECT COUNT(*) FROM one_time_keys WHERE user_id = 1 AND device_id = ?1", params![device_id], |row| row.get(0))
1264 .expect("otk count");
1265 assert_eq!(otk_count, 0);
1266 let fallback_count: i64 = conn
1267 .query_row("SELECT COUNT(*) FROM fallback_keys WHERE user_id = 1 AND device_id = ?1", params![device_id], |row| row.get(0))
1268 .expect("fallback count");
1269 assert_eq!(fallback_count, 0);
1270 let to_device_count: i64 = conn
1271 .query_row("SELECT COUNT(*) FROM to_device_messages WHERE recipient_device_id = ?1", params![device_id], |row| row.get(0))
1272 .expect("to-device count");
1273 assert_eq!(to_device_count, 0);
1274
1275 let after = count_device_list_changes(&conn, 1);
1276 assert_eq!(after, before + 1, "exactly one device_list_changes row must be appended");
1277
1278 assert_eq!(
1279 delete_device_by_credential(&mut conn, CredentialKind::Bearer, "tok-hash-1", T0).expect("second revoke is a no-op"),
1280 None
1281 );
1282 }
1283
1284 #[test]
1287 fn add_one_time_keys_is_idempotent_for_identical_json_and_refuses_changed_json() {
1288 let mut conn = test_conn();
1289 let key = ("signed_curve25519:AAAAAQ".to_string(), "signed_curve25519".to_string(), r#"{"key":"k1"}"#.to_string());
1290 add_one_time_keys(&mut conn, 1, "DEV1", &[key.clone()]).expect("first add");
1291 add_one_time_keys(&mut conn, 1, "DEV1", &[key.clone()]).expect("identical resubmission is a no-op");
1292
1293 let count: i64 = conn
1294 .query_row("SELECT COUNT(*) FROM one_time_keys WHERE user_id = 1 AND device_id = 'DEV1'", [], |row| row.get(0))
1295 .expect("count");
1296 assert_eq!(count, 1);
1297
1298 let changed = ("signed_curve25519:AAAAAQ".to_string(), "signed_curve25519".to_string(), r#"{"key":"k2-different"}"#.to_string());
1299 let err = add_one_time_keys(&mut conn, 1, "DEV1", &[changed]).unwrap_err();
1300 assert!(matches!(err, MatrixKeysStoreError::OneTimeKeyConflict(ref id) if id == "signed_curve25519:AAAAAQ"));
1301 }
1302
1303 #[test]
1306 fn to_device_delete_up_to_leaves_later_messages() {
1307 let mut conn = test_conn();
1308 let s1 = enqueue_to_device(&mut conn, 2, &[(1, "DEV1".to_string(), "m.a".to_string(), "{}".to_string())]).expect("send 1");
1309 let s2 = enqueue_to_device(&mut conn, 2, &[(1, "DEV1".to_string(), "m.b".to_string(), "{}".to_string())]).expect("send 2");
1310 assert!(s2 > s1);
1311
1312 let deleted = delete_to_device_up_to(&conn, 1, "DEV1", s1).expect("delete up to s1");
1313 assert_eq!(deleted, 1);
1314
1315 let remaining = to_device_for(&conn, 1, "DEV1", 0, 10).expect("remaining");
1316 assert_eq!(remaining.len(), 1);
1317 assert_eq!(remaining[0].stream_id, s2);
1318 }
1319
1320 #[test]
1323 fn device_keys_change_logs_a_device_list_change() {
1324 let mut conn = test_conn();
1325 let before = count_device_list_changes(&conn, 1);
1326 upsert_device_keys(&mut conn, 1, "DEV1", "[\"m.olm.v1\"]", "{}", "{}", T0).expect("upload keys");
1327 let after = count_device_list_changes(&conn, 1);
1328 assert_eq!(after, before + 1);
1329
1330 let rows = device_keys_for(&conn, &[1]).expect("query");
1331 assert_eq!(rows.len(), 1);
1332 assert_eq!(rows[0].device_id, "DEV1");
1333
1334 assert_eq!(device_keys_for(&conn, &[]).expect("empty input"), Vec::new());
1335 }
1336
1337 #[test]
1340 fn put_backup_sessions_refuses_a_stale_version() {
1341 let mut conn = test_conn();
1342 let v1 = create_backup_version(&conn, 1, "m.megolm_backup.v1", "{}", T0).expect("v1");
1343 let v2 = create_backup_version(&conn, 1, "m.megolm_backup.v1", "{}", T0).expect("v2");
1344 assert!(v2 > v1);
1345
1346 let err = put_backup_sessions(&mut conn, 1, v1, &[("!room:x".to_string(), "sess1".to_string(), "{}".to_string())], T0).unwrap_err();
1347 assert!(matches!(err, MatrixKeysStoreError::WrongBackupVersion));
1348
1349 put_backup_sessions(&mut conn, 1, v2, &[("!room:x".to_string(), "sess1".to_string(), "{}".to_string())], T0).expect("current version accepted");
1350 }
1351
1352 #[test]
1355 fn delete_device_cascades_keys_and_pending_to_device() {
1356 let mut conn = test_conn();
1357 let device_id = create_device(&conn, 1, CredentialKind::Web, "sess-1", T0).expect("create device");
1358 upsert_device_keys(&mut conn, 1, &device_id, "[]", "{}", "{}", T0).expect("device keys");
1359 add_one_time_keys(
1360 &mut conn,
1361 1,
1362 &device_id,
1363 &[("signed_curve25519:AAAAAQ".to_string(), "signed_curve25519".to_string(), "{}".to_string())],
1364 )
1365 .expect("otk");
1366 upsert_fallback_key(&conn, 1, &device_id, "signed_curve25519", "signed_curve25519:FB", "{}", T0).expect("fallback");
1367 enqueue_to_device(&mut conn, 9, &[(1, device_id.clone(), "m.text".to_string(), "{}".to_string())]).expect("to-device");
1368
1369 let before = count_device_list_changes(&conn, 1);
1370 let deleted = delete_device(&mut conn, 1, &device_id, T0).expect("delete");
1371 assert!(deleted);
1372
1373 assert_eq!(get_device(&conn, 1, &device_id).expect("get"), None);
1374 assert!(device_keys_for(&conn, &[1]).expect("keys").is_empty());
1375 assert_eq!(count_one_time_keys(&conn, 1, &device_id).expect("otk count").len(), 0);
1376 assert!(unused_fallback_key_types(&conn, 1, &device_id).expect("fallback").is_empty());
1377 assert!(to_device_for(&conn, 1, &device_id, 0, 10).expect("to-device").is_empty());
1378
1379 let after = count_device_list_changes(&conn, 1);
1380 assert_eq!(after, before + 1);
1381
1382 let deleted_again = delete_device(&mut conn, 1, &device_id, T0).expect("second delete is a no-op");
1383 assert!(!deleted_again);
1384 }
1385
1386 #[test]
1389 fn create_device_mints_a_fresh_device_id_and_touch_device_updates_last_seen() {
1390 let conn = test_conn();
1391 let d1 = create_device(&conn, 1, CredentialKind::Bearer, "tok-a", T0).expect("create 1");
1392 let d2 = create_device(&conn, 1, CredentialKind::Bearer, "tok-b", T0).expect("create 2");
1393 assert_ne!(d1, d2, "two different credentials must mint different device ids");
1394
1395 touch_device(&conn, 1, &d1, "2026-09-24T01:00:00+00:00").expect("touch");
1396 let row = get_device(&conn, 1, &d1).expect("get").expect("row exists");
1397 assert_eq!(row.last_seen_at, "2026-09-24T01:00:00+00:00");
1398 assert_eq!(row.created_at, T0, "created_at must not move on touch");
1399 }
1400
1401 #[test]
1402 fn device_for_credential_finds_the_row_created_by_create_device() {
1403 let conn = test_conn();
1404 let device_id = create_device(&conn, 5, CredentialKind::Web, "sess-xyz", T0).expect("create");
1405 let found = device_for_credential(&conn, CredentialKind::Web, "sess-xyz").expect("lookup").expect("row exists");
1406 assert_eq!(found.user_id, 5);
1407 assert_eq!(found.device_id, device_id);
1408
1409 assert_eq!(device_for_credential(&conn, CredentialKind::Bearer, "sess-xyz").expect("wrong kind"), None);
1410 }
1411
1412 #[test]
1413 fn set_device_display_name_updates_only_the_named_device() {
1414 let conn = test_conn();
1415 let d1 = create_device(&conn, 1, CredentialKind::Bearer, "a", T0).expect("d1");
1416 let d2 = create_device(&conn, 1, CredentialKind::Bearer, "b", T0).expect("d2");
1417
1418 assert!(set_device_display_name(&conn, 1, &d1, Some("My Phone")).expect("set"));
1419 assert!(!set_device_display_name(&conn, 1, "nonexistent", Some("x")).expect("missing device is a no-op returning false"));
1420
1421 let devices = list_devices(&conn, 1).expect("list");
1422 assert_eq!(devices.len(), 2);
1423 let named = devices.iter().find(|d| d.device_id == d1).expect("d1 present");
1424 assert_eq!(named.display_name.as_deref(), Some("My Phone"));
1425 let unnamed = devices.iter().find(|d| d.device_id == d2).expect("d2 present");
1426 assert_eq!(unnamed.display_name, None);
1427 }
1428
1429 #[test]
1430 fn devices_for_reaper_lists_every_device_with_its_credential() {
1431 let conn = test_conn();
1432 create_device(&conn, 1, CredentialKind::Bearer, "tok-1", T0).expect("d1");
1433 create_device(&conn, 2, CredentialKind::Web, "sess-2", T0).expect("d2");
1434
1435 let mut all = devices_for_reaper(&conn).expect("reaper list");
1436 all.sort_by_key(|(user_id, ..)| *user_id);
1437 assert_eq!(all.len(), 2);
1438 assert_eq!(all[0].0, 1);
1439 assert_eq!(all[0].2, CredentialKind::Bearer);
1440 assert_eq!(all[0].3, "tok-1");
1441 assert_eq!(all[1].0, 2);
1442 assert_eq!(all[1].2, CredentialKind::Web);
1443 }
1444
1445 #[test]
1446 fn count_one_time_keys_groups_by_algorithm() {
1447 let mut conn = test_conn();
1448 add_one_time_keys(
1449 &mut conn,
1450 1,
1451 "DEV1",
1452 &[
1453 ("signed_curve25519:A".to_string(), "signed_curve25519".to_string(), "{}".to_string()),
1454 ("signed_curve25519:B".to_string(), "signed_curve25519".to_string(), "{}".to_string()),
1455 ("other_algo:C".to_string(), "other_algo".to_string(), "{}".to_string()),
1456 ],
1457 )
1458 .expect("add");
1459
1460 let counts = count_one_time_keys(&conn, 1, "DEV1").expect("count");
1461 assert_eq!(counts.get("signed_curve25519"), Some(&2));
1462 assert_eq!(counts.get("other_algo"), Some(&1));
1463 }
1464
1465 #[test]
1466 fn cross_signing_upsert_round_trips_and_logs_a_device_list_change() {
1467 let mut conn = test_conn();
1468 let before = count_device_list_changes(&conn, 1);
1469 upsert_cross_signing_key(&mut conn, 1, CrossSigningUsage::Master, r#"{"keys":{}}"#, T0).expect("master");
1470 upsert_cross_signing_key(&mut conn, 1, CrossSigningUsage::SelfSigning, r#"{"keys":{}}"#, T0).expect("self signing");
1471
1472 let after = count_device_list_changes(&conn, 1);
1473 assert_eq!(after, before + 2);
1474
1475 let keys = cross_signing_keys_for(&conn, &[1]).expect("query");
1476 assert_eq!(keys.len(), 2);
1477 assert!(keys.iter().any(|k| k.usage == CrossSigningUsage::Master));
1478 assert!(keys.iter().any(|k| k.usage == CrossSigningUsage::SelfSigning));
1479
1480 assert_eq!(cross_signing_keys_for(&conn, &[]).expect("empty input"), Vec::new());
1481 }
1482
1483 #[test]
1484 fn add_signatures_and_signatures_for_round_trip() {
1485 let mut conn = test_conn();
1486 add_signatures(&mut conn, &[(1, 2, "DEVICEX".to_string(), r#"{"sig":"abc"}"#.to_string(), T0.to_string())]).expect("add");
1487
1488 let sigs = signatures_for(&conn, 2, "DEVICEX").expect("query");
1489 assert_eq!(sigs.len(), 1);
1490 assert_eq!(sigs[0].signer_user_id, 1);
1491 assert_eq!(sigs[0].signature_json, r#"{"sig":"abc"}"#);
1492
1493 assert!(signatures_for(&conn, 2, "OTHER").expect("no match").is_empty());
1494 }
1495
1496 #[test]
1497 fn backup_version_lifecycle_create_get_update_delete() {
1498 let conn = test_conn();
1499 let version = create_backup_version(&conn, 1, "m.megolm_backup.v1", r#"{"a":1}"#, T0).expect("create");
1500
1501 assert_eq!(current_backup_version(&conn, 1).expect("current").expect("row").version, version);
1502
1503 assert!(update_backup_version_auth_data(&conn, 1, version, r#"{"a":2}"#).expect("update"));
1504 let updated = get_backup_version(&conn, 1, version).expect("get").expect("row");
1505 assert_eq!(updated.auth_data, r#"{"a":2}"#);
1506 assert_eq!(updated.etag, 1);
1507
1508 assert!(delete_backup_version(&conn, 1, version).expect("delete"));
1509 assert_eq!(current_backup_version(&conn, 1).expect("current after delete"), None);
1510 assert!(!delete_backup_version(&conn, 1, version).expect("second delete is a no-op"));
1511 }
1512
1513 #[test]
1514 fn backup_sessions_put_get_delete_and_etag_bumps() {
1515 let mut conn = test_conn();
1516 let version = create_backup_version(&conn, 1, "m.megolm_backup.v1", "{}", T0).expect("create");
1517
1518 put_backup_sessions(
1519 &mut conn,
1520 1,
1521 version,
1522 &[
1523 ("!room1:x".to_string(), "sessA".to_string(), r#"{"d":1}"#.to_string()),
1524 ("!room1:x".to_string(), "sessB".to_string(), r#"{"d":2}"#.to_string()),
1525 ("!room2:x".to_string(), "sessC".to_string(), r#"{"d":3}"#.to_string()),
1526 ],
1527 T0,
1528 )
1529 .expect("put");
1530
1531 let (count, etag_after_put) = backup_count_and_etag(&conn, 1, version).expect("count+etag");
1532 assert_eq!(count, 3);
1533 assert_eq!(etag_after_put, 1);
1534
1535 let room1_sessions = get_backup_sessions(&conn, 1, version, Some("!room1:x"), None).expect("room1");
1536 assert_eq!(room1_sessions.len(), 2);
1537
1538 let one = get_backup_sessions(&conn, 1, version, Some("!room1:x"), Some("sessA")).expect("one");
1539 assert_eq!(one.len(), 1);
1540 assert_eq!(one[0].session_data, r#"{"d":1}"#);
1541
1542 let deleted = delete_backup_sessions(&mut conn, 1, version, Some("!room1:x"), None).expect("delete room1");
1543 assert_eq!(deleted, 2);
1544 let (count_after, etag_after_delete) = backup_count_and_etag(&conn, 1, version).expect("count+etag after delete");
1545 assert_eq!(count_after, 1);
1546 assert_eq!(etag_after_delete, 2);
1547 }
1548
1549 #[test]
1550 fn device_list_changes_between_is_distinct_and_bounded() {
1551 let mut conn = test_conn();
1552 log_device_list_change(&mut conn, 1, T0).expect("change 1");
1553 let boundary = crate::store::max_stream_id(&conn).expect("boundary");
1554 log_device_list_change(&mut conn, 1, T0).expect("change 2 for the same user");
1555 log_device_list_change(&mut conn, 2, T0).expect("change for a different user");
1556
1557 let mut changed = device_list_changes_between(&conn, boundary, crate::store::max_stream_id(&conn).expect("max")).expect("query");
1558 changed.sort();
1559 assert_eq!(changed, vec![1, 2]);
1560 }
1561
1562 #[test]
1563 fn cross_signing_key_for_finds_exactly_the_named_usage() {
1564 let mut conn = test_conn();
1565 upsert_cross_signing_key(&mut conn, 1, CrossSigningUsage::Master, r#"{"usage":["master"]}"#, T0).expect("master");
1566
1567 let master = cross_signing_key_for(&conn, 1, CrossSigningUsage::Master).expect("query").expect("row exists");
1568 assert_eq!(master.usage, CrossSigningUsage::Master);
1569 assert_eq!(cross_signing_key_for(&conn, 1, CrossSigningUsage::SelfSigning).expect("query"), None);
1570 assert_eq!(cross_signing_key_for(&conn, 2, CrossSigningUsage::Master).expect("different user"), None);
1571 }
1572
1573 #[test]
1574 fn enqueue_to_device_deduped_is_idempotent_per_txn() {
1575 let mut conn = test_conn();
1576 let messages = [(1_i64, "DEV1".to_string(), "m.room_key".to_string(), r#"{"k":1}"#.to_string())];
1577
1578 let first = enqueue_to_device_deduped(&mut conn, 9, "SENDER_DEV", "txn-1", &messages, T0).expect("first send");
1579 assert_eq!(first, ToDeviceDedupOutcome::New);
1580 assert_eq!(to_device_for(&conn, 1, "DEV1", 0, 10).expect("after first").len(), 1);
1581
1582 let second = enqueue_to_device_deduped(&mut conn, 9, "SENDER_DEV", "txn-1", &messages, T0).expect("repeat send");
1583 assert_eq!(second, ToDeviceDedupOutcome::AlreadySent);
1584 assert_eq!(to_device_for(&conn, 1, "DEV1", 0, 10).expect("after repeat").len(), 1, "a repeated txn_id must not enqueue a second copy");
1585 }
1586
1587 #[test]
1588 fn enqueue_to_device_fans_out_and_is_scoped_to_the_recipient_device() {
1589 let mut conn = test_conn();
1590 let last = enqueue_to_device(
1591 &mut conn,
1592 9,
1593 &[
1594 (1, "DEV1".to_string(), "m.room_key".to_string(), r#"{"k":1}"#.to_string()),
1595 (1, "DEV2".to_string(), "m.room_key".to_string(), r#"{"k":1}"#.to_string()),
1596 ],
1597 )
1598 .expect("enqueue");
1599 assert_eq!(crate::store::max_stream_id(&conn).expect("max"), last);
1600
1601 let for_dev1 = to_device_for(&conn, 1, "DEV1", 0, 10).expect("dev1");
1602 assert_eq!(for_dev1.len(), 1);
1603 let for_dev2 = to_device_for(&conn, 1, "DEV2", 0, 10).expect("dev2");
1604 assert_eq!(for_dev2.len(), 1);
1605 assert_ne!(for_dev1[0].stream_id, for_dev2[0].stream_id);
1606
1607 let empty_batch = enqueue_to_device(&mut conn, 9, &[]).expect("empty batch is a no-op");
1608 assert_eq!(empty_batch, last, "an empty batch must not mint a fresh stream id");
1609 }
1610}