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