1use std::{borrow::Borrow, collections::HashMap, iter::once, sync::Arc, time::Instant};
2
3use futures::{FutureExt, StreamExt, TryFutureExt};
4use ruma::{
5 CanonicalJsonObject, EventId, OwnedEventId, RoomId, RoomVersionId, ServerName,
6 events::StateEventType, room_version_rules::RoomVersionRules,
7};
8use tuwunel_core::{
9 Result, debug, debug_info, err, implement, is_equal_to,
10 matrix::{Event, EventTypeExt, PduEvent, StateKey, pdu::check_rules, room_version},
11 trace,
12 utils::{
13 BoolExt,
14 stream::{BroadbandExt, ReadyExt},
15 },
16 warn,
17};
18
19use super::{
20 backoff::{Context, Disposition, UPGRADE_RETRY},
21 policy_server::PolicyCheck,
22 state_local_build::{WalkMode, compare_shadow},
23};
24use crate::rooms::{
25 state::{RoomMutexGuard, Trigger, prune_goal},
26 state_compressor::{CompressedState, HashSetCompressStateEvent},
27 state_res::auth_check,
28 timeline::RawPduId,
29};
30
31#[cfg(test)]
32mod tests;
33
34#[derive(Clone, Copy)]
37enum ResolvedVia {
38 Derived,
39 Memo,
40 Local,
41 Fetch,
42}
43
44#[derive(Clone, Copy)]
46enum Standing {
47 Evaluate(bool),
50
51 Withheld,
53}
54
55#[implement(super::Service)]
56#[tracing::instrument(
57 name = "upgrade",
58 level = "debug",
59 ret(level = "debug"),
60 skip_all,
61 fields(lev = %recursion_level)
62)]
63#[expect(clippy::too_many_arguments)]
64pub(super) async fn upgrade_outlier_to_timeline_pdu(
65 &self,
66 origin: &ServerName,
67 room_id: &RoomId,
68 incoming_pdu: PduEvent,
69 mut pdu_json: CanonicalJsonObject,
70 room_version: &RoomVersionId,
71 recursion_level: usize,
72 create_event_id: &EventId,
73) -> Result<Option<(RawPduId, bool)>> {
74 if let Ok(pdu_id) = self
76 .services
77 .timeline
78 .get_pdu_id(incoming_pdu.event_id())
79 .await
80 {
81 debug!(?pdu_id, "Exists.");
82 return Ok(Some((pdu_id, false)));
83 }
84
85 trace!("Upgrading to timeline pdu");
86
87 let timer = Instant::now();
88 let room_rules = room_version::rules(room_version)?;
89
90 trace!(format = ?room_rules.event_format, "Checking format");
91 check_rules(&pdu_json, &room_rules.event_format)?;
92
93 let Standing::Evaluate(cleared) = self
94 .soft_fail_standing(&incoming_pdu, &room_rules, &mut pdu_json)
95 .await?
96 else {
97 return Ok(None);
98 };
99
100 let (state_at_incoming_event, resolved_via) = self
101 .resolve_state_at_incoming_event(
102 origin,
103 room_id,
104 &incoming_pdu,
105 room_version,
106 recursion_level,
107 create_event_id,
108 )
109 .await?;
110
111 self.auth_check_outlier_pdu(room_id, &incoming_pdu, &room_rules, &state_at_incoming_event)
112 .await?;
113
114 let soft_fail_pre = !cleared
115 && self
116 .compute_soft_fail(&incoming_pdu, &room_rules, &mut pdu_json)
117 .await?;
118
119 trace!("Locking the room");
122 let state_lock = self.services.state.mutex.lock(room_id).await;
123
124 let soft_fail_current_state = !self
126 .current_state_auth_passes(room_id, &incoming_pdu, &room_rules)
127 .await;
128
129 let soft_fail = soft_fail_pre || soft_fail_current_state;
130
131 let mut extremities = self
132 .compute_remaining_extremities(room_id, &incoming_pdu)
133 .await;
134
135 let config = &self.services.server.config;
136 let max = config.forward_extremities_max;
137 let len = extremities
138 .len()
139 .saturating_add(usize::from(!soft_fail));
140
141 if max > 0 && len > max {
142 let goal = prune_goal(
143 len,
144 max,
145 config.forward_extremities_emergency_max,
146 config.forward_extremities_prune_batch,
147 );
148
149 let summary = self
150 .services
151 .state
152 .prune_forward_extremities(room_id, &mut extremities, goal, Trigger::Receive)
153 .await;
154
155 debug!(?summary, "Pruned forward extremities over the cap.");
156 }
157
158 trace!("Compressing state...");
159 let state_ids_compressed: Arc<CompressedState> = self
160 .services
161 .state_compressor
162 .compress_state_events(
163 state_at_incoming_event
164 .iter()
165 .map(|(ssk, eid)| (ssk, eid.borrow())),
166 )
167 .collect()
168 .map(Arc::new)
169 .await;
170
171 if matches!(resolved_via, ResolvedVia::Local | ResolvedVia::Fetch) && !soft_fail {
172 self.cache_resolved_state(room_id, incoming_pdu.event_id(), state_ids_compressed.clone())
173 .await;
174 }
175
176 if incoming_pdu.state_key().is_some() && !soft_fail {
179 self.resolve_and_force_state_after(
180 room_id,
181 room_version,
182 &incoming_pdu,
183 &state_at_incoming_event,
184 &state_lock,
185 )
186 .boxed()
187 .await?;
188 }
189
190 trace!("Appending pdu to timeline");
193
194 let incoming_extremity = once(incoming_pdu.event_id()).filter(|_| !soft_fail);
196
197 let extremities = extremities
198 .iter()
199 .map(Borrow::borrow)
200 .chain(incoming_extremity);
201
202 let pdu_id = self
203 .services
204 .timeline
205 .append_incoming_pdu(
206 &incoming_pdu,
207 pdu_json,
208 extremities,
209 state_ids_compressed,
210 soft_fail,
211 &state_lock,
212 )
213 .await?;
214
215 debug_assert!(
216 pdu_id.is_some() || soft_fail,
217 "Ok(None) returned by timeline for soft-failed PDU's"
218 );
219
220 if soft_fail {
221 self.services
222 .pdu_metadata
223 .mark_event_soft_failed(incoming_pdu.event_id());
224
225 self.record_outcome(Context::Upgrade, incoming_pdu.event_id(), Disposition::Transient);
226
227 drop(state_lock);
228 warn!(
229 event_id = %incoming_pdu.event_id(),
230 redact_or_policy = soft_fail_pre,
231 current_state = soft_fail_current_state,
232 elapsed = ?timer.elapsed(),
233 "Event was soft failed.",
234 );
235
236 return Ok(None);
237 }
238
239 drop(state_lock);
240
241 if cleared {
242 self.services
243 .pdu_metadata
244 .clear_event_soft_failed(incoming_pdu.event_id());
245
246 self.record_success(Context::Upgrade, incoming_pdu.event_id())
247 .await;
248 }
249
250 debug_info!(
251 elapsed = ?timer.elapsed(),
252 "Accepted",
253 );
254
255 Ok(pdu_id.zip(Some(true)))
256}
257
258#[implement(super::Service)]
259#[tracing::instrument(
260 name = "auth_current",
261 level = "debug",
262 skip_all,
263 fields(%room_id, event_id = %incoming_pdu.event_id()),
264)]
265async fn current_state_auth_passes(
266 &self,
267 room_id: &RoomId,
268 incoming_pdu: &PduEvent,
269 room_rules: &RoomVersionRules,
270) -> bool {
271 trace!("Gathering current-state auth events.");
272 let Ok(auth_events) = self
273 .services
274 .state
275 .get_auth_events(
276 room_id,
277 incoming_pdu.kind(),
278 incoming_pdu.sender(),
279 incoming_pdu.state_key(),
280 incoming_pdu.content(),
281 &room_rules.authorization,
282 true,
283 )
284 .await
285 .inspect_err(|error| {
286 warn!(
287 auth_leg = "current_state",
288 event_id = %incoming_pdu.event_id(),
289 %room_id,
290 %error,
291 "Current-state auth event lookup failed; soft-failing event.",
292 );
293 })
294 else {
295 return false;
296 };
297
298 let state_fetch = async |k: StateEventType, s: StateKey| {
299 auth_events
300 .get(&k.with_state_key(s.as_str()))
301 .map(ToOwned::to_owned)
302 .ok_or_else(|| err!(Request(NotFound("state event not found"))))
303 };
304
305 let event_fetch = async |event_id: OwnedEventId| self.event_fetch(&event_id).await;
306
307 trace!("Performing current-state auth check.");
308 auth_check(room_rules, incoming_pdu, &event_fetch, &state_fetch)
309 .await
310 .inspect_err(|error| {
311 warn!(
312 auth_leg = "current_state",
313 event_id = %incoming_pdu.event_id(),
314 %room_id,
315 %error,
316 "Current-state auth check failed; soft-failing event.",
317 );
318 })
319 .is_ok()
320}
321
322#[implement(super::Service)]
329async fn soft_fail_standing(
330 &self,
331 incoming_pdu: &PduEvent,
332 room_rules: &RoomVersionRules,
333 pdu_json: &mut CanonicalJsonObject,
334) -> Result<Standing> {
335 let event_id = incoming_pdu.event_id();
336
337 if !self
338 .services
339 .pdu_metadata
340 .is_event_soft_failed(event_id)
341 .await
342 {
343 return Ok(Standing::Evaluate(false));
344 }
345
346 if self
347 .is_suppressed(Context::Upgrade, event_id, UPGRADE_RETRY)
348 .await
349 .is_deny()
350 {
351 debug!(%event_id, "Soft failed; deferring re-evaluation.");
352 return Ok(Standing::Withheld);
353 }
354
355 if self
356 .compute_soft_fail(incoming_pdu, room_rules, pdu_json)
357 .await?
358 {
359 self.record_outcome(Context::Upgrade, event_id, Disposition::Transient);
360
361 debug!(%event_id, "Still soft failed.");
362 return Ok(Standing::Withheld);
363 }
364
365 Ok(Standing::Evaluate(true))
366}
367
368#[implement(super::Service)]
369async fn resolve_state_at_incoming_event(
370 &self,
371 origin: &ServerName,
372 room_id: &RoomId,
373 incoming_pdu: &PduEvent,
374 room_version: &RoomVersionId,
375 recursion_level: usize,
376 create_event_id: &EventId,
377) -> Result<(HashMap<u64, OwnedEventId>, ResolvedVia)> {
378 trace!("Resolving state at event");
382
383 let state_at_incoming_event = if incoming_pdu.prev_events().count() == 1 {
384 self.state_at_incoming_degree_one(incoming_pdu)
385 .await?
386 } else {
387 self.state_at_incoming_resolved(incoming_pdu, room_id, room_version)
388 .boxed()
389 .await?
390 };
391
392 if let Some(state) = state_at_incoming_event {
393 return Ok((state, ResolvedVia::Derived));
394 }
395
396 if let Some(state) = self
397 .cached_resolved_state(incoming_pdu.event_id())
398 .await
399 {
400 return Ok((state, ResolvedVia::Memo));
401 }
402
403 let config = &self.services.server.config;
404 let mode = if config.resolve_state_locally_shadow {
405 WalkMode::Shadow
406 } else {
407 WalkMode::Active
408 };
409
410 let enabled = config.resolve_state_locally && config.resolve_state_locally_max > 0;
411
412 let local = enabled
413 .then_async(|| {
414 self.state_at_incoming_local(
415 room_id,
416 incoming_pdu,
417 room_version,
418 create_event_id,
419 mode,
420 )
421 .boxed()
422 })
423 .await
424 .transpose()?
425 .flatten();
426
427 let shadow = match (mode, local) {
428 | (WalkMode::Active, Some(state)) => return Ok((state, ResolvedVia::Local)),
429 | (WalkMode::Shadow, local) => local,
430 | (WalkMode::Active, None) => None,
431 };
432
433 let state = self
434 .fetch_state(
435 origin,
436 room_id,
437 incoming_pdu.event_id(),
438 room_version,
439 recursion_level,
440 create_event_id,
441 )
442 .boxed()
443 .await?
444 .expect("fetch_state always resolves state to some");
445
446 if let Some(local) = shadow {
447 compare_shadow(room_id, incoming_pdu.event_id(), &local, &state);
448 }
449
450 Ok((state, ResolvedVia::Fetch))
451}
452
453#[implement(super::Service)]
454async fn auth_check_outlier_pdu(
455 &self,
456 room_id: &RoomId,
457 incoming_pdu: &PduEvent,
458 room_rules: &RoomVersionRules,
459 state_at_incoming_event: &HashMap<u64, OwnedEventId>,
460) -> Result {
461 let state_fetch = async |k: StateEventType, s: StateKey| {
464 let shortstatekey = self
465 .services
466 .short
467 .get_shortstatekey(&k, s.as_str())
468 .await?;
469
470 let event_id = state_at_incoming_event
471 .get(&shortstatekey)
472 .ok_or_else(|| {
473 err!(Request(NotFound(
474 "shortstatekey {shortstatekey:?} not found for ({k:?},{s:?})"
475 )))
476 })?;
477
478 self.services.timeline.get_pdu(event_id).await
479 };
480
481 let event_fetch = async |event_id: OwnedEventId| self.event_fetch(&event_id).await;
482
483 trace!("Performing positional auth check.");
484 auth_check(room_rules, incoming_pdu, &event_fetch, &state_fetch)
485 .inspect_err(|error| {
486 warn!(
487 auth_leg = "positional",
488 event_id = %incoming_pdu.event_id(),
489 %room_id,
490 %error,
491 "Positional auth check failed.",
492 );
493 })
494 .await?;
495
496 Ok(())
497}
498
499#[implement(super::Service)]
500async fn compute_soft_fail(
501 &self,
502 incoming_pdu: &PduEvent,
503 room_rules: &RoomVersionRules,
504 pdu_json: &mut CanonicalJsonObject,
505) -> Result<bool> {
506 trace!("Performing soft-fail check");
508 let soft_fail_redact = match incoming_pdu.redacts_id(room_rules) {
509 | None => false,
510 | Some(redact_id) =>
511 !self
512 .services
513 .state_accessor
514 .user_can_redact(&redact_id, incoming_pdu.sender(), incoming_pdu.room_id(), true)
515 .await?,
516 };
517
518 Ok(soft_fail_redact
520 || matches!(
521 self.verify_or_fetch_inbound_policy_signature(pdu_json, incoming_pdu)
522 .await,
523 PolicyCheck::Invalid,
524 ))
525}
526
527#[implement(super::Service)]
528async fn compute_remaining_extremities(
529 &self,
530 room_id: &RoomId,
531 incoming_pdu: &PduEvent,
532) -> Vec<OwnedEventId> {
533 trace!("Calculating extremities");
536 let extremities: Vec<_> = self
537 .services
538 .state
539 .get_forward_extremities(room_id)
540 .ready_filter(|&event_id| {
541 !incoming_pdu
543 .prev_events()
544 .any(is_equal_to!(event_id))
545 })
546 .map(ToOwned::to_owned)
547 .broad_filter_map(async |event_id| {
548 self.services
550 .pdu_metadata
551 .is_event_referenced(room_id, &event_id)
552 .await
553 .eq(&false)
554 .then_some(event_id)
555 })
556 .collect()
557 .await;
558
559 debug!(
560 retained = extremities.len(),
561 prev_events = incoming_pdu.prev_events().count(),
562 "Retained extremities checked against prev_events.",
563 );
564
565 extremities
566}
567
568#[implement(super::Service)]
569async fn resolve_and_force_state_after(
570 &self,
571 room_id: &RoomId,
572 room_version: &RoomVersionId,
573 incoming_pdu: &PduEvent,
574 state_at_incoming_event: &HashMap<u64, OwnedEventId>,
575 state_lock: &RoomMutexGuard,
576) -> Result {
577 let mut state_after = state_at_incoming_event.clone();
579 if let Some(state_key) = incoming_pdu.state_key() {
580 let event_id = incoming_pdu.event_id();
581 let event_type = incoming_pdu.kind();
582 let shortstatekey = self
583 .services
584 .short
585 .get_or_create_shortstatekey(&event_type.to_string().into(), state_key)
586 .await;
587
588 state_after.insert(shortstatekey, event_id.to_owned());
589 debug!(
591 ?event_id,
592 ?event_type,
593 ?state_key,
594 ?shortstatekey,
595 state_after = state_after.len(),
596 "Adding event to state."
597 );
598 }
599
600 trace!("Resolving new room state.");
601 let new_room_state = self
602 .resolve_state(room_id, room_version, state_after)
603 .boxed()
604 .await?;
605
606 trace!("Saving resolved state.");
608 let HashSetCompressStateEvent { shortstatehash, added, removed } = self
609 .services
610 .state_compressor
611 .save_state(room_id, new_room_state)
612 .await?;
613
614 debug!(
615 ?shortstatehash,
616 added = added.len(),
617 removed = removed.len(),
618 "Forcing new room state."
619 );
620 self.services
621 .state
622 .force_state(room_id, shortstatehash, added, removed, state_lock)
623 .await?;
624
625 Ok(())
626}