Skip to main content

tuwunel_service/rooms/state_res/resolve/
iterative_auth_check.rs

1use 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/// Perform the iterative auth checks to the given list of events.
23///
24/// Definition in the specification:
25///
26/// The iterative auth checks algorithm takes as input an initial room state and
27/// a sorted list of state events, and constructs a new room state by iterating
28/// through the event list and applying the state event to the room state if the
29/// state event is allowed by the authorization rules. If the state event is not
30/// allowed by the authorization rules, then the event is ignored. If a
31/// (event_type, state_key) key that is required for checking the authorization
32/// rules is not present in the state, then the appropriate state event from the
33/// event’s auth_events is used if the auth event is not rejected.
34///
35/// ## Arguments
36///
37/// * `rules` - The authorization rules for the current room version.
38/// * `events` - The sorted state events to apply to the `partial_state`.
39/// * `state` - The current state that was partially resolved for the room.
40/// * `fetch_event` - Function to fetch an event in the room given its event ID.
41///
42/// ## Returns
43///
44/// Returns the partially resolved state, or an `Err(_)` if one of the state
45/// events in the room has an unexpected format.
46#[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	// If the `m.room.create` event is not in the auth events, we need to add it,
131	// because it's always part of the state and required in the auth rules.
132	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}