Skip to main content

tuwunel_service/pusher/
notification.rs

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
17/// Per-thread unread counts: `(notification, highlight)` keyed by thread root.
18type ThreadCounts = BTreeMap<OwnedEventId, (u64, u64)>;
19
20/// Per-thread last-read counts keyed by thread root. Used by sync v3 to
21/// gate emission of `unread_thread_notifications` to threads whose read
22/// cursor advanced within the sync window.
23type ThreadLastReads = BTreeMap<OwnedEventId, u64>;
24
25/// Reset the room's main-timeline notification counts.
26///
27/// The last-read stamp gates sync output; callers dispatch the badge refresh
28/// after every reset.
29#[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	// The increment path is a read-modify-write under this lock; an unlocked
60	// zero could land inside it and be overwritten by the stale sum.
61	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/// Reset counts for a single thread within a room.
72///
73/// The last-read stamp gates sync output.
74#[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/// Clear all per-thread notification state for this user and room.
100///
101/// The `Interfix` prefix excludes the main row. The notification-count sweep
102/// runs under the increment mutex so a concurrent read-modify-write cannot
103/// resurrect a cleared row.
104#[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/// Dispatcher: route a receipt's `ReceiptThread` to the matching reset path.
136///
137/// `Unthreaded` clears all room and thread counts; `Main` clears only the
138/// main-timeline counts; `Thread(id)` clears just that thread unless the
139/// acknowledged event is the thread root. `None` denotes a non-receipt reset.
140#[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/// Return the user's account-wide unread notification count.
179///
180/// Joined main and thread rows contribute to a saturating total.
181#[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/// Per-thread `(notification, highlight)` counts for one room and user.
220/// `Interfix` excludes the legacy 2-tuple main row from the scan; only
221/// 3-tuple `(user, room, root)` rows match.
222#[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/// Per-thread last-read counts for one room and user. `Interfix` keeps the
283/// scan to 3-tuple `(room, user, root)` rows; the legacy 2-tuple main row
284/// is excluded by construction and lives behind `last_notification_read`.
285#[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}