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