Skip to main content

tuwunel_service/rooms/state_res/test_utils/
fetch.rs

1use 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}