tuwunel_service/rooms/short/
mod.rs1use std::{borrow::Borrow, sync::Arc};
8
9use futures::{FutureExt, Stream, StreamExt, pin_mut};
10use ruma::{EventId, OwnedEventId, OwnedRoomId, RoomId, events::StateEventType};
11use serde::Deserialize;
12pub use tuwunel_core::matrix::{ShortEventId, ShortId, ShortRoomId, ShortStateKey};
16use tuwunel_core::{
17 Err, Result, err, implement,
18 matrix::StateKey,
19 utils,
20 utils::{
21 IterStream, MutexMap,
22 hash::sha256::Digest,
23 stream::{ReadyExt, WidebandExt},
24 },
25};
26use tuwunel_database::{Deserialized, Get, Map, Qry, Txn};
27
28pub struct Service {
33 db: Data,
34 creating: Creating,
35 services: Arc<crate::services::OnceServices>,
36}
37
38struct Data {
39 eventid_shorteventid: Arc<Map>,
40 shorteventid_eventid: Arc<Map>,
41 statekey_shortstatekey: Arc<Map>,
42 shortstatekey_statekey: Arc<Map>,
43 roomid_shortroomid: Arc<Map>,
44 statehash_shortstatehash: Arc<Map>,
45}
46
47#[derive(Default)]
54struct Creating {
55 shorteventid: MutexMap<OwnedEventId, ()>,
56 shortstatekey: MutexMap<(StateEventType, StateKey), ()>,
57 shortstatehash: MutexMap<Digest, ()>,
58 shortroomid: MutexMap<OwnedRoomId, ()>,
59}
60
61pub type ShortStateHash = ShortId;
65
66impl crate::Service for Service {
67 fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
68 Ok(Arc::new(Self {
69 db: Data {
70 eventid_shorteventid: args.db["eventid_shorteventid"].clone(),
71 shorteventid_eventid: args.db["shorteventid_eventid"].clone(),
72 statekey_shortstatekey: args.db["statekey_shortstatekey"].clone(),
73 shortstatekey_statekey: args.db["shortstatekey_statekey"].clone(),
74 roomid_shortroomid: args.db["roomid_shortroomid"].clone(),
75 statehash_shortstatehash: args.db["statehash_shortstatehash"].clone(),
76 },
77 creating: Creating::default(),
78 services: args.services.clone(),
79 }))
80 }
81
82 fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
83}
84
85#[implement(Service)]
94pub async fn get_or_create_shorteventid(&self, event_id: &EventId) -> ShortEventId {
95 if let Ok(shorteventid) = self.get_shorteventid(event_id).await {
96 return shorteventid;
97 }
98
99 self.create_shorteventid(event_id).await
100}
101
102#[implement(Service)]
113pub fn multi_get_or_create_shorteventid<'a, I>(
114 &'a self,
115 event_ids: I,
116) -> impl Stream<Item = ShortEventId> + Send + '_
117where
118 I: Iterator<Item = &'a EventId> + Clone + Send + 'a,
119{
120 event_ids
121 .clone()
122 .stream()
123 .get(&self.db.eventid_shorteventid)
124 .zip(event_ids.into_iter().stream())
125 .wide_then(async |(result, event_id)| match result {
126 | Ok(ref short) => utils::u64_from_u8(short),
127 | Err(_) => self.create_shorteventid(event_id).await,
128 })
129}
130
131#[implement(Service)]
132async fn create_shorteventid(&self, event_id: &EventId) -> ShortEventId {
133 let _lock = self.creating.shorteventid.lock(event_id).await;
134
135 if let Ok(shorteventid) = self.get_shorteventid(event_id).await {
136 return shorteventid;
137 }
138
139 let short = self.services.globals.next_count();
140 let mut txn = self.services.db.txn();
141
142 txn.insert_raw(&self.db.shorteventid_eventid, (*short).to_be_bytes(), event_id);
143 txn.insert_raw(&self.db.eventid_shorteventid, event_id, (*short).to_be_bytes());
144 txn.execute();
145
146 *short
147}
148
149#[implement(Service)]
153pub async fn get_shorteventid(&self, event_id: &EventId) -> Result<ShortEventId> {
154 self.db
155 .eventid_shorteventid
156 .get(event_id)
157 .await
158 .deserialized()
159}
160
161#[implement(Service)]
170pub async fn get_or_create_shortstatekey(
171 &self,
172 event_type: &StateEventType,
173 state_key: &str,
174) -> ShortStateKey {
175 if let Ok(shortstatekey) = self
176 .get_shortstatekey(event_type, state_key)
177 .await
178 {
179 return shortstatekey;
180 }
181
182 self.create_shortstatekey(event_type, state_key)
183 .await
184}
185
186#[implement(Service)]
187async fn create_shortstatekey(
188 &self,
189 event_type: &StateEventType,
190 state_key: &str,
191) -> ShortStateKey {
192 let owned_key = (event_type.clone(), StateKey::from_str(state_key));
193 let _lock = self.creating.shortstatekey.lock(&owned_key).await;
194
195 if let Ok(shortstatekey) = self
196 .get_shortstatekey(event_type, state_key)
197 .await
198 {
199 return shortstatekey;
200 }
201
202 let key = (event_type, state_key);
203 let shortstatekey = self.services.globals.next_count();
204 let mut txn = self.services.db.txn();
205
206 txn.put(&self.db.shortstatekey_statekey, *shortstatekey, key);
207 txn.put(&self.db.statekey_shortstatekey, key, *shortstatekey);
208 txn.execute();
209
210 *shortstatekey
211}
212
213#[implement(Service)]
218pub async fn get_shortstatekey(
219 &self,
220 event_type: &StateEventType,
221 state_key: &str,
222) -> Result<ShortStateKey> {
223 let key = (event_type, state_key);
224 self.db
225 .statekey_shortstatekey
226 .qry(&key)
227 .await
228 .deserialized()
229}
230
231#[implement(Service)]
236pub async fn get_eventid_from_short<Id>(&self, shorteventid: ShortEventId) -> Result<Id>
237where
238 Id: for<'de> Deserialize<'de> + Send + Sized + ToOwned,
239 <Id as ToOwned>::Owned: Borrow<EventId>,
240{
241 const BUFSIZE: usize = size_of::<ShortEventId>();
242
243 self.db
244 .shorteventid_eventid
245 .aqry::<BUFSIZE, _>(&shorteventid)
246 .await
247 .deserialized()
248 .map_err(|e| err!(Database("Failed to find EventId from short {shorteventid:?}: {e:?}")))
249}
250
251#[implement(Service)]
255pub fn multi_get_eventid_from_short<'a, Id, S>(
256 &'a self,
257 shorteventid: S,
258) -> impl Stream<Item = Result<Id>> + Send + 'a
259where
260 S: Stream<Item = ShortEventId> + Send + 'a,
261 Id: for<'de> Deserialize<'de> + Send + Sized + ToOwned + 'a,
262 <Id as ToOwned>::Owned: Borrow<EventId>,
263{
264 shorteventid
265 .qry(&self.db.shorteventid_eventid)
266 .map(Deserialized::deserialized)
267}
268
269#[implement(Service)]
274pub async fn get_statekey_from_short(
275 &self,
276 shortstatekey: ShortStateKey,
277) -> Result<(StateEventType, StateKey)> {
278 const BUFSIZE: usize = size_of::<ShortStateKey>();
279
280 self.db
281 .shortstatekey_statekey
282 .aqry::<BUFSIZE, _>(&shortstatekey)
283 .await
284 .deserialized()
285 .map_err(|e| {
286 err!(Database(
287 "Failed to find (StateEventType, state_key) from short {shortstatekey:?}: {e:?}"
288 ))
289 })
290}
291
292#[implement(Service)]
296pub fn multi_get_statekey_from_short<'a, S>(
297 &'a self,
298 shortstatekey: S,
299) -> impl Stream<Item = Result<(StateEventType, StateKey)>> + Send + 'a
300where
301 S: Stream<Item = ShortStateKey> + Send + 'a,
302{
303 shortstatekey
304 .qry(&self.db.shortstatekey_statekey)
305 .map(Deserialized::deserialized)
306}
307
308#[implement(Service)]
322pub async fn get_or_create_shortstatehash<F>(
323 &self,
324 state_hash: &Digest,
325 write_statediff: F,
326) -> Result<(ShortStateHash, bool)>
327where
328 F: FnOnce(&mut Txn, ShortStateHash) -> Result,
329{
330 if let Ok(shortstatehash) = self.get_shortstatehash(state_hash).await {
331 return Ok((shortstatehash, true));
332 }
333
334 self.create_shortstatehash(state_hash, write_statediff)
335 .await
336}
337
338#[implement(Service)]
339async fn create_shortstatehash<F>(
340 &self,
341 state_hash: &Digest,
342 write_statediff: F,
343) -> Result<(ShortStateHash, bool)>
344where
345 F: FnOnce(&mut Txn, ShortStateHash) -> Result,
346{
347 let _lock = self
348 .creating
349 .shortstatehash
350 .lock(state_hash)
351 .await;
352
353 if let Ok(shortstatehash) = self.get_shortstatehash(state_hash).await {
354 return Ok((shortstatehash, true));
355 }
356
357 let shortstatehash = self.services.globals.next_count();
358 let mut txn = self.services.db.txn();
359
360 txn.insert_raw(
361 &self.db.statehash_shortstatehash,
362 state_hash,
363 (*shortstatehash).to_be_bytes(),
364 );
365 write_statediff(&mut txn, *shortstatehash)?;
366 txn.execute();
367
368 Ok((*shortstatehash, false))
369}
370
371#[implement(Service)]
375pub async fn get_shortstatehash(&self, state_hash: &Digest) -> Result<ShortStateHash> {
376 self.db
377 .statehash_shortstatehash
378 .get(state_hash)
379 .await
380 .deserialized()
381}
382
383#[implement(Service)]
387pub async fn get_shortroomid(&self, room_id: &RoomId) -> Result<ShortRoomId> {
388 self.db
389 .roomid_shortroomid
390 .get(room_id)
391 .await
392 .deserialized()
393}
394
395#[implement(Service)]
400pub async fn get_roomid_from_short(&self, shortroomid_: ShortRoomId) -> Result<OwnedRoomId> {
401 let stream = self.iter_shortroomids();
402
403 pin_mut!(stream);
404 stream
405 .ready_find(|&(_, shortroomid)| shortroomid == shortroomid_)
406 .map(|found| found.map(|(room_id, _)| room_id.to_owned()))
407 .await
408 .ok_or_else(|| err!(Database("Failed to find RoomId from {shortroomid_:?}")))
409}
410
411#[implement(Service)]
417pub fn iter_shortroomids(&self) -> impl Stream<Item = (&RoomId, ShortRoomId)> + Send + '_ {
418 self.db
419 .roomid_shortroomid
420 .stream()
421 .ready_filter_map(Result::ok)
422}
423
424#[implement(Service)]
433pub async fn get_or_create_shortroomid(&self, room_id: &RoomId) -> ShortRoomId {
434 if let Ok(shortroomid) = self.get_shortroomid(room_id).await {
435 return shortroomid;
436 }
437
438 self.create_shortroomid(room_id).await
439}
440
441#[implement(Service)]
442async fn create_shortroomid(&self, room_id: &RoomId) -> ShortRoomId {
443 const BUFSIZE: usize = size_of::<ShortRoomId>();
444
445 let _lock = self.creating.shortroomid.lock(room_id).await;
446
447 if let Ok(shortroomid) = self.get_shortroomid(room_id).await {
448 return shortroomid;
449 }
450
451 let short = self.services.globals.next_count();
452
453 debug_assert!(size_of_val(&*short) == BUFSIZE, "buffer requirement changed");
454
455 self.db
456 .roomid_shortroomid
457 .raw_aput::<BUFSIZE, _, _>(room_id, *short);
458
459 *short
460}
461
462#[implement(Service)]
467pub async fn delete_shortroomid(&self, room_id: &RoomId) -> Result {
468 if self
469 .db
470 .roomid_shortroomid
471 .exists(room_id)
472 .await
473 .is_ok()
474 {
475 self.db.roomid_shortroomid.remove(room_id);
476 Ok(())
477 } else {
478 Err!(Database("not found"))
479 }
480}