1use std::{
2 net::IpAddr,
3 time::{Duration, SystemTime},
4};
5
6use futures::{FutureExt, Stream, StreamExt, future::join};
7use ruma::{
8 DeviceId, MilliSecondsSinceUnixEpoch, OwnedDeviceId, OwnedUserId, UserId,
9 api::client::device::Device, events::AnyToDeviceEvent, serde::Raw,
10};
11use serde_json::json;
12use tuwunel_core::{
13 Err, Result, at, implement, trace,
14 utils::{
15 BoolExt, ReadyExt, random_string,
16 stream::{IterStream, TryIgnore},
17 string::to_small_string,
18 time::{
19 duration_since_epoch, timepoint_from_epoch, timepoint_from_now, timepoint_has_passed,
20 },
21 },
22};
23use tuwunel_database::{Cbor, Deserialized, Ignore, Interfix, Json, Txn};
24
25use super::DeviceListChange;
26
27const DEVICE_ID_LENGTH: usize = 10;
29
30pub const TOKEN_LENGTH: usize = 32;
32
33#[implement(super::Service)]
38#[tracing::instrument(level = "info", skip(self, access_token))]
39pub async fn create_device(
40 &self,
41 user_id: &UserId,
42 device_id: Option<&DeviceId>,
43 (access_token, expires_in): (Option<&str>, Option<Duration>),
44 refresh_token: Option<&str>,
45 initial_device_display_name: Option<&str>,
46 client_ip: Option<IpAddr>,
47) -> Result<OwnedDeviceId> {
48 let device_id = resolve_device_id(device_id);
49
50 if !self.exists(user_id).await {
51 return Err!(Request(InvalidParam(error!(
52 "Called create_device for non-existent user {user_id}"
53 ))));
54 }
55
56 if self
57 .is_cross_signing_key_id(user_id, device_id.as_str())
58 .await?
59 {
60 return Err!(Request(Forbidden("Device ID matches a cross-signing key ID.")));
61 }
62
63 let notify = true;
64 self.put_device_metadata(user_id, notify, &Device {
65 device_id: device_id.clone(),
66 display_name: initial_device_display_name.map(Into::into),
67 last_seen_ts: Some(MilliSecondsSinceUnixEpoch::now()),
68 last_seen_ip: client_ip.map(to_small_string),
69 })
70 .await;
71
72 if let Some(access_token) = access_token {
73 self.set_access_token(user_id, &device_id, access_token, expires_in, refresh_token)
74 .await?;
75 }
76
77 Ok(device_id)
78}
79
80fn resolve_device_id(device_id: Option<&DeviceId>) -> OwnedDeviceId {
81 device_id
83 .filter(|device_id| !device_id.as_str().is_empty())
84 .map(ToOwned::to_owned)
85 .unwrap_or_else(|| OwnedDeviceId::from(random_string(DEVICE_ID_LENGTH)))
86}
87
88#[implement(super::Service)]
93#[tracing::instrument(level = "info", skip(self))]
94pub async fn remove_device(&self, user_id: &UserId, device_id: &DeviceId) {
95 self.remove_tokens(user_id, device_id).await;
97
98 self.db
100 .todeviceid_events
101 .del_prefix(&(user_id, device_id, Interfix))
102 .await;
103
104 self.services
106 .pusher
107 .get_device_pushkeys(user_id, device_id)
108 .map(Vec::into_iter)
109 .map(IterStream::stream)
110 .flatten_stream()
111 .for_each(async |pushkey| {
112 self.services
113 .pusher
114 .delete_pusher(user_id, &pushkey)
115 .await;
116 })
117 .await;
118
119 self.remove_dehydrated_device(user_id, Some(device_id))
121 .await
122 .ok();
123
124 self.remove_one_time_keys(user_id, device_id)
125 .await;
126
127 self.db
129 .userdeviceidalgorithm_fallback
130 .del_prefix(&(user_id, device_id, Interfix))
131 .await;
132
133 let event_type = format!("org.matrix.msc3890.local_notification_settings.{device_id}").into();
135 self.services
136 .account_data
137 .delete(None, user_id, event_type)
138 .await
139 .ok();
140
141 let userdeviceid = (user_id, device_id);
142
143 self.db.keyid_key.del(userdeviceid);
144 self.db.userdeviceid_metadata.del(userdeviceid);
145 self.db.oidcdevice_userdeviceid.del(userdeviceid);
146
147 self.mark_device_key_update(user_id, DeviceListChange::Deleted(device_id))
148 .await;
149}
150
151#[implement(super::Service)]
153pub fn all_device_ids<'a>(
154 &'a self,
155 user_id: &'a UserId,
156) -> impl Stream<Item = &DeviceId> + Send + 'a {
157 let prefix = (user_id, Interfix);
158 self.db
159 .userdeviceid_metadata
160 .keys_prefix(&prefix)
161 .ignore_err()
162 .map(|(_, device_id): (Ignore, &DeviceId)| device_id)
163}
164
165#[implement(super::Service)]
167#[tracing::instrument(level = "trace", skip(self, token))]
168pub async fn find_from_token(
169 &self,
170 token: &str,
171) -> Result<(OwnedUserId, OwnedDeviceId, Option<SystemTime>)> {
172 self.db
173 .token_userdeviceid
174 .get(token)
175 .await
176 .deserialized()
177 .and_then(|(user_id, device_id, expires_at): (_, _, Option<u64>)| {
178 let expires_at = expires_at
179 .map(Duration::from_secs)
180 .map(timepoint_from_epoch)
181 .transpose()?;
182
183 Ok((user_id, device_id, expires_at))
184 })
185}
186
187#[implement(super::Service)]
188#[tracing::instrument(level = "debug", skip(self))]
189pub async fn remove_tokens(&self, user_id: &UserId, device_id: &DeviceId) {
190 let remove_access = self
191 .remove_access_token(user_id, device_id)
192 .map(Result::ok);
193
194 let remove_refresh = self
195 .remove_refresh_token(user_id, device_id)
196 .map(Result::ok);
197
198 join(remove_access, remove_refresh).await;
199}
200
201#[implement(super::Service)]
203#[tracing::instrument(level = "debug", skip(self))]
204pub async fn set_access_token(
205 &self,
206 user_id: &UserId,
207 device_id: &DeviceId,
208 access_token: &str,
209 expires_in: Option<Duration>,
210 refresh_token: Option<&str>,
211) -> Result {
212 assert!(
213 access_token.len() >= TOKEN_LENGTH,
214 "Caller must supply an access_token >= {TOKEN_LENGTH} chars."
215 );
216
217 if let Some(refresh_token) = refresh_token {
218 self.set_refresh_token(user_id, device_id, refresh_token)
219 .await?;
220 }
221
222 let expires_at = expires_in
223 .map(timepoint_from_now)
224 .transpose()?
225 .map(duration_since_epoch)
226 .as_ref()
227 .map(Duration::as_secs);
228
229 let userdeviceid = (user_id, device_id);
230
231 let previous = self
233 .db
234 .userdeviceid_token
235 .qry(&userdeviceid)
236 .await
237 .deserialized::<String>()
238 .ok();
239
240 let mut txn = self.services.db.txn();
241
242 if let Some(previous) = previous.as_deref() {
243 let key = (user_id, device_id, previous);
244
245 txn.put_raw(&self.db.userdeviceidtoken_index, key, []);
246 }
247
248 let key = (user_id, device_id, access_token);
249 let value = (user_id, device_id, expires_at);
250
251 txn.raw_put(&self.db.token_userdeviceid, access_token, value);
252 txn.put_raw(&self.db.userdeviceidtoken_index, key, []);
253 txn.put_raw(&self.db.userdeviceid_token, userdeviceid, access_token);
254
255 txn.execute();
256
257 Ok(())
258}
259
260#[implement(super::Service)]
263pub async fn remove_access_token(&self, user_id: &UserId, device_id: &DeviceId) -> Result {
264 let prefix = (user_id, device_id, Interfix);
265 self.db
266 .userdeviceidtoken_index
267 .keys_prefix(&prefix)
268 .ignore_err()
269 .ready_for_each(|(_, _, token): (Ignore, Ignore, &str)| {
270 self.db.token_userdeviceid.remove(token);
271 self.db
272 .userdeviceidtoken_index
273 .del((user_id, device_id, token));
274 })
275 .await;
276
277 let token = self
279 .db
280 .userdeviceid_token
281 .qry(&(user_id, device_id))
282 .await
283 .deserialized::<String>()
284 .ok();
285
286 let mut txn = self.services.db.txn();
287
288 if let Some(token) = token.as_deref() {
289 txn.del_raw(&self.db.token_userdeviceid, token);
290 }
291
292 txn.del(&self.db.userdeviceid_token, (user_id, device_id));
293
294 txn.execute();
295
296 Ok(())
297}
298
299#[implement(super::Service)]
302pub async fn remove_access_token_value(&self, access_token: &str) {
303 let owner = self
304 .db
305 .token_userdeviceid
306 .get(access_token)
307 .await
308 .deserialized::<(OwnedUserId, OwnedDeviceId, Option<u64>)>()
309 .ok();
310
311 let mut txn = self.services.db.txn();
312
313 if let Some((user_id, device_id, _)) = owner {
314 let user_device_token = (&*user_id, &*device_id, access_token);
315
316 txn.del(&self.db.userdeviceidtoken_index, user_device_token);
317 }
318
319 txn.del_raw(&self.db.token_userdeviceid, access_token);
320
321 txn.execute();
322}
323
324#[implement(super::Service)]
325pub fn generate_access_token(&self, expires: bool) -> (String, Option<Duration>) {
326 let access_token = random_string(TOKEN_LENGTH);
327 let expires_in = expires
328 .then_some(self.services.server.config.access_token_ttl)
329 .map(Duration::from_secs);
330
331 (access_token, expires_in)
332}
333
334#[implement(super::Service)]
336#[tracing::instrument(level = "debug", skip(self))]
337pub async fn set_refresh_token(
338 &self,
339 user_id: &UserId,
340 device_id: &DeviceId,
341 refresh_token: &str,
342) -> Result {
343 debug_assert!(refresh_token.starts_with("refresh_"), "refresh_token missing prefix");
344
345 let config = &self.services.server.config;
346 let ttl = config.refresh_token_ttl;
347 let idle_only = config.refresh_token_idle_only;
348
349 let prior_expires_at: Option<SystemTime> = (ttl != 0 && !idle_only)
351 .then_async(|| self.find_refresh_token_expires_at(user_id, device_id))
352 .await
353 .flatten();
354
355 let spent: Option<String> = self
358 .db
359 .userdeviceid_refresh
360 .qry(&(user_id, device_id))
361 .await
362 .deserialized()
363 .ok();
364
365 self.remove_refresh_token(user_id, device_id)
367 .await
368 .ok();
369
370 let expires_at = match (ttl, prior_expires_at) {
371 | (0, _) => None,
372 | (_, Some(prior)) => Some(prior),
373 | (ttl, None) => Some(timepoint_from_now(Duration::from_secs(ttl))?),
374 };
375
376 let expires_at_secs = expires_at
377 .map(duration_since_epoch)
378 .as_ref()
379 .map(Duration::as_secs);
380
381 let userdeviceid = (user_id, device_id);
382 let value = (user_id, device_id, expires_at_secs);
383 let mut txn = self.services.db.txn();
384
385 txn.raw_put(&self.db.token_userdeviceid, refresh_token, value);
386 txn.put_raw(&self.db.userdeviceid_refresh, userdeviceid, refresh_token);
387
388 if let Some(spent) = spent {
391 let spent_at = duration_since_epoch(SystemTime::now()).as_secs();
392 let value = (user_id, device_id, refresh_token, spent_at);
393
394 txn.raw_put(&self.db.spentrefresh_userdeviceid, &*spent, value);
395 txn.put_raw(&self.db.userdeviceid_spentrefresh, userdeviceid, &*spent);
396 }
397
398 txn.execute();
399
400 Ok(())
401}
402
403#[implement(super::Service)]
407async fn find_refresh_token_expires_at(
408 &self,
409 user_id: &UserId,
410 device_id: &DeviceId,
411) -> Option<SystemTime> {
412 let userdeviceid = (user_id, device_id);
413 let old_token: String = self
414 .db
415 .userdeviceid_refresh
416 .qry(&userdeviceid)
417 .await
418 .deserialized()
419 .ok()?;
420
421 let (_, _, expires_at_secs): (Ignore, Ignore, Option<u64>) = self
422 .db
423 .token_userdeviceid
424 .get(&old_token)
425 .await
426 .deserialized()
427 .ok()?;
428
429 expires_at_secs
430 .map(Duration::from_secs)
431 .map(timepoint_from_epoch)?
432 .ok()
433}
434
435#[implement(super::Service)]
438pub async fn remove_refresh_token(&self, user_id: &UserId, device_id: &DeviceId) -> Result {
439 let userdeviceid = (user_id, device_id);
440 let refresh_token = self
441 .db
442 .userdeviceid_refresh
443 .qry(&userdeviceid)
444 .await;
445
446 let mut txn = self.services.db.txn();
447
448 if let Ok(refresh_token) = refresh_token {
449 txn.del_raw(&self.db.token_userdeviceid, &refresh_token);
450 }
451
452 txn.del(&self.db.userdeviceid_refresh, userdeviceid);
453
454 self.forget_spent_refresh_token(user_id, device_id, &mut txn)
455 .await;
456
457 txn.execute();
458
459 Ok(())
460}
461
462#[implement(super::Service)]
465async fn forget_spent_refresh_token(
466 &self,
467 user_id: &UserId,
468 device_id: &DeviceId,
469 txn: &mut Txn,
470) {
471 let userdeviceid = (user_id, device_id);
472
473 if let Ok(spent) = self
474 .db
475 .userdeviceid_spentrefresh
476 .qry(&userdeviceid)
477 .await
478 {
479 txn.del_raw(&self.db.spentrefresh_userdeviceid, &spent);
480 }
481
482 txn.del(&self.db.userdeviceid_spentrefresh, userdeviceid);
483}
484
485pub enum RefreshToken {
488 Current {
490 user_id: OwnedUserId,
491 device_id: OwnedDeviceId,
492 expires_at: Option<SystemTime>,
493 },
494
495 Replayed {
500 user_id: OwnedUserId,
501 device_id: OwnedDeviceId,
502 current: String,
503 grace: bool,
504 },
505
506 Unknown,
508}
509
510#[implement(super::Service)]
512pub async fn classify_refresh_token(&self, presented: &str) -> RefreshToken {
513 if let Ok((user_id, device_id, expires_at)) = self.find_from_token(presented).await {
516 let current: Option<String> = self
517 .db
518 .userdeviceid_refresh
519 .qry(&(&user_id, &device_id))
520 .await
521 .deserialized()
522 .ok();
523
524 if current.as_deref() == Some(presented) {
525 return RefreshToken::Current { user_id, device_id, expires_at };
526 }
527 }
528
529 let Ok((user_id, device_id, successor, spent_at)) = self
532 .db
533 .spentrefresh_userdeviceid
534 .get(presented)
535 .await
536 .deserialized::<(OwnedUserId, OwnedDeviceId, String, u64)>()
537 else {
538 return RefreshToken::Unknown;
539 };
540
541 let current: Option<String> = self
542 .db
543 .userdeviceid_refresh
544 .qry(&(&user_id, &device_id))
545 .await
546 .deserialized()
547 .ok();
548
549 let grace_window = self
550 .services
551 .server
552 .config
553 .refresh_token_reuse_grace;
554 let elapsed = duration_since_epoch(SystemTime::now())
555 .as_secs()
556 .saturating_sub(spent_at);
557
558 let grace = grace_window != 0
559 && elapsed <= grace_window
560 && current.as_deref() == Some(successor.as_str());
561
562 RefreshToken::Replayed {
563 user_id,
564 device_id,
565 current: successor,
566 grace,
567 }
568}
569
570#[must_use]
571pub fn generate_refresh_token() -> String { format!("refresh_{}", random_string(TOKEN_LENGTH)) }
572
573#[implement(super::Service)]
574pub fn add_to_device_event(
575 &self,
576 sender: &UserId,
577 target_user_id: &UserId,
578 target_device_id: &DeviceId,
579 event_type: &str,
580 content: &serde_json::Value,
581) -> u64 {
582 let count = self.services.globals.next_count();
583
584 let key = (target_user_id, target_device_id, *count);
585 self.db.todeviceid_events.put(
586 key,
587 Json(json!({
588 "type": event_type,
589 "sender": sender,
590 "content": content,
591 })),
592 );
593
594 trace!(
595 %target_user_id,
596 %target_device_id,
597 count = *count,
598 %event_type,
599 %sender,
600 "to_device write",
601 );
602
603 *count
604}
605
606#[implement(super::Service)]
607pub fn get_to_device_events<'a>(
608 &'a self,
609 user_id: &'a UserId,
610 device_id: &'a DeviceId,
611 since: Option<u64>,
612 to: Option<u64>,
613) -> impl Stream<Item = (u64, Raw<AnyToDeviceEvent>)> + Send + 'a {
614 type Key<'a> = (&'a UserId, &'a DeviceId, u64);
615
616 let from = (user_id, device_id, since.map_or(0, |since| since.saturating_add(1)));
617
618 self.db
619 .todeviceid_events
620 .stream_from(&from)
621 .ignore_err()
622 .ready_take_while(move |((user_id_, device_id_, count), _): &(Key<'_>, _)| {
623 user_id == *user_id_ && device_id == *device_id_ && to.is_none_or(|to| *count <= to)
624 })
625 .map(|((_, _, count), event)| (count, event))
626}
627
628#[implement(super::Service)]
629pub async fn remove_to_device_events<Until>(
630 &self,
631 user_id: &UserId,
632 device_id: &DeviceId,
633 until: Until,
634) where
635 Until: Into<Option<u64>> + Send,
636{
637 type Key<'a> = (&'a UserId, &'a DeviceId, u64);
638
639 let until = until.into().unwrap_or(u64::MAX);
640 let from = (user_id, device_id, until);
641 self.db
642 .todeviceid_events
643 .rev_keys_from(&from)
644 .ignore_err()
645 .ready_take_while(move |(user_id_, device_id_, _): &Key<'_>| {
646 user_id == *user_id_ && device_id == *device_id_
647 })
648 .ready_for_each(|key: Key<'_>| {
649 self.db.todeviceid_events.del(key);
650 })
651 .await;
652}
653
654#[implement(super::Service)]
655pub async fn update_device_last_seen(
656 &self,
657 user_id: &UserId,
658 device_id: &DeviceId,
659 last_seen_ip: Option<IpAddr>,
660 last_seen_ts: Option<MilliSecondsSinceUnixEpoch>,
661) -> Result {
662 let mut device = self
663 .get_device_metadata(user_id, device_id)
664 .await?;
665
666 if let Some(last_seen_ip) = last_seen_ip.map(to_small_string) {
667 device.last_seen_ip.replace(last_seen_ip);
668 }
669
670 device
671 .last_seen_ts
672 .replace(last_seen_ts.unwrap_or_else(MilliSecondsSinceUnixEpoch::now));
673
674 self.put_device_metadata(user_id, false, &device)
675 .await;
676
677 Ok(())
678}
679
680#[implement(super::Service)]
684#[tracing::instrument(level = "trace", skip(self, device))]
685pub async fn put_device_metadata(&self, user_id: &UserId, notify: bool, device: &Device) {
686 let key = (user_id, &device.device_id);
687 self.db
688 .userdeviceid_metadata
689 .put(key, Json(device));
690
691 if notify {
692 self.mark_device_key_update(user_id, DeviceListChange::Device(&device.device_id))
693 .boxed() .await;
695 }
696}
697
698#[implement(super::Service)]
700pub async fn get_device_metadata(
701 &self,
702 user_id: &UserId,
703 device_id: &DeviceId,
704) -> Result<Device> {
705 self.db
706 .userdeviceid_metadata
707 .qry(&(user_id, device_id))
708 .await
709 .deserialized()
710 .inspect(|device: &Device| {
711 debug_assert_eq!(&device.device_id, device_id, "device_id mismatch");
712 })
713}
714
715#[implement(super::Service)]
716pub async fn device_exists(&self, user_id: &UserId, device_id: &DeviceId) -> bool {
717 self.db
718 .userdeviceid_metadata
719 .contains(&(user_id, device_id))
720 .await
721}
722
723#[implement(super::Service)]
724pub async fn is_oidc_device(&self, user_id: &UserId, device_id: &DeviceId) -> bool {
725 self.db
726 .oidcdevice_userdeviceid
727 .contains(&(user_id, device_id))
728 .await
729}
730
731#[implement(super::Service)]
734pub async fn get_oidc_device_idp(
735 &self,
736 user_id: &UserId,
737 device_id: &DeviceId,
738) -> Option<String> {
739 self.db
740 .oidcdevice_userdeviceid
741 .qry(&(user_id, device_id))
742 .await
743 .deserialized::<Json<String>>()
744 .ok()
745 .map(|Json(idp)| idp)
746 .filter(|idp| !idp.is_empty())
747}
748
749#[implement(super::Service)]
750pub fn mark_oidc_device(&self, user_id: &UserId, device_id: &DeviceId, idp_id: &str) {
751 self.db
752 .oidcdevice_userdeviceid
753 .put((user_id, device_id), Json(idp_id));
754}
755
756#[expect(clippy::must_use_candidate)]
759#[implement(super::Service)]
760pub fn allow_cross_signing_replacement(&self, user_id: &UserId) -> SystemTime {
761 let duration = Duration::from_mins(10);
762 let expires = timepoint_from_now(duration).expect("failed to create timepoint from now");
763
764 self.db
765 .oidccskeybypass_userid
766 .raw_put(user_id, Cbor(expires));
767
768 expires
769}
770
771#[implement(super::Service)]
773pub async fn can_replace_cross_signing_keys(&self, user_id: &UserId) -> bool {
774 let Ok(expires): Result<SystemTime, _> = self
775 .db
776 .oidccskeybypass_userid
777 .get(user_id)
778 .await
779 .deserialized::<Cbor<_>>()
780 .map(at!(0))
781 else {
782 return false;
783 };
784
785 if !timepoint_has_passed(expires) {
786 return true;
787 }
788
789 self.db.oidccskeybypass_userid.remove(user_id);
790 false
791}
792
793#[implement(super::Service)]
794pub async fn get_devicelist_version(&self, user_id: &UserId) -> Result<u64> {
795 self.db
796 .userid_devicelistversion
797 .get(user_id)
798 .await
799 .deserialized()
800}
801
802#[implement(super::Service)]
803pub fn all_devices_metadata<'a>(
804 &'a self,
805 user_id: &'a UserId,
806) -> impl Stream<Item = Device> + Send + 'a {
807 let key = (user_id, Interfix);
808 self.db
809 .userdeviceid_metadata
810 .stream_prefix(&key)
811 .ignore_err()
812 .map(|(_, val): (Ignore, Device)| val)
813}
814
815#[cfg(test)]
816mod tests {
817 use super::*;
818
819 #[test]
820 fn absent_device_id_is_generated() {
821 let device_id = resolve_device_id(None);
822
823 assert_eq!(device_id.as_str().len(), DEVICE_ID_LENGTH);
824 }
825
826 #[test]
827 fn empty_device_id_is_generated() {
828 let device_id = resolve_device_id(Some("".into()));
829
830 assert_eq!(device_id.as_str().len(), DEVICE_ID_LENGTH);
831 }
832
833 #[test]
834 fn provided_device_id_is_preserved() {
835 let device_id = resolve_device_id(Some("HELLOWORLD".into()));
836
837 assert_eq!(device_id.as_str(), "HELLOWORLD");
838 }
839}