tuwunel_service/rooms/state_accessor/
state.rs1use std::{ops::Deref, sync::Arc};
7
8use futures::{
9 FutureExt, Stream, StreamExt, TryFutureExt, TryStreamExt, future::try_join, pin_mut,
10};
11use ruma::{
12 OwnedEventId, UserId,
13 events::{
14 StateEventType,
15 room::member::{MembershipState, RoomMemberEventContent},
16 },
17};
18use serde::Deserialize;
19use tuwunel_core::{
20 Result, at, err, implement,
21 matrix::{Event, Pdu, StateKey},
22 pair_of,
23 utils::{
24 result::FlatOk,
25 stream::{BroadbandExt, IterStream, ReadyExt, TryBroadbandExt, TryIgnore, TryTools},
26 },
27};
28
29use crate::rooms::{
30 short::{ShortEventId, ShortStateHash, ShortStateKey},
31 state_compressor::{CompressedState, compress_state_event, parse_compressed_state_event},
32};
33
34#[implement(super::Service)]
39#[inline]
40pub async fn user_was_joined(&self, shortstatehash: ShortStateHash, user_id: &UserId) -> bool {
41 matches!(
42 self.user_membership(shortstatehash, user_id)
43 .await,
44 MembershipState::Join
45 )
46}
47
48#[implement(super::Service)]
53#[inline]
54pub async fn user_was_invited(&self, shortstatehash: ShortStateHash, user_id: &UserId) -> bool {
55 let membership = self
56 .user_membership(shortstatehash, user_id)
57 .await;
58
59 matches!(membership, MembershipState::Join | MembershipState::Invite)
60}
61
62#[implement(super::Service)]
67pub async fn user_membership(
68 &self,
69 shortstatehash: ShortStateHash,
70 user_id: &UserId,
71) -> MembershipState {
72 self.state_get_content(shortstatehash, &StateEventType::RoomMember, user_id.as_str())
73 .await
74 .map_or(MembershipState::Leave, |c: RoomMemberEventContent| c.membership)
75}
76
77#[implement(super::Service)]
83pub async fn user_membership_at_pdu(&self, user_id: &UserId, pdu: &Pdu) -> MembershipState {
84 if let Some(membership) = pdu.membership_for(user_id) {
85 return membership;
86 }
87
88 let Ok(shortstatehash) = self
89 .services
90 .state
91 .pdu_shortstatehash(pdu.event_id())
92 .await
93 else {
94 return MembershipState::Leave;
95 };
96
97 self.user_membership(shortstatehash, user_id)
98 .await
99}
100
101#[implement(super::Service)]
106pub async fn state_get_content<T>(
107 &self,
108 shortstatehash: ShortStateHash,
109 event_type: &StateEventType,
110 state_key: &str,
111) -> Result<T>
112where
113 T: for<'de> Deserialize<'de> + Send,
114{
115 self.state_get(shortstatehash, event_type, state_key)
116 .await
117 .and_then(|event| event.get_content())
118}
119
120#[implement(super::Service)]
125pub async fn state_contains(
126 &self,
127 shortstatehash: ShortStateHash,
128 event_type: &StateEventType,
129 state_key: &str,
130) -> bool {
131 let Ok(shortstatekey) = self
132 .services
133 .short
134 .get_shortstatekey(event_type, state_key)
135 .await
136 else {
137 return false;
138 };
139
140 self.state_contains_shortstatekey(shortstatehash, shortstatekey)
141 .await
142}
143
144#[implement(super::Service)]
149pub async fn state_contains_type(
150 &self,
151 shortstatehash: ShortStateHash,
152 event_type: &StateEventType,
153) -> bool {
154 let state_keys = self.state_keys(shortstatehash, event_type);
155
156 pin_mut!(state_keys);
157 state_keys.next().await.is_some()
158}
159
160#[implement(super::Service)]
165pub async fn state_contains_shortstatekey(
166 &self,
167 shortstatehash: ShortStateHash,
168 shortstatekey: ShortStateKey,
169) -> bool {
170 let start = compress_state_event(shortstatekey, 0);
171 let end = compress_state_event(shortstatekey, u64::MAX);
172
173 self.load_full_state(shortstatehash)
174 .map_ok(|full_state| full_state.range(start..=end).next().copied())
175 .await
176 .flat_ok()
177 .is_some()
178}
179
180#[implement(super::Service)]
185pub async fn state_get(
186 &self,
187 shortstatehash: ShortStateHash,
188 event_type: &StateEventType,
189 state_key: &str,
190) -> Result<Pdu> {
191 let event_id: OwnedEventId = self
192 .state_get_id(shortstatehash, event_type, state_key)
193 .await?;
194
195 self.services.timeline.get_pdu(&event_id).await
196}
197
198#[implement(super::Service)]
202pub async fn state_get_id(
203 &self,
204 shortstatehash: ShortStateHash,
205 event_type: &StateEventType,
206 state_key: &str,
207) -> Result<OwnedEventId> {
208 let shorteventid = self
209 .state_get_shortid(shortstatehash, event_type, state_key)
210 .await?;
211
212 self.services
213 .short
214 .get_eventid_from_short(shorteventid)
215 .await
216}
217
218#[implement(super::Service)]
223pub async fn state_get_shortid(
224 &self,
225 shortstatehash: ShortStateHash,
226 event_type: &StateEventType,
227 state_key: &str,
228) -> Result<ShortEventId> {
229 let shortstatekey = self
230 .services
231 .short
232 .get_shortstatekey(event_type, state_key)
233 .await?;
234
235 let start = compress_state_event(shortstatekey, 0);
236 let end = compress_state_event(shortstatekey, u64::MAX);
237 self.load_full_state(shortstatehash)
238 .map_ok(|full_state| {
239 full_state
240 .range(start..=end)
241 .next()
242 .copied()
243 .map(parse_compressed_state_event)
244 .map(at!(1))
245 .ok_or(err!(Request(NotFound("Not found in room state"))))
246 })
247 .await?
248}
249
250#[implement(super::Service)]
255pub fn state_type_pdus<'a>(
256 &'a self,
257 shortstatehash: ShortStateHash,
258 event_type: &'a StateEventType,
259) -> impl Stream<Item = impl Event> + Send + 'a {
260 self.state_keys_with_ids(shortstatehash, event_type)
261 .map(at!(1))
262 .broad_filter_map(async |event_id: OwnedEventId| {
263 self.services
264 .timeline
265 .get_pdu(&event_id)
266 .await
267 .ok()
268 })
269}
270
271#[implement(super::Service)]
276pub fn state_keys_with_ids<'a>(
277 &'a self,
278 shortstatehash: ShortStateHash,
279 event_type: &'a StateEventType,
280) -> impl Stream<Item = (StateKey, OwnedEventId)> + Send + 'a {
281 self.state_keys_with_shortids(shortstatehash, event_type)
282 .unzip()
283 .map(|(state_keys, shorteventids): (Vec<_>, Vec<_>)| {
284 self.services
285 .short
286 .multi_get_eventid_from_short(shorteventids.into_iter().stream())
287 .zip(state_keys.into_iter().stream())
288 .ready_filter_map(|(eid, sk)| eid.map(move |eid| (sk, eid)).ok())
289 })
290 .flatten_stream()
291}
292
293#[implement(super::Service)]
298pub fn state_keys_with_shortids<'a>(
299 &'a self,
300 shortstatehash: ShortStateHash,
301 event_type: &'a StateEventType,
302) -> impl Stream<Item = (StateKey, ShortEventId)> + Send + 'a {
303 self.state_full_shortids(shortstatehash)
304 .ignore_err()
305 .unzip()
306 .map(move |(shortstatekeys, shorteventids): (Vec<_>, Vec<_>)| {
307 self.services
308 .short
309 .multi_get_statekey_from_short(shortstatekeys.into_iter().stream())
310 .zip(shorteventids.into_iter().stream())
311 .ready_filter_map(|(res, id)| res.map(|res| (res, id)).ok())
312 .ready_filter_map(move |((event_type_, state_key), event_id)| {
313 event_type_
314 .eq(event_type)
315 .then_some((state_key, event_id))
316 })
317 })
318 .flatten_stream()
319}
320
321#[implement(super::Service)]
326pub fn state_keys<'a>(
327 &'a self,
328 shortstatehash: ShortStateHash,
329 event_type: &'a StateEventType,
330) -> impl Stream<Item = StateKey> + Send + 'a {
331 let short_ids = self
332 .state_full_shortids(shortstatehash)
333 .ignore_err()
334 .map(at!(0));
335
336 self.services
337 .short
338 .multi_get_statekey_from_short(short_ids)
339 .ready_filter_map(Result::ok)
340 .ready_filter_map(move |(event_type_, state_key)| {
341 event_type_.eq(event_type).then_some(state_key)
342 })
343}
344
345#[implement(super::Service)]
350#[inline]
351pub fn state_removed(
352 &self,
353 shortstatehash: pair_of!(ShortStateHash),
354) -> impl Stream<Item = (ShortStateKey, ShortEventId)> + Send + '_ {
355 self.state_added((shortstatehash.1, shortstatehash.0))
356}
357
358#[implement(super::Service)]
363pub fn state_added(
364 &self,
365 shortstatehash: pair_of!(ShortStateHash),
366) -> impl Stream<Item = (ShortStateKey, ShortEventId)> + Send + '_ {
367 let a = self.load_full_state(shortstatehash.0);
368 let b = self.load_full_state(shortstatehash.1);
369 try_join(a, b)
370 .map_ok(|(a, b)| b.difference(&a).copied().collect::<Vec<_>>())
371 .map_ok(IterStream::try_stream)
372 .try_flatten_stream()
373 .ignore_err()
374 .map(parse_compressed_state_event)
375}
376
377#[implement(super::Service)]
382pub fn state_full(
383 &self,
384 shortstatehash: ShortStateHash,
385) -> impl Stream<Item = ((StateEventType, StateKey), impl Event)> + Send + '_ {
386 self.state_full_pdus(shortstatehash)
387 .ready_filter_map(|pdu| {
388 Some(((pdu.kind().to_cow_str().into(), pdu.state_key()?.into()), pdu))
389 })
390}
391
392#[implement(super::Service)]
397pub fn state_full_pdus(
398 &self,
399 shortstatehash: ShortStateHash,
400) -> impl Stream<Item = impl Event> + Send + '_ {
401 let short_ids = self
402 .state_full_shortids(shortstatehash)
403 .ignore_err()
404 .map(at!(1));
405
406 self.services
407 .short
408 .multi_get_eventid_from_short(short_ids)
409 .ready_filter_map(Result::ok)
410 .broad_filter_map(async |event_id: OwnedEventId| {
411 self.services
412 .timeline
413 .get_pdu(&event_id)
414 .await
415 .ok()
416 })
417}
418
419#[implement(super::Service)]
424pub fn state_full_pdus_strict(
425 &self,
426 shortstatehash: ShortStateHash,
427) -> impl Stream<Item = Result<impl Event>> + Send + '_ {
428 self.state_full_ids_strict(shortstatehash)
429 .broad_and_then(async |(_, event_id)| self.services.timeline.get_pdu(&event_id).await)
430}
431
432#[implement(super::Service)]
437pub fn state_full_ids(
438 &self,
439 shortstatehash: ShortStateHash,
440) -> impl Stream<Item = (ShortStateKey, OwnedEventId)> + Send + '_ {
441 self.state_full_shortids(shortstatehash)
442 .ignore_err()
443 .unzip()
444 .map(|(shortstatekeys, shorteventids): (Vec<_>, Vec<_>)| {
445 self.services
446 .short
447 .multi_get_eventid_from_short(shorteventids.into_iter().stream())
448 .zip(shortstatekeys.into_iter().stream())
449 .ready_filter_map(|(eid, ssk)| eid.ok().map(|eid| (ssk, eid)))
450 })
451 .flatten_stream()
452}
453
454#[implement(super::Service)]
459pub fn state_full_ids_strict(
460 &self,
461 shortstatehash: ShortStateHash,
462) -> impl Stream<Item = Result<(ShortStateKey, OwnedEventId)>> + Send + '_ {
463 self.state_full_shortids(shortstatehash)
464 .try_unzip::<Vec<_>, Vec<_>>()
465 .and_then(async move |(shortstatekeys, shorteventids)| {
466 self.services
467 .short
468 .multi_get_eventid_from_short(shorteventids.into_iter().stream())
469 .zip(shortstatekeys.into_iter().stream())
470 .map(|(event_id, shortstatekey)| {
471 event_id.map(|event_id| (shortstatekey, event_id))
472 })
473 .try_collect::<Vec<_>>()
474 .await
475 })
476 .map_ok(Vec::into_iter)
477 .map_ok(IterStream::try_stream)
478 .try_flatten_stream()
479}
480
481#[implement(super::Service)]
486pub fn state_full_shortids(
487 &self,
488 shortstatehash: ShortStateHash,
489) -> impl Stream<Item = Result<(ShortStateKey, ShortEventId)>> + Send + '_ {
490 self.load_full_state(shortstatehash)
491 .map_ok(|full_state| {
492 full_state
493 .deref()
494 .iter()
495 .copied()
496 .map(parse_compressed_state_event)
497 .collect()
498 })
499 .map_ok(Vec::into_iter)
500 .map_ok(IterStream::try_stream)
501 .try_flatten_stream()
502}
503
504#[implement(super::Service)]
505#[tracing::instrument(name = "load", level = "debug", skip(self))]
506async fn load_full_state(&self, shortstatehash: ShortStateHash) -> Result<Arc<CompressedState>> {
507 self.services
508 .state_compressor
509 .load_shortstatehash_info(shortstatehash)
510 .map_err(|e| err!(Database("Missing state IDs: {e}")))
511 .map_ok(|vec| {
512 vec.last()
513 .expect("at least one layer")
514 .full_state
515 .clone()
516 })
517 .await
518}