1use 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
45pub 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#[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#[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)]
148pub 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; 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() })
245 .collect::<Vec<_>>()
246 .map(IntoIterator::into_iter)
247 .map(Itertools::sorted_unstable)
248 .map(Itertools::dedup)
249 .map(Iterator::collect)
250 .boxed() .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 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 #[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 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}