Skip to main content

tuwunel_service/rooms/event_handler/
upgrade_outlier_pdu.rs

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/// How the state at the incoming event was obtained, deciding the memo write
35/// at the upgrade site.
36#[derive(Clone, Copy)]
37enum ResolvedVia {
38	Derived,
39	Memo,
40	Local,
41	Fetch,
42}
43
44/// Outcome of re-examining an event that already carries a soft-fail marker.
45#[derive(Clone, Copy)]
46enum Standing {
47	/// Evaluate the event. `true` when a re-check already cleared a lapsed
48	/// marker, so the soft-fail computation is settled.
49	Evaluate(bool),
50
51	/// The standing verdict holds; leave the event withheld.
52	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	// Skip the PDU if we already have it as a timeline event
75	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	// 13. Use state resolution to find new room state
120	// We start looking at current room state now, so lets lock the room
121	trace!("Locking the room");
122	let state_lock = self.services.state.mutex.lock(room_id).await;
123
124	// 14. Check if the event passes auth based on the current room state.
125	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	// A soft-failed event is not a forward extremity, so it never drives the
177	// room's current state; only an accepted state event resolves forward.
178	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	// We use the `state_at_event` instead of `state_after` so we accurately
191	// represent the state for this event.
192	trace!("Appending pdu to timeline");
193
194	// Incoming event will be referenced in prev_events unless soft-failed.
195	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/// Re-examines an event that already carries a soft-fail marker.
323///
324/// The marker is a standing verdict rather than a permanent rejection, so it
325/// lapses on the upgrade backoff and the event is weighed again. Asking before
326/// state resolution keeps a still-refused event cheap to decline, since a
327/// cached policy answer needs no round trip.
328#[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	// 10. Fetch missing state and auth chain events by calling /state_ids at
379	//     backwards extremities doing all the checks in this list starting at 1.
380	//     These are not timeline events.
381	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	// 11. Check the auth of the event passes based on the state of the event
462
463	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	// Soft fail check before doing state res
507	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	// MSC4284: soft-fail when the policy server rejects the event.
519	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	// Now we calculate the set of extremities this room has after the incoming
534	// event has been applied. We start with the previous extremities (aka leaves)
535	trace!("Calculating extremities");
536	let extremities: Vec<_> = self
537		.services
538		.state
539		.get_forward_extremities(room_id)
540		.ready_filter(|&event_id| {
541			// Remove any that are referenced by this incoming event's prev_events
542			!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			// Only keep those extremities were not referenced yet
549			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	// We also add state after incoming event to the fork states
578	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		// Now it's the state after the event.
590		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	// Set the new room state to the resolved state
607	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}