Skip to main content

tuwunel_api/client/keys/
claim_keys.rs

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
37/// # `POST /_matrix/client/r0/keys/claim`
38///
39/// Claims one-time keys
40pub(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			// MSC3983: claim from appservices before marking local fallback keys used.
109			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			// Omit a depleted user entirely; Synapse returns no entry, not an empty map.
138			(!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}