Skip to main content

tuwunel_api/client/keys/
get_keys.rs

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
54/// # `POST /_matrix/client/r0/keys/query`
55///
56/// Get end-to-end encryption keys for the given users.
57///
58/// - Always fetches users from other servers over federation
59/// - Gets master keys, self-signing keys, user signing keys and device keys.
60/// - The master and self-signing keys contain signatures that the user is
61///   allowed to see
62pub(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, // Always allow local users to see device names of other local users
74	)
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
348/// Merges signatures from our cached copy of the user's master key (if any)
349/// onto the remote-supplied master key, persists the merged copy to our
350/// database, and returns the merged Raw value for the response.
351async 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	// Don't notify: a notification would trigger another key request resulting
376	// in an endless loop.
377	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}