tuwunel_service/server_keys/
request.rs1use std::{collections::BTreeMap, convert::identity, fmt::Debug};
8
9use futures::{FutureExt, StreamExt, TryFutureExt};
10use ruma::{
11 OwnedServerName, OwnedServerSigningKeyId, ServerName, ServerSigningKeyId,
12 api::federation::discovery::{
13 ServerSigningKeys, get_remote_server_keys,
14 get_remote_server_keys_batch::{self, v2::QueryCriteria},
15 get_server_keys,
16 },
17};
18use tuwunel_core::{
19 Err, Result, error, implement, info, trace,
20 utils::stream::{IterStream, ReadyExt, TryBroadbandExt, TryReadyExt},
21};
22
23#[implement(super::Service)]
28pub(super) async fn batch_notary_request<'a, S, K>(
29 &self,
30 notary: &ServerName,
31 batch: S,
32) -> Result<Vec<ServerSigningKeys>>
33where
34 S: Iterator<Item = (&'a ServerName, K)> + Send,
35 K: Iterator<Item = &'a ServerSigningKeyId> + Send,
36{
37 use get_remote_server_keys_batch::v2::Request;
38 type RumaBatch = BTreeMap<OwnedServerName, BTreeMap<OwnedServerSigningKeyId, QueryCriteria>>;
39
40 let criteria = QueryCriteria {
41 minimum_valid_until_ts: Some(self.minimum_valid_ts()),
42 };
43
44 let mut server_keys = batch.fold(RumaBatch::new(), |mut batch, (server, key_ids)| {
45 batch
46 .entry(server.into())
47 .or_default()
48 .extend(key_ids.map(|key_id| (key_id.into(), criteria.clone())));
49
50 batch
51 });
52
53 let total_keys = server_keys
54 .values()
55 .flat_map(|ids| ids.iter())
56 .count();
57
58 debug_assert!(total_keys > 0, "empty batch request to notary");
59
60 let batch_max = self
61 .services
62 .server
63 .config
64 .trusted_server_batch_size;
65
66 let batch_concurrency = self
67 .services
68 .server
69 .config
70 .trusted_server_batch_concurrency;
71
72 let batches: Vec<_> = server_keys
73 .keys()
74 .rev()
75 .step_by(batch_max.saturating_sub(1))
76 .skip(1)
77 .chain(server_keys.keys().next())
78 .cloned()
79 .collect();
80
81 batches
82 .iter()
83 .stream()
84 .enumerate()
85 .map(|(i, batch)| {
86 let request = Request {
87 server_keys: server_keys.split_off(batch),
88 };
89
90 if request.server_keys.is_empty() {
91 return None;
92 }
93
94 trace!(
95 %i, %notary, ?batch,
96 remaining = ?server_keys,
97 requesting = ?request.server_keys.keys(),
98 "Request to notary server."
99 );
100
101 info!(
102 %notary,
103 remaining = %server_keys.len(),
104 requesting = %request.server_keys.len(),
105 "Sending request to notary server..."
106 );
107
108 Some(Ok(request))
109 })
110 .ready_filter_map(identity)
111 .broadn_and_then(batch_concurrency, |request| {
112 self.services
113 .federation
114 .execute_synapse(notary, request)
115 })
116 .ready_try_fold(Vec::new(), |mut results, response| {
117 let response = response
118 .server_keys
119 .into_iter()
120 .map(|key| key.deserialize())
121 .filter_map(Result::ok);
122
123 trace!(
124 %notary, ?response,
125 "Response from notary server."
126 );
127
128 results.extend(response);
129
130 info!(
131 "Received {0} keys out of {1} from notary server so far...",
132 results.len(),
133 total_keys,
134 );
135
136 Ok(results)
137 })
138 .inspect_err(|e| {
139 error!(
140 ?notary, %batch_max, %batch_concurrency, %total_keys,
141 "Requesting keys from notary server failed: {e}",
142 );
143 })
144 .boxed()
145 .await
146}
147
148#[implement(super::Service)]
153pub async fn notary_request(
154 &self,
155 notary: &ServerName,
156 target: &ServerName,
157) -> Result<impl Iterator<Item = ServerSigningKeys> + Clone + Debug + Send + use<>> {
158 use get_remote_server_keys::v2::Request;
159
160 let request = Request {
161 server_name: target.into(),
162 minimum_valid_until_ts: self.minimum_valid_ts(),
163 };
164
165 let response = self
166 .services
167 .federation
168 .execute(notary, request)
169 .await?
170 .server_keys
171 .into_iter()
172 .map(|key| key.deserialize())
173 .filter_map(Result::ok);
174
175 Ok(response)
176}
177
178#[implement(super::Service)]
183pub async fn server_request(&self, target: &ServerName) -> Result<ServerSigningKeys> {
184 use get_server_keys::v2::Request;
185
186 let server_signing_key = self
187 .services
188 .federation
189 .execute(target, Request::new())
190 .await
191 .map(|response| response.server_key)
192 .and_then(|key| key.deserialize().map_err(Into::into))?;
193
194 if server_signing_key.server_name != target {
195 return Err!(BadServerResponse(debug_warn!(
196 requested = ?target,
197 response = ?server_signing_key.server_name,
198 "Server responded with bogus server_name"
199 )));
200 }
201
202 Ok(server_signing_key)
203}