1use std::{collections::BTreeMap, fmt::Debug};
2
3use futures::{StreamExt, future::join3, stream::select};
4use ruma::{EventId, OwnedEventId, RoomId, UserId, events::receipt::ReceiptThread};
5use serde::Serialize;
6use tuwunel_core::{
7 Result, implement, trace,
8 utils::{
9 stream::{BroadbandExt, ReadyExt, TryIgnore},
10 u64_from_u8,
11 },
12};
13use tuwunel_database::{
14 Deserialized, Ignore, IgnoreAll, Interfix, KeyBuf, deserialize_from_slice as deserialize_key,
15};
16
17type ThreadCounts = BTreeMap<OwnedEventId, (u64, u64)>;
19
20type ThreadLastReads = BTreeMap<OwnedEventId, u64>;
24
25#[implement(super::Service)]
30#[tracing::instrument(level = "debug", skip(self))]
31pub async fn reset_notification_counts(&self, user_id: &UserId, room_id: &RoomId) {
32 let count = self.services.globals.next_count();
33
34 let userroom_id = (user_id, room_id);
35
36 self.reset_notification_count(room_id, user_id, userroom_id)
37 .await;
38
39 self.db
40 .userroomid_highlightcount
41 .put(userroom_id, 0_u64);
42
43 let roomuser_id = (room_id, user_id);
44 self.db
45 .roomuserid_lastnotificationread
46 .put(roomuser_id, *count);
47
48 let removed = self.clear_suppressed_room(user_id, room_id);
49 if removed > 0 {
50 trace!(?user_id, ?room_id, removed, "Cleared suppressed push events after read");
51 }
52}
53
54#[implement(super::Service)]
55async fn reset_notification_count<K>(&self, room_id: &RoomId, user_id: &UserId, key: K)
56where
57 K: Serialize + Debug + Send + Sync,
58{
59 let _lock = self
62 .notification_increment_mutex
63 .lock(&(room_id.to_owned(), user_id.to_owned()))
64 .await;
65
66 self.db
67 .userroomid_notificationcount
68 .put(key, 0_u64);
69}
70
71#[implement(super::Service)]
75#[tracing::instrument(level = "debug", skip(self))]
76pub async fn reset_thread_notification_counts(
77 &self,
78 user_id: &UserId,
79 room_id: &RoomId,
80 thread_root: &EventId,
81) {
82 let count = self.services.globals.next_count();
83
84 let userroom_thread = (user_id, room_id, thread_root);
85
86 self.reset_notification_count(room_id, user_id, userroom_thread)
87 .await;
88
89 self.db
90 .userroomid_highlightcount
91 .put(userroom_thread, 0_u64);
92
93 let roomuser_thread = (room_id, user_id, thread_root);
94 self.db
95 .roomuserid_lastnotificationread
96 .put(roomuser_thread, *count);
97}
98
99#[implement(super::Service)]
105#[tracing::instrument(level = "debug", skip(self))]
106pub async fn clear_all_thread_notification_counts(&self, user_id: &UserId, room_id: &RoomId) {
107 let userroom_prefix = (user_id, room_id, Interfix);
108 let roomuser_prefix = (room_id, user_id, Interfix);
109
110 let highlights = self
111 .db
112 .userroomid_highlightcount
113 .del_prefix(&userroom_prefix);
114
115 let last_reads = self
116 .db
117 .roomuserid_lastnotificationread
118 .del_prefix(&roomuser_prefix);
119
120 let notifications = async {
121 let _lock = self
122 .notification_increment_mutex
123 .lock(&(room_id.to_owned(), user_id.to_owned()))
124 .await;
125
126 self.db
127 .userroomid_notificationcount
128 .del_prefix(&userroom_prefix)
129 .await;
130 };
131
132 join3(notifications, highlights, last_reads).await;
133}
134
135#[implement(super::Service)]
141pub async fn reset_notification_counts_for_thread(
142 &self,
143 user_id: &UserId,
144 room_id: &RoomId,
145 acknowledged: Option<&EventId>,
146 thread: &ReceiptThread,
147) {
148 match thread {
149 | ReceiptThread::Thread(root) if acknowledged == Some(root) => {},
150 | ReceiptThread::Main =>
151 self.reset_notification_counts(user_id, room_id)
152 .await,
153 | ReceiptThread::Thread(root) =>
154 self.reset_thread_notification_counts(user_id, room_id, root)
155 .await,
156 | _ => {
157 self.reset_notification_counts(user_id, room_id)
158 .await;
159
160 self.clear_all_thread_notification_counts(user_id, room_id)
161 .await;
162 },
163 }
164}
165
166#[implement(super::Service)]
167#[tracing::instrument(level = "debug", skip(self), ret(level = "trace"))]
168pub async fn notification_count(&self, user_id: &UserId, room_id: &RoomId) -> u64 {
169 let key = (user_id, room_id);
170 self.db
171 .userroomid_notificationcount
172 .qry(&key)
173 .await
174 .deserialized()
175 .unwrap_or(0)
176}
177
178#[implement(super::Service)]
182#[tracing::instrument(level = "trace", skip(self), ret)]
183pub async fn global_notification_count(&self, user_id: &UserId) -> u64 {
184 self.db
185 .userroomid_notificationcount
186 .stream_prefix_raw(&(user_id, Interfix))
187 .ignore_err()
188 .ready_filter_map(|(key, count)| {
189 let count = u64_from_u8(count);
190
191 (count > 0).then(|| (KeyBuf::from(key), count))
192 })
193 .broad_filter_map(async |(key, count)| {
194 let (_, room_id, _): (Ignore, &RoomId, IgnoreAll) =
195 deserialize_key(&key).expect("notification count key");
196
197 self.services
198 .state_cache
199 .is_joined(user_id, room_id)
200 .await
201 .then_some(count)
202 })
203 .ready_fold(0_u64, u64::saturating_add)
204 .await
205}
206
207#[implement(super::Service)]
208#[tracing::instrument(level = "debug", skip(self), ret(level = "trace"))]
209pub async fn highlight_count(&self, user_id: &UserId, room_id: &RoomId) -> u64 {
210 let key = (user_id, room_id);
211 self.db
212 .userroomid_highlightcount
213 .qry(&key)
214 .await
215 .deserialized()
216 .unwrap_or(0)
217}
218
219#[implement(super::Service)]
223#[tracing::instrument(level = "debug", skip(self))]
224pub async fn thread_notification_counts(
225 &self,
226 user_id: &UserId,
227 room_id: &RoomId,
228) -> ThreadCounts {
229 let prefix = (user_id, room_id, Interfix);
230 let notifications = self
231 .db
232 .userroomid_notificationcount
233 .stream_prefix(&prefix)
234 .ignore_err()
235 .map(notification_kv);
236
237 let highlights = self
238 .db
239 .userroomid_highlightcount
240 .stream_prefix(&prefix)
241 .ignore_err()
242 .map(highlight_kv);
243
244 select(notifications, highlights)
245 .ready_fold(ThreadCounts::default(), merge_thread_count)
246 .await
247}
248
249fn notification_kv(
250 (key, notifications): ((&UserId, &RoomId, OwnedEventId), u64),
251) -> (OwnedEventId, (u64, u64)) {
252 (key.2, (notifications, 0))
253}
254
255fn highlight_kv(
256 (key, highlights): ((&UserId, &RoomId, OwnedEventId), u64),
257) -> (OwnedEventId, (u64, u64)) {
258 (key.2, (0, highlights))
259}
260
261fn merge_thread_count(
262 mut counts: ThreadCounts,
263 (root, (notifications, highlights)): (OwnedEventId, (u64, u64)),
264) -> ThreadCounts {
265 let entry = counts.entry(root).or_default();
266 entry.0 = entry.0.saturating_add(notifications);
267 entry.1 = entry.1.saturating_add(highlights);
268 counts
269}
270
271#[implement(super::Service)]
272#[tracing::instrument(level = "debug", skip(self), ret(level = "trace"))]
273pub async fn last_notification_read(&self, user_id: &UserId, room_id: &RoomId) -> Result<u64> {
274 let key = (room_id, user_id);
275 self.db
276 .roomuserid_lastnotificationread
277 .qry(&key)
278 .await
279 .deserialized()
280}
281
282#[implement(super::Service)]
286#[tracing::instrument(level = "debug", skip(self))]
287pub async fn thread_last_notification_reads(
288 &self,
289 user_id: &UserId,
290 room_id: &RoomId,
291) -> ThreadLastReads {
292 let prefix = (room_id, user_id, Interfix);
293 self.db
294 .roomuserid_lastnotificationread
295 .stream_prefix(&prefix)
296 .ignore_err()
297 .map(|((_, _, root), count): ((Ignore, Ignore, OwnedEventId), u64)| (root, count))
298 .collect()
299 .await
300}
301
302#[implement(super::Service)]
303pub async fn delete_room_notification_read(&self, room_id: &RoomId) -> Result {
304 let key = (room_id, Interfix);
305 self.db
306 .roomuserid_lastnotificationread
307 .keys_prefix_raw(&key)
308 .ignore_err()
309 .ready_for_each(|key| {
310 trace!("Removing key: {key:?}");
311 self.db
312 .roomuserid_lastnotificationread
313 .remove(key);
314 })
315 .await;
316
317 Ok(())
318}