tuwunel_service/rooms/event_handler/
resolve_state.rs1use std::{
2 borrow::Borrow,
3 collections::HashMap,
4 sync::{
5 Arc,
6 atomic::{AtomicBool, Ordering},
7 },
8};
9
10use futures::{FutureExt, Stream, StreamExt, TryFutureExt, TryStreamExt};
11use ruma::{EventId, OwnedEventId, RoomId, RoomVersionId};
12use serde::Deserialize;
13use tuwunel_core::{
14 Err, Result, err, error, implement,
15 matrix::room_version,
16 trace,
17 utils::stream::{IterStream, ReadyExt, TryWidebandExt, WidebandExt},
18};
19
20use crate::rooms::{
21 state_compressor::CompressedState,
22 state_res::{self, AuthSet, FetchEvent, StateMap},
23 timeline,
24};
25
26#[derive(Clone, Copy)]
27struct Strict<'a> {
28 timeline: &'a timeline::Service,
29 complete: Option<&'a AtomicBool>,
30}
31
32#[implement(super::Service)]
33#[tracing::instrument(
34 name = "state",
35 level = "debug",
36 skip_all,
37 fields(
38 incoming = ?incoming_state.len()
39 ),
40)]
41pub async fn resolve_state(
42 &self,
43 room_id: &RoomId,
44 room_version: &RoomVersionId,
45 incoming_state: HashMap<u64, OwnedEventId>,
46) -> Result<Arc<CompressedState>> {
47 trace!("Loading current room state ids");
48 let current_sstatehash = self
49 .services
50 .state
51 .get_room_shortstatehash(room_id)
52 .map_err(|e| err!(Database(error!("No state for {room_id:?}: {e:?}"))))
53 .await?;
54
55 let current_state_ids: HashMap<_, _> = self
56 .services
57 .state_accessor
58 .state_full_ids(current_sstatehash)
59 .collect()
60 .await;
61
62 trace!("Loading fork states");
63 let fork_states = [current_state_ids, incoming_state];
64 let auth_chains = fork_states
65 .iter()
66 .try_stream()
67 .wide_and_then(|state| {
68 self.services
71 .auth_chain
72 .event_ids_iter(room_id, room_version, state.values().map(Borrow::borrow))
73 .try_collect()
74 .map_ok(AuthSet::from_distinct)
75 })
76 .ready_filter_map(Result::ok);
77
78 let fork_states = fork_states
79 .iter()
80 .stream()
81 .wide_then(|fork_state| {
82 let shortstatekeys = fork_state.keys().copied().stream();
83 let event_ids = fork_state.values().cloned().stream();
84 self.services
85 .short
86 .multi_get_statekey_from_short(shortstatekeys)
87 .zip(event_ids)
88 .ready_filter_map(|(ty_sk, id)| Some((ty_sk.ok()?, id)))
89 .collect::<StateMap<OwnedEventId>>()
90 });
91
92 trace!("Resolving state");
93 let state = self
94 .state_resolution(room_id, room_version, fork_states, auth_chains, None)
95 .await?;
96
97 trace!("State resolution done.");
98 let state_events: Vec<_> = state
99 .iter()
100 .stream()
101 .wide_then(|((event_type, state_key), event_id)| {
102 self.services
103 .short
104 .get_or_create_shortstatekey(event_type, state_key)
105 .map(move |shortstatekey| (shortstatekey, event_id))
106 })
107 .collect()
108 .await;
109
110 trace!("Compressing state...");
111 let new_room_state: CompressedState = self
112 .services
113 .state_compressor
114 .compress_state_events(
115 state_events
116 .iter()
117 .map(|(ssk, eid)| (ssk, (*eid).borrow())),
118 )
119 .collect()
120 .await;
121
122 Ok(Arc::new(new_room_state))
123}
124
125#[implement(super::Service)]
126#[tracing::instrument(name = "resolve", level = "debug", skip_all)]
127pub(super) async fn state_resolution<StateSets, AuthSets>(
128 &self,
129 _room_id: &RoomId,
130 room_version: &RoomVersionId,
131 state_sets: StateSets,
132 auth_chains: AuthSets,
133 complete: Option<&AtomicBool>,
134) -> Result<StateMap<OwnedEventId>>
135where
136 StateSets: Stream<Item = StateMap<OwnedEventId>> + Send,
137 AuthSets: Stream<Item = AuthSet<OwnedEventId>> + Send,
138{
139 let fetch = Strict {
140 timeline: &self.services.timeline,
141 complete,
142 };
143
144 state_res::resolve(
145 &room_version::rules(room_version)?,
146 state_sets,
147 auth_chains,
148 fetch,
149 self.services.server.config.hydra_backports,
150 )
151 .inspect_err(|error| {
152 if let Some(complete) = complete {
153 complete.store(false, Ordering::Relaxed);
154 }
155
156 error!(?error, "State resolution failed.");
157 })
158 .await
159}
160
161impl FetchEvent for Strict<'_> {
162 async fn get<T>(self, event_id: &EventId) -> Result<T>
163 where
164 T: for<'de> Deserialize<'de> + Send,
165 {
166 FetchEvent::get(self.timeline, event_id)
167 .map_err(|error| match self.complete.filter(|_| error.is_not_found()) {
168 | None => error,
169 | Some(complete) => {
170 complete.store(false, Ordering::Relaxed);
171
172 err!(Database("State resolution references missing event {event_id}."))
173 },
174 })
175 .await
176 }
177
178 async fn exists(self, event_id: &EventId) -> Result<bool> {
179 let found = FetchEvent::exists(self.timeline, event_id).await?;
180
181 match self.complete.filter(|_| !found) {
182 | None => Ok(found),
183 | Some(complete) => {
184 complete.store(false, Ordering::Relaxed);
185
186 Err!(Database("State resolution references missing event {event_id}."))
187 },
188 }
189 }
190}