1use std::{fmt::Debug, sync::Arc};
2
3use futures::{Stream, StreamExt, stream::iter};
4use ruma::{OwnedServerName, ServerName, UserId};
5use tuwunel_core::{
6 Error, Result, at, utils,
7 utils::{ReadyExt, stream::TryIgnore},
8};
9use tuwunel_database::{Database, Deserialized, Map, Txn};
10
11use super::{
12 Destination, EduBuf, SendingEvent, TAG_BADGE_REFRESH, TAG_DEVICE_LIST_CHANGED, TAG_TO_DEVICE,
13};
14
15pub(super) type OutgoingItem = (Key, SendingEvent, Destination);
16pub(super) type SendingItem = (Key, SendingEvent);
17pub(super) type QueueItem = (Key, SendingEvent);
18pub(super) type Key = Vec<u8>;
19
20pub struct Data {
21 servercurrentevent_data: Arc<Map>,
22 servernameevent_data: Arc<Map>,
23 servername_educount: Arc<Map>,
24 pub(super) db: Arc<Database>,
25 services: Arc<crate::services::OnceServices>,
26}
27
28impl Data {
29 pub(super) fn new(args: &crate::Args<'_>) -> Self {
30 let db = &args.db;
31 Self {
32 servercurrentevent_data: db["servercurrentevent_data"].clone(),
33 servernameevent_data: db["servernameevent_data"].clone(),
34 servername_educount: db["servername_educount"].clone(),
35 db: args.db.clone(),
36 services: args.services.clone(),
37 }
38 }
39
40 #[inline]
41 pub(super) fn delete_active_request(&self, key: &[u8]) {
42 self.servercurrentevent_data.remove(key);
43 }
44
45 pub(super) async fn delete_all_active_requests_for(&self, destination: &Destination) {
46 let prefix = destination.get_prefix();
47 self.servercurrentevent_data
48 .raw_keys_prefix(&prefix)
49 .ignore_err()
50 .ready_for_each(|key| self.servercurrentevent_data.remove(key))
51 .await;
52 }
53
54 pub(super) async fn delete_all_requests_for(&self, destination: &Destination) {
55 let prefix = destination.get_prefix();
56 self.servercurrentevent_data
57 .raw_keys_prefix(&prefix)
58 .ignore_err()
59 .ready_for_each(|key| self.servercurrentevent_data.remove(key))
60 .await;
61
62 self.servernameevent_data
63 .raw_keys_prefix(&prefix)
64 .ignore_err()
65 .ready_for_each(|key| self.servernameevent_data.remove(key))
66 .await;
67 }
68
69 pub(super) fn mark_as_active<'a, I>(&self, events: I)
70 where
71 I: Iterator<Item = &'a QueueItem>,
72 {
73 events
74 .filter(|(key, _)| !key.is_empty())
75 .fold(self.db.txn(), |mut txn, (key, val)| {
76 txn.insert_raw(&self.servercurrentevent_data, key, val.value_bytes());
77 txn.del_raw(&self.servernameevent_data, key);
78 txn
79 })
80 .execute();
81 }
82
83 pub(super) fn persist_active_edus(&self, server: &ServerName, edus: &[EduBuf]) {
86 let prefix = Destination::Federation(server.to_owned()).get_prefix();
87
88 let items = edus.iter().map(|edu| {
89 let mut key = prefix.clone();
90 let count = self.services.globals.next_count();
91 let count = count.to_be_bytes();
92 key.extend(&count);
93
94 (key, edu.as_slice())
95 });
96
97 Txn::insert(&self.servercurrentevent_data, items).execute();
98 }
99
100 #[inline]
101 pub fn active_requests(&self) -> impl Stream<Item = OutgoingItem> + Send + '_ {
102 self.servercurrentevent_data
103 .raw_stream()
104 .ignore_err()
105 .map(|(key, val)| {
106 let (dest, event) =
107 parse_servercurrentevent(key, val).expect("invalid servercurrentevent");
108
109 (key.to_vec(), event, dest)
110 })
111 }
112
113 #[inline]
114 pub fn active_requests_for(
115 &self,
116 destination: &Destination,
117 ) -> impl Stream<Item = SendingItem> + Send + '_ + use<'_> {
118 let prefix = destination.get_prefix();
119 self.servercurrentevent_data
120 .raw_stream_from(&prefix)
121 .ignore_err()
122 .ready_take_while(move |(key, _)| key.starts_with(&prefix))
123 .map(|(key, val)| {
124 let (_, event) =
125 parse_servercurrentevent(key, val).expect("invalid servercurrentevent");
126
127 (key.to_vec(), event)
128 })
129 }
130
131 pub(super) fn queue_requests<'a, I>(&self, requests: I) -> Vec<Vec<u8>>
132 where
133 I: Iterator<Item = (&'a SendingEvent, &'a Destination)> + Clone + Debug + Send,
134 {
135 let keys: Vec<_> = requests
136 .clone()
137 .map(|(event, dest)| {
138 let mut key = dest.get_prefix();
139 if let SendingEvent::Pdu(value) = event {
140 key.extend(value.as_ref());
141 } else {
142 let count = self.services.globals.next_count();
143 let count = count.to_be_bytes();
144 key.extend(&count);
145 }
146
147 key
148 })
149 .collect();
150
151 let items = keys
152 .iter()
153 .map(Vec::as_slice)
154 .zip(requests.map(at!(0)))
155 .map(|(key, event)| (key, event.value_bytes()));
156
157 Txn::insert(&self.servernameevent_data, items).execute();
158
159 keys
160 }
161
162 pub(super) fn retain_queued<'a, I>(
167 &'a self,
168 events: I,
169 ) -> impl Stream<Item = QueueItem> + Send + 'a
170 where
171 I: IntoIterator<Item = QueueItem> + Send + 'a,
172 I::IntoIter: Send,
173 {
174 iter(events).filter_map(async |item| {
175 let key = &item.0;
176 let exists = async || {
177 self.servernameevent_data
178 .exists(key)
179 .await
180 .is_ok()
181 };
182
183 (key.is_empty() || exists().await).then_some(item)
184 })
185 }
186
187 pub fn queued_requests(
188 &self,
189 destination: &Destination,
190 ) -> impl Stream<Item = QueueItem> + Send + '_ + use<'_> {
191 let prefix = destination.get_prefix();
192 self.servernameevent_data
193 .raw_stream_from(&prefix)
194 .ignore_err()
195 .ready_take_while(move |(key, _)| key.starts_with(&prefix))
196 .map(|(key, val)| {
197 let (_, event) =
198 parse_servercurrentevent(key, val).expect("invalid servercurrentevent");
199
200 (key.to_vec(), event)
201 })
202 }
203
204 pub(super) fn queued_badge_refresh_destinations(
208 &self,
209 ) -> impl Stream<Item = Destination> + Send + '_ {
210 self.servernameevent_data
211 .raw_stream_from(b"$")
212 .ignore_err()
213 .ready_take_while(|(key, _)| key.starts_with(b"$"))
214 .ready_filter_map(|(key, val)| {
215 (val == [TAG_BADGE_REFRESH]).then(|| {
216 parse_servercurrentevent(key, val)
217 .expect("invalid servercurrentevent")
218 .0
219 })
220 })
221 }
222
223 pub(super) fn set_latest_educount(&self, server_name: &ServerName, last_count: u64) {
224 self.servername_educount
225 .raw_put(server_name, last_count);
226 }
227
228 pub async fn get_latest_educount(&self, server_name: &ServerName) -> u64 {
229 self.servername_educount
230 .get(server_name)
231 .await
232 .deserialized()
233 .unwrap_or(0)
234 }
235}
236
237pub(super) fn parse_servercurrentevent(
238 key: &[u8],
239 value: &[u8],
240) -> Result<(Destination, SendingEvent)> {
241 Ok::<_, Error>(if key.starts_with(b"+") {
243 let mut parts = key[1..].splitn(2, |&b| b == 0xFF);
244
245 let server = parts
246 .next()
247 .expect("splitn always returns one element");
248 let event = parts
249 .next()
250 .ok_or_else(|| Error::bad_database("Invalid bytes in servercurrentpdus."))?;
251
252 let server = utils::string_from_bytes(server).map_err(|_| {
253 Error::bad_database("Invalid server bytes in server_currenttransaction")
254 })?;
255
256 let decoded = match value {
257 | [] => SendingEvent::Pdu(event.into()),
258 | [TAG_TO_DEVICE, ..] => SendingEvent::ToDevice(value.into()),
259 | [TAG_DEVICE_LIST_CHANGED, ..] => SendingEvent::DeviceListChanged(value.into()),
260 | _ => SendingEvent::Edu(value.into()),
261 };
262
263 (Destination::Appservice(server), decoded)
264 } else if key.starts_with(b"$") {
265 let mut parts = key[1..].splitn(3, |&b| b == 0xFF);
266
267 let user = parts
268 .next()
269 .expect("splitn always returns one element");
270 let user_string = utils::str_from_bytes(user)
271 .map_err(|_| Error::bad_database("Invalid user string in servercurrentevent"))?;
272 let user_id = UserId::parse(user_string)
273 .map_err(|_| Error::bad_database("Invalid user id in servercurrentevent"))?;
274
275 let pushkey = parts
276 .next()
277 .ok_or_else(|| Error::bad_database("Invalid bytes in servercurrentpdus."))?;
278 let pushkey_string = utils::string_from_bytes(pushkey)
279 .map_err(|_| Error::bad_database("Invalid pushkey in servercurrentevent"))?;
280
281 let event = parts
282 .next()
283 .ok_or_else(|| Error::bad_database("Invalid bytes in servercurrentpdus."))?;
284
285 (Destination::Push(user_id, pushkey_string), match value {
286 | [] => SendingEvent::Pdu(event.into()),
287 | [tag] if *tag == TAG_BADGE_REFRESH => SendingEvent::BadgeRefresh,
288 | _ => SendingEvent::Edu(value.into()),
289 })
290 } else {
291 let mut parts = key.splitn(2, |&b| b == 0xFF);
292
293 let server = parts
294 .next()
295 .expect("splitn always returns one element");
296 let event = parts
297 .next()
298 .ok_or_else(|| Error::bad_database("Invalid bytes in servercurrentpdus."))?;
299
300 let server = utils::string_from_bytes(server).map_err(|_| {
301 Error::bad_database("Invalid server bytes in server_currenttransaction")
302 })?;
303
304 (
305 Destination::Federation(OwnedServerName::parse(&server).map_err(|_| {
306 Error::bad_database("Invalid server string in server_currenttransaction")
307 })?),
308 if value.is_empty() {
309 SendingEvent::Pdu(event.into())
310 } else {
311 SendingEvent::Edu(value.into())
312 },
313 )
314 })
315}