tuwunel_service/rooms/state_res/resolve/
power_sort.rs1use 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#[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 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 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#[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#[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#[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}