tuwunel_service/rooms/state_res/resolve/
mainline_sort.rs1use std::collections::HashMap;
2
3use futures::{Stream, StreamExt, TryFutureExt, TryStreamExt, stream::try_unfold};
4use ruma::{EventId, OwnedEventId, events::TimelineEventType};
5use tuwunel_core::{
6 Error, Result, at,
7 matrix::{Event, PduEvent, event_id::RandomState},
8 result::NotFound,
9 trace,
10 utils::stream::{BroadbandExt, TryReadyExt},
11};
12
13use super::super::FetchEvent;
14
15type Positions<'a> = HashMap<&'a EventId, usize, RandomState>;
17
18#[tracing::instrument(
44 level = "debug",
45 skip_all,
46 fields(
47 power_levels = power_level_event_id
48 .as_deref()
49 .map(EventId::as_str)
50 .unwrap_or_default(),
51 )
52)]
53pub(super) async fn mainline_sort<'a, RemainingEvents>(
54 power_level_event_id: Option<OwnedEventId>,
55 events: RemainingEvents,
56 fetch: impl FetchEvent,
57) -> Result<Vec<OwnedEventId>>
58where
59 RemainingEvents: Stream<Item = &'a EventId> + Send,
60{
61 let mainline: Vec<_> = try_unfold(power_level_event_id, async |power_level_event_id| {
63 let Some(power_level_event_id) = power_level_event_id else {
64 return Ok::<_, Error>(None);
65 };
66
67 let power_level_event = fetch
68 .get::<PduEvent>(&power_level_event_id)
69 .await?;
70
71 let this_event_id = power_level_event.event_id().to_owned();
72 let next_event_id = get_power_levels_auth_event(&power_level_event, fetch)
73 .map_ok(|event| {
74 event
75 .as_ref()
76 .map(Event::event_id)
77 .map(ToOwned::to_owned)
78 })
79 .await?;
80
81 trace!(?this_event_id, ?next_event_id, "mainline descent",);
82
83 Ok(Some((this_event_id, next_event_id)))
84 })
85 .try_collect()
86 .await?;
87
88 let positions: Positions<'_> = mainline
89 .iter()
90 .rev()
91 .map(AsRef::as_ref)
92 .enumerate()
93 .map(|(position, event_id)| (event_id, position))
94 .collect();
95
96 events
97 .map(ToOwned::to_owned)
98 .broad_then(async |event_id| {
99 let Some(event) = fetch
100 .get::<PduEvent>(&event_id)
101 .await
102 .optional()?
103 else {
104 return Ok(None);
105 };
106
107 let origin_server_ts = event.origin_server_ts();
108 let Some(position) = mainline_position(Some(event), &positions, fetch)
109 .await
110 .optional()?
111 else {
112 return Ok(None);
113 };
114
115 Ok(Some((event_id, (position, origin_server_ts))))
116 })
117 .ready_try_filter_map(Result::Ok)
118 .inspect_ok(|(event_id, (position, origin_server_ts))| {
119 trace!(position, ?origin_server_ts, ?event_id, "mainline position");
120 })
121 .try_collect()
122 .map_ok(|mut events: Vec<_>| {
123 events.sort_by(|a, b| {
124 let (a_pos, a_ots) = &a.1;
125 let (b_pos, b_ots) = &b.1;
126 a_pos
127 .cmp(b_pos)
128 .then(a_ots.cmp(b_ots))
129 .then(a.cmp(b))
130 });
131
132 events.into_iter().map(at!(0)).collect()
133 })
134 .await
135}
136
137#[tracing::instrument(
150 name = "position",
151 level = "trace",
152 ret(level = "trace"),
153 skip_all,
154 fields(
155 mainline = positions.len(),
156 event = ?current_event.as_ref().map(Event::event_id).map(ToOwned::to_owned),
157 )
158)]
159async fn mainline_position(
160 mut current_event: Option<PduEvent>,
161 positions: &Positions<'_>,
162 fetch: impl FetchEvent,
163) -> Result<usize> {
164 while let Some(event) = current_event {
165 trace!(
166 event_id = ?event.event_id(),
167 "mainline position search",
168 );
169
170 if let Some(position) = positions.get(event.event_id()) {
174 return Ok(position.saturating_add(1));
175 }
176
177 current_event = get_power_levels_auth_event(&event, fetch).await?;
179 }
180
181 Ok(0)
184}
185
186#[tracing::instrument(level = "trace", skip_all)]
187async fn get_power_levels_auth_event(
188 event: &PduEvent,
189 fetch: impl FetchEvent,
190) -> Result<Option<PduEvent>> {
191 for auth_event_id in event.auth_events() {
193 let auth_event: PduEvent = fetch.get(auth_event_id).await?;
194
195 if auth_event.is_type_and_state_key(&TimelineEventType::RoomPowerLevels, "") {
196 return Ok(Some(auth_event));
197 }
198 }
199
200 Ok(None)
201}