1use std::collections::BTreeMap;
2
3use axum::extract::State;
4use futures::{StreamExt, future::join};
5use ruma::{
6 OneTimeKeyAlgorithm, OwnedDeviceId, OwnedOneTimeKeyId, OwnedUserId, ServerName, UserId,
7 api::{client::keys::claim_keys, federation},
8 encryption::OneTimeKey,
9 serde::Raw,
10};
11use serde_json::json;
12use tuwunel_core::{
13 Result, debug_warn,
14 utils::{
15 BoolExt, IterStream,
16 stream::{BroadbandExt, ReadyExt},
17 },
18};
19use tuwunel_service::Services;
20
21use super::FailureMap;
22use crate::Ruma;
23
24#[derive(Default)]
25struct Claims {
26 one_time_keys: OneTimeKeyMap,
27 failures: FailureMap,
28}
29
30type RequestClaims = BTreeMap<OwnedUserId, Algorithms>;
31type ServerClaims<'a> = BTreeMap<&'a ServerName, RequestClaims>;
32type LocalClaim<'a> = (&'a UserId, &'a Algorithms);
33type Algorithms = BTreeMap<OwnedDeviceId, OneTimeKeyAlgorithm>;
34type OneTimeKeys = BTreeMap<OwnedOneTimeKeyId, Raw<OneTimeKey>>;
35type OneTimeKeyMap = BTreeMap<OwnedUserId, BTreeMap<OwnedDeviceId, OneTimeKeys>>;
36
37pub(crate) async fn claim_keys_route(
41 State(services): State<crate::State>,
42 body: Ruma<claim_keys::v3::Request>,
43) -> Result<claim_keys::v3::Response> {
44 claim_keys_helper(&services, &body.one_time_keys).await
45}
46
47pub(crate) async fn claim_keys_helper(
48 services: &Services,
49 one_time_keys_input: &RequestClaims,
50) -> Result<claim_keys::v3::Response> {
51 let (local_users, remote_users): (Vec<_>, Vec<_>) = one_time_keys_input
52 .iter()
53 .map(|(uid, map)| (uid.as_ref(), map))
54 .partition(|(user_id, _)| services.globals.user_is_local(user_id));
55
56 let server: ServerClaims<'_> =
57 remote_users
58 .into_iter()
59 .fold(BTreeMap::new(), |mut acc, (user_id, map)| {
60 acc.entry(user_id.server_name())
61 .or_default()
62 .insert(user_id.to_owned(), map.clone());
63 acc
64 });
65
66 let local = collect_local_one_time_keys(services, &local_users);
67 let federation = collect_federation_one_time_keys(services, server);
68
69 let (local, federation) = join(local, federation).await;
70 let merged = local.merge(federation);
71
72 Ok(claim_keys::v3::Response {
73 failures: merged.failures,
74 one_time_keys: merged.one_time_keys,
75 })
76}
77
78async fn collect_local_one_time_keys(services: &Services, users: &[LocalClaim<'_>]) -> Claims {
79 let one_time_keys = users
80 .iter()
81 .copied()
82 .stream()
83 .broad_filter_map(async |(user_id, requested)| {
84 let (mut device_keys, mut needed) = requested
85 .iter()
86 .stream()
87 .fold(
88 (BTreeMap::new(), BTreeMap::new()),
89 async |(mut device_keys, mut needed), (device_id, algorithm)| {
90 match services
91 .users
92 .take_one_time_key(user_id, device_id, algorithm)
93 .await
94 {
95 | Ok(key) => {
96 device_keys.insert(device_id.clone(), [key].into());
97 },
98 | Err(_) => {
99 needed.insert(device_id.clone(), algorithm.clone());
100 },
101 }
102
103 (device_keys, needed)
104 },
105 )
106 .await;
107
108 let claimed = needed
110 .is_empty()
111 .is_false()
112 .then_async(|| services.appservice.claim_keys(user_id, &needed))
113 .await
114 .unwrap_or_default();
115
116 for (device_id, keys) in claimed {
117 needed.remove(&device_id);
118 device_keys.insert(device_id, keys);
119 }
120
121 let device_keys = needed
122 .into_iter()
123 .stream()
124 .fold(device_keys, async |mut device_keys, (device_id, algorithm)| {
125 if let Ok(key) = services
126 .users
127 .take_fallback_key(user_id, &device_id, &algorithm)
128 .await
129 {
130 device_keys.insert(device_id, [key].into());
131 }
132
133 device_keys
134 })
135 .await;
136
137 (!device_keys.is_empty()).then(|| (user_id.to_owned(), device_keys))
139 })
140 .collect()
141 .await;
142
143 Claims { one_time_keys, ..Default::default() }
144}
145
146async fn collect_federation_one_time_keys(
147 services: &Services,
148 server: ServerClaims<'_>,
149) -> Claims {
150 server
151 .into_iter()
152 .stream()
153 .broad_then(async |(server, one_time_keys)| {
154 let failed = || Claims {
155 failures: [(server.to_string(), json!({}))].into(),
156 ..Default::default()
157 };
158
159 let request = federation::keys::claim_keys::v1::Request { one_time_keys };
160
161 match services
162 .federation
163 .execute_keys(server, request)
164 .await
165 .inspect_err(
166 |e| debug_warn!(%server, "claim_keys federation request failed: {e}"),
167 ) {
168 | Err(_e) => failed(),
169 | Ok(keys) => Claims {
170 one_time_keys: keys.one_time_keys,
171 failures: Default::default(),
172 },
173 }
174 })
175 .ready_fold(Claims::default(), Claims::merge)
176 .await
177}
178
179impl Claims {
180 fn merge(mut self, other: Self) -> Self {
181 self.one_time_keys.extend(other.one_time_keys);
182 self.failures.extend(other.failures);
183 self
184 }
185}