Skip to main content

tuwunel_service/rooms/auth_chain/
mod.rs

1//! Resolves and caches transitive authorization-event ancestry.
2//!
3//! Traversals operate on short event IDs, enforce room membership on cache
4//! misses, and store complete chains for repeated starting sets. Public stream
5//! adapters translate the results back to event IDs with tolerant or explicit
6//! completeness reporting.
7
8use std::{
9	collections::{BTreeSet, HashSet},
10	iter::once,
11	sync::{
12		Arc,
13		atomic::{AtomicBool, Ordering},
14	},
15	time::Instant,
16};
17
18use async_trait::async_trait;
19use futures::{
20	FutureExt, Stream, StreamExt, TryFutureExt, pin_mut,
21	stream::{FuturesUnordered, unfold},
22};
23use ruma::{
24	EventId, OwnedEventId, OwnedRoomId, RoomId, RoomVersionId,
25	room_version_rules::RoomVersionRules,
26};
27use serde::Deserialize;
28use tuwunel_core::{
29	Err, Result, at, debug, debug_error, err, implement,
30	itertools::Itertools,
31	matrix::room_version,
32	pdu::AuthEvents,
33	smallvec::SmallVec,
34	trace,
35	utils::{
36		BoolExt, IterStream,
37		stream::{BroadbandExt, ReadyExt, TryExpect, automatic_width},
38	},
39	validated, warn,
40};
41use tuwunel_database::Map;
42
43use crate::rooms::short::ShortEventId;
44
45/// Computes transitive authorization chains and maintains their persistent cache.
46///
47/// Starting events are distributed across a fixed number of buckets before
48/// traversal. Only complete cache-miss walks are written. Existing cache hits
49/// are trusted without rechecking completeness, so historical or
50/// downgrade-written rows can still supply partial chains until they expire or
51/// are cleared. Combined bucket results are sorted and deduplicated by short
52/// event ID rather than graph order.
53pub struct Service {
54	services: Arc<crate::services::OnceServices>,
55	db: Data,
56}
57
58struct Data {
59	authchainkey_authchain: Arc<Map>,
60}
61
62type Bucket<'a> = BTreeSet<(ShortEventId, &'a EventId)>;
63type CacheKey = SmallVec<[ShortEventId; 1]>;
64
65#[async_trait]
66impl crate::Service for Service {
67	fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
68		Ok(Arc::new(Self {
69			services: args.services.clone(),
70			db: Data {
71				authchainkey_authchain: args.db["authchainkey_authchain"].clone(),
72			},
73		}))
74	}
75
76	async fn clear_cache(&self) { self.db.authchainkey_authchain.clear().await; }
77
78	fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
79}
80
81/// Streams the transitive auth ancestors of a set of starting events.
82///
83/// The starting events themselves are not included unless encountered as
84/// ancestors. Traversal and reverse-mapping failures are filtered from this
85/// tolerant interface, so the stream can yield a partial chain without error.
86#[implement(Service)]
87pub fn event_ids_iter<'a, I>(
88	&'a self,
89	room_id: &'a RoomId,
90	room_version: &'a RoomVersionId,
91	starting_events: I,
92) -> impl Stream<Item = Result<OwnedEventId>> + Send + 'a
93where
94	I: Iterator<Item = &'a EventId> + Clone + ExactSizeIterator + Send + 'a,
95{
96	self.get_auth_chain(room_id, room_version, starting_events)
97		.map_ok(|chain| {
98			self.services
99				.short
100				.multi_get_eventid_from_short(chain.into_iter().stream())
101				.ready_filter(Result::is_ok)
102		})
103		.try_flatten_stream()
104}
105
106/// Streams an auth chain and reports whether its observed inputs were complete.
107///
108/// The caller initializes `complete` to `true` and checks it after consuming
109/// the stream. Cached chains leave it unchanged. Any polled reverse-mapping or
110/// ancestor-walk failure sets it to `false`, so the stream must be fully
111/// consumed before the flag is inspected.
112#[implement(Service)]
113pub fn event_ids_iter_strict<'a, I>(
114	&'a self,
115	room_id: &'a RoomId,
116	room_version: &'a RoomVersionId,
117	starting_events: I,
118	complete: &'a AtomicBool,
119) -> impl Stream<Item = Result<OwnedEventId>> + Send + 'a
120where
121	I: Iterator<Item = &'a EventId> + Clone + ExactSizeIterator + Send + 'a,
122{
123	self.get_auth_chain_strict(room_id, room_version, starting_events, complete)
124		.map_ok(move |chain| {
125			self.services
126				.short
127				.multi_get_eventid_from_short(chain.into_iter().stream())
128				.inspect(move |result| {
129					if result.is_err() {
130						complete.store(false, Ordering::Relaxed);
131					}
132				})
133				.ready_filter(Result::is_ok)
134		})
135		.try_flatten_stream()
136}
137
138#[implement(Service)]
139#[tracing::instrument(
140	name = "auth_chain",
141	level = "debug",
142	skip_all,
143	fields(
144		%room_id,
145		starting_events = %starting_events.clone().count(),
146	)
147)]
148/// Returns sorted short IDs for the transitive auth ancestors of starting events.
149///
150/// Starting events are excluded unless reached again as ancestors, and missing
151/// ancestry can produce a partial successful result. Complete per-bucket
152/// results are cached under roomless keys derived from each bucket's exact set
153/// of starting short IDs; uncached traversal still rejects events from another
154/// room.
155pub async fn get_auth_chain<'a, I>(
156	&'a self,
157	room_id: &RoomId,
158	room_version: &RoomVersionId,
159	starting_events: I,
160) -> Result<Vec<ShortEventId>>
161where
162	I: Iterator<Item = &'a EventId> + Clone + ExactSizeIterator + Send + 'a,
163{
164	let complete = AtomicBool::new(true);
165
166	self.get_auth_chain_inner(room_id, room_version, starting_events, &complete)
167		.await
168}
169
170#[implement(Service)]
171#[tracing::instrument(
172	name = "auth_chain",
173	level = "debug",
174	skip_all,
175	fields(
176		%room_id,
177		starting_events = %starting_events.clone().count(),
178	)
179)]
180async fn get_auth_chain_strict<'a, I>(
181	&'a self,
182	room_id: &RoomId,
183	room_version: &RoomVersionId,
184	starting_events: I,
185	complete: &AtomicBool,
186) -> Result<Vec<ShortEventId>>
187where
188	I: Iterator<Item = &'a EventId> + Clone + ExactSizeIterator + Send + 'a,
189{
190	self.get_auth_chain_inner(room_id, room_version, starting_events, complete)
191		.inspect_err(|_| complete.store(false, Ordering::Relaxed))
192		.await
193}
194
195#[implement(Service)]
196async fn get_auth_chain_inner<'a, I>(
197	&'a self,
198	room_id: &RoomId,
199	room_version: &RoomVersionId,
200	starting_events: I,
201	complete: &AtomicBool,
202) -> Result<Vec<ShortEventId>>
203where
204	I: Iterator<Item = &'a EventId> + Clone + ExactSizeIterator + Send + 'a,
205{
206	const NUM_BUCKETS: usize = 50; //TODO: change possible w/o disrupting db?
207	const BUCKET: Bucket<'_> = BTreeSet::new();
208
209	let started = Instant::now();
210	let room_rules = room_version::rules(room_version)?;
211	let starting_events_count = starting_events.clone().count();
212	let starting_ids = self
213		.services
214		.short
215		.multi_get_or_create_shorteventid(starting_events.clone())
216		.zip(starting_events.stream());
217
218	pin_mut!(starting_ids);
219	let mut buckets = [BUCKET; NUM_BUCKETS];
220	while let Some((short, starting_event)) = starting_ids.next().await {
221		let bucket: usize = short.try_into()?;
222		let bucket: usize = validated!(bucket % NUM_BUCKETS);
223		buckets[bucket].insert((short, starting_event));
224	}
225
226	debug!(
227		starting_events = starting_events_count,
228		elapsed = ?started.elapsed(),
229		"start",
230	);
231
232	let full_auth_chain: Vec<ShortEventId> = buckets
233		.iter()
234		.stream()
235		.flat_map_unordered(automatic_width(), |starting_events| {
236			self.get_chunk_auth_chain(
237				room_id,
238				&started,
239				starting_events.iter().copied(),
240				&room_rules,
241				complete,
242			)
243			.boxed() // Unpin for flat_map_unordered
244		})
245		.collect::<Vec<_>>()
246		.map(IntoIterator::into_iter)
247		.map(Itertools::sorted_unstable)
248		.map(Itertools::dedup)
249		.map(Iterator::collect)
250		.boxed() // erase region
251		.await;
252
253	debug!(
254		chain_length = ?full_auth_chain.len(),
255		elapsed = ?started.elapsed(),
256		"done",
257	);
258
259	Ok(full_auth_chain)
260}
261
262#[implement(Service)]
263#[tracing::instrument(
264	name = "outer",
265	level = "trace",
266	skip_all,
267	fields(
268		starting_events = %starting_events.clone().count(),
269	)
270)]
271fn get_chunk_auth_chain<'a, I>(
272	&'a self,
273	room_id: &'a RoomId,
274	started: &'a Instant,
275	starting_events: I,
276	room_rules: &'a RoomVersionRules,
277	complete: &'a AtomicBool,
278) -> impl Stream<Item = ShortEventId> + Send + 'a
279where
280	I: Iterator<Item = (ShortEventId, &'a EventId)> + Clone + Send + Sync + 'a,
281{
282	self.get_cached_auth_chain(starting_events.clone().map(at!(0)))
283		.map_ok(IntoIterator::into_iter)
284		.map_ok(IterStream::try_stream)
285		.or_else(async move |_| {
286			let chain = self
287				.build_chunk_auth_chain(room_id, started, starting_events, room_rules, complete)
288				.await;
289
290			Ok(chain.into_iter().try_stream())
291		})
292		.try_flatten_stream()
293		.map_expect("either cache hit or cache miss yields a chain")
294}
295
296#[implement(Service)]
297async fn build_chunk_auth_chain<'a, I>(
298	&'a self,
299	room_id: &'a RoomId,
300	started: &'a Instant,
301	starting_events: I,
302	room_rules: &'a RoomVersionRules,
303	complete: &'a AtomicBool,
304) -> Vec<ShortEventId>
305where
306	I: Iterator<Item = (ShortEventId, &'a EventId)> + Clone + Send + Sync + 'a,
307{
308	let chunk_complete = AtomicBool::new(true);
309
310	let build_chain = async |(shortid, event_id): (ShortEventId, &'a EventId)| {
311		if let Ok(cached) = self.get_cached_auth_chain(once(shortid)).await {
312			return cached;
313		}
314
315		let event_complete = AtomicBool::new(true);
316		let auth_chain: Vec<_> = self
317			.get_event_auth_chain(room_id, event_id, room_rules, &event_complete)
318			.collect()
319			.await;
320
321		match event_complete.load(Ordering::Relaxed) {
322			| true => self.put_cached_auth_chain(once(shortid), auth_chain.as_slice()),
323			| false => {
324				chunk_complete.store(false, Ordering::Relaxed);
325				complete.store(false, Ordering::Relaxed);
326			},
327		}
328
329		debug!(
330			?event_id,
331			elapsed = ?started.elapsed(),
332			"Cache missed event"
333		);
334
335		auth_chain
336	};
337
338	let chunk_chain: Vec<_> = starting_events
339		.clone()
340		.stream()
341		.broad_then(build_chain)
342		.collect::<Vec<_>>()
343		.map(IntoIterator::into_iter)
344		.map(Iterator::flatten)
345		.map(Itertools::sorted_unstable)
346		.map(Itertools::dedup)
347		.map(Iterator::collect)
348		.await;
349
350	match chunk_complete.load(Ordering::Relaxed) {
351		| false => debug!(
352			elapsed = ?started.elapsed(),
353			"Incomplete chunk not cached",
354		),
355		| true => {
356			self.put_cached_auth_chain(starting_events.map(at!(0)), chunk_chain.as_slice());
357
358			debug!(
359				chunk_chain_length = ?chunk_chain.len(),
360				elapsed = ?started.elapsed(),
361				"Cache missed chunk",
362			);
363		},
364	}
365
366	chunk_chain
367}
368
369#[implement(Service)]
370#[tracing::instrument(name = "inner", level = "trace", skip_all)]
371fn get_event_auth_chain<'a>(
372	&'a self,
373	room_id: &'a RoomId,
374	event_id: &'a EventId,
375	room_rules: &'a RoomVersionRules,
376	complete: &'a AtomicBool,
377) -> impl Stream<Item = ShortEventId> + Send + 'a {
378	self.get_event_auth_chain_ids(room_id, event_id, room_rules, complete)
379		.broad_then(async move |auth_event| {
380			self.services
381				.short
382				.get_or_create_shorteventid(&auth_event)
383				.await
384		})
385}
386
387#[implement(Service)]
388#[tracing::instrument(
389	name = "inner_ids",
390	level = "trace",
391	skip_all,
392	fields(%event_id)
393)]
394fn get_event_auth_chain_ids<'a>(
395	&'a self,
396	room_id: &'a RoomId,
397	event_id: &'a EventId,
398	room_rules: &'a RoomVersionRules,
399	complete: &'a AtomicBool,
400) -> impl Stream<Item = OwnedEventId> + Send + 'a {
401	struct State<Fut> {
402		todo: FuturesUnordered<Fut>,
403		seen: HashSet<OwnedEventId>,
404		implied: Option<OwnedEventId>,
405	}
406
407	// MSC4291 rooms imply the create event rather than naming it in auth_events.
408	let create_event_id = room_rules
409		.authorization
410		.room_create_event_id_as_room_id
411		.and_then(|| room_id.as_event_id().ok())
412		.filter(|create_event_id| create_event_id.ne(&event_id));
413
414	let starting_events = self.get_event_auth_event_ids(room_id, event_id.to_owned());
415
416	let state = State {
417		todo: once(starting_events).collect(),
418		seen: create_event_id.iter().cloned().collect(),
419		implied: create_event_id,
420	};
421
422	let eval = |auth_events: AuthEvents, mut state: State<_>| {
423		let implied = state.implied.take();
424		let push = |auth_event: &OwnedEventId| {
425			trace!(todo = state.todo.len(), ?auth_event, "push");
426			state
427				.todo
428				.push(self.get_event_auth_event_ids(room_id, auth_event.clone()));
429		};
430
431		let seen = |auth_event: OwnedEventId| {
432			state
433				.seen
434				.insert(auth_event.clone())
435				.then_some(auth_event)
436		};
437
438		let unseen = auth_events
439			.into_iter()
440			.filter_map(seen)
441			.inspect(push);
442
443		let out = implied
444			.into_iter()
445			.chain(unseen)
446			.collect::<AuthEvents>()
447			.into_iter()
448			.stream();
449
450		(out, state)
451	};
452
453	// unfold requires FnMut returning a nameable future; an async closure
454	// capturing eval does not satisfy it.
455	#[expect(closure_returning_async_block)]
456	unfold(state, move |mut state| async move {
457		match state.todo.next().await {
458			| None => None,
459			| Some(Ok(auth_events)) => Some(eval(auth_events, state)),
460			| Some(Err(e)) => {
461				complete.store(false, Ordering::Relaxed);
462
463				// A missing ancestor is a normal backfill gap and is skipped;
464				// any other error is corrupt data and ends the walk.
465				e.is_not_found()
466					.then(move || (AuthEvents::new().into_iter().stream(), state))
467			},
468		}
469	})
470	.flatten()
471}
472
473#[implement(Service)]
474#[tracing::instrument(
475	name = "cache_put",
476	level = "debug",
477	skip_all,
478	fields(
479		key_len = key.clone().count(),
480		chain_len = auth_chain.len(),
481	)
482)]
483fn put_cached_auth_chain<I>(&self, key: I, auth_chain: &[ShortEventId])
484where
485	I: Iterator<Item = ShortEventId> + Clone + Send,
486{
487	let key = key.collect::<CacheKey>();
488
489	debug_assert!(!key.is_empty(), "auth_chain key must not be empty");
490
491	self.db
492		.authchainkey_authchain
493		.put(key.as_slice(), auth_chain);
494}
495
496#[implement(Service)]
497#[tracing::instrument(
498	name = "cache_get",
499	level = "trace",
500	err(level = "trace"),
501	skip_all,
502	fields(
503		key_len = %key.clone().count()
504	),
505)]
506async fn get_cached_auth_chain<I>(&self, key: I) -> Result<Vec<ShortEventId>>
507where
508	I: Iterator<Item = ShortEventId> + Clone + Send,
509{
510	let key = key.collect::<CacheKey>();
511
512	if key.is_empty() {
513		return Ok(Vec::new());
514	}
515
516	let chain = self
517		.db
518		.authchainkey_authchain
519		.qry(key.as_slice())
520		.map_err(|_| err!(Request(NotFound("auth_chain not cached"))))
521		.await?;
522
523	if !chain.len().is_multiple_of(size_of::<u64>()) {
524		return Err!(Request(NotFound("malformed auth_chain cache")));
525	}
526
527	let chain = chain
528		.as_chunks::<{ size_of::<u64>() }>()
529		.0
530		.iter()
531		.copied()
532		.map(u64::from_be_bytes)
533		.collect();
534
535	Ok(chain)
536}
537
538#[implement(Service)]
539#[tracing::instrument(
540	name = "auth_events",
541	level = "trace",
542	ret(level = "trace"),
543	err(level = "trace"),
544	skip_all,
545	fields(%event_id)
546)]
547async fn get_event_auth_event_ids<'a>(
548	&'a self,
549	room_id: &'a RoomId,
550	event_id: OwnedEventId,
551) -> Result<AuthEvents> {
552	#[derive(Deserialize)]
553	struct Pdu {
554		auth_events: AuthEvents,
555		room_id: OwnedRoomId,
556	}
557
558	let pdu: Pdu = self
559		.services
560		.timeline
561		.get(&event_id)
562		.inspect_err(|e| {
563			debug_error!(?event_id, ?room_id, "auth chain event: {e}");
564		})
565		.await?;
566
567	if pdu.room_id != room_id {
568		return Err!(Request(Forbidden(error!(
569			?event_id,
570			?room_id,
571			wrong_room_id = ?pdu.room_id,
572			"auth event for incorrect room",
573		))));
574	}
575
576	Ok(pdu.auth_events)
577}