tuwunel_service/rooms/state_res/
resolve.rs1#[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
40pub type StateMap<Id> = BTreeMap<TypeStateKey, Id>;
43
44#[derive(Clone)]
49pub struct AuthSet<Id>(Vec<Id>);
50
51pub type ConflictMap<Id> = StateMap<ConflictVec<Id>>;
53
54type ConflictVec<Id> = SmallVec<[Id; 2]>;
58
59type ConflictedSet = HashSet<OwnedEventId, RandomState>;
61
62impl<Id> AuthSet<Id> {
63 #[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#[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 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 let full_conflicted_set =
159 full_conflicted_set(rules, conflicted_states, auth_sets, fetch, hydra_backports).await?;
160
161 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() .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 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 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 let sorted_remaining_events = have_remaining_events
226 .then_async(move || mainline_sort(power_event.cloned(), remaining_events, fetch))
227 .boxed(); 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 let mut resolved_state =
245 iterative_auth_check(rules, sorted_remaining_events, partially_resolved_state, fetch)
246 .boxed()
247 .await?;
248
249 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 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(); 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}