Skip to main content

tuwunel_service/rooms/state_res/resolve/
power_sort.rs

1use std::{collections::HashMap, iter::once};
2
3use futures::{StreamExt, TryFutureExt, TryStreamExt, stream::FuturesUnordered};
4use ruma::{
5	EventId, OwnedEventId,
6	events::{TimelineEventType, room::power_levels::UserPowerLevel},
7	room_version_rules::RoomVersionRules,
8};
9use serde::Deserialize;
10use tuwunel_core::{
11	Result, err,
12	matrix::{Event, PduEvent},
13	result::NotFound,
14	utils::stream::{IterStream, TryBroadbandExt, TryReadyExt},
15};
16
17use super::{
18	super::{
19		FetchEvent,
20		events::{
21			PowerEvent, RoomCreateEvent, RoomPowerLevelsEvent, RoomPowerLevelsIntField,
22			power_levels::RoomPowerLevelsEventOptionExt,
23		},
24		fetch_event::AuthRefs,
25		topological_sort,
26		topological_sort::ReferencedIds,
27	},
28	ConflictedSet,
29};
30
31/// Enlarge the given list of conflicted power events by adding the events in
32/// their auth chain that are in the full conflicted set, and sort it using
33/// reverse topological power ordering.
34///
35/// ## Arguments
36///
37/// * `full_conflicted_set` - The full conflicted set.
38///
39/// * `rules` - The authorization rules for the current room version.
40///
41/// * `fetch` - Handle for reading events in the room by event ID.
42///
43/// ## Returns
44///
45/// Returns the ordered list of event IDs from earliest to latest.
46#[tracing::instrument(
47	level = "debug",
48	skip_all,
49	fields(
50		conflicted = full_conflicted_set.len(),
51	)
52)]
53pub(super) async fn power_sort(
54	rules: &RoomVersionRules,
55	full_conflicted_set: &ConflictedSet,
56	fetch: impl FetchEvent,
57) -> Result<Vec<OwnedEventId>> {
58	// A representation of the DAG, a map of event ID to its list of auth events
59	// that are in the full conflicted set. Fill the graph.
60	let graph = full_conflicted_set
61		.iter()
62		.try_stream()
63		.broad_and_then(async |id| {
64			is_power_event_id(id, fetch)
65				.map_ok(|is_power| is_power.then(|| id.clone()))
66				.await
67		})
68		.ready_try_filter_map(Result::Ok)
69		.enumerate()
70		.map(|(i, event_id)| event_id.map(|event_id| (i, event_id)))
71		.try_fold(HashMap::new(), |graph, (i, event_id)| {
72			add_event_auth_chain(full_conflicted_set, graph, event_id, fetch, i)
73		})
74		.await?;
75
76	// The map of event ID to the power level of the sender of the event.
77	// Get the power level of the sender of each event in the graph.
78	let event_to_power_level: HashMap<_, _> = graph
79		.keys()
80		.try_stream()
81		.map_ok(AsRef::as_ref)
82		.broad_and_then(|event_id| {
83			power_level_for_sender(event_id, rules, fetch)
84				.map_ok(move |sender_power| (event_id.to_owned(), sender_power))
85		})
86		.try_collect()
87		.await?;
88
89	let query = async |event_id: OwnedEventId| {
90		let power_level = *event_to_power_level
91			.get(&event_id)
92			.ok_or_else(|| err!(Request(NotFound("Missing PL event: {event_id}"))))?;
93
94		let event: PduEvent = fetch.get(&event_id).await?;
95
96		Ok((power_level, event.origin_server_ts()))
97	};
98
99	topological_sort(graph, &query).await
100}
101
102/// Add the event with the given event ID and all the events in its auth chain
103/// that are in the full conflicted set to the graph.
104///
105/// Missing events are skipped. Other fetch errors preserve their original
106/// kind.
107#[tracing::instrument(
108	name = "auth_chain",
109	level = "trace",
110	skip_all,
111	fields(
112		graph = graph.len(),
113		?event_id,
114		%i,
115	)
116)]
117pub(super) async fn add_event_auth_chain(
118	full_conflicted_set: &ConflictedSet,
119	mut graph: HashMap<OwnedEventId, ReferencedIds>,
120	event_id: OwnedEventId,
121	fetch: impl FetchEvent,
122	i: usize,
123) -> Result<HashMap<OwnedEventId, ReferencedIds>> {
124	let mut todo: FuturesUnordered<_> = once(fetch_optional_event(event_id, fetch)).collect();
125
126	while let Some(event) = todo.next().await {
127		let (event_id, event): (_, Option<AuthRefs>) = event?;
128		let Some(event) = event else {
129			continue;
130		};
131
132		graph.entry(event_id.clone()).or_default();
133
134		for auth_event_id in event
135			.auth_events
136			.into_iter()
137			.filter(|auth_event_id| full_conflicted_set.contains(auth_event_id))
138		{
139			if !graph.contains_key(&auth_event_id) {
140				todo.push(fetch_optional_event(auth_event_id.clone(), fetch));
141			}
142
143			let references = graph
144				.get_mut(&event_id)
145				.expect("event_id present in graph");
146
147			if !references.contains(&auth_event_id) {
148				references.push(auth_event_id);
149			}
150		}
151	}
152
153	Ok(graph)
154}
155
156/// Finds the power level for an event sender.
157///
158/// This is only valid for topological sorting. Missing events retain the
159/// default-power behavior; other dependency failures are returned unchanged.
160#[tracing::instrument(
161	name = "sender_power",
162	level = "trace",
163	skip_all,
164	fields(
165		?event_id,
166	)
167)]
168pub(super) async fn power_level_for_sender(
169	event_id: &EventId,
170	rules: &RoomVersionRules,
171	fetch: impl FetchEvent,
172) -> Result<UserPowerLevel> {
173	let event: Option<PduEvent> = fetch.get(event_id).await.optional()?;
174
175	let hydra_room_id = rules
176		.authorization
177		.room_create_event_id_as_room_id;
178
179	let mut create_event = None;
180	let mut power_levels_event = None;
181	if hydra_room_id && let Some(event) = event.as_ref() {
182		let create_id = event.room_id().as_event_id()?;
183		let fetched: PduEvent = fetch.get(&create_id).await?;
184
185		_ = create_event.insert(RoomCreateEvent::new(fetched));
186	}
187
188	for auth_event_id in event
189		.as_ref()
190		.map(Event::auth_events)
191		.into_iter()
192		.flatten()
193	{
194		use TimelineEventType::{RoomCreate, RoomPowerLevels};
195
196		let Some(auth_event) = fetch
197			.get::<PduEvent>(auth_event_id)
198			.await
199			.optional()?
200		else {
201			continue;
202		};
203
204		if !hydra_room_id && auth_event.is_type_and_state_key(&RoomCreate, "") {
205			_ = create_event.get_or_insert_with(|| RoomCreateEvent::new(auth_event));
206		} else if auth_event.is_type_and_state_key(&RoomPowerLevels, "") {
207			_ = power_levels_event.get_or_insert_with(|| RoomPowerLevelsEvent::new(auth_event));
208		}
209
210		if power_levels_event.is_some() && create_event.is_some() {
211			break;
212		}
213	}
214
215	let creators = create_event
216		.as_ref()
217		.map(|event| event.creators(&rules.authorization))
218		.transpose()?;
219
220	if let Some((event, creators)) = event.as_ref().zip(creators) {
221		power_levels_event.user_power_level(event.sender(), creators, &rules.authorization)
222	} else {
223		power_levels_event
224			.get_as_int_or_default(RoomPowerLevelsIntField::UsersDefault, &rules.authorization)
225			.map(Into::into)
226	}
227}
228
229/// Whether the given event ID belongs to a power event.
230///
231/// See the docs of `is_power_event()` for the definition of a power event.
232#[tracing::instrument(
233	name = "is_power_event",
234	level = "trace",
235	skip_all,
236	fields(
237		?event_id,
238	)
239)]
240pub(super) async fn is_power_event_id(
241	event_id: &EventId,
242	fetch: impl FetchEvent,
243) -> Result<bool> {
244	Ok(fetch
245		.get(event_id)
246		.await
247		.optional()?
248		.is_some_and(|PowerEvent(is_power)| is_power))
249}
250
251async fn fetch_optional_event<T>(
252	event_id: OwnedEventId,
253	fetch: impl FetchEvent,
254) -> Result<(OwnedEventId, Option<T>)>
255where
256	T: for<'de> Deserialize<'de> + Send,
257{
258	let event = fetch.get::<T>(&event_id).await.optional()?;
259
260	Ok((event_id, event))
261}