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};
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/// How the state at the incoming event was obtained, deciding the memo write
36/// at the upgrade site.
37#[derive(Clone, Copy)]
38enum ResolvedVia {
39	Derived,
40	Memo,
41	Local,
42	Fetch,
43}
44
45/// Outcome of re-examining an event that already carries a soft-fail marker.
46#[derive(Clone, Copy)]
47enum Standing {
48	/// Evaluate the event. `true` when a re-check already cleared a lapsed
49	/// marker, so the soft-fail computation is settled.
50	Evaluate(bool),
51
52	/// The standing verdict holds; leave the event withheld.
53	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	// Skip the PDU if we already have it as a timeline event
80	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	// 13. Use state resolution to find new room state
125	// We start looking at current room state now, so lets lock the room
126	trace!("Locking the room");
127	let state_lock = self.services.state.mutex.lock(room_id).await;
128
129	// 14. Check if the event passes auth based on the current room state.
130	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	// Memoize only complete local or fetched state after positional auth succeeds;
177	// soft-failed events cannot establish reusable state.
178	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	// A soft-failed event is not a forward extremity, so it never drives the
184	// room's current state; only an accepted state event resolves forward.
185	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() // cold arm: state event
194		.await?;
195	}
196
197	// We use the `state_at_event` instead of `state_after` so we accurately
198	// represent the state for this event.
199	trace!("Appending pdu to timeline");
200
201	// Incoming event will be referenced in prev_events unless soft-failed.
202	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/// Re-examines an event that already carries a soft-fail marker.
338///
339/// The marker is a standing verdict rather than a permanent rejection, so it
340/// lapses on the upgrade backoff and the event is weighed again. Asking before
341/// state resolution keeps a still-refused event cheap to decline, since a
342/// cached policy answer needs no round trip.
343#[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	// 10. Fetch missing state and auth chain events by calling /state_ids at
394	//     backwards extremities doing all the checks in this list starting at 1.
395	//     These are not timeline events.
396	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() // cold arm: multiple prev events
404			.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() // cold arm: federation fallback
459		.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	// Every event holding a `shorteventid_shortstatehash` row passed spec check 5
474	// (auth against the state at its own position) as a hard reject; soft failure
475	// (spec check 6) still writes the row, so soft-failed events are valid fold
476	// inputs while positionally rejected events never gain a row.
477
478	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	// Soft fail check before doing state res
506	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	// MSC4284: soft-fail when the policy server rejects the event.
518	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	// Now we calculate the set of extremities this room has after the incoming
533	// event has been applied. We start with the previous extremities (aka leaves)
534	trace!("Calculating extremities");
535	let extremities: Vec<_> = self
536		.services
537		.state
538		.get_forward_extremities(room_id)
539		.ready_filter(|&event_id| {
540			// Remove any that are referenced by this incoming event's prev_events
541			!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			// Only keep those extremities were not referenced yet
548			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	// We also add state after incoming event to the fork states
577	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		// Now it's the state after the event.
589		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() // size firewall
603		.await?;
604
605	// Set the new room state to the resolved state
606	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}