tuwunel_api/client/keys/
claim_keys.rs1use 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
41pub(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 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 (!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}