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::{
8		client::keys::claim_keys,
9		federation::keys::claim_keys::v1::{
10			Request as FederationRequest, Response as FederationResponse,
11		},
12	},
13	encryption::OneTimeKey,
14	serde::Raw,
15};
16use tuwunel_core::{
17	Result,
18	utils::{IterStream, stream::BroadbandExt},
19};
20use tuwunel_service::{
21	Services,
22	federation::feds::{OutcomeExt, fanout_with},
23};
24
25use super::{FailureMap, federation_failures, federation_opts};
26use crate::Ruma;
27
28#[derive(Default)]
29struct Claims {
30	one_time_keys: OneTimeKeyMap,
31	failures: FailureMap,
32}
33
34type RequestClaims = BTreeMap<OwnedUserId, Algorithms>;
35type ServerClaims<'a> = BTreeMap<&'a ServerName, RequestClaims>;
36type LocalClaim<'a> = (&'a UserId, &'a Algorithms);
37type Algorithms = BTreeMap<OwnedDeviceId, OneTimeKeyAlgorithm>;
38type OneTimeKeys = BTreeMap<OwnedOneTimeKeyId, Raw<OneTimeKey>>;
39type OneTimeKeyMap = BTreeMap<OwnedUserId, BTreeMap<OwnedDeviceId, OneTimeKeys>>;
40
41/// # `POST /_matrix/client/r0/keys/claim`
42///
43/// Claims one-time keys
44pub(crate) async fn claim_keys_route(
45	State(services): State<crate::State>,
46	body: Ruma<claim_keys::v3::Request>,
47) -> Result<claim_keys::v3::Response> {
48	claim_keys_helper(&services, &body.one_time_keys).await
49}
50
51pub(crate) async fn claim_keys_helper(
52	services: &Services,
53	one_time_keys_input: &RequestClaims,
54) -> Result<claim_keys::v3::Response> {
55	let (local_users, remote_users): (Vec<_>, Vec<_>) = one_time_keys_input
56		.iter()
57		.map(|(uid, map)| (uid.as_ref(), map))
58		.partition(|(user_id, _)| services.globals.user_is_local(user_id));
59
60	let server: ServerClaims<'_> =
61		remote_users
62			.into_iter()
63			.fold(BTreeMap::new(), |mut acc, (user_id, map)| {
64				acc.entry(user_id.server_name())
65					.or_default()
66					.insert(user_id.to_owned(), map.clone());
67				acc
68			});
69
70	let local = collect_local_one_time_keys(services, &local_users);
71	let federation = collect_federation_one_time_keys(services, server);
72
73	let (local, federation) = join(local, federation).await;
74	let merged = local.merge(federation);
75
76	Ok(claim_keys::v3::Response {
77		failures: merged.failures,
78		one_time_keys: merged.one_time_keys,
79	})
80}
81
82async fn collect_local_one_time_keys(services: &Services, users: &[LocalClaim<'_>]) -> Claims {
83	let take_one_time_key = async |(user_id, device_id, algorithm)| {
84		let key = services
85			.users
86			.take_one_time_key(user_id, device_id, algorithm)
87			.await
88			.ok();
89
90		// MSC2732: serve the fallback key when the one-time pool is empty.
91		let key = match key {
92			| Some(key) => Some(key),
93			| None => services
94				.users
95				.take_fallback_key(user_id, device_id, algorithm)
96				.await
97				.ok(),
98		};
99
100		key.map(|key| (device_id.to_owned(), [key].into()))
101	};
102
103	let one_time_keys = users
104		.iter()
105		.copied()
106		.stream()
107		.broad_filter_map(async |(user_id, requested)| {
108			let device_keys: BTreeMap<_, _> = requested
109				.iter()
110				.stream()
111				.map(|(device_id, algorithm)| (user_id, device_id.as_ref(), algorithm))
112				.filter_map(take_one_time_key)
113				.collect()
114				.await;
115
116			// Omit a depleted user entirely; Synapse returns no entry, not an empty map.
117			(!device_keys.is_empty()).then(|| (user_id.to_owned(), device_keys))
118		})
119		.collect()
120		.await;
121
122	Claims { one_time_keys, ..Default::default() }
123}
124
125async fn collect_federation_one_time_keys(
126	services: &Services,
127	server: ServerClaims<'_>,
128) -> Claims {
129	let requests = server
130		.into_iter()
131		.stream()
132		.map(|(server, one_time_keys)| (server.to_owned(), FederationRequest { one_time_keys }));
133
134	let outcomes = fanout_with(
135		requests,
136		async |server, request| {
137			services
138				.federation
139				.execute_keys(&server, request)
140				.await
141		},
142		federation_opts(services),
143	);
144
145	let (claims, faults) = outcomes
146		.merge(Claims::default(), |mut claims, response: FederationResponse| {
147			claims
148				.one_time_keys
149				.extend(response.one_time_keys);
150
151			claims
152		})
153		.await;
154
155	Claims {
156		failures: federation_failures("claim_keys", faults).collect(),
157		..claims
158	}
159}
160
161impl Claims {
162	fn merge(mut self, other: Self) -> Self {
163		self.one_time_keys.extend(other.one_time_keys);
164		self.failures.extend(other.failures);
165		self
166	}
167}