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 Ok(stream_id)
779}
780
781pub fn log_device_list_change(conn: &mut Connection, user_id: i64, now: &str) -> rusqlite::Result<i64> {
787 let tx = conn.transaction()?;
788 let stream_id = log_device_list_change_tx(&tx, user_id, now)?;
789 tx.commit()?;
790 Ok(stream_id)
791}
792
793pub fn device_list_changes_between(conn: &Connection, from_exclusive: i64, to_inclusive: i64) -> rusqlite::Result<Vec<i64>> {
798 let mut stmt = conn.prepare("SELECT DISTINCT user_id FROM device_list_changes WHERE stream_id > ?1 AND stream_id <= ?2")?;
799 let rows = stmt.query_map(params![from_exclusive, to_inclusive], |row| row.get(0))?;
800 rows.collect()
801}
802
803#[derive(Debug, Clone, PartialEq)]
808pub struct CrossSigningKey {
809 pub user_id: i64,
810 pub usage: CrossSigningUsage,
811 pub key_json: String,
812 pub uploaded_at: String,
813}
814
815const CROSS_SIGNING_KEY_SELECT_COLUMNS: &str = "user_id, usage, key_json, uploaded_at";
816
817fn cross_signing_key_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<CrossSigningKey> {
818 let usage_raw: String = row.get(1)?;
819 Ok(CrossSigningKey {
820 user_id: row.get(0)?,
821 usage: decode_enum(1, "usage", &usage_raw, CrossSigningUsage::from_wire_name)?,
822 key_json: row.get(2)?,
823 uploaded_at: row.get(3)?,
824 })
825}
826
827pub fn upsert_cross_signing_key(conn: &mut Connection, user_id: i64, usage: CrossSigningUsage, key_json: &str, now: &str) -> rusqlite::Result<()> {
831 let tx = conn.transaction()?;
832 tx.execute(
833 "INSERT INTO cross_signing_keys (user_id, usage, key_json, uploaded_at) VALUES (?1, ?2, ?3, ?4)
834 ON CONFLICT(user_id, usage) DO UPDATE SET key_json = excluded.key_json, uploaded_at = excluded.uploaded_at",
835 params![user_id, usage.as_str(), key_json, now],
836 )?;
837 log_device_list_change_tx(&tx, user_id, now)?;
838 tx.commit()?;
839 Ok(())
840}
841
842pub fn cross_signing_key_for(conn: &Connection, user_id: i64, usage: CrossSigningUsage) -> rusqlite::Result<Option<CrossSigningKey>> {
849 conn.query_row(
850 &format!("SELECT {CROSS_SIGNING_KEY_SELECT_COLUMNS} FROM cross_signing_keys WHERE user_id = ?1 AND usage = ?2"),
851 params![user_id, usage.as_str()],
852 cross_signing_key_from_row,
853 )
854 .optional()
855}
856
857pub fn cross_signing_keys_for(conn: &Connection, user_ids: &[i64]) -> rusqlite::Result<Vec<CrossSigningKey>> {
860 if user_ids.is_empty() {
861 return Ok(Vec::new());
862 }
863 let placeholders = vec!["?"; user_ids.len()].join(",");
864 let sql = format!("SELECT {CROSS_SIGNING_KEY_SELECT_COLUMNS} FROM cross_signing_keys WHERE user_id IN ({placeholders})");
865 let mut stmt = conn.prepare(&sql)?;
866 let bound: Vec<&dyn rusqlite::ToSql> = user_ids.iter().map(|id| id as &dyn rusqlite::ToSql).collect();
867 let rows = stmt.query_map(bound.as_slice(), cross_signing_key_from_row)?;
868 rows.collect()
869}
870
871#[derive(Debug, Clone, PartialEq)]
872pub struct CrossSigningSignature {
873 pub id: i64,
874 pub signer_user_id: i64,
875 pub target_user_id: i64,
876 pub target_key_id: String,
877 pub signature_json: String,
878 pub uploaded_at: String,
879}
880
881const CROSS_SIGNING_SIGNATURE_SELECT_COLUMNS: &str = "id, signer_user_id, target_user_id, target_key_id, signature_json, uploaded_at";
882
883fn cross_signing_signature_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<CrossSigningSignature> {
884 Ok(CrossSigningSignature {
885 id: row.get(0)?,
886 signer_user_id: row.get(1)?,
887 target_user_id: row.get(2)?,
888 target_key_id: row.get(3)?,
889 signature_json: row.get(4)?,
890 uploaded_at: row.get(5)?,
891 })
892}
893
894pub fn add_signatures(conn: &mut Connection, signatures: &[(i64, i64, String, String, String)]) -> rusqlite::Result<()> {
898 let tx = conn.transaction()?;
899 for (signer_user_id, target_user_id, target_key_id, signature_json, uploaded_at) in signatures {
900 tx.execute(
901 "INSERT INTO cross_signing_signatures (signer_user_id, target_user_id, target_key_id, signature_json, uploaded_at)
902 VALUES (?1, ?2, ?3, ?4, ?5)",
903 params![signer_user_id, target_user_id, target_key_id, signature_json, uploaded_at],
904 )?;
905 }
906 tx.commit()?;
907 Ok(())
908}
909
910pub fn signatures_for(conn: &Connection, target_user_id: i64, target_key_id: &str) -> rusqlite::Result<Vec<CrossSigningSignature>> {
912 let mut stmt = conn.prepare(&format!(
913 "SELECT {CROSS_SIGNING_SIGNATURE_SELECT_COLUMNS} FROM cross_signing_signatures WHERE target_user_id = ?1 AND target_key_id = ?2"
914 ))?;
915 let rows = stmt.query_map(params![target_user_id, target_key_id], cross_signing_signature_from_row)?;
916 rows.collect()
917}
918
919#[derive(Debug, Clone, PartialEq)]
924pub struct KeyBackupVersion {
925 pub version: i64,
926 pub user_id: i64,
927 pub algorithm: String,
928 pub auth_data: String,
929 pub etag: i64,
930 pub is_deleted: bool,
931 pub created_at: String,
932}
933
934const KEY_BACKUP_VERSION_SELECT_COLUMNS: &str = "version, user_id, algorithm, auth_data, etag, is_deleted, created_at";
935
936fn key_backup_version_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<KeyBackupVersion> {
937 Ok(KeyBackupVersion {
938 version: row.get(0)?,
939 user_id: row.get(1)?,
940 algorithm: row.get(2)?,
941 auth_data: row.get(3)?,
942 etag: row.get(4)?,
943 is_deleted: row.get(5)?,
944 created_at: row.get(6)?,
945 })
946}
947
948pub fn create_backup_version(conn: &Connection, user_id: i64, algorithm: &str, auth_data: &str, now: &str) -> rusqlite::Result<i64> {
952 conn.execute(
953 "INSERT INTO key_backup_versions (user_id, algorithm, auth_data, etag, is_deleted, created_at)
954 VALUES (?1, ?2, ?3, 0, 0, ?4)",
955 params![user_id, algorithm, auth_data, now],
956 )?;
957 Ok(conn.last_insert_rowid())
958}
959
960pub fn current_backup_version(conn: &Connection, user_id: i64) -> rusqlite::Result<Option<KeyBackupVersion>> {
963 conn.query_row(
964 &format!(
965 "SELECT {KEY_BACKUP_VERSION_SELECT_COLUMNS} FROM key_backup_versions
966 WHERE user_id = ?1 AND is_deleted = 0 ORDER BY version DESC LIMIT 1"
967 ),
968 params![user_id],
969 key_backup_version_from_row,
970 )
971 .optional()
972}
973
974pub fn get_backup_version(conn: &Connection, user_id: i64, version: i64) -> rusqlite::Result<Option<KeyBackupVersion>> {
979 conn.query_row(
980 &format!("SELECT {KEY_BACKUP_VERSION_SELECT_COLUMNS} FROM key_backup_versions WHERE user_id = ?1 AND version = ?2"),
981 params![user_id, version],
982 key_backup_version_from_row,
983 )
984 .optional()
985}
986
987pub fn update_backup_version_auth_data(conn: &Connection, user_id: i64, version: i64, auth_data: &str) -> rusqlite::Result<bool> {
990 let changed = conn.execute(
991 "UPDATE key_backup_versions SET auth_data = ?1, etag = etag + 1 WHERE user_id = ?2 AND version = ?3 AND is_deleted = 0",
992 params![auth_data, user_id, version],
993 )?;
994 Ok(changed > 0)
995}
996
997pub fn delete_backup_version(conn: &Connection, user_id: i64, version: i64) -> rusqlite::Result<bool> {
1001 let changed = conn.execute(
1002 "UPDATE key_backup_versions SET is_deleted = 1, etag = etag + 1 WHERE user_id = ?1 AND version = ?2 AND is_deleted = 0",
1003 params![user_id, version],
1004 )?;
1005 Ok(changed > 0)
1006}
1007
1008#[derive(Debug, Clone, PartialEq)]
1009pub struct KeyBackupSession {
1010 pub user_id: i64,
1011 pub version: i64,
1012 pub room_id: String,
1013 pub session_id: String,
1014 pub session_data: String,
1015 pub updated_at: String,
1016}
1017
1018const KEY_BACKUP_SESSION_SELECT_COLUMNS: &str = "user_id, version, room_id, session_id, session_data, updated_at";
1019
1020fn key_backup_session_from_row(row: &rusqlite::Row<'_>) -> rusqlite::Result<KeyBackupSession> {
1021 Ok(KeyBackupSession {
1022 user_id: row.get(0)?,
1023 version: row.get(1)?,
1024 room_id: row.get(2)?,
1025 session_id: row.get(3)?,
1026 session_data: row.get(4)?,
1027 updated_at: row.get(5)?,
1028 })
1029}
1030
1031pub fn put_backup_sessions(
1036 conn: &mut Connection,
1037 user_id: i64,
1038 version: i64,
1039 sessions: &[(String, String, String)],
1040 now: &str,
1041) -> Result<(), MatrixKeysStoreError> {
1042 let tx = conn.transaction()?;
1043 let current_version: Option<i64> = tx.query_row(
1044 "SELECT MAX(version) FROM key_backup_versions WHERE user_id = ?1 AND is_deleted = 0",
1045 params![user_id],
1046 |row| row.get(0),
1047 )?;
1048 if current_version != Some(version) {
1049 return Err(MatrixKeysStoreError::WrongBackupVersion);
1050 }
1051 for (room_id, session_id, session_data) in sessions {
1052 tx.execute(
1053 "INSERT INTO key_backup_sessions (user_id, version, room_id, session_id, session_data, updated_at)
1054 VALUES (?1, ?2, ?3, ?4, ?5, ?6)
1055 ON CONFLICT(user_id, version, room_id, session_id) DO UPDATE SET
1056 session_data = excluded.session_data, updated_at = excluded.updated_at",
1057 params![user_id, version, room_id, session_id, session_data, now],
1058 )?;
1059 }
1060 tx.execute(
1061 "UPDATE key_backup_versions SET etag = etag + 1 WHERE user_id = ?1 AND version = ?2",
1062 params![user_id, version],
1063 )?;
1064 tx.commit()?;
1065 Ok(())
1066}
1067
1068pub fn get_backup_sessions(
1073 conn: &Connection,
1074 user_id: i64,
1075 version: i64,
1076 room_id: Option<&str>,
1077 session_id: Option<&str>,
1078) -> rusqlite::Result<Vec<KeyBackupSession>> {
1079 match (room_id, session_id) {
1080 (Some(room_id), Some(session_id)) => {
1081 let mut stmt = conn.prepare(&format!(
1082 "SELECT {KEY_BACKUP_SESSION_SELECT_COLUMNS} FROM key_backup_sessions
1083 WHERE user_id = ?1 AND version = ?2 AND room_id = ?3 AND session_id = ?4"
1084 ))?;
1085 let rows = stmt.query_map(params![user_id, version, room_id, session_id], key_backup_session_from_row)?;
1086 rows.collect()
1087 }
1088 (Some(room_id), None) => {
1089 let mut stmt = conn.prepare(&format!(
1090 "SELECT {KEY_BACKUP_SESSION_SELECT_COLUMNS} FROM key_backup_sessions
1091 WHERE user_id = ?1 AND version = ?2 AND room_id = ?3"
1092 ))?;
1093 let rows = stmt.query_map(params![user_id, version, room_id], key_backup_session_from_row)?;
1094 rows.collect()
1095 }
1096 (None, _) => {
1097 let mut stmt = conn.prepare(&format!(
1098 "SELECT {KEY_BACKUP_SESSION_SELECT_COLUMNS} FROM key_backup_sessions WHERE user_id = ?1 AND version = ?2"
1099 ))?;
1100 let rows = stmt.query_map(params![user_id, version], key_backup_session_from_row)?;
1101 rows.collect()
1102 }
1103 }
1104}
1105
1106pub fn delete_backup_sessions(
1110 conn: &mut Connection,
1111 user_id: i64,
1112 version: i64,
1113 room_id: Option<&str>,
1114 session_id: Option<&str>,
1115) -> rusqlite::Result<usize> {
1116 let tx = conn.transaction()?;
1117 let deleted = match (room_id, session_id) {
1118 (Some(room_id), Some(session_id)) => tx.execute(
1119 "DELETE FROM key_backup_sessions WHERE user_id = ?1 AND version = ?2 AND room_id = ?3 AND session_id = ?4",
1120 params![user_id, version, room_id, session_id],
1121 )?,
1122 (Some(room_id), None) => tx.execute(
1123 "DELETE FROM key_backup_sessions WHERE user_id = ?1 AND version = ?2 AND room_id = ?3",
1124 params![user_id, version, room_id],
1125 )?,
1126 (None, _) => tx.execute(
1127 "DELETE FROM key_backup_sessions WHERE user_id = ?1 AND version = ?2",
1128 params![user_id, version],
1129 )?,
1130 };
1131 if deleted > 0 {
1132 tx.execute(
1133 "UPDATE key_backup_versions SET etag = etag + 1 WHERE user_id = ?1 AND version = ?2",
1134 params![user_id, version],
1135 )?;
1136 }
1137 tx.commit()?;
1138 Ok(deleted)
1139}
1140
1141pub fn backup_count_and_etag(conn: &Connection, user_id: i64, version: i64) -> rusqlite::Result<(i64, i64)> {
1144 let count: i64 = conn.query_row(
1145 "SELECT COUNT(*) FROM key_backup_sessions WHERE user_id = ?1 AND version = ?2",
1146 params![user_id, version],
1147 |row| row.get(0),
1148 )?;
1149 let etag: i64 = conn.query_row(
1150 "SELECT etag FROM key_backup_versions WHERE user_id = ?1 AND version = ?2",
1151 params![user_id, version],
1152 |row| row.get(0),
1153 )?;
1154 Ok((count, etag))
1155}
1156
1157#[cfg(test)]
1158mod tests {
1159 use super::*;
1160
1161 const T0: &str = "2026-09-24T00:00:00+00:00";
1162
1163 fn test_conn() -> Connection {
1164 let conn = Connection::open_in_memory().expect("in-memory sqlite");
1165 crate::store::create_matrix_schema(&conn).expect("matrix schema (stream_counter lives there)");
1166 create_matrix_keys_schema(&conn).expect("matrix keys schema");
1167 conn
1168 }
1169
1170 fn count_device_list_changes(conn: &Connection, user_id: i64) -> i64 {
1171 conn.query_row("SELECT COUNT(*) FROM device_list_changes WHERE user_id = ?1", params![user_id], |row| row.get(0))
1172 .expect("count device_list_changes")
1173 }
1174
1175 #[test]
1178 fn claim_one_time_key_deletes_it_so_a_second_claim_gets_a_different_key_or_none() {
1179 let mut conn = test_conn();
1180 add_one_time_keys(
1181 &mut conn,
1182 1,
1183 "DEV1",
1184 &[
1185 ("signed_curve25519:AAAAAQ".to_string(), "signed_curve25519".to_string(), r#"{"key":"k1"}"#.to_string()),
1186 ("signed_curve25519:AAAAAg".to_string(), "signed_curve25519".to_string(), r#"{"key":"k2"}"#.to_string()),
1187 ],
1188 )
1189 .expect("add otks");
1190
1191 let first = claim_one_time_key(&mut conn, 1, "DEV1", "signed_curve25519").expect("claim 1").expect("has a key");
1192 let second = claim_one_time_key(&mut conn, 1, "DEV1", "signed_curve25519").expect("claim 2").expect("has a different key");
1193 assert_ne!(first.0, second.0, "the two claims must return different key ids");
1194
1195 let third = claim_one_time_key(&mut conn, 1, "DEV1", "signed_curve25519").expect("claim 3");
1196 assert_eq!(third, None, "no one-time keys or fallback keys remain");
1197 }
1198
1199 #[test]
1202 fn claim_falls_back_to_a_fallback_key_without_deleting_it_and_marks_it_used() {
1203 let mut conn = test_conn();
1204 upsert_fallback_key(&conn, 1, "DEV1", "signed_curve25519", "signed_curve25519:FALLBACK", r#"{"key":"fb"}"#, T0).expect("upsert fallback");
1205
1206 let claimed = claim_one_time_key(&mut conn, 1, "DEV1", "signed_curve25519").expect("claim").expect("fallback returned");
1207 assert_eq!(claimed.0, "signed_curve25519:FALLBACK");
1208 assert_eq!(claimed.1, r#"{"key":"fb"}"#);
1209
1210 let claimed_again = claim_one_time_key(&mut conn, 1, "DEV1", "signed_curve25519").expect("claim again").expect("fallback still there");
1211 assert_eq!(claimed_again.0, "signed_curve25519:FALLBACK", "a fallback key is never deleted on claim");
1212
1213 let unused = unused_fallback_key_types(&conn, 1, "DEV1").expect("unused types");
1214 assert!(unused.is_empty(), "the fallback key must be marked used after its first claim");
1215 }
1216
1217 #[test]
1220 fn device_unused_fallback_key_types_excludes_a_used_one() {
1221 let mut conn = test_conn();
1222 upsert_fallback_key(&conn, 1, "DEV1", "signed_curve25519", "signed_curve25519:FB1", r#"{"key":"fb1"}"#, T0).expect("fallback 1");
1223 upsert_fallback_key(&conn, 1, "DEV1", "olm_curve25519", "olm_curve25519:FB2", r#"{"key":"fb2"}"#, T0).expect("fallback 2");
1224
1225 let before = unused_fallback_key_types(&conn, 1, "DEV1").expect("before claim");
1226 assert_eq!(before.len(), 2);
1227
1228 claim_one_time_key(&mut conn, 1, "DEV1", "signed_curve25519").expect("claim marks it used");
1229
1230 let after = unused_fallback_key_types(&conn, 1, "DEV1").expect("after claim");
1231 assert_eq!(after, vec!["olm_curve25519".to_string()]);
1232 }
1233
1234 #[test]
1238 fn delete_device_by_credential_deletes_the_device_and_its_key_rows_and_logs_a_device_list_change() {
1239 let mut conn = test_conn();
1240 let device_id = create_device(&conn, 1, CredentialKind::Bearer, "tok-hash-1", T0).expect("create device");
1241 upsert_device_keys(&mut conn, 1, &device_id, "[]", "{}", "{}", T0).expect("upload device keys");
1242 add_one_time_keys(
1243 &mut conn,
1244 1,
1245 &device_id,
1246 &[("signed_curve25519:AAAAAQ".to_string(), "signed_curve25519".to_string(), "{}".to_string())],
1247 )
1248 .expect("otk");
1249 upsert_fallback_key(&conn, 1, &device_id, "signed_curve25519", "signed_curve25519:FB", "{}", T0).expect("fallback");
1250 enqueue_to_device(&mut conn, 2, &[(1, device_id.clone(), "m.text".to_string(), "{}".to_string())]).expect("to-device");
1251
1252 let before = count_device_list_changes(&conn, 1);
1253 let deleted = delete_device_by_credential(&mut conn, CredentialKind::Bearer, "tok-hash-1", T0).expect("revoke");
1254 assert_eq!(deleted, Some((1, device_id.clone())));
1255
1256 assert_eq!(get_device(&conn, 1, &device_id).expect("get"), None);
1257 let keys_count: i64 = conn
1258 .query_row("SELECT COUNT(*) FROM device_keys WHERE user_id = 1 AND device_id = ?1", params![device_id], |row| row.get(0))
1259 .expect("keys count");
1260 assert_eq!(keys_count, 0);
1261 let otk_count: i64 = conn
1262 .query_row("SELECT COUNT(*) FROM one_time_keys WHERE user_id = 1 AND device_id = ?1", params![device_id], |row| row.get(0))
1263 .expect("otk count");
1264 assert_eq!(otk_count, 0);
1265 let fallback_count: i64 = conn
1266 .query_row("SELECT COUNT(*) FROM fallback_keys WHERE user_id = 1 AND device_id = ?1", params![device_id], |row| row.get(0))
1267 .expect("fallback count");
1268 assert_eq!(fallback_count, 0);
1269 let to_device_count: i64 = conn
1270 .query_row("SELECT COUNT(*) FROM to_device_messages WHERE recipient_device_id = ?1", params![device_id], |row| row.get(0))
1271 .expect("to-device count");
1272 assert_eq!(to_device_count, 0);
1273
1274 let after = count_device_list_changes(&conn, 1);
1275 assert_eq!(after, before + 1, "exactly one device_list_changes row must be appended");
1276
1277 assert_eq!(
1278 delete_device_by_credential(&mut conn, CredentialKind::Bearer, "tok-hash-1", T0).expect("second revoke is a no-op"),
1279 None
1280 );
1281 }
1282
1283 #[test]
1286 fn add_one_time_keys_is_idempotent_for_identical_json_and_refuses_changed_json() {
1287 let mut conn = test_conn();
1288 let key = ("signed_curve25519:AAAAAQ".to_string(), "signed_curve25519".to_string(), r#"{"key":"k1"}"#.to_string());
1289 add_one_time_keys(&mut conn, 1, "DEV1", &[key.clone()]).expect("first add");
1290 add_one_time_keys(&mut conn, 1, "DEV1", &[key.clone()]).expect("identical resubmission is a no-op");
1291
1292 let count: i64 = conn
1293 .query_row("SELECT COUNT(*) FROM one_time_keys WHERE user_id = 1 AND device_id = 'DEV1'", [], |row| row.get(0))
1294 .expect("count");
1295 assert_eq!(count, 1);
1296
1297 let changed = ("signed_curve25519:AAAAAQ".to_string(), "signed_curve25519".to_string(), r#"{"key":"k2-different"}"#.to_string());
1298 let err = add_one_time_keys(&mut conn, 1, "DEV1", &[changed]).unwrap_err();
1299 assert!(matches!(err, MatrixKeysStoreError::OneTimeKeyConflict(ref id) if id == "signed_curve25519:AAAAAQ"));
1300 }
1301
1302 #[test]
1305 fn to_device_delete_up_to_leaves_later_messages() {
1306 let mut conn = test_conn();
1307 let s1 = enqueue_to_device(&mut conn, 2, &[(1, "DEV1".to_string(), "m.a".to_string(), "{}".to_string())]).expect("send 1");
1308 let s2 = enqueue_to_device(&mut conn, 2, &[(1, "DEV1".to_string(), "m.b".to_string(), "{}".to_string())]).expect("send 2");
1309 assert!(s2 > s1);
1310
1311 let deleted = delete_to_device_up_to(&conn, 1, "DEV1", s1).expect("delete up to s1");
1312 assert_eq!(deleted, 1);
1313
1314 let remaining = to_device_for(&conn, 1, "DEV1", 0, 10).expect("remaining");
1315 assert_eq!(remaining.len(), 1);
1316 assert_eq!(remaining[0].stream_id, s2);
1317 }
1318
1319 #[test]
1322 fn device_keys_change_logs_a_device_list_change() {
1323 let mut conn = test_conn();
1324 let before = count_device_list_changes(&conn, 1);
1325 upsert_device_keys(&mut conn, 1, "DEV1", "[\"m.olm.v1\"]", "{}", "{}", T0).expect("upload keys");
1326 let after = count_device_list_changes(&conn, 1);
1327 assert_eq!(after, before + 1);
1328
1329 let rows = device_keys_for(&conn, &[1]).expect("query");
1330 assert_eq!(rows.len(), 1);
1331 assert_eq!(rows[0].device_id, "DEV1");
1332
1333 assert_eq!(device_keys_for(&conn, &[]).expect("empty input"), Vec::new());
1334 }
1335
1336 #[test]
1339 fn put_backup_sessions_refuses_a_stale_version() {
1340 let mut conn = test_conn();
1341 let v1 = create_backup_version(&conn, 1, "m.megolm_backup.v1", "{}", T0).expect("v1");
1342 let v2 = create_backup_version(&conn, 1, "m.megolm_backup.v1", "{}", T0).expect("v2");
1343 assert!(v2 > v1);
1344
1345 let err = put_backup_sessions(&mut conn, 1, v1, &[("!room:x".to_string(), "sess1".to_string(), "{}".to_string())], T0).unwrap_err();
1346 assert!(matches!(err, MatrixKeysStoreError::WrongBackupVersion));
1347
1348 put_backup_sessions(&mut conn, 1, v2, &[("!room:x".to_string(), "sess1".to_string(), "{}".to_string())], T0).expect("current version accepted");
1349 }
1350
1351 #[test]
1354 fn delete_device_cascades_keys_and_pending_to_device() {
1355 let mut conn = test_conn();
1356 let device_id = create_device(&conn, 1, CredentialKind::Web, "sess-1", T0).expect("create device");
1357 upsert_device_keys(&mut conn, 1, &device_id, "[]", "{}", "{}", T0).expect("device keys");
1358 add_one_time_keys(
1359 &mut conn,
1360 1,
1361 &device_id,
1362 &[("signed_curve25519:AAAAAQ".to_string(), "signed_curve25519".to_string(), "{}".to_string())],
1363 )
1364 .expect("otk");
1365 upsert_fallback_key(&conn, 1, &device_id, "signed_curve25519", "signed_curve25519:FB", "{}", T0).expect("fallback");
1366 enqueue_to_device(&mut conn, 9, &[(1, device_id.clone(), "m.text".to_string(), "{}".to_string())]).expect("to-device");
1367
1368 let before = count_device_list_changes(&conn, 1);
1369 let deleted = delete_device(&mut conn, 1, &device_id, T0).expect("delete");
1370 assert!(deleted);
1371
1372 assert_eq!(get_device(&conn, 1, &device_id).expect("get"), None);
1373 assert!(device_keys_for(&conn, &[1]).expect("keys").is_empty());
1374 assert_eq!(count_one_time_keys(&conn, 1, &device_id).expect("otk count").len(), 0);
1375 assert!(unused_fallback_key_types(&conn, 1, &device_id).expect("fallback").is_empty());
1376 assert!(to_device_for(&conn, 1, &device_id, 0, 10).expect("to-device").is_empty());
1377
1378 let after = count_device_list_changes(&conn, 1);
1379 assert_eq!(after, before + 1);
1380
1381 let deleted_again = delete_device(&mut conn, 1, &device_id, T0).expect("second delete is a no-op");
1382 assert!(!deleted_again);
1383 }
1384
1385 #[test]
1388 fn create_device_mints_a_fresh_device_id_and_touch_device_updates_last_seen() {
1389 let conn = test_conn();
1390 let d1 = create_device(&conn, 1, CredentialKind::Bearer, "tok-a", T0).expect("create 1");
1391 let d2 = create_device(&conn, 1, CredentialKind::Bearer, "tok-b", T0).expect("create 2");
1392 assert_ne!(d1, d2, "two different credentials must mint different device ids");
1393
1394 touch_device(&conn, 1, &d1, "2026-09-24T01:00:00+00:00").expect("touch");
1395 let row = get_device(&conn, 1, &d1).expect("get").expect("row exists");
1396 assert_eq!(row.last_seen_at, "2026-09-24T01:00:00+00:00");
1397 assert_eq!(row.created_at, T0, "created_at must not move on touch");
1398 }
1399
1400 #[test]
1401 fn device_for_credential_finds_the_row_created_by_create_device() {
1402 let conn = test_conn();
1403 let device_id = create_device(&conn, 5, CredentialKind::Web, "sess-xyz", T0).expect("create");
1404 let found = device_for_credential(&conn, CredentialKind::Web, "sess-xyz").expect("lookup").expect("row exists");
1405 assert_eq!(found.user_id, 5);
1406 assert_eq!(found.device_id, device_id);
1407
1408 assert_eq!(device_for_credential(&conn, CredentialKind::Bearer, "sess-xyz").expect("wrong kind"), None);
1409 }
1410
1411 #[test]
1412 fn set_device_display_name_updates_only_the_named_device() {
1413 let conn = test_conn();
1414 let d1 = create_device(&conn, 1, CredentialKind::Bearer, "a", T0).expect("d1");
1415 let d2 = create_device(&conn, 1, CredentialKind::Bearer, "b", T0).expect("d2");
1416
1417 assert!(set_device_display_name(&conn, 1, &d1, Some("My Phone")).expect("set"));
1418 assert!(!set_device_display_name(&conn, 1, "nonexistent", Some("x")).expect("missing device is a no-op returning false"));
1419
1420 let devices = list_devices(&conn, 1).expect("list");
1421 assert_eq!(devices.len(), 2);
1422 let named = devices.iter().find(|d| d.device_id == d1).expect("d1 present");
1423 assert_eq!(named.display_name.as_deref(), Some("My Phone"));
1424 let unnamed = devices.iter().find(|d| d.device_id == d2).expect("d2 present");
1425 assert_eq!(unnamed.display_name, None);
1426 }
1427
1428 #[test]
1429 fn devices_for_reaper_lists_every_device_with_its_credential() {
1430 let conn = test_conn();
1431 create_device(&conn, 1, CredentialKind::Bearer, "tok-1", T0).expect("d1");
1432 create_device(&conn, 2, CredentialKind::Web, "sess-2", T0).expect("d2");
1433
1434 let mut all = devices_for_reaper(&conn).expect("reaper list");
1435 all.sort_by_key(|(user_id, ..)| *user_id);
1436 assert_eq!(all.len(), 2);
1437 assert_eq!(all[0].0, 1);
1438 assert_eq!(all[0].2, CredentialKind::Bearer);
1439 assert_eq!(all[0].3, "tok-1");
1440 assert_eq!(all[1].0, 2);
1441 assert_eq!(all[1].2, CredentialKind::Web);
1442 }
1443
1444 #[test]
1445 fn count_one_time_keys_groups_by_algorithm() {
1446 let mut conn = test_conn();
1447 add_one_time_keys(
1448 &mut conn,
1449 1,
1450 "DEV1",
1451 &[
1452 ("signed_curve25519:A".to_string(), "signed_curve25519".to_string(), "{}".to_string()),
1453 ("signed_curve25519:B".to_string(), "signed_curve25519".to_string(), "{}".to_string()),
1454 ("other_algo:C".to_string(), "other_algo".to_string(), "{}".to_string()),
1455 ],
1456 )
1457 .expect("add");
1458
1459 let counts = count_one_time_keys(&conn, 1, "DEV1").expect("count");
1460 assert_eq!(counts.get("signed_curve25519"), Some(&2));
1461 assert_eq!(counts.get("other_algo"), Some(&1));
1462 }
1463
1464 #[test]
1465 fn cross_signing_upsert_round_trips_and_logs_a_device_list_change() {
1466 let mut conn = test_conn();
1467 let before = count_device_list_changes(&conn, 1);
1468 upsert_cross_signing_key(&mut conn, 1, CrossSigningUsage::Master, r#"{"keys":{}}"#, T0).expect("master");
1469 upsert_cross_signing_key(&mut conn, 1, CrossSigningUsage::SelfSigning, r#"{"keys":{}}"#, T0).expect("self signing");
1470
1471 let after = count_device_list_changes(&conn, 1);
1472 assert_eq!(after, before + 2);
1473
1474 let keys = cross_signing_keys_for(&conn, &[1]).expect("query");
1475 assert_eq!(keys.len(), 2);
1476 assert!(keys.iter().any(|k| k.usage == CrossSigningUsage::Master));
1477 assert!(keys.iter().any(|k| k.usage == CrossSigningUsage::SelfSigning));
1478
1479 assert_eq!(cross_signing_keys_for(&conn, &[]).expect("empty input"), Vec::new());
1480 }
1481
1482 #[test]
1483 fn add_signatures_and_signatures_for_round_trip() {
1484 let mut conn = test_conn();
1485 add_signatures(&mut conn, &[(1, 2, "DEVICEX".to_string(), r#"{"sig":"abc"}"#.to_string(), T0.to_string())]).expect("add");
1486
1487 let sigs = signatures_for(&conn, 2, "DEVICEX").expect("query");
1488 assert_eq!(sigs.len(), 1);
1489 assert_eq!(sigs[0].signer_user_id, 1);
1490 assert_eq!(sigs[0].signature_json, r#"{"sig":"abc"}"#);
1491
1492 assert!(signatures_for(&conn, 2, "OTHER").expect("no match").is_empty());
1493 }
1494
1495 #[test]
1496 fn backup_version_lifecycle_create_get_update_delete() {
1497 let conn = test_conn();
1498 let version = create_backup_version(&conn, 1, "m.megolm_backup.v1", r#"{"a":1}"#, T0).expect("create");
1499
1500 assert_eq!(current_backup_version(&conn, 1).expect("current").expect("row").version, version);
1501
1502 assert!(update_backup_version_auth_data(&conn, 1, version, r#"{"a":2}"#).expect("update"));
1503 let updated = get_backup_version(&conn, 1, version).expect("get").expect("row");
1504 assert_eq!(updated.auth_data, r#"{"a":2}"#);
1505 assert_eq!(updated.etag, 1);
1506
1507 assert!(delete_backup_version(&conn, 1, version).expect("delete"));
1508 assert_eq!(current_backup_version(&conn, 1).expect("current after delete"), None);
1509 assert!(!delete_backup_version(&conn, 1, version).expect("second delete is a no-op"));
1510 }
1511
1512 #[test]
1513 fn backup_sessions_put_get_delete_and_etag_bumps() {
1514 let mut conn = test_conn();
1515 let version = create_backup_version(&conn, 1, "m.megolm_backup.v1", "{}", T0).expect("create");
1516
1517 put_backup_sessions(
1518 &mut conn,
1519 1,
1520 version,
1521 &[
1522 ("!room1:x".to_string(), "sessA".to_string(), r#"{"d":1}"#.to_string()),
1523 ("!room1:x".to_string(), "sessB".to_string(), r#"{"d":2}"#.to_string()),
1524 ("!room2:x".to_string(), "sessC".to_string(), r#"{"d":3}"#.to_string()),
1525 ],
1526 T0,
1527 )
1528 .expect("put");
1529
1530 let (count, etag_after_put) = backup_count_and_etag(&conn, 1, version).expect("count+etag");
1531 assert_eq!(count, 3);
1532 assert_eq!(etag_after_put, 1);
1533
1534 let room1_sessions = get_backup_sessions(&conn, 1, version, Some("!room1:x"), None).expect("room1");
1535 assert_eq!(room1_sessions.len(), 2);
1536
1537 let one = get_backup_sessions(&conn, 1, version, Some("!room1:x"), Some("sessA")).expect("one");
1538 assert_eq!(one.len(), 1);
1539 assert_eq!(one[0].session_data, r#"{"d":1}"#);
1540
1541 let deleted = delete_backup_sessions(&mut conn, 1, version, Some("!room1:x"), None).expect("delete room1");
1542 assert_eq!(deleted, 2);
1543 let (count_after, etag_after_delete) = backup_count_and_etag(&conn, 1, version).expect("count+etag after delete");
1544 assert_eq!(count_after, 1);
1545 assert_eq!(etag_after_delete, 2);
1546 }
1547
1548 #[test]
1549 fn device_list_changes_between_is_distinct_and_bounded() {
1550 let mut conn = test_conn();
1551 log_device_list_change(&mut conn, 1, T0).expect("change 1");
1552 let boundary = crate::store::max_stream_id(&conn).expect("boundary");
1553 log_device_list_change(&mut conn, 1, T0).expect("change 2 for the same user");
1554 log_device_list_change(&mut conn, 2, T0).expect("change for a different user");
1555
1556 let mut changed = device_list_changes_between(&conn, boundary, crate::store::max_stream_id(&conn).expect("max")).expect("query");
1557 changed.sort();
1558 assert_eq!(changed, vec![1, 2]);
1559 }
1560
1561 #[test]
1562 fn cross_signing_key_for_finds_exactly_the_named_usage() {
1563 let mut conn = test_conn();
1564 upsert_cross_signing_key(&mut conn, 1, CrossSigningUsage::Master, r#"{"usage":["master"]}"#, T0).expect("master");
1565
1566 let master = cross_signing_key_for(&conn, 1, CrossSigningUsage::Master).expect("query").expect("row exists");
1567 assert_eq!(master.usage, CrossSigningUsage::Master);
1568 assert_eq!(cross_signing_key_for(&conn, 1, CrossSigningUsage::SelfSigning).expect("query"), None);
1569 assert_eq!(cross_signing_key_for(&conn, 2, CrossSigningUsage::Master).expect("different user"), None);
1570 }
1571
1572 #[test]
1573 fn enqueue_to_device_deduped_is_idempotent_per_txn() {
1574 let mut conn = test_conn();
1575 let messages = [(1_i64, "DEV1".to_string(), "m.room_key".to_string(), r#"{"k":1}"#.to_string())];
1576
1577 let first = enqueue_to_device_deduped(&mut conn, 9, "SENDER_DEV", "txn-1", &messages, T0).expect("first send");
1578 assert_eq!(first, ToDeviceDedupOutcome::New);
1579 assert_eq!(to_device_for(&conn, 1, "DEV1", 0, 10).expect("after first").len(), 1);
1580
1581 let second = enqueue_to_device_deduped(&mut conn, 9, "SENDER_DEV", "txn-1", &messages, T0).expect("repeat send");
1582 assert_eq!(second, ToDeviceDedupOutcome::AlreadySent);
1583 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");
1584 }
1585
1586 #[test]
1587 fn enqueue_to_device_fans_out_and_is_scoped_to_the_recipient_device() {
1588 let mut conn = test_conn();
1589 let last = enqueue_to_device(
1590 &mut conn,
1591 9,
1592 &[
1593 (1, "DEV1".to_string(), "m.room_key".to_string(), r#"{"k":1}"#.to_string()),
1594 (1, "DEV2".to_string(), "m.room_key".to_string(), r#"{"k":1}"#.to_string()),
1595 ],
1596 )
1597 .expect("enqueue");
1598 assert_eq!(crate::store::max_stream_id(&conn).expect("max"), last);
1599
1600 let for_dev1 = to_device_for(&conn, 1, "DEV1", 0, 10).expect("dev1");
1601 assert_eq!(for_dev1.len(), 1);
1602 let for_dev2 = to_device_for(&conn, 1, "DEV2", 0, 10).expect("dev2");
1603 assert_eq!(for_dev2.len(), 1);
1604 assert_ne!(for_dev1[0].stream_id, for_dev2[0].stream_id);
1605
1606 let empty_batch = enqueue_to_device(&mut conn, 9, &[]).expect("empty batch is a no-op");
1607 assert_eq!(empty_batch, last, "an empty batch must not mint a fresh stream id");
1608 }
1609}