1use std::collections::{BTreeMap, HashMap};
2
3use axum::extract::State;
4use futures::{
5 FutureExt, StreamExt,
6 future::{join, join4},
7};
8use ruma::{
9 CanonicalJsonObject, CanonicalJsonValue, DeviceId, OwnedDeviceId, OwnedUserId, ServerName,
10 UserId,
11 api::{
12 client::{device::Device, keys::get_keys},
13 federation::keys::get_keys::v1::{
14 Request as FederationRequest, Response as FederationResponse,
15 },
16 },
17 encryption::{CrossSigningKey, DeviceKeys},
18 serde::Raw,
19};
20use serde_json::value::to_raw_value;
21use tuwunel_core::{
22 Result, debug_warn, implement,
23 utils::{
24 BoolExt, IterStream,
25 future::TryExtExt,
26 json,
27 stream::{BroadbandExt, ReadyExt},
28 },
29};
30use tuwunel_service::{
31 Services,
32 federation::feds::{Outcome, OutcomeExt, fanout_with},
33 users::parse_master_key,
34};
35
36use super::{FailureMap, federation_failures, federation_opts};
37use crate::Ruma;
38
39#[derive(Default)]
40struct Keys {
41 device_keys: DeviceKeyMap,
42 master_keys: CrossSigningKeys,
43 self_signing_keys: CrossSigningKeys,
44 user_signing_keys: CrossSigningKeys,
45 failures: FailureMap,
46}
47
48type DeviceLists = BTreeMap<OwnedUserId, Vec<OwnedDeviceId>>;
49type DeviceKeyMap = BTreeMap<OwnedUserId, BTreeMap<OwnedDeviceId, Raw<DeviceKeys>>>;
50type ServerDevices<'a> = HashMap<&'a ServerName, DeviceLists>;
51type LocalDeviceUser<'a> = (&'a UserId, &'a Vec<OwnedDeviceId>);
52type CrossSigningKeys = BTreeMap<OwnedUserId, Raw<CrossSigningKey>>;
53
54pub(crate) async fn get_keys_route(
63 State(services): State<crate::State>,
64 body: Ruma<get_keys::v3::Request>,
65) -> Result<get_keys::v3::Response> {
66 let sender_user = body.sender_user();
67
68 get_keys_helper(
69 &services,
70 Some(sender_user),
71 &body.device_keys,
72 |u| u == sender_user,
73 true, )
75 .await
76}
77
78pub(crate) async fn get_keys_helper<F>(
79 services: &Services,
80 sender_user: Option<&UserId>,
81 device_keys_input: &DeviceLists,
82 allowed_signatures: F,
83 include_display_names: bool,
84) -> Result<get_keys::v3::Response>
85where
86 F: Fn(&UserId) -> bool + Send + Sync,
87{
88 let (local_users, remote_users): (Vec<LocalDeviceUser<'_>>, Vec<_>) = device_keys_input
89 .iter()
90 .map(|(uid, dids)| (uid.as_ref(), dids))
91 .partition(|(user_id, _)| services.globals.user_is_local(user_id));
92
93 let server: ServerDevices<'_> =
94 remote_users
95 .into_iter()
96 .fold(HashMap::new(), |mut acc, (user_id, device_ids)| {
97 acc.entry(user_id.server_name())
98 .or_default()
99 .insert(user_id.to_owned(), device_ids.clone());
100 acc
101 });
102
103 let local = collect_local_keys(
104 services,
105 &local_users,
106 sender_user,
107 &allowed_signatures,
108 include_display_names,
109 );
110
111 let federation = collect_federation_keys(services, server, sender_user, &allowed_signatures);
112
113 let (local, federation) = join(local, federation).await;
114 Ok(local.merge(federation).into_response())
115}
116
117async fn collect_local_keys<F>(
118 services: &Services,
119 users: &[LocalDeviceUser<'_>],
120 sender_user: Option<&UserId>,
121 allowed_signatures: &F,
122 include_display_names: bool,
123) -> Keys
124where
125 F: Fn(&UserId) -> bool + Send + Sync,
126{
127 users
128 .iter()
129 .copied()
130 .stream()
131 .broad_then(async |(user_id, device_ids)| {
132 collect_local_user_keys(
133 services,
134 user_id,
135 device_ids,
136 sender_user,
137 allowed_signatures,
138 include_display_names,
139 )
140 .await
141 })
142 .ready_fold(Keys::default(), Keys::merge)
143 .await
144}
145
146async fn collect_local_user_keys<F>(
147 services: &Services,
148 user_id: &UserId,
149 device_ids: &[OwnedDeviceId],
150 sender_user: Option<&UserId>,
151 allowed_signatures: &F,
152 include_display_names: bool,
153) -> Keys
154where
155 F: Fn(&UserId) -> bool + Send + Sync,
156{
157 let device_keys =
158 collect_local_device_keys(services, user_id, device_ids, include_display_names);
159
160 let master_key = services
161 .users
162 .get_master_key(sender_user, user_id, allowed_signatures)
163 .ok();
164
165 let self_signing_key = services
166 .users
167 .get_self_signing_key(sender_user, user_id, allowed_signatures)
168 .ok();
169
170 let user_signing_key = (Some(user_id) == sender_user)
171 .then_async(|| services.users.get_user_signing_key(user_id).ok())
172 .map(Option::flatten);
173
174 let (device_keys, master_key, self_signing_key, user_signing_key) =
175 join4(device_keys, master_key, self_signing_key, user_signing_key).await;
176
177 let owned = || user_id.to_owned();
178 Keys {
179 device_keys: BTreeMap::from([(owned(), device_keys)]),
180 master_keys: master_key
181 .map(|k| (owned(), k))
182 .into_iter()
183 .collect(),
184
185 self_signing_keys: self_signing_key
186 .map(|k| (owned(), k))
187 .into_iter()
188 .collect(),
189
190 user_signing_keys: user_signing_key
191 .map(|k| (owned(), k))
192 .into_iter()
193 .collect(),
194
195 ..Default::default()
196 }
197}
198
199async fn collect_local_device_keys(
200 services: &Services,
201 user_id: &UserId,
202 device_ids: &[OwnedDeviceId],
203 include_display_names: bool,
204) -> BTreeMap<OwnedDeviceId, Raw<DeviceKeys>> {
205 let stream = if device_ids.is_empty() {
206 services
207 .users
208 .all_device_ids(user_id)
209 .map(ToOwned::to_owned)
210 .left_stream()
211 } else {
212 device_ids.iter().cloned().stream().right_stream()
213 };
214
215 stream
216 .broad_filter_map(async |device_id| {
217 get_local_device_keys(services, user_id, &device_id, include_display_names)
218 .await
219 .map(|keys| (device_id, keys))
220 })
221 .collect()
222 .await
223}
224
225async fn get_local_device_keys(
226 services: &Services,
227 user_id: &UserId,
228 device_id: &DeviceId,
229 include_display_names: bool,
230) -> Option<Raw<DeviceKeys>> {
231 let mut keys = services
232 .users
233 .get_device_keys(user_id, device_id)
234 .await
235 .ok()?;
236
237 let metadata = services
238 .users
239 .get_device_metadata(user_id, device_id)
240 .await
241 .inspect_err(|e| debug_warn!(?user_id, ?device_id, "device metadata missing: {e}"))
242 .ok()?;
243
244 add_unsigned_device_display_name(&mut keys, metadata, include_display_names)
245 .inspect_err(|e| debug_warn!(?user_id, ?device_id, "invalid device keys: {e}"))
246 .ok()?;
247
248 Some(keys)
249}
250
251async fn collect_federation_keys<F>(
252 services: &Services,
253 server: ServerDevices<'_>,
254 sender_user: Option<&UserId>,
255 allowed_signatures: &F,
256) -> Keys
257where
258 F: Fn(&UserId) -> bool + Send + Sync,
259{
260 let requests = server
261 .into_iter()
262 .stream()
263 .map(|(server, device_keys)| (server.to_owned(), FederationRequest { device_keys }));
264
265 let outcomes = fanout_with(
266 requests,
267 async |server, request| {
268 services
269 .federation
270 .execute_keys(&server, request)
271 .await
272 },
273 federation_opts(services),
274 )
275 .broad_then(async |outcome| {
276 process_federation_outcome(services, sender_user, allowed_signatures, outcome).await
277 });
278
279 let (keys, faults) = outcomes.merge(Keys::default(), Keys::merge).await;
280
281 Keys {
282 failures: federation_failures("get_keys", faults).collect(),
283 ..keys
284 }
285}
286
287async fn process_federation_outcome<F>(
288 services: &Services,
289 sender_user: Option<&UserId>,
290 allowed_signatures: &F,
291 outcome: Outcome<FederationResponse>,
292) -> Outcome<Keys>
293where
294 F: Fn(&UserId) -> bool + Send + Sync,
295{
296 let Outcome { origin, elapsed, result } = outcome;
297 let result = match result {
298 | Err(fault) => Err(fault),
299 | Ok(response) => {
300 let keys =
301 process_federation_response(services, sender_user, allowed_signatures, response)
302 .await;
303
304 Ok(keys)
305 },
306 };
307
308 Outcome { origin, elapsed, result }
309}
310
311async fn process_federation_response<F>(
312 services: &Services,
313 sender_user: Option<&UserId>,
314 allowed_signatures: &F,
315 response: FederationResponse,
316) -> Keys
317where
318 F: Fn(&UserId) -> bool + Send + Sync,
319{
320 let FederationResponse {
321 master_keys,
322 self_signing_keys,
323 device_keys,
324 } = response;
325
326 let master_keys = master_keys
327 .into_iter()
328 .stream()
329 .broad_filter_map(async |(user, master_key)| {
330 merge_remote_master_key(services, sender_user, allowed_signatures, &user, master_key)
331 .await
332 .inspect_err(|e| debug_warn!(?user, "skipping master key from federation: {e}"))
333 .map(|raw| (user, raw))
334 .ok()
335 })
336 .collect()
337 .await;
338
339 Keys {
340 device_keys,
341 master_keys,
342 self_signing_keys,
343 user_signing_keys: BTreeMap::new(),
344 failures: BTreeMap::new(),
345 }
346}
347
348async fn merge_remote_master_key<F>(
352 services: &Services,
353 sender_user: Option<&UserId>,
354 allowed_signatures: &F,
355 user: &UserId,
356 master_key_raw: Raw<CrossSigningKey>,
357) -> Result<Raw<CrossSigningKey>>
358where
359 F: Fn(&UserId) -> bool + Send + Sync,
360{
361 let (master_key_id, mut master_key) = parse_master_key(user, &master_key_raw)?;
362 let our_raw = services
363 .users
364 .get_key(&master_key_id, sender_user, user, allowed_signatures)
365 .await;
366
367 if let Ok(our_raw) = our_raw
368 && let Ok((_, mut ours)) = parse_master_key(user, &our_raw)
369 {
370 master_key.signatures.append(&mut ours.signatures);
371 }
372
373 let raw = json::to_raw(&master_key)?;
374
375 services
378 .users
379 .add_cross_signing_keys(user, &Some(raw.clone()), &None, &None, false)
380 .await?;
381
382 Ok(raw)
383}
384
385fn add_unsigned_device_display_name(
386 keys: &mut Raw<DeviceKeys>,
387 metadata: Device,
388 include_display_names: bool,
389) -> Result {
390 let Some(display_name) = metadata.display_name else {
391 return Ok(());
392 };
393
394 let mut object = keys.deserialize_as_unchecked::<CanonicalJsonObject>()?;
395
396 if let CanonicalJsonValue::Object(unsigned) = object
397 .entry("unsigned".into())
398 .or_insert_with(|| CanonicalJsonObject::default().into())
399 {
400 let display_name = if include_display_names {
401 CanonicalJsonValue::String(display_name.to_string())
402 } else {
403 CanonicalJsonValue::String(metadata.device_id.into())
404 };
405
406 unsigned.insert("device_display_name".into(), display_name);
407 }
408
409 *keys = Raw::from_json(to_raw_value(&object)?);
410
411 Ok(())
412}
413
414#[implement(Keys)]
415fn merge(mut self, other: Self) -> Self {
416 self.failures.extend(other.failures);
417 self.device_keys.extend(other.device_keys);
418 self.master_keys.extend(other.master_keys);
419 self.self_signing_keys
420 .extend(other.self_signing_keys);
421 self.user_signing_keys
422 .extend(other.user_signing_keys);
423 self
424}
425
426#[implement(Keys)]
427fn into_response(self) -> get_keys::v3::Response {
428 get_keys::v3::Response {
429 failures: self.failures,
430 device_keys: self.device_keys,
431 master_keys: self.master_keys,
432 self_signing_keys: self.self_signing_keys,
433 user_signing_keys: self.user_signing_keys,
434 }
435}