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#[derive(Clone, Copy, Debug, Eq, PartialEq)]
42pub enum DeviceListChange<D> {
43 Device(D),
45
46 Deleted(D),
48
49 CrossSigning,
51
52 Resync,
54}
55
56#[derive(Debug, Eq, PartialEq)]
60pub struct DeviceListRecord {
61 pub change: DeviceListChange<OwnedDeviceId>,
63
64 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#[derive(Debug, Deserialize, Serialize)]
100struct FallbackEntry {
101 key_id: OwnedOneTimeKeyId,
102 key: Raw<OneTimeKey>,
103 used: bool,
104}
105
106type 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 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 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 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 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 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
413fn 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#[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#[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 {
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 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#[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 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 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#[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
1306fn 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 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}