tuwunel_service/rooms/state_res/resolve/
iterative_auth_check.rs1use futures::{Stream, StreamExt, TryFutureExt, TryStreamExt};
2use ruma::{
3 EventId, OwnedEventId, events::TimelineEventType, room_version_rules::RoomVersionRules,
4};
5use tuwunel_core::{
6 Result, debug_warn, err, error,
7 matrix::{Event, EventTypeExt, PduEvent, StateKey, TypeStateKey},
8 smallvec::SmallVec,
9 utils::stream::{IterStream, ReadyExt, TryReadyExt, TryWidebandExt},
10};
11
12use super::{
13 super::{
14 AuthCheckOutcome, FetchEvent, auth_types_for_event, check_state_dependent_auth_rules,
15 event_auth::classify_auth_error,
16 },
17 StateMap,
18};
19
20type AuthEvents = SmallVec<[(TypeStateKey, PduEvent); 4]>;
21
22#[tracing::instrument(
47 name = "iterative_auth",
48 level = "debug",
49 skip_all,
50 fields(
51 states = ?state.len(),
52 )
53)]
54pub(super) async fn iterative_auth_check<'b, SortedPowerEvents, Fetch>(
55 rules: &RoomVersionRules,
56 events: SortedPowerEvents,
57 state: StateMap<OwnedEventId>,
58 fetch: Fetch,
59) -> Result<StateMap<OwnedEventId>>
60where
61 SortedPowerEvents: Stream<Item = &'b EventId> + Send,
62 Fetch: FetchEvent,
63{
64 events
65 .map(Ok)
66 .wide_and_then(async |event_id| {
67 let event: PduEvent = fetch.get(event_id).await?;
68 let state_key = event.state_key().map(StateKey::from);
69
70 Ok(state_key.map(|state_key| (event_id, state_key, event)))
71 })
72 .ready_try_filter_map(Result::<_>::Ok)
73 .try_fold(state, |state, (event_id, state_key, event)| {
74 auth_check(rules, state, event_id, state_key, event, fetch)
75 })
76 .await
77}
78
79#[tracing::instrument(
80 name = "check",
81 level = "trace",
82 skip_all,
83 fields(
84 %event_id,
85 ?state_key,
86 )
87)]
88async fn auth_check<Fetch>(
89 rules: &RoomVersionRules,
90 mut state: StateMap<OwnedEventId>,
91 event_id: &EventId,
92 state_key: StateKey,
93 event: PduEvent,
94 fetch: Fetch,
95) -> Result<StateMap<OwnedEventId>>
96where
97 Fetch: FetchEvent,
98{
99 let Ok(auth_types) = auth_types_for_event(
100 event.event_type(),
101 event.sender(),
102 Some(&state_key),
103 event.content(),
104 &rules.authorization,
105 true,
106 )
107 .inspect_err(|e| error!("failed to get auth types for event: {e}")) else {
108 return Ok(state);
109 };
110
111 let auth_types_events = auth_types
112 .stream()
113 .ready_filter_map(|key| {
114 state
115 .get(&key)
116 .map(move |auth_event_id| (auth_event_id, key))
117 })
118 .filter_map(async |(id, key)| {
119 fetch_auth_event(id, fetch)
120 .await
121 .map(|event| event.map(|event| (key, event)))
122 })
123 .ready_try_filter_map(|(key, auth_event)| {
124 Ok(auth_event
125 .rejected()
126 .eq(&false)
127 .then_some((key, auth_event)))
128 });
129
130 let also_need_create_event = *event.event_type() != TimelineEventType::RoomCreate
133 && rules
134 .authorization
135 .room_create_event_id_as_room_id;
136
137 let also_create_id: Option<OwnedEventId> = also_need_create_event
138 .then(|| event.room_id().as_event_id().ok())
139 .flatten();
140
141 let auth_events = event
142 .auth_events()
143 .chain(also_create_id.as_deref().into_iter())
144 .stream()
145 .filter_map(async |id| fetch_auth_event(id, fetch).await)
146 .ready_try_filter_map(|auth_event| {
147 let state_key = auth_event
148 .state_key()
149 .ok_or_else(|| err!(Request(InvalidParam("Missing state_key"))))?;
150
151 let key_val = auth_event
152 .rejected()
153 .eq(&false)
154 .then_some((auth_event.event_type().with_state_key(state_key), auth_event));
155
156 Ok(key_val)
157 });
158
159 let auth_events = auth_events
160 .chain(auth_types_events)
161 .try_collect()
162 .map_ok(|mut vec: AuthEvents| {
163 vec.sort_by(|a, b| a.0.cmp(&b.0));
164 vec.reverse();
165 vec.dedup_by(|a, b| a.0.eq(&b.0));
166 vec
167 })
168 .await?;
169
170 let outcome =
171 match check_state_dependent_auth_rules(rules, &event, auth_events.as_slice()).await {
172 | Ok(()) => AuthCheckOutcome::Allow,
173 | Err(error) => classify_auth_error(error)?,
174 };
175
176 match outcome {
177 | AuthCheckOutcome::Allow => {
178 let key = event.event_type().with_state_key(state_key);
179
180 state.insert(key, event_id.to_owned());
181 },
182 | AuthCheckOutcome::Deny(error) => {
183 debug_warn!(
184 %event_id,
185 sender = %event.sender(),
186 event_type = ?event.event_type(),
187 ?state_key,
188 %error,
189 "event failed auth check"
190 );
191 },
192 }
193
194 Ok(state)
195}
196
197async fn fetch_auth_event<Fetch>(id: &EventId, fetch: Fetch) -> Option<Result<PduEvent>>
198where
199 Fetch: FetchEvent,
200{
201 match fetch.get::<PduEvent>(id).await {
202 | Ok(event) => Some(Ok(event)),
203 | Err(error) if error.is_not_found() => {
204 debug_warn!(%id, %error, "missing auth event");
205 None
206 },
207 | Err(error) => Some(Err(error)),
208 }
209}