Skip to main content

tuwunel_service/rooms/state_res/
resolve.rs

1#[cfg(test)]
2mod tests;
3
4mod auth_difference;
5mod conflicted_subgraph;
6mod iterative_auth_check;
7mod mainline_sort;
8mod power_sort;
9mod split_conflicted;
10
11use std::{
12	collections::{BTreeMap, HashSet},
13	ops::Deref,
14	vec::IntoIter,
15};
16
17use futures::{FutureExt, Stream, StreamExt, TryFutureExt, TryStreamExt};
18use ruma::{OwnedEventId, events::StateEventType, room_version_rules::RoomVersionRules};
19use tuwunel_core::{
20	Result, debug,
21	itertools::Itertools,
22	matrix::{TypeStateKey, event_id::RandomState},
23	smallvec::SmallVec,
24	trace,
25	utils::{
26		BoolExt,
27		stream::{BroadbandExt, IterStream, ReadyExt},
28	},
29};
30
31use self::{
32	auth_difference::auth_difference, conflicted_subgraph::conflicted_subgraph_dfs,
33	iterative_auth_check::iterative_auth_check, mainline_sort::mainline_sort,
34	power_sort::power_sort, split_conflicted::split_conflicted_state,
35};
36use super::FetchEvent;
37#[cfg(test)]
38use super::test_utils;
39
40/// A mapping of event type and state_key to some value `T`, usually an
41/// `EventId`.
42pub type StateMap<Id> = BTreeMap<TypeStateKey, Id>;
43
44/// Full recursive auth chain for one candidate [`StateMap`].
45///
46/// Values are distinct and immutable after construction. Their order is
47/// arbitrary, and consumers must not depend on it.
48#[derive(Clone)]
49pub struct AuthSet<Id>(Vec<Id>);
50
51/// Conflicting event ids for each contested state key.
52pub type ConflictMap<Id> = StateMap<ConflictVec<Id>>;
53
54/// Event ids contesting one state key.
55///
56/// Two forks disputing a key is the modal conflict, so two ids stay inline.
57type ConflictVec<Id> = SmallVec<[Id; 2]>;
58
59/// The full conflicted set (arbitrary order).
60type ConflictedSet = HashSet<OwnedEventId, RandomState>;
61
62impl<Id> AuthSet<Id> {
63	/// Creates an auth set from distinct identifiers.
64	///
65	/// The caller must ensure `ids` contains no duplicates. Duplicates are
66	/// not checked, so hot paths avoid redundant work.
67	#[inline]
68	#[must_use]
69	pub(crate) fn from_distinct(ids: Vec<Id>) -> Self { Self(ids) }
70}
71
72impl<Id> Default for AuthSet<Id> {
73	fn default() -> Self { Self(Vec::new()) }
74}
75
76impl<Id: Ord> FromIterator<Id> for AuthSet<Id> {
77	fn from_iter<I: IntoIterator<Item = Id>>(iter: I) -> Self {
78		Self::from_distinct(
79			iter.into_iter()
80				.sorted_unstable()
81				.dedup()
82				.collect(),
83		)
84	}
85}
86
87impl<Id> IntoIterator for AuthSet<Id> {
88	type IntoIter = IntoIter<Id>;
89	type Item = Id;
90
91	fn into_iter(self) -> Self::IntoIter { self.0.into_iter() }
92}
93
94/// Apply the [state resolution] algorithm introduced in room version 2 to
95/// resolve the state of a room.
96///
97/// ## Arguments
98///
99/// * `rules` - The rules to apply for the version of the current room.
100///
101/// * `state_maps` - The incoming states to resolve. Each `StateMap` represents
102///   a possible fork in the state of a room.
103///
104/// * `auth_sets` - The list of full recursive sets of `auth_events` for each
105///   event in the `state_maps`. Inputs must not contain duplicates.
106///
107/// * `fetch_event` - Function to fetch an event in the room given its event ID.
108///
109/// ## Invariants
110///
111/// The caller of `resolve` must ensure that all the events are from the same
112/// room.
113///
114/// ## Returns
115///
116/// The resolved room state.
117///
118/// [state resolution]: https://spec.matrix.org/latest/rooms/v2/#state-resolution
119#[tracing::instrument(level = "debug", skip_all)]
120pub async fn resolve<States, AuthSets, Fetch>(
121	rules: &RoomVersionRules,
122	state_maps: States,
123	auth_sets: AuthSets,
124	fetch: Fetch,
125	hydra_backports: bool,
126) -> Result<StateMap<OwnedEventId>>
127where
128	States: Stream<Item = StateMap<OwnedEventId>> + Send,
129	AuthSets: Stream<Item = AuthSet<OwnedEventId>> + Send,
130	Fetch: FetchEvent,
131{
132	// Split the unconflicted state map and the conflicted state set.
133	let (unconflicted_state, conflicted_states) = split_conflicted_state(state_maps).await;
134
135	debug!(
136		unconflicted = unconflicted_state.len(),
137		conflicted_states = conflicted_states.len(),
138		conflicted_events = conflicted_states
139			.values()
140			.fold(0_usize, |a, s| a.saturating_add(s.len())),
141		"unresolved states"
142	);
143
144	trace!(
145		?unconflicted_state,
146		?conflicted_states,
147		unconflicted = unconflicted_state.len(),
148		conflicted_states = conflicted_states.len(),
149		"unresolved states"
150	);
151
152	if conflicted_states.is_empty() {
153		return Ok(unconflicted_state.into_iter().collect());
154	}
155
156	// 0. The full conflicted set is the union of the conflicted state set and the
157	//    auth difference. Don't honor events that don't exist.
158	let full_conflicted_set =
159		full_conflicted_set(rules, conflicted_states, auth_sets, fetch, hydra_backports).await?;
160
161	// 1. Select the set X of all power events that appear in the full conflicted
162	//    set. For each such power event P, enlarge X by adding the events in the
163	//    auth chain of P which also belong to the full conflicted set. Sort X into
164	//    a list using the reverse topological power ordering.
165	let sorted_power_set: Vec<_> = power_sort(rules, &full_conflicted_set, fetch)
166		.inspect_ok(|list| debug!(count = list.len(), "sorted power events"))
167		.inspect_ok(|list| trace!(?list, "sorted power events"))
168		.boxed() // size firewall
169		.await?;
170
171	let power_set_event_ids: Vec<_> = sorted_power_set
172		.iter()
173		.sorted_unstable()
174		.collect();
175
176	let sorted_power_set = sorted_power_set
177		.iter()
178		.stream()
179		.map(AsRef::as_ref);
180
181	let begin_with_empty_state_map = rules
182		.state_res
183		.v2_rules()
184		.is_some_and(|r| r.begin_iterative_auth_checks_with_empty_state_map)
185		|| hydra_backports;
186
187	let initial_state = begin_with_empty_state_map
188		.is_false()
189		.then(|| unconflicted_state.clone())
190		.unwrap_or_default();
191
192	// 2. Apply the iterative auth checks algorithm, starting from the unconflicted
193	//    state map, to the list of events from the previous step to get a partially
194	//    resolved state.
195	// query-depth firewall
196	let partially_resolved_state =
197		iterative_auth_check(rules, sorted_power_set, initial_state, fetch)
198			.boxed()
199			.inspect_ok(|map| debug!(count = map.len(), "partially resolved power state"))
200			.inspect_ok(|map| trace!(?map, "partially resolved power state"))
201			.await?;
202
203	// This "epochs" power level event
204	let power_ty_sk = (StateEventType::RoomPowerLevels, "".into());
205	let power_event = partially_resolved_state.get(&power_ty_sk);
206	debug!(event_id = ?power_event, "epoch power event");
207
208	let remaining_events: Vec<_> = full_conflicted_set
209		.into_iter()
210		.filter(|id| power_set_event_ids.binary_search(&id).is_err())
211		.collect();
212
213	debug!(count = remaining_events.len(), "remaining events");
214	trace!(list = ?remaining_events, "remaining events");
215
216	let have_remaining_events = !remaining_events.is_empty();
217	let remaining_events = remaining_events
218		.iter()
219		.stream()
220		.map(AsRef::as_ref);
221
222	// 3. Take all remaining events that weren’t picked in step 1 and order them by
223	//    the mainline ordering based on the power level in the partially resolved
224	//    state obtained in step 2.
225	let sorted_remaining_events = have_remaining_events
226		.then_async(move || mainline_sort(power_event.cloned(), remaining_events, fetch))
227		.boxed(); // size firewall
228
229	let sorted_remaining_events = sorted_remaining_events
230		.await
231		.unwrap_or(Ok(Vec::new()))?;
232
233	debug!(count = sorted_remaining_events.len(), "sorted remaining events");
234	trace!(list = ?sorted_remaining_events, "sorted remaining events");
235
236	let sorted_remaining_events = sorted_remaining_events
237		.iter()
238		.stream()
239		.map(AsRef::as_ref);
240
241	// 4. Apply the iterative auth checks algorithm on the partial resolved state
242	//    and the list of events from the previous step.
243	// query-depth firewall
244	let mut resolved_state =
245		iterative_auth_check(rules, sorted_remaining_events, partially_resolved_state, fetch)
246			.boxed()
247			.await?;
248
249	// 5. Update the result by replacing any event with the event with the same key
250	//    from the unconflicted state map, if such an event exists, to get the final
251	//    resolved state.
252	resolved_state.extend(unconflicted_state);
253
254	debug!(resolved_state = resolved_state.len(), "resolved state");
255	trace!(?resolved_state, "resolved state");
256
257	Ok(resolved_state)
258}
259
260#[tracing::instrument(
261	name = "conflicted",
262	level = "debug",
263	skip_all,
264	fields(
265		states = conflicted_states.len(),
266		events = conflicted_states.values().flatten().count()
267	),
268)]
269async fn full_conflicted_set<AuthSets>(
270	rules: &RoomVersionRules,
271	conflicted_states: ConflictMap<OwnedEventId>,
272	auth_sets: AuthSets,
273	fetch: impl FetchEvent,
274	hydra_backports: bool,
275) -> Result<ConflictedSet>
276where
277	AuthSets: Stream<Item = AuthSet<OwnedEventId>> + Send,
278{
279	let consider_conflicted_subgraph = rules
280		.state_res
281		.v2_rules()
282		.is_some_and(|rules| rules.consider_conflicted_state_subgraph)
283		|| hydra_backports;
284
285	let conflicted_state_set: Vec<_> = conflicted_states
286		.values()
287		.flatten()
288		.sorted_unstable()
289		.dedup()
290		.collect();
291
292	// Since `org.matrix.hydra.11`, fetch the conflicted state subgraph.
293	let conflicted_subgraph = consider_conflicted_subgraph
294		.then_async(async || conflicted_subgraph_dfs(&conflicted_state_set, fetch))
295		.map(Option::into_iter)
296		.map(IterStream::stream)
297		.flatten_stream()
298		.flatten()
299		.boxed(); // erase region
300
301	let conflicted_state_ids = conflicted_state_set
302		.iter()
303		.map(Deref::deref)
304		.cloned()
305		.stream();
306
307	auth_difference(auth_sets)
308		.chain(conflicted_state_ids)
309		.broad_then(async |id| {
310			let exists = fetch.exists(&id).await?;
311
312			Ok(exists.then_some(id))
313		})
314		.ready_filter_map(Result::transpose)
315		.chain(conflicted_subgraph)
316		.try_collect::<ConflictedSet>()
317		.inspect_ok(|set| debug!(count = set.len(), "full conflicted set"))
318		.inspect_ok(|set| trace!(?set, "full conflicted set"))
319		.await
320}