Skip to main content

tuwunel_service/users/
keys.rs

1use std::{collections::BTreeMap, mem, ops::Deref, sync::Arc};
2
3use futures::{
4	Stream, StreamExt, TryFutureExt,
5	future::{join, try_join3},
6	pin_mut,
7};
8use ruma::{
9	AnyKeyName, DeviceId, KeyId, OneTimeKeyAlgorithm, OneTimeKeyId, OneTimeKeyName,
10	OwnedDeviceId, OwnedKeyId, OwnedOneTimeKeyId, OwnedRoomId, OwnedServerName, RoomId,
11	SigningKeyId, UInt, UserId,
12	encryption::{CrossSigningKey, DeviceKeys, OneTimeKey},
13	serde::{Base64, Raw, base64::Standard},
14	signatures::{
15		VerificationError, to_canonical_json_string_for_signing, verify_canonical_json_bytes,
16	},
17};
18use serde::{Deserialize, Serialize};
19use tuwunel_core::{
20	Err, Error, Result,
21	debug::INFO_SPAN_LEVEL,
22	debug_error, debug_warn, err, implement,
23	smallvec::SmallVec,
24	utils::{
25		BoolExt, IterStream, ReadyExt,
26		result::LogErr,
27		stream::{BroadbandExt, TryIgnore},
28		to_canonical_object,
29	},
30};
31use tuwunel_database::{
32	Deserialized, Ignore, Interfix, Json, KeyBuf, Map, Txn, deserialize_from_slice, serialize_key,
33};
34
35type Servers = SmallVec<[OwnedServerName; 1]>;
36type Signatures = SmallVec<[(String, String); 1]>;
37
38/// Mutation announced to clients and federation peers.
39///
40/// Device identifiers are borrowed at mutation sites and owned in stored records.
41#[derive(Clone, Copy, Debug, Eq, PartialEq)]
42pub enum DeviceListChange<D> {
43	/// A device was created or its keys or metadata changed.
44	Device(D),
45
46	/// A device was removed.
47	Deleted(D),
48
49	/// Cross-signing keys or their signatures changed.
50	CrossSigning,
51
52	/// Peers must fetch the user's device-list snapshot.
53	Resync,
54}
55
56/// Durable detail for one local user's device-list notification.
57///
58/// Cross-signing changes retain the current stream position without advancing it.
59#[derive(Debug, Eq, PartialEq)]
60pub struct DeviceListRecord {
61	/// Mutation associated with the notification count.
62	pub change: DeviceListChange<OwnedDeviceId>,
63
64	/// Per-user stream position after the mutation.
65	pub stream_id: u64,
66}
67
68enum DeviceListChangeKind {
69	Resync = 0,
70	Device = 1,
71	Deleted = 2,
72	CrossSigning = 3,
73}
74
75#[derive(Clone, Copy, Debug, Eq, PartialEq)]
76enum KeyRole {
77	Device,
78	CrossSigningRoot,
79	SelfSigning,
80	UserSigning,
81}
82
83#[derive(Clone, Copy, Debug, Eq, PartialEq)]
84enum SignatureWrite {
85	Merge,
86	ReplaceSender,
87}
88
89#[derive(Clone, Copy, Debug, Eq, PartialEq)]
90enum SignatureAction {
91	Ignore,
92	Reject,
93	Write(KeyRole, SignatureWrite),
94}
95
96/// MSC2732: row stored under `(user, device, algorithm)` in
97/// `userdeviceidalgorithm_fallback`. Fallback keys are not deleted on
98/// claim; the row is rewritten with `used = true`.
99#[derive(Debug, Deserialize, Serialize)]
100struct FallbackEntry {
101	key_id: OwnedOneTimeKeyId,
102	key: Raw<OneTimeKey>,
103	used: bool,
104}
105
106/// Row-key shape of `onetimekeyid4225_otk`: per-device pool keyed by
107/// upload-order count for MSC4225 ordering.
108type OtkRowKey<'a> = (&'a UserId, &'a DeviceId, u64, &'a OneTimeKeyId);
109
110#[implement(super::Service)]
111pub async fn add_one_time_keys(
112	&self,
113	user_id: &UserId,
114	device_id: &DeviceId,
115	keys: &BTreeMap<OwnedOneTimeKeyId, Raw<OneTimeKey>>,
116	limit: usize,
117) -> Result {
118	let mut txn = self.services.db.txn();
119	// Hold the oldest permit so the retirement frontier cannot pass this batch
120	// before commit.
121	let mut oldest_count = None;
122	let mut last_count = None;
123
124	for (id, key) in keys.iter().take(limit) {
125		let Ok(Some(count)) = self
126			.add_one_time_key(user_id, device_id, id, key, &mut txn)
127			.await
128		else {
129			continue;
130		};
131
132		last_count = Some(*count);
133		oldest_count = oldest_count.or(Some(count));
134	}
135
136	if let Some(count) = last_count {
137		txn.raw_put(&self.db.userid_lastonetimekeyupdate, user_id, count);
138	}
139
140	txn.execute();
141	drop(oldest_count);
142
143	Ok(())
144}
145
146#[implement(super::Service)]
147pub async fn add_one_time_key(
148	&self,
149	user_id: &UserId,
150	device_id: &DeviceId,
151	one_time_key_key: &KeyId<OneTimeKeyAlgorithm, OneTimeKeyName>,
152	one_time_key_value: &Raw<OneTimeKey>,
153	txn: &mut Txn,
154) -> Result<Option<impl Deref<Target = u64> + Send + use<>>> {
155	let Some(otk) = self.db.onetimekeyid4225_otk.as_ref() else {
156		return Err!(Database("one-time-key column unavailable"));
157	};
158
159	if !self.device_exists(user_id, device_id).await {
160		return Err!(Database(error!(
161			?user_id,
162			?device_id,
163			"User does not exist or device has no metadata."
164		)));
165	}
166
167	if let Err(e) = one_time_key_value
168		.deserialize()
169		.map_err(Into::into)
170	{
171		debug_error!(
172			?one_time_key_key,
173			?one_time_key_value,
174			"Invalid one time key JSON submitted by client, skipping: {e}"
175		);
176
177		return Err(e);
178	}
179
180	// Racy dedup: two concurrent uploads of the same id can both pass this
181	// check and produce duplicate rows that persist until aged out by prune.
182	let prefix = (user_id, device_id, Interfix);
183	let already_present = otk
184		.keys_prefix(&prefix)
185		.ignore_err()
186		.ready_any(|(.., id): OtkRowKey<'_>| id == one_time_key_key)
187		.await;
188
189	if already_present {
190		return Ok(None);
191	}
192
193	let count = self.services.globals.next_count();
194
195	// MSC4225: RocksDB iterates the (user, device) prefix in count_be ascending
196	// order, so /keys/claim issues one-time keys in the order they were uploaded.
197	txn.put(
198		otk,
199		(user_id, device_id, *count, one_time_key_key.as_str()),
200		Json(one_time_key_value),
201	);
202
203	Ok(Some(count))
204}
205
206#[implement(super::Service)]
207pub async fn add_fallback_keys<'a, Keys>(
208	&self,
209	user_id: &UserId,
210	device_id: &DeviceId,
211	keys: Keys,
212) -> Result
213where
214	Keys: Iterator<Item = (&'a OneTimeKeyId, &'a Raw<OneTimeKey>)> + Send + 'a,
215{
216	let mut txn = self.services.db.txn();
217	// Hold the oldest permit so the retirement frontier cannot pass this batch
218	// before commit.
219	let mut oldest_count = None;
220	let mut last_count = None;
221
222	for (id, key) in keys {
223		let Ok(count) = self
224			.add_fallback_key(user_id, device_id, id, key, &mut txn)
225			.await
226		else {
227			continue;
228		};
229
230		last_count = Some(*count);
231		oldest_count = oldest_count.or(Some(count));
232	}
233
234	if let Some(count) = last_count {
235		txn.raw_put(&self.db.userid_lastonetimekeyupdate, user_id, count);
236	}
237
238	txn.execute();
239	drop(oldest_count);
240
241	Ok(())
242}
243
244#[implement(super::Service)]
245pub async fn add_fallback_key(
246	&self,
247	user_id: &UserId,
248	device_id: &DeviceId,
249	one_time_key_key: &KeyId<OneTimeKeyAlgorithm, OneTimeKeyName>,
250	one_time_key_value: &Raw<OneTimeKey>,
251	txn: &mut Txn,
252) -> Result<impl Deref<Target = u64> + Send + use<>> {
253	if !self.device_exists(user_id, device_id).await {
254		return Err!(Database(error!(
255			?user_id,
256			?device_id,
257			"User does not exist or device has no metadata."
258		)));
259	}
260
261	if let Err(e) = one_time_key_value
262		.deserialize()
263		.map_err(Into::into)
264	{
265		debug_error!(
266			?one_time_key_key,
267			?one_time_key_value,
268			"Invalid fallback key JSON submitted by client, skipping: {e}"
269		);
270
271		return Err(e);
272	}
273
274	let entry = FallbackEntry {
275		key_id: one_time_key_key.to_owned(),
276		key: one_time_key_value.clone(),
277		used: false,
278	};
279
280	let key = (user_id, device_id, one_time_key_key.algorithm());
281	let count = self.services.globals.next_count();
282
283	txn.put(&self.db.userdeviceidalgorithm_fallback, key, Json(&entry));
284
285	Ok(count)
286}
287
288#[implement(super::Service)]
289pub async fn take_fallback_key(
290	&self,
291	user_id: &UserId,
292	device_id: &DeviceId,
293	algorithm: &OneTimeKeyAlgorithm,
294) -> Result<(OwnedKeyId<OneTimeKeyAlgorithm, OneTimeKeyName>, Raw<OneTimeKey>)> {
295	let key = (user_id, device_id, algorithm);
296	let entry: FallbackEntry = self
297		.db
298		.userdeviceidalgorithm_fallback
299		.qry(&key)
300		.await
301		.deserialized::<Json<_>>()
302		.map(|Json(entry)| entry)
303		.map_err(|_| err!(Request(NotFound("No fallback key found"))))?;
304
305	let updated = FallbackEntry { used: true, ..entry };
306	self.db
307		.userdeviceidalgorithm_fallback
308		.put(key, Json(&updated));
309
310	Ok((updated.key_id, updated.key))
311}
312
313#[implement(super::Service)]
314pub fn unused_fallback_key_algorithms<'a>(
315	&'a self,
316	user_id: &'a UserId,
317	device_id: &'a DeviceId,
318) -> impl Stream<Item = OneTimeKeyAlgorithm> + Send + 'a {
319	type KeyVal = ((Ignore, Ignore, OneTimeKeyAlgorithm), Json<FallbackEntry>);
320
321	let prefix = (user_id, device_id);
322	self.db
323		.userdeviceidalgorithm_fallback
324		.stream_prefix(&prefix)
325		.ignore_err()
326		.ready_filter_map(|((_, _, algorithm), Json(entry)): KeyVal| {
327			entry.used.is_false().then_some(algorithm)
328		})
329}
330
331#[implement(super::Service)]
332pub async fn last_one_time_keys_update(&self, user_id: &UserId) -> u64 {
333	self.db
334		.userid_lastonetimekeyupdate
335		.get(user_id)
336		.await
337		.deserialized()
338		.unwrap_or(0)
339}
340
341#[implement(super::Service)]
342pub async fn take_one_time_key(
343	&self,
344	user_id: &UserId,
345	device_id: &DeviceId,
346	key_algorithm: &OneTimeKeyAlgorithm,
347) -> Result<(OwnedKeyId<OneTimeKeyAlgorithm, OneTimeKeyName>, Raw<OneTimeKey>)> {
348	let Some(otk) = self.db.onetimekeyid4225_otk.as_ref() else {
349		return Err!(Request(NotFound("No one-time-key found")));
350	};
351
352	let update_count = self.services.globals.next_count();
353	self.db
354		.userid_lastonetimekeyupdate
355		.insert(user_id, update_count.to_be_bytes());
356
357	let prefix = (user_id, device_id, Interfix);
358	let one_time_keys = otk
359		.stream_prefix(&prefix)
360		.ignore_err()
361		.ready_filter(|(row, _): &(OtkRowKey<'_>, &[u8])| row.3.algorithm() == *key_algorithm);
362
363	pin_mut!(one_time_keys);
364	let ((user_id, device_id, count, id), val) = one_time_keys
365		.next()
366		.await
367		.ok_or_else(|| err!(Request(NotFound("No one-time-key found"))))?;
368
369	otk.del((user_id, device_id, count, id));
370
371	Ok((id.into(), serde_json::from_slice(val)?))
372}
373
374#[implement(super::Service)]
375pub async fn count_one_time_keys(
376	&self,
377	user_id: &UserId,
378	device_id: &DeviceId,
379) -> BTreeMap<OneTimeKeyAlgorithm, UInt> {
380	let Some(otk) = self.db.onetimekeyid4225_otk.as_ref() else {
381		// Without the MSC4225 column this node cannot observe the authoritative
382		// pool, so preserve "unknown" instead of falsely reporting zero keys.
383		return BTreeMap::new();
384	};
385
386	let prefix = (user_id, device_id, Interfix);
387	let algorithm_counts: BTreeMap<OneTimeKeyAlgorithm, UInt> = otk
388		.keys_prefix(&prefix)
389		.ignore_err()
390		.ready_fold(BTreeMap::new(), |mut acc, (.., id): OtkRowKey<'_>| {
391			let count: &mut UInt = acc.entry(id.algorithm()).or_default();
392			*count = count.saturating_add(1_u32.into());
393			acc
394		})
395		.await;
396
397	let total = algorithm_counts
398		.values()
399		.copied()
400		.map(TryInto::try_into)
401		.filter_map(Result::ok)
402		.fold(0_usize, usize::saturating_add);
403
404	let limit = self.services.config.one_time_key_limit;
405	if let Some(excess) = total.checked_sub(limit).filter(|&n| n > 0) {
406		self.prune_one_time_keys(user_id, device_id, excess)
407			.await;
408	}
409
410	complete_one_time_key_counts(algorithm_counts)
411}
412
413/// Keep zero-count algorithms visible to clients after an OTK pool is drained.
414///
415/// An empty map is omitted from `/sync` by ruma. Some clients interpret an
416/// omitted count as "unknown" and therefore do not replenish a
417/// previously-uploaded Olm account. Only `signed_curve25519` is seeded,
418/// matching Synapse; clients do not maintain unsigned curve25519 keys.
419fn complete_one_time_key_counts(
420	mut counts: BTreeMap<OneTimeKeyAlgorithm, UInt>,
421) -> BTreeMap<OneTimeKeyAlgorithm, UInt> {
422	counts
423		.entry(OneTimeKeyAlgorithm::SignedCurve25519)
424		.or_default();
425	counts
426}
427
428/// MSC4225: drop the `excess` oldest rows for this `(user, device)`. Forward
429/// iteration over the prefix runs in count_be ascending order, so
430/// `take(excess)` yields the earliest-uploaded rows.
431#[implement(super::Service)]
432pub async fn prune_one_time_keys(&self, user_id: &UserId, device_id: &DeviceId, excess: usize) {
433	let Some(otk) = self.db.onetimekeyid4225_otk.as_ref() else {
434		return;
435	};
436
437	let prefix = (user_id, device_id, Interfix);
438	otk.keys_prefix(&prefix)
439		.ignore_err()
440		.take(excess)
441		.ready_for_each(|row: OtkRowKey<'_>| {
442			otk.del(row);
443		})
444		.await;
445}
446
447/// Drops every one-time key still outstanding for this `(user, device)`.
448///
449/// Device removal has to reach the pool as well as the MSC2732 fallback key,
450/// or a claim against the deleted device still succeeds from the rows left
451/// behind.
452#[implement(super::Service)]
453pub async fn remove_one_time_keys(&self, user_id: &UserId, device_id: &DeviceId) {
454	let Some(otk) = self.db.onetimekeyid4225_otk.as_ref() else {
455		return;
456	};
457
458	otk.del_prefix(&(user_id, device_id, Interfix))
459		.await;
460}
461
462#[implement(super::Service)]
463pub async fn add_device_keys(
464	&self,
465	user_id: &UserId,
466	device_id: &DeviceId,
467	device_keys: &Raw<DeviceKeys>,
468) {
469	let key = (user_id, device_id);
470
471	self.db.keyid_key.put(key, Json(device_keys));
472	self.mark_device_key_update(user_id, DeviceListChange::Device(device_id))
473		.await;
474}
475
476#[implement(super::Service)]
477pub async fn add_cross_signing_keys(
478	&self,
479	user_id: &UserId,
480	master_key: &Option<Raw<CrossSigningKey>>,
481	self_signing_key: &Option<Raw<CrossSigningKey>>,
482	user_signing_key: &Option<Raw<CrossSigningKey>>,
483	notify: bool,
484) -> Result {
485	// TODO: Check signatures
486	{
487		let master_key_key = master_key
488			.as_ref()
489			.map(|master_key| parse_master_key(user_id, master_key).map(|(key, _)| key))
490			.transpose()?;
491
492		let self_signing_key_key = self_signing_key
493			.as_ref()
494			.map(|self_signing_key| parse_self_signing_key(user_id, self_signing_key))
495			.transpose()?;
496
497		let user_signing_key_id = user_signing_key
498			.as_ref()
499			.map(parse_user_signing_key)
500			.transpose()?;
501
502		let mut txn = self.services.db.txn();
503
504		if let Some((master_key, master_key_key)) =
505			master_key.as_ref().zip(master_key_key.as_ref())
506		{
507			txn.insert_raw(
508				&self.db.keyid_key,
509				master_key_key,
510				master_key.json().get().as_bytes(),
511			);
512			txn.insert_raw(&self.db.userid_masterkeyid, user_id.as_bytes(), master_key_key);
513		}
514
515		if let Some((self_signing_key, self_signing_key_key)) = self_signing_key
516			.as_ref()
517			.zip(self_signing_key_key.as_ref())
518		{
519			txn.insert_raw(
520				&self.db.keyid_key,
521				self_signing_key_key,
522				self_signing_key.json().get(),
523			);
524			txn.insert_raw(
525				&self.db.userid_selfsigningkeyid,
526				user_id.as_bytes(),
527				self_signing_key_key,
528			);
529		}
530
531		if let Some((user_signing_key, user_signing_key_id)) = user_signing_key
532			.as_ref()
533			.zip(user_signing_key_id.as_ref())
534		{
535			let user_signing_key_key = (user_id, user_signing_key_id);
536
537			txn.put_raw(
538				&self.db.keyid_key,
539				user_signing_key_key,
540				user_signing_key.json().get().as_bytes(),
541			);
542
543			txn.raw_put(&self.db.userid_usersigningkeyid, user_id, user_signing_key_key);
544		}
545
546		txn.execute();
547	};
548
549	if notify {
550		self.mark_device_key_update(user_id, DeviceListChange::CrossSigning)
551			.await;
552	}
553
554	Ok(())
555}
556
557fn parse_self_signing_key(
558	user_id: &UserId,
559	self_signing_key: &Raw<CrossSigningKey>,
560) -> Result<KeyBuf> {
561	let mut self_signing_key_ids = self_signing_key
562		.deserialize()
563		.map_err(|e| err!(Request(InvalidParam("Invalid self signing key: {e:?}"))))?
564		.keys
565		.into_values();
566
567	let self_signing_key_id = self_signing_key_ids
568		.next()
569		.ok_or_else(|| err!(Request(InvalidParam("Self signing key contained no key."))))?;
570
571	if self_signing_key_ids.next().is_some() {
572		return Err!(Request(InvalidParam("Self signing key contained more than one key.")));
573	}
574
575	serialize_key((user_id, self_signing_key_id))
576}
577
578#[implement(super::Service)]
579pub async fn sign_key(
580	&self,
581	target_id: &UserId,
582	key_id: &str,
583	signatures: Signatures,
584	sender_id: &UserId,
585) -> Result {
586	let key = (target_id, key_id);
587
588	let mut target_key: serde_json::Value = self
589		.db
590		.keyid_key
591		.qry(&key)
592		.await
593		.map_err(|error| match error {
594			| error if error.is_not_found() =>
595				err!(Request(NotFound("Tried to sign nonexistent key"))),
596			| error => error,
597		})?
598		.deserialized()
599		.map_err(|e| err!(Database(debug_warn!("key in keyid_key is invalid: {e:?}"))))?;
600
601	let target_role = self
602		.uploaded_key_role(target_id, key_id)
603		.await?
604		.ok_or_else(|| err!(Request(NotFound("Unknown device"))))?;
605
606	if !key_matches_role(&target_key, target_id, key_id, target_role) {
607		return Err!(Request(NotFound("Unknown device")));
608	}
609
610	let same_user = sender_id == target_id;
611	let mut canonical = None;
612	let mut accepted = false;
613	let mut changed = false;
614
615	for (signature_id, signature) in signatures {
616		let signature_key_id = <&SigningKeyId<AnyKeyName>>::try_from(signature_id.as_str())
617			.map_err(|source| VerificationError::ParseIdentifier {
618				identifier_type: "signing key ID",
619				source,
620			})?;
621
622		let signer_role = self
623			.uploaded_key_role(sender_id, signature_key_id.key_name().as_str())
624			.await?;
625
626		let (signer_role, write) = match signature_action(same_user, target_role, signer_role) {
627			| SignatureAction::Ignore => continue,
628			| SignatureAction::Reject => return Err!(Request(NotFound("Unknown device"))),
629			| SignatureAction::Write(signer_role, write) => (signer_role, write),
630		};
631
632		if canonical.is_none() {
633			canonical = Some(canonical_key(&target_key)?);
634		}
635
636		let canonical = canonical
637			.as_deref()
638			.ok_or_else(|| err!(Database("canonical key was not initialized")))?;
639
640		let signature = self
641			.verify_key_signature(
642				sender_id,
643				signature_key_id,
644				signer_role,
645				&signature,
646				canonical.as_bytes(),
647			)
648			.map_ok(|signature| signature.encode())
649			.await?;
650
651		changed |= match write {
652			| SignatureWrite::Merge =>
653				insert_signatures(&mut target_key, sender_id, [(signature_id, signature)])?,
654			| SignatureWrite::ReplaceSender =>
655				replace_signatures(&mut target_key, sender_id, (signature_id, signature))?,
656		};
657
658		accepted = true;
659	}
660
661	if !accepted {
662		return Err!(Request(NotFound("Unknown device")));
663	}
664
665	if !changed {
666		return Ok(());
667	}
668
669	let key = (target_id, key_id);
670	self.db.keyid_key.put(key, Json(target_key));
671
672	if same_user {
673		let change = match target_role {
674			| KeyRole::Device => DeviceListChange::Device(key_id.into()),
675			| _ => DeviceListChange::CrossSigning,
676		};
677
678		self.mark_device_key_update(target_id, change)
679			.await;
680	} else {
681		let count = self.services.globals.next_count();
682
683		self.db
684			.keychangeid_userid
685			.put_raw((sender_id, *count), target_id);
686	}
687
688	Ok(())
689}
690
691#[implement(super::Service)]
692async fn uploaded_key_role(&self, user_id: &UserId, key_id: &str) -> Result<Option<KeyRole>> {
693	let row_key = serialize_key((user_id, key_id))?;
694	let device_id: &DeviceId = key_id.into();
695
696	let (device, pointers) = join(
697		self.device_exists(user_id, device_id),
698		self.cross_signing_pointers(user_id, row_key.as_slice()),
699	)
700	.await;
701
702	let [root, self_signing, user_signing] = pointers?;
703
704	Ok(key_role([device, root, self_signing, user_signing]))
705}
706
707#[implement(super::Service)]
708pub(super) async fn is_cross_signing_key_id(
709	&self,
710	user_id: &UserId,
711	key_id: &str,
712) -> Result<bool> {
713	// Cross-signing rows share the device key row space under the bare public key.
714	let row_key = serialize_key((user_id, key_id))?;
715
716	self.cross_signing_pointers(user_id, row_key.as_slice())
717		.map_ok(|pointers| pointers.contains(&true))
718		.await
719}
720
721#[implement(super::Service)]
722async fn cross_signing_pointers(&self, user_id: &UserId, row_key: &[u8]) -> Result<[bool; 3]> {
723	try_join3(
724		pointer_matches(&self.db.userid_masterkeyid, user_id, row_key),
725		pointer_matches(&self.db.userid_selfsigningkeyid, user_id, row_key),
726		pointer_matches(&self.db.userid_usersigningkeyid, user_id, row_key),
727	)
728	.map_ok(|(root, self_signing, user_signing)| [root, self_signing, user_signing])
729	.await
730}
731
732#[tracing::instrument(
733	level = "trace",
734	skip_all,
735	fields(
736		user = %user_id,
737	)
738)]
739async fn pointer_matches(map: &Arc<Map>, user_id: &UserId, row_key: &[u8]) -> Result<bool> {
740	match map.get(user_id).await {
741		| Ok(pointer) => Ok(&*pointer == row_key),
742		| Err(error) if error.is_not_found() => Ok(false),
743		| Err(error) => Err(error),
744	}
745}
746
747fn key_role([device, root, self_signing, user_signing]: [bool; 4]) -> Option<KeyRole> {
748	match (device, root, self_signing, user_signing) {
749		| (true, false, false, false) => Some(KeyRole::Device),
750		| (false, true, false, false) => Some(KeyRole::CrossSigningRoot),
751		| (false, false, true, false) => Some(KeyRole::SelfSigning),
752		| (false, false, false, true) => Some(KeyRole::UserSigning),
753		| _ => None,
754	}
755}
756
757fn key_matches_role(
758	key: &serde_json::Value,
759	user_id: &UserId,
760	key_id: &str,
761	role: KeyRole,
762) -> bool {
763	if key
764		.get("user_id")
765		.and_then(serde_json::Value::as_str)
766		!= Some(user_id.as_str())
767	{
768		return false;
769	}
770
771	match role {
772		| KeyRole::Device =>
773			key.get("device_id")
774				.and_then(serde_json::Value::as_str)
775				== Some(key_id),
776		| KeyRole::CrossSigningRoot | KeyRole::SelfSigning | KeyRole::UserSigning =>
777			key.get("device_id").is_none()
778				&& key
779					.get("keys")
780					.and_then(serde_json::Value::as_object)
781					.is_some_and(|keys| {
782						keys.len() == 1
783							&& keys
784								.values()
785								.any(|value| value.as_str() == Some(key_id))
786					}),
787	}
788}
789
790fn signature_action(
791	same_user: bool,
792	target: KeyRole,
793	signer: Option<KeyRole>,
794) -> SignatureAction {
795	match (same_user, target, signer) {
796		| (true, KeyRole::CrossSigningRoot, Some(KeyRole::Device)) =>
797			SignatureAction::Write(KeyRole::Device, SignatureWrite::Merge),
798		| (true, KeyRole::CrossSigningRoot, _) => SignatureAction::Reject,
799		| (true, KeyRole::Device, Some(KeyRole::SelfSigning)) =>
800			SignatureAction::Write(KeyRole::SelfSigning, SignatureWrite::Merge),
801		| (false, KeyRole::CrossSigningRoot, Some(KeyRole::UserSigning)) =>
802			SignatureAction::Write(KeyRole::UserSigning, SignatureWrite::ReplaceSender),
803		| _ => SignatureAction::Ignore,
804	}
805}
806
807fn canonical_key(key: &serde_json::Value) -> Result<String> {
808	let key = to_canonical_object(key)?;
809
810	Ok(to_canonical_json_string_for_signing(&key)?)
811}
812
813#[implement(super::Service)]
814#[tracing::instrument(
815	level = "trace",
816	skip_all,
817	fields(
818		sender = %sender_id,
819		signing_key_id = %key_id,
820		?role,
821	)
822)]
823async fn verify_key_signature(
824	&self,
825	sender_id: &UserId,
826	key_id: &SigningKeyId<AnyKeyName>,
827	role: KeyRole,
828	signature: &str,
829	canonical: &[u8],
830) -> Result<Base64> {
831	let signing_key: serde_json::Value = self
832		.db
833		.keyid_key
834		.qry(&(sender_id, key_id.key_name().as_str()))
835		.map_err(|error| match error {
836			| error if error.is_not_found() =>
837				VerificationError::NoPublicKeysForEntity(sender_id.to_string()).into(),
838			| error => error,
839		})
840		.await?
841		.deserialized()
842		.map_err(|e| err!(Database(debug_warn!("key in keyid_key is invalid: {e:?}"))))?;
843
844	if !key_matches_role(&signing_key, sender_id, key_id.key_name().as_str(), role) {
845		return Err(VerificationError::NoPublicKeysForEntity(sender_id.to_string()).into());
846	}
847
848	let public_key = signing_key
849		.get("keys")
850		.and_then(|keys| keys.get(key_id.as_str()))
851		.and_then(serde_json::Value::as_str)
852		.ok_or_else(|| {
853			Error::from(VerificationError::NoPublicKeysForEntity(sender_id.to_string()))
854		})?;
855
856	verify_signature(sender_id, key_id, public_key, signature, canonical)
857}
858
859fn verify_signature(
860	sender_id: &UserId,
861	key_id: &SigningKeyId<AnyKeyName>,
862	public_key: &str,
863	signature: &str,
864	canonical: &[u8],
865) -> Result<Base64> {
866	let public_key = Base64::<Standard>::parse(public_key)
867		.map_err(|_| VerificationError::NoPublicKeysForEntity(sender_id.to_string()))?;
868
869	let signature = Base64::<Standard>::parse(signature).map_err(|source| {
870		VerificationError::InvalidBase64Signature {
871			path: format!("signatures.{sender_id}.{key_id}"),
872			source,
873		}
874	})?;
875
876	verify_canonical_json_bytes(
877		&key_id.algorithm(),
878		public_key.as_bytes(),
879		signature.as_bytes(),
880		canonical,
881	)?;
882
883	Ok(signature)
884}
885
886fn insert_signatures(
887	key: &mut serde_json::Value,
888	sender_id: &UserId,
889	additional: impl IntoIterator<Item = (String, String)>,
890) -> Result<bool> {
891	let signatures = signatures_map(key)?;
892
893	let signatures = signatures
894		.entry(sender_id.to_string())
895		.or_insert_with(|| serde_json::Map::new().into())
896		.as_object_mut()
897		.ok_or_else(|| {
898			err!(Database(debug_warn!("signature data in keyid_key for a user is invalid.")))
899		})?;
900
901	let changed = additional
902		.into_iter()
903		.fold(false, |changed, (key_id, signature)| {
904			let entry_changed = signatures
905				.get(&key_id)
906				.and_then(serde_json::Value::as_str)
907				!= Some(&signature);
908
909			if entry_changed {
910				signatures.insert(key_id, signature.into());
911			}
912
913			changed | entry_changed
914		});
915
916	Ok(changed)
917}
918
919fn signatures_map(
920	key: &mut serde_json::Value,
921) -> Result<&mut serde_json::Map<String, serde_json::Value>> {
922	let key = key
923		.as_object_mut()
924		.ok_or_else(|| err!(Database(debug_warn!("key in keyid_key is not an object."))))?;
925
926	key.entry("signatures")
927		.or_insert_with(|| serde_json::Map::new().into())
928		.as_object_mut()
929		.ok_or_else(|| {
930			err!(Database(debug_warn!("key in keyid_key has invalid signatures field.")))
931		})
932}
933
934fn replace_signatures(
935	key: &mut serde_json::Value,
936	sender_id: &UserId,
937	(key_id, signature): (String, String),
938) -> Result<bool> {
939	let signatures = signatures_map(key)?;
940	let replacement = serde_json::Map::from_iter([(key_id, signature.into())]);
941
942	match signatures.get_mut(sender_id.as_str()) {
943		| Some(signatures)
944			if signatures
945				.as_object()
946				.is_some_and(|signatures| signatures == &replacement) =>
947			return Ok(false),
948		| Some(signatures) => *signatures = replacement.into(),
949		| None => {
950			signatures.insert(sender_id.to_string(), replacement.into());
951		},
952	}
953
954	Ok(true)
955}
956
957#[implement(super::Service)]
958#[inline]
959pub fn keys_changed<'a>(
960	&'a self,
961	user_id: &'a UserId,
962	from: u64,
963	to: Option<u64>,
964) -> impl Stream<Item = &UserId> + Send + 'a {
965	self.keys_changed_user_or_room(user_id.as_str(), from, to)
966		.map(|(user_id, ..)| user_id)
967}
968
969#[implement(super::Service)]
970#[inline]
971pub fn room_keys_changed<'a>(
972	&'a self,
973	room_id: &'a RoomId,
974	from: u64,
975	to: Option<u64>,
976) -> impl Stream<Item = (&UserId, u64)> + Send + 'a {
977	self.keys_changed_user_or_room(room_id.as_str(), from, to)
978}
979
980#[implement(super::Service)]
981fn keys_changed_user_or_room<'a>(
982	&'a self,
983	user_or_room_id: &'a str,
984	from: u64,
985	to: Option<u64>,
986) -> impl Stream<Item = (&UserId, u64)> + Send + 'a {
987	type KeyVal<'a> = ((&'a str, u64), &'a UserId);
988
989	let to = to.unwrap_or(u64::MAX);
990	let start = (user_or_room_id, from.saturating_add(1));
991	self.db
992		.keychangeid_userid
993		.stream_from(&start)
994		.ignore_err()
995		.ready_take_while(move |((prefix, count), _): &KeyVal<'_>| {
996			*prefix == user_or_room_id && *count <= to
997		})
998		.map(|((_, count), user_id): KeyVal<'_>| (user_id, count))
999}
1000
1001/// Records a device-list change for a user and pushes it to interested peers.
1002///
1003/// The change is durable before any peer is contacted, and the unbounded sender
1004/// channels refuse only once closed, so a refused flush is a shutdown race
1005/// rather than a fault.
1006#[implement(super::Service)]
1007#[tracing::instrument(
1008	name = "device_key_update"
1009	level = INFO_SPAN_LEVEL,
1010	skip_all,
1011	fields(%user_id),
1012)]
1013pub async fn mark_device_key_update(
1014	&self,
1015	user_id: &UserId,
1016	change: DeviceListChange<&DeviceId>,
1017) {
1018	let update_all_rooms = !self
1019		.services
1020		.config
1021		.device_key_update_encrypted_rooms_only;
1022
1023	let all_or_is_encrypted = async |room_id: &RoomId| {
1024		update_all_rooms
1025			|| self
1026				.services
1027				.state_accessor
1028				.is_encrypted_room(room_id)
1029				.await
1030	};
1031
1032	let local = self.services.globals.user_is_local(user_id);
1033	let guard = local
1034		.then_async(async || self.device_list_mutex.lock(user_id).await)
1035		.await;
1036
1037	let count = self.services.globals.next_count();
1038
1039	if local {
1040		let current = self
1041			.get_devicelist_version(user_id)
1042			.await
1043			.unwrap_or(0);
1044
1045		let stream_id = current
1046			.checked_add(u64::from(!matches!(change, DeviceListChange::CrossSigning)))
1047			.expect("device-list stream exhausted");
1048
1049		if !matches!(change, DeviceListChange::CrossSigning) {
1050			self.db
1051				.userid_devicelistversion
1052				.raw_put(user_id, stream_id);
1053		}
1054
1055		let (kind, device_id) = match change {
1056			| DeviceListChange::Resync => (DeviceListChangeKind::Resync, ""),
1057			| DeviceListChange::Device(id) => (DeviceListChangeKind::Device, id.as_str()),
1058			| DeviceListChange::Deleted(id) => (DeviceListChangeKind::Deleted, id.as_str()),
1059			| DeviceListChange::CrossSigning => (DeviceListChangeKind::CrossSigning, ""),
1060		};
1061
1062		// The value layout is (kind: u8, stream_id: u64, device_id: str).
1063		self.db
1064			.keychangeid_devicechange
1065			.put(*count, (u8::from(kind), stream_id, device_id));
1066	}
1067
1068	let user_key = (user_id, *count);
1069
1070	self.db
1071		.keychangeid_userid
1072		.put_raw(user_key, user_id);
1073
1074	self.services
1075		.state_cache
1076		.rooms_joined(user_id)
1077		.filter(|room_id| all_or_is_encrypted(*room_id))
1078		.ready_for_each(|room_id| {
1079			let room_key = (room_id, *count);
1080			self.db
1081				.keychangeid_userid
1082				.put_raw(room_key, user_id);
1083		})
1084		.await;
1085
1086	drop(guard);
1087
1088	self.services
1089		.sending
1090		.send_device_list_appservices(user_id, *count)
1091		.await
1092		.log_err()
1093		.ok();
1094
1095	if !local {
1096		return;
1097	}
1098
1099	// device_list_update EDUs reach remote servers only on a sender flush.
1100	let mut servers: Servers = self
1101		.services
1102		.state_cache
1103		.rooms_joined(user_id)
1104		.filter(|room_id| all_or_is_encrypted(*room_id))
1105		.map(ToOwned::to_owned)
1106		.broad_then(async |room_id: OwnedRoomId| {
1107			self.services
1108				.state_cache
1109				.remote_room_servers(&room_id)
1110				.map(ToOwned::to_owned)
1111				.collect()
1112				.await
1113		})
1114		.flat_map(|servers: Vec<OwnedServerName>| servers.into_iter().stream())
1115		.collect()
1116		.await;
1117
1118	servers.sort_unstable();
1119	servers.dedup();
1120
1121	self.services
1122		.sending
1123		.flush_servers(servers.iter().map(|server| &**server).stream())
1124		.await
1125		.inspect_err(|e| debug_warn!(%e, "device list flush refused"))
1126		.ok();
1127}
1128
1129/// Reads the mutation recorded for a local device-list notification.
1130///
1131/// Missing records return `NotFound`; unknown kind tags request a snapshot resync.
1132#[implement(super::Service)]
1133#[tracing::instrument(level = "trace", skip(self))]
1134pub async fn device_list_change(&self, count: u64) -> Result<DeviceListRecord> {
1135	let row = self
1136		.db
1137		.keychangeid_devicechange
1138		.qry(&count)
1139		.await?;
1140
1141	let (kind, stream_id, device_id): (u8, u64, &str) = deserialize_from_slice(&row)?;
1142	let change =
1143		match DeviceListChangeKind::try_from(kind).unwrap_or(DeviceListChangeKind::Resync) {
1144			| DeviceListChangeKind::Resync => DeviceListChange::Resync,
1145			| DeviceListChangeKind::Device => DeviceListChange::Device(device_id.into()),
1146			| DeviceListChangeKind::Deleted => DeviceListChange::Deleted(device_id.into()),
1147			| DeviceListChangeKind::CrossSigning => DeviceListChange::CrossSigning,
1148		};
1149
1150	Ok(DeviceListRecord { change, stream_id })
1151}
1152
1153impl From<DeviceListChangeKind> for u8 {
1154	fn from(kind: DeviceListChangeKind) -> Self {
1155		match kind {
1156			| DeviceListChangeKind::Resync => 0,
1157			| DeviceListChangeKind::Device => 1,
1158			| DeviceListChangeKind::Deleted => 2,
1159			| DeviceListChangeKind::CrossSigning => 3,
1160		}
1161	}
1162}
1163
1164impl TryFrom<u8> for DeviceListChangeKind {
1165	type Error = ();
1166
1167	fn try_from(kind: u8) -> Result<Self, Self::Error> {
1168		match kind {
1169			| 0 => Ok(Self::Resync),
1170			| 1 => Ok(Self::Device),
1171			| 2 => Ok(Self::Deleted),
1172			| 3 => Ok(Self::CrossSigning),
1173			| _ => Err(()),
1174		}
1175	}
1176}
1177
1178#[implement(super::Service)]
1179pub async fn get_device_keys<'a>(
1180	&'a self,
1181	user_id: &'a UserId,
1182	device_id: &DeviceId,
1183) -> Result<Raw<DeviceKeys>> {
1184	let key_id = (user_id, device_id);
1185	self.db
1186		.keyid_key
1187		.qry(&key_id)
1188		.await
1189		.deserialized()
1190}
1191
1192#[implement(super::Service)]
1193pub async fn get_key<F>(
1194	&self,
1195	key_id: &[u8],
1196	sender_user: Option<&UserId>,
1197	user_id: &UserId,
1198	allowed_signatures: &F,
1199) -> Result<Raw<CrossSigningKey>>
1200where
1201	F: Fn(&UserId) -> bool + Send + Sync,
1202{
1203	let key: serde_json::Value = self
1204		.db
1205		.keyid_key
1206		.get(key_id)
1207		.await
1208		.deserialized()?;
1209
1210	let cleaned = clean_signatures(key, sender_user, user_id, allowed_signatures)?;
1211	let raw_value = serde_json::value::to_raw_value(&cleaned)?;
1212
1213	Ok(Raw::from_json(raw_value))
1214}
1215
1216#[implement(super::Service)]
1217pub async fn get_master_key<F>(
1218	&self,
1219	sender_user: Option<&UserId>,
1220	user_id: &UserId,
1221	allowed_signatures: &F,
1222) -> Result<Raw<CrossSigningKey>>
1223where
1224	F: Fn(&UserId) -> bool + Send + Sync,
1225{
1226	let key_id = self.db.userid_masterkeyid.get(user_id).await?;
1227
1228	self.get_key(&key_id, sender_user, user_id, allowed_signatures)
1229		.await
1230}
1231
1232#[implement(super::Service)]
1233pub async fn get_self_signing_key<F>(
1234	&self,
1235	sender_user: Option<&UserId>,
1236	user_id: &UserId,
1237	allowed_signatures: &F,
1238) -> Result<Raw<CrossSigningKey>>
1239where
1240	F: Fn(&UserId) -> bool + Send + Sync,
1241{
1242	let key_id = self
1243		.db
1244		.userid_selfsigningkeyid
1245		.get(user_id)
1246		.await?;
1247
1248	self.get_key(&key_id, sender_user, user_id, allowed_signatures)
1249		.await
1250}
1251
1252#[implement(super::Service)]
1253pub async fn get_user_signing_key(&self, user_id: &UserId) -> Result<Raw<CrossSigningKey>> {
1254	self.db
1255		.userid_usersigningkeyid
1256		.get(user_id)
1257		.and_then(|key_id| self.db.keyid_key.get(&*key_id))
1258		.await
1259		.deserialized()
1260}
1261
1262pub fn parse_master_key(
1263	user_id: &UserId,
1264	master_key: &Raw<CrossSigningKey>,
1265) -> Result<(Vec<u8>, CrossSigningKey)> {
1266	let mut prefix = user_id.as_bytes().to_vec();
1267	prefix.push(0xFF);
1268
1269	let master_key = master_key
1270		.deserialize()
1271		.map_err(|_| err!(Request(InvalidParam("Invalid master key"))))?;
1272
1273	let mut master_key_ids = master_key.keys.values();
1274	let master_key_id = master_key_ids
1275		.next()
1276		.ok_or(err!(Request(InvalidParam("Master key contained no key."))))?;
1277
1278	if master_key_ids.next().is_some() {
1279		return Err!(Request(InvalidParam("Master key contained more than one key.")));
1280	}
1281
1282	let mut master_key_key = prefix.clone();
1283	master_key_key.extend_from_slice(master_key_id.as_bytes());
1284
1285	Ok((master_key_key, master_key))
1286}
1287
1288pub(super) fn parse_user_signing_key(user_signing_key: &Raw<CrossSigningKey>) -> Result<String> {
1289	let mut user_signing_key_ids = user_signing_key
1290		.deserialize()
1291		.map_err(|_| err!(Request(InvalidParam("Invalid user signing key"))))?
1292		.keys
1293		.into_values();
1294
1295	let user_signing_key_id = user_signing_key_ids
1296		.next()
1297		.ok_or(err!(Request(InvalidParam("User signing key contained no key."))))?;
1298
1299	if user_signing_key_ids.next().is_some() {
1300		return Err!(Request(InvalidParam("User signing key contained more than one key.")));
1301	}
1302
1303	Ok(user_signing_key_id)
1304}
1305
1306/// Ensure that a user only sees signatures from themselves and the target user
1307fn clean_signatures<F>(
1308	mut cross_signing_key: serde_json::Value,
1309	sender_user: Option<&UserId>,
1310	user_id: &UserId,
1311	allowed_signatures: &F,
1312) -> Result<serde_json::Value>
1313where
1314	F: Fn(&UserId) -> bool + Send + Sync,
1315{
1316	if let Some(signatures) = cross_signing_key
1317		.get_mut("signatures")
1318		.and_then(|v| v.as_object_mut())
1319	{
1320		// Don't allocate for the full size of the current signatures, but require
1321		// at most one resize if nothing is dropped
1322		let new_capacity = signatures.len() / 2;
1323		for (user, signature) in
1324			mem::replace(signatures, serde_json::Map::with_capacity(new_capacity))
1325		{
1326			let sid = <&UserId>::try_from(user.as_str())
1327				.map_err(|e| err!(Database("Invalid user ID in database: {e}")))?;
1328
1329			if sender_user == Some(user_id) || sid == user_id || allowed_signatures(sid) {
1330				signatures.insert(user, signature);
1331			}
1332		}
1333	}
1334
1335	Ok(cross_signing_key)
1336}
1337
1338#[cfg(test)]
1339mod tests {
1340	use ruma::{
1341		signatures::{Ed25519KeyPair, KeyPair},
1342		user_id,
1343	};
1344	use tuwunel_database::serialize_to_vec;
1345
1346	use super::*;
1347
1348	#[test]
1349	fn device_list_record_codec() {
1350		for device_id in ["DEVICE", ""] {
1351			let record = (1_u8, 42_u64, device_id);
1352			let bytes = serialize_to_vec(record).expect("serialize device change");
1353
1354			assert_eq!(bytes.len(), 1 + 1 + 8 + 1 + device_id.len());
1355			assert_eq!(bytes[0], 1);
1356			assert_eq!(bytes[1], 0xFF);
1357			assert_eq!(&bytes[2..10], &42_u64.to_be_bytes());
1358			assert_eq!(bytes[10], 0xFF);
1359			assert_eq!(&bytes[11..], device_id.as_bytes());
1360			let decoded: (u8, u64, &str) =
1361				deserialize_from_slice(&bytes).expect("decode device change");
1362
1363			assert_eq!(decoded, record);
1364		}
1365	}
1366
1367	fn signature_fixture() -> (String, String, Vec<u8>) {
1368		let der = Ed25519KeyPair::generate();
1369		let keypair = Ed25519KeyPair::from_der(&der, "DEVICE".to_owned())
1370			.expect("key pair should be generated");
1371
1372		let key = serde_json::json!({
1373			"user_id": "@alice:example.com",
1374			"device_id": "DEVICE",
1375			"keys": { "ed25519:DEVICE": "public-key" },
1376		});
1377
1378		let canonical = canonical_key(&key)
1379			.expect("signing JSON should serialize")
1380			.into_bytes();
1381
1382		let signature = keypair.sign(&canonical).base64();
1383		let public_key = Base64::<Standard, _>::new(keypair.public_key()).encode();
1384
1385		(public_key, signature, canonical)
1386	}
1387
1388	#[test]
1389	fn verifies_canonical_signature_bytes() {
1390		let sender_id = user_id!("@alice:example.com");
1391		let key_id = <&SigningKeyId<AnyKeyName>>::try_from("ed25519:DEVICE")
1392			.expect("signature key ID should parse");
1393
1394		let (public_key, signature, canonical) = signature_fixture();
1395
1396		verify_signature(sender_id, key_id, &public_key, &signature, &canonical)
1397			.expect("signature should verify");
1398	}
1399
1400	#[test]
1401	fn canonicalizes_stored_key_for_verification() {
1402		let sender_id = user_id!("@alice:example.com");
1403		let key_id = <&SigningKeyId<AnyKeyName>>::try_from("ed25519:DEVICE")
1404			.expect("signature key ID should parse");
1405
1406		let der = Ed25519KeyPair::generate();
1407		let keypair = Ed25519KeyPair::from_der(&der, "DEVICE".to_owned())
1408			.expect("key pair should be generated");
1409
1410		let mut stored_key = serde_json::json!({
1411			"user_id": sender_id,
1412			"device_id": "DEVICE",
1413			"keys": { "ed25519:DEVICE": "public-key" },
1414		});
1415
1416		let canonical = canonical_key(&stored_key).expect("stored key should canonicalize");
1417		let signature = keypair.sign(canonical.as_bytes()).base64();
1418		let public_key = Base64::<Standard, _>::new(keypair.public_key()).encode();
1419
1420		stored_key["signatures"] = serde_json::json!({
1421			"@bob:example.com": { "ed25519:BOB": "bob-signature" },
1422		});
1423
1424		stored_key["unsigned"] = serde_json::json!({ "server_data": "ignored" });
1425		let canonical_with_metadata =
1426			canonical_key(&stored_key).expect("stored key with metadata should canonicalize");
1427
1428		assert_eq!(canonical_with_metadata, canonical);
1429		verify_signature(
1430			sender_id,
1431			key_id,
1432			&public_key,
1433			&signature,
1434			canonical_with_metadata.as_bytes(),
1435		)
1436		.expect("signature over the stored key should verify");
1437	}
1438
1439	#[test]
1440	fn rejects_signature_over_different_key() {
1441		let sender_id = user_id!("@alice:example.com");
1442		let key_id = <&SigningKeyId<AnyKeyName>>::try_from("ed25519:DEVICE")
1443			.expect("signature key ID should parse");
1444
1445		let (public_key, signature, _) = signature_fixture();
1446
1447		let error = verify_signature(sender_id, key_id, &public_key, &signature, b"{}")
1448			.expect_err("signature over another object should fail");
1449
1450		assert!(matches!(error, Error::Signatures(_)));
1451	}
1452
1453	#[test]
1454	fn rejects_malformed_signature_base64() {
1455		let sender_id = user_id!("@alice:example.com");
1456		let key_id = <&SigningKeyId<AnyKeyName>>::try_from("ed25519:DEVICE")
1457			.expect("signature key ID should parse");
1458
1459		let (public_key, _, canonical) = signature_fixture();
1460
1461		let error = verify_signature(sender_id, key_id, &public_key, "not base64?", &canonical)
1462			.expect_err("malformed signature base64 should fail");
1463
1464		assert!(matches!(
1465			error,
1466			Error::Signatures(VerificationError::InvalidBase64Signature { .. })
1467		));
1468	}
1469
1470	#[test]
1471	fn classifies_only_unambiguous_key_roles() {
1472		for mask in 0_u8..16 {
1473			let matches = [mask & 1 != 0, mask & 2 != 0, mask & 4 != 0, mask & 8 != 0];
1474			let expected = match mask {
1475				| 1 => Some(KeyRole::Device),
1476				| 2 => Some(KeyRole::CrossSigningRoot),
1477				| 4 => Some(KeyRole::SelfSigning),
1478				| 8 => Some(KeyRole::UserSigning),
1479				| _ => None,
1480			};
1481
1482			assert_eq!(key_role(matches), expected, "role mask {mask:04b}");
1483		}
1484	}
1485
1486	#[test]
1487	fn binds_key_roles_to_owner_and_row_shape() {
1488		let user_id = user_id!("@alice:example.com");
1489		let device = serde_json::json!({
1490			"user_id": user_id,
1491			"device_id": "DEVICE",
1492			"keys": { "ed25519:DEVICE": "device-public-key" },
1493		});
1494
1495		let cross_signing = serde_json::json!({
1496			"user_id": user_id,
1497			"usage": ["untrusted-value"],
1498			"keys": { "ed25519:ROOT": "root-public-key" },
1499		});
1500
1501		let multiple_keys = serde_json::json!({
1502			"user_id": user_id,
1503			"keys": {
1504				"ed25519:ROOT": "root-public-key",
1505				"ed25519:OTHER": "other-public-key",
1506			},
1507		});
1508
1509		assert!(key_matches_role(&device, user_id, "DEVICE", KeyRole::Device));
1510		assert!(!key_matches_role(&device, user_id, "DEVICE", KeyRole::CrossSigningRoot));
1511		assert!(key_matches_role(
1512			&cross_signing,
1513			user_id,
1514			"root-public-key",
1515			KeyRole::CrossSigningRoot
1516		));
1517		assert!(!key_matches_role(
1518			&multiple_keys,
1519			user_id,
1520			"root-public-key",
1521			KeyRole::CrossSigningRoot
1522		));
1523		assert!(!key_matches_role(
1524			&cross_signing,
1525			user_id,
1526			"different-public-key",
1527			KeyRole::CrossSigningRoot
1528		));
1529
1530		assert!(!key_matches_role(
1531			&cross_signing,
1532			user_id!("@bob:example.com"),
1533			"root-public-key",
1534			KeyRole::CrossSigningRoot
1535		));
1536	}
1537
1538	#[test]
1539	fn accepts_only_supported_signature_scopes() {
1540		let roles = [
1541			KeyRole::Device,
1542			KeyRole::CrossSigningRoot,
1543			KeyRole::SelfSigning,
1544			KeyRole::UserSigning,
1545		];
1546		let cases = [false, true].into_iter().flat_map(|same_user| {
1547			roles.into_iter().flat_map(move |target| {
1548				roles
1549					.into_iter()
1550					.map(Some)
1551					.chain([None])
1552					.map(move |signer| (same_user, target, signer))
1553			})
1554		});
1555
1556		for (same_user, target, signer) in cases {
1557			let expected = match (same_user, target, signer) {
1558				| (true, KeyRole::CrossSigningRoot, Some(KeyRole::Device)) =>
1559					SignatureAction::Write(KeyRole::Device, SignatureWrite::Merge),
1560				| (true, KeyRole::CrossSigningRoot, _) => SignatureAction::Reject,
1561				| (true, KeyRole::Device, Some(KeyRole::SelfSigning)) =>
1562					SignatureAction::Write(KeyRole::SelfSigning, SignatureWrite::Merge),
1563				| (false, KeyRole::CrossSigningRoot, Some(KeyRole::UserSigning)) =>
1564					SignatureAction::Write(KeyRole::UserSigning, SignatureWrite::ReplaceSender),
1565				| _ => SignatureAction::Ignore,
1566			};
1567
1568			assert_eq!(
1569				signature_action(same_user, target, signer),
1570				expected,
1571				"same_user={same_user}, target={target:?}, signer={signer:?}"
1572			);
1573		}
1574	}
1575
1576	#[test]
1577	fn replacing_signatures_bounds_key_rotation_history() {
1578		let sender_id = user_id!("@alice:example.com");
1579		let mut key = serde_json::json!({
1580			"signatures": {
1581				"@alice:example.com": { "ed25519:OLD": "old-signature" },
1582				"@bob:example.com": { "ed25519:BOB": "bob-signature" },
1583			},
1584		});
1585
1586		let changed = replace_signatures(
1587			&mut key,
1588			sender_id,
1589			("ed25519:NEW".to_owned(), "new-signature".to_owned()),
1590		)
1591		.expect("signature replacement should succeed");
1592
1593		assert!(changed);
1594
1595		let changed = replace_signatures(
1596			&mut key,
1597			sender_id,
1598			("ed25519:NEW".to_owned(), "new-signature".to_owned()),
1599		)
1600		.expect("idempotent signature replacement should succeed");
1601
1602		assert!(!changed);
1603
1604		let sender_signatures = key["signatures"][sender_id.as_str()]
1605			.as_object()
1606			.expect("sender signatures should remain an object");
1607
1608		assert_eq!(sender_signatures.len(), 1);
1609		assert_eq!(sender_signatures["ed25519:NEW"], "new-signature");
1610		assert_eq!(key["signatures"]["@bob:example.com"]["ed25519:BOB"], "bob-signature");
1611	}
1612
1613	#[test]
1614	fn insert_signatures_creates_missing_map() {
1615		let sender_id = user_id!("@alice:example.com");
1616		let mut key = serde_json::json!({
1617			"user_id": sender_id,
1618			"keys": { "ed25519:ALICE": "ALICE" },
1619		});
1620		let signatures = [("ed25519:ALICE".to_owned(), "alice-signature".to_owned())];
1621
1622		let changed = insert_signatures(&mut key, sender_id, signatures)
1623			.expect("signature insertion should succeed");
1624
1625		assert!(changed);
1626
1627		let signatures = [("ed25519:ALICE".to_owned(), "alice-signature".to_owned())];
1628		let changed = insert_signatures(&mut key, sender_id, signatures)
1629			.expect("idempotent signature insertion should succeed");
1630
1631		assert!(!changed);
1632
1633		assert_eq!(key["signatures"][sender_id.as_str()]["ed25519:ALICE"], "alice-signature");
1634	}
1635
1636	#[test]
1637	fn insert_signatures_preserves_existing_signers() {
1638		let sender_id = user_id!("@alice:example.com");
1639		let mut key = serde_json::json!({
1640			"user_id": sender_id,
1641			"keys": { "ed25519:ROOT": "root-public-key" },
1642			"signatures": {
1643				"@alice:example.com": { "ed25519:OLD": "old-signature" },
1644				"@bob:example.com": { "ed25519:BOB": "bob-signature" },
1645			},
1646		});
1647		let signatures = [
1648			("ed25519:ALICE1".to_owned(), "alice-signature-1".to_owned()),
1649			("ed25519:ALICE2".to_owned(), "alice-signature-2".to_owned()),
1650		];
1651
1652		let changed = insert_signatures(&mut key, sender_id, signatures)
1653			.expect("signature insertion should succeed");
1654
1655		assert!(changed);
1656
1657		let expected = serde_json::json!({
1658			"user_id": sender_id,
1659			"keys": { "ed25519:ROOT": "root-public-key" },
1660			"signatures": {
1661				"@alice:example.com": {
1662					"ed25519:OLD": "old-signature",
1663					"ed25519:ALICE1": "alice-signature-1",
1664					"ed25519:ALICE2": "alice-signature-2",
1665				},
1666				"@bob:example.com": { "ed25519:BOB": "bob-signature" },
1667			},
1668		});
1669
1670		assert_eq!(key, expected);
1671	}
1672
1673	#[test]
1674	fn empty_one_time_key_counts_include_signed_zero() {
1675		let counts = complete_one_time_key_counts(BTreeMap::new());
1676
1677		assert_eq!(counts.len(), 1);
1678		assert_eq!(counts.get(&OneTimeKeyAlgorithm::SignedCurve25519), Some(&UInt::from(0_u32)));
1679	}
1680
1681	#[test]
1682	fn existing_one_time_key_counts_are_preserved() {
1683		let mut counts = BTreeMap::new();
1684		counts.insert(OneTimeKeyAlgorithm::from("curve25519"), UInt::from(11_u32));
1685		counts.insert(OneTimeKeyAlgorithm::SignedCurve25519, UInt::from(17_u32));
1686
1687		let counts = complete_one_time_key_counts(counts);
1688
1689		assert_eq!(
1690			counts.get(&OneTimeKeyAlgorithm::from("curve25519")),
1691			Some(&UInt::from(11_u32))
1692		);
1693		assert_eq!(counts.get(&OneTimeKeyAlgorithm::SignedCurve25519), Some(&UInt::from(17_u32)));
1694	}
1695}