tuwunel_service/rooms/state_res/test_utils/
fetch.rs1use std::{collections::HashMap, future::ready, hash::BuildHasher};
2
3use ruma::{EventId, OwnedEventId, events::StateEventType};
4use serde::Deserialize;
5use tuwunel_core::{
6 Result,
7 matrix::{PduEvent, StateKey},
8};
9
10#[cfg(test)]
11use super::TestStateMap;
12use super::{
13 super::{FetchEvent, FetchState},
14 event_not_found,
15};
16
17impl<F, Fut> FetchEvent for &F
18where
19 F: Fn(OwnedEventId) -> Fut + Sync,
20 Fut: Future<Output = Result<PduEvent>> + Send,
21{
22 async fn get<T>(self, event_id: &EventId) -> Result<T>
23 where
24 T: for<'de> Deserialize<'de> + Send,
25 {
26 let event = self(event_id.to_owned()).await?;
27
28 project(&event)
29 }
30
31 async fn exists(self, event_id: &EventId) -> Result<bool> {
32 match self(event_id.to_owned()).await {
33 | Ok(_) => Ok(true),
34 | Err(error) if error.is_not_found() => Ok(false),
35 | Err(error) => Err(error),
36 }
37 }
38}
39
40impl<E, X, EFut, XFut> FetchEvent for (&E, &X)
41where
42 E: Fn(OwnedEventId) -> EFut + Sync,
43 X: Fn(OwnedEventId) -> XFut + Sync,
44 EFut: Future<Output = Result<PduEvent>> + Send,
45 XFut: Future<Output = Result<bool>> + Send,
46{
47 fn get<T>(self, event_id: &EventId) -> impl Future<Output = Result<T>> + Send
48 where
49 T: for<'de> Deserialize<'de> + Send,
50 {
51 FetchEvent::get(self.0, event_id)
52 }
53
54 fn exists(self, event_id: &EventId) -> impl Future<Output = Result<bool>> + Send {
55 self.1(event_id.to_owned())
56 }
57}
58
59impl<S: BuildHasher + Sync> FetchEvent for &HashMap<OwnedEventId, PduEvent, S> {
60 fn get<T>(self, event_id: &EventId) -> impl Future<Output = Result<T>> + Send
61 where
62 T: for<'de> Deserialize<'de> + Send,
63 {
64 ready(
65 HashMap::get(self, event_id)
66 .ok_or_else(|| event_not_found(event_id))
67 .and_then(project),
68 )
69 }
70
71 fn exists(self, event_id: &EventId) -> impl Future<Output = Result<bool>> + Send {
72 ready(Ok(self.contains_key(event_id)))
73 }
74}
75
76impl<S: BuildHasher + Sync> FetchEvent for &HashMap<OwnedEventId, Vec<u8>, S> {
77 fn get<T>(self, event_id: &EventId) -> impl Future<Output = Result<T>> + Send
78 where
79 T: for<'de> Deserialize<'de> + Send,
80 {
81 ready(
82 HashMap::get(self, event_id)
83 .ok_or_else(|| event_not_found(event_id))
84 .and_then(|row| serde_json::from_slice(row).map_err(Into::into)),
85 )
86 }
87
88 fn exists(self, event_id: &EventId) -> impl Future<Output = Result<bool>> + Send {
89 ready(Ok(self.contains_key(event_id)))
90 }
91}
92
93impl<F, Fut> FetchState for &F
94where
95 F: Fn(StateEventType, StateKey) -> Fut + Sync,
96 Fut: Future<Output = Result<PduEvent>> + Send,
97{
98 type Pdu = PduEvent;
99
100 fn get(
101 self,
102 event_type: StateEventType,
103 state_key: StateKey,
104 ) -> impl Future<Output = Result<Self::Pdu>> + Send {
105 self(event_type, state_key)
106 }
107}
108
109#[cfg(test)]
110impl<'a> FetchState for &'a TestStateMap {
111 type Pdu = &'a PduEvent;
112
113 fn get(
114 self,
115 event_type: StateEventType,
116 state_key: StateKey,
117 ) -> impl Future<Output = Result<Self::Pdu>> + Send {
118 ready(TestStateMap::get(self, &event_type, state_key.as_str()))
119 }
120}
121
122fn project<T: for<'de> Deserialize<'de>>(event: &PduEvent) -> Result<T> {
123 let row = serde_json::to_vec(event)?;
124
125 serde_json::from_slice(&row).map_err(Into::into)
126}