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