Skip to main content

tuwunel_service/server_keys/
acquire.rs

1//! Best-effort acquisition of missing federation verify keys.
2//!
3//! Requests are deduplicated by server and key ID, checked against local
4//! storage, and then sent to origins and trusted notaries in configured order.
5//! Missing keys are logged rather than returned to callers as an error.
6
7use std::{
8	borrow::Borrow,
9	collections::{BTreeMap, BTreeSet},
10	time::Duration,
11};
12
13use futures::{StreamExt, stream::FuturesUnordered};
14use ruma::{
15	CanonicalJsonObject, OwnedServerName, OwnedServerSigningKeyId, ServerName,
16	ServerSigningKeyId, api::federation::discovery::ServerSigningKeys, serde::Raw,
17};
18use serde_json::value::RawValue as RawJsonValue;
19use tokio::time::{Instant, timeout_at};
20use tuwunel_core::{
21	debug, debug_error, debug_warn, error, implement, info, result::FlatOk, trace, warn,
22};
23
24use super::key_exists;
25
26type Batch = BTreeMap<OwnedServerName, Vec<OwnedServerSigningKeyId>>;
27
28/// Acquires the signing keys referenced by a collection of raw events.
29///
30/// Signature maps are grouped and deduplicated before acquisition. Events with
31/// malformed JSON or unreadable signature fields are silently skipped.
32#[implement(super::Service)]
33pub async fn acquire_events_pubkeys<'a, I>(&self, events: I)
34where
35	I: Iterator<Item = &'a Box<RawJsonValue>> + Send,
36{
37	type Batch = BTreeMap<OwnedServerName, BTreeSet<OwnedServerSigningKeyId>>;
38	type Signatures = BTreeMap<OwnedServerName, BTreeMap<OwnedServerSigningKeyId, String>>;
39
40	let mut batch = Batch::new();
41	events
42		.cloned()
43		.map(Raw::<CanonicalJsonObject>::from_json)
44		.map(|event| event.get_field::<Signatures>("signatures"))
45		.filter_map(FlatOk::flat_ok)
46		.flat_map(IntoIterator::into_iter)
47		.for_each(|(server, sigs)| {
48			batch
49				.entry(server)
50				.or_default()
51				.extend(sigs.into_keys());
52		});
53
54	let batch = batch
55		.iter()
56		.map(|(server, keys)| (server.borrow(), keys.iter().map(Borrow::borrow)));
57
58	self.acquire_pubkeys(batch).await;
59}
60
61/// Best-effort acquires a batch of server signing keys into the local cache.
62///
63/// Local storage is checked first, then origins and trusted notaries are tried
64/// according to configuration. The method returns no status; keys still missing
65/// after all allowed sources are logged.
66#[implement(super::Service)]
67pub async fn acquire_pubkeys<'a, S, K>(&self, batch: S)
68where
69	S: Iterator<Item = (&'a ServerName, K)> + Send + Clone,
70	K: Iterator<Item = &'a ServerSigningKeyId> + Send + Clone,
71{
72	let notary_only = self
73		.services
74		.config
75		.only_query_trusted_key_servers;
76
77	let notary_first_always = self
78		.services
79		.config
80		.query_trusted_key_servers_first;
81
82	let notary_first_on_join = self
83		.services
84		.config
85		.query_trusted_key_servers_first_on_join;
86
87	let requested_servers = batch.clone().count();
88	let requested_keys = batch
89		.clone()
90		.flat_map(|(_, key_ids)| key_ids)
91		.count();
92
93	debug!("acquire {requested_keys} keys from {requested_servers}");
94
95	let mut missing = self.acquire_locals(batch).await;
96	let mut missing_keys = keys_count(&missing);
97	let mut missing_servers = missing.len();
98	if missing_servers == 0 {
99		return;
100	}
101
102	info!("{missing_keys} keys for {missing_servers} servers will be acquired");
103
104	if notary_first_always || notary_first_on_join {
105		missing = self.acquire_notary(missing.into_iter()).await;
106		missing_keys = keys_count(&missing);
107		missing_servers = missing.len();
108		if missing_keys == 0 {
109			return;
110		}
111
112		warn!(
113			"missing {missing_keys} keys for {missing_servers} servers from all notaries first"
114		);
115	}
116
117	if !notary_only {
118		missing = self.acquire_origins(missing.into_iter()).await;
119		missing_keys = keys_count(&missing);
120		missing_servers = missing.len();
121		if missing_keys == 0 {
122			return;
123		}
124
125		debug_warn!("missing {missing_keys} keys for {missing_servers} servers unreachable");
126	}
127
128	if !notary_first_always && !notary_first_on_join {
129		missing = self.acquire_notary(missing.into_iter()).await;
130		missing_keys = keys_count(&missing);
131		missing_servers = missing.len();
132		if missing_keys == 0 {
133			return;
134		}
135
136		debug_warn!(
137			"still missing {missing_keys} keys for {missing_servers} servers from all notaries."
138		);
139	}
140
141	if missing_keys > 0 {
142		warn!(
143			"did not obtain {missing_keys} keys for {missing_servers} servers out of \
144			 {requested_keys} total keys for {requested_servers} total servers."
145		);
146	}
147
148	for (server, key_ids) in missing {
149		debug_warn!(?server, ?key_ids, "missing");
150	}
151}
152
153#[implement(super::Service)]
154async fn acquire_locals<'a, S, K>(&self, batch: S) -> Batch
155where
156	S: Iterator<Item = (&'a ServerName, K)> + Send,
157	K: Iterator<Item = &'a ServerSigningKeyId> + Send,
158{
159	let mut missing = Batch::new();
160	for (server, key_ids) in batch {
161		for key_id in key_ids {
162			if !self.verify_key_exists(server, key_id).await {
163				missing
164					.entry(server.into())
165					.or_default()
166					.push(key_id.into());
167			}
168		}
169	}
170
171	missing
172}
173
174#[implement(super::Service)]
175async fn acquire_origins<I>(&self, batch: I) -> Batch
176where
177	I: Iterator<Item = (OwnedServerName, Vec<OwnedServerSigningKeyId>)> + Send,
178{
179	let timeout = Instant::now()
180		.checked_add(Duration::from_secs(45))
181		.expect("timeout overflows");
182
183	let mut requests: FuturesUnordered<_> = batch
184		.map(|(origin, key_ids)| self.acquire_origin(origin, key_ids, timeout))
185		.collect();
186
187	let mut missing = Batch::new();
188	while let Some((origin, key_ids)) = requests.next().await {
189		if !key_ids.is_empty() {
190			missing.insert(origin, key_ids);
191		}
192	}
193
194	missing
195}
196
197#[implement(super::Service)]
198async fn acquire_origin(
199	&self,
200	origin: OwnedServerName,
201	mut key_ids: Vec<OwnedServerSigningKeyId>,
202	timeout: Instant,
203) -> (OwnedServerName, Vec<OwnedServerSigningKeyId>) {
204	match timeout_at(timeout, self.server_request(&origin)).await {
205		| Err(e) => debug_warn!(?origin, "timed out: {e}"),
206		| Ok(Err(e)) => debug_error!(?origin, "{e}"),
207		| Ok(Ok(server_keys)) => {
208			trace!(
209				%origin,
210				?key_ids,
211				?server_keys,
212				"received server_keys"
213			);
214
215			self.add_signing_keys(server_keys.clone()).await;
216			key_ids.retain(|key_id| !key_exists(&server_keys, key_id));
217		},
218	}
219
220	(origin, key_ids)
221}
222
223#[implement(super::Service)]
224async fn acquire_notary<I>(&self, batch: I) -> Batch
225where
226	I: Iterator<Item = (OwnedServerName, Vec<OwnedServerSigningKeyId>)> + Send,
227{
228	let mut missing: Batch = batch.collect();
229	for notary in &self.services.config.trusted_servers {
230		let missing_keys = keys_count(&missing);
231		let missing_servers = missing.len();
232		debug!(
233			"Asking notary {notary} for {missing_keys} missing keys from {missing_servers} \
234			 servers"
235		);
236
237		let batch = missing
238			.iter()
239			.map(|(server, keys)| (server.borrow(), keys.iter().map(Borrow::borrow)));
240
241		match self.batch_notary_request(notary, batch).await {
242			| Err(e) => error!("Failed to contact notary {notary:?}: {e}"),
243			| Ok(results) =>
244				for server_keys in results {
245					self.acquire_notary_result(&mut missing, server_keys)
246						.await;
247				},
248		}
249	}
250
251	missing
252}
253
254#[implement(super::Service)]
255async fn acquire_notary_result(&self, missing: &mut Batch, server_keys: ServerSigningKeys) {
256	let server = &server_keys.server_name;
257	self.add_signing_keys(server_keys.clone()).await;
258
259	if let Some(key_ids) = missing.get_mut(server) {
260		key_ids.retain(|key_id| !key_exists(&server_keys, key_id));
261		if key_ids.is_empty() {
262			missing.remove(server);
263		}
264	}
265}
266
267fn keys_count(batch: &Batch) -> usize {
268	batch
269		.values()
270		.flat_map(|key_ids| key_ids.iter())
271		.count()
272}