Skip to main content

tuwunel_service/pusher/
mod.rs

1mod append;
2mod badge;
3mod notification;
4mod request;
5mod send;
6mod suppressed;
7#[cfg(test)]
8mod tests;
9
10use std::{
11	collections::BTreeMap,
12	sync::{Arc, LazyLock},
13};
14
15use futures::{Stream, StreamExt, TryFutureExt, future::join};
16use ruma::{
17	DeviceId, OwnedDeviceId, OwnedEventId, OwnedRoomId, OwnedUserId, RoomId, UserId,
18	api::client::push::{Pusher, PusherKind, set_pusher::v3::PusherAction},
19	events::{
20		AnySyncTimelineEvent, GlobalAccountDataEventType, push_rules::PushRulesEvent,
21		room::power_levels::RoomPowerLevels,
22	},
23	push::{Action, FlattenedJson, PushConditionPowerLevelsCtx, PushConditionRoomCtx, Ruleset},
24	serde::Raw,
25	uint,
26};
27use serde::Deserialize;
28use tracing::Level;
29use tuwunel_core::{
30	Err, Result, err, implement,
31	matrix::Event,
32	utils::{
33		MutexMap,
34		future::TryExtExt,
35		result::ErrLog,
36		stream::{BroadbandExt, IterStream, ReadyExt, TryIgnore, WidebandExt},
37	},
38};
39use tuwunel_database::{Database, Deserialized, Ignore, Interfix, Json, Map};
40use url::Url;
41
42use self::badge::SentBadges;
43pub use self::{
44	append::Notified,
45	suppressed::{SuppressedPushes, SuppressedRooms},
46};
47
48/// The events an event relates to, keyed by relation type, for MSC3664.
49type RelatedEvents = BTreeMap<String, FlattenedJson>;
50
51/// The inputs for evaluating one user's push rules against one event.
52pub struct Evaluate<'a, 'b> {
53	/// The user whose rules are evaluated, and who is notified.
54	pub user: &'b UserId,
55	/// The user's ruleset, which outlives the borrowed actions returned.
56	pub ruleset: &'a Ruleset,
57	/// The room's power levels, without which some conditions never match.
58	pub power_levels: Option<&'b RoomPowerLevels>,
59	/// The event being evaluated, in the format the conditions flatten.
60	pub pdu: &'b Raw<AnySyncTimelineEvent>,
61	/// The room the event was sent to.
62	pub room_id: &'b RoomId,
63	/// The events it relates to, or `None` where they are not resolved at all.
64	pub related_events: Option<&'b Arc<RelatedEvents>>,
65}
66
67/// MSC3664 matches a reply as if it carried a relation of this type.
68const IN_REPLY_TO: &str = "m.in_reply_to";
69
70/// Shared by every event that resolves no relations, which is every event
71/// while MSC3664 evaluation is disabled.
72static NO_RELATED_EVENTS: LazyLock<Arc<RelatedEvents>> = LazyLock::new(Arc::default);
73
74#[derive(Deserialize)]
75struct ExtractRelatesTo {
76	#[serde(rename = "m.relates_to")]
77	relates_to: RelatesTo,
78}
79
80#[derive(Deserialize)]
81struct RelatesTo {
82	rel_type: Option<String>,
83
84	event_id: Option<OwnedEventId>,
85
86	#[serde(rename = "m.in_reply_to")]
87	in_reply_to: Option<InReplyTo>,
88}
89
90#[derive(Deserialize)]
91struct InReplyTo {
92	event_id: OwnedEventId,
93}
94
95pub struct Service {
96	services: Arc<crate::services::OnceServices>,
97	notification_increment_mutex: MutexMap<(OwnedRoomId, OwnedUserId), ()>,
98	highlight_increment_mutex: MutexMap<(OwnedRoomId, OwnedUserId), ()>,
99	db: Data,
100	suppressed: suppressed::SuppressedQueue,
101	sent_badges: SentBadges,
102}
103
104struct Data {
105	db: Arc<Database>,
106	senderkey_pusher: Arc<Map>,
107	pushkey_deviceid: Arc<Map>,
108	useridcount_notification: Arc<Map>,
109	userroomid_highlightcount: Arc<Map>,
110	userroomid_notificationcount: Arc<Map>,
111	roomuserid_lastnotificationread: Arc<Map>,
112}
113
114impl crate::Service for Service {
115	fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
116		Ok(Arc::new(Self {
117			services: args.services.clone(),
118			notification_increment_mutex: MutexMap::new(),
119			highlight_increment_mutex: MutexMap::new(),
120			db: Data {
121				db: args.db.clone(),
122				senderkey_pusher: args.db["senderkey_pusher"].clone(),
123				pushkey_deviceid: args.db["pushkey_deviceid"].clone(),
124				useridcount_notification: args.db["useridcount_notification"].clone(),
125				userroomid_highlightcount: args.db["userroomid_highlightcount"].clone(),
126				userroomid_notificationcount: args.db["userroomid_notificationcount"].clone(),
127				roomuserid_lastnotificationread: args.db["roomuserid_lastnotificationread"]
128					.clone(),
129			},
130			suppressed: suppressed::SuppressedQueue::default(),
131			sent_badges: SentBadges::default(),
132		}))
133	}
134
135	fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
136}
137
138#[implement(Service)]
139pub async fn set_pusher(
140	&self,
141	sender: &UserId,
142	sender_device: &DeviceId,
143	pusher: &PusherAction,
144) -> Result {
145	match pusher {
146		| PusherAction::Delete(ids) =>
147			self.set_pusher_delete(sender, ids.pushkey.as_str())
148				.await,
149		| PusherAction::Post(data) =>
150			self.set_pusher_post(sender, sender_device, pusher, &data.pusher)?,
151	}
152
153	Ok(())
154}
155
156#[implement(Service)]
157async fn set_pusher_delete(&self, sender: &UserId, pushkey: &str) {
158	self.delete_pusher(sender, pushkey).await;
159}
160
161#[implement(Service)]
162fn set_pusher_post(
163	&self,
164	sender: &UserId,
165	sender_device: &DeviceId,
166	action: &PusherAction,
167	pusher: &Pusher,
168) -> Result {
169	let pushkey = pusher.ids.pushkey.as_str();
170
171	if pushkey.len() > 512 {
172		return Err!(Request(InvalidParam("Push key length cannot be greater than 512 bytes.")));
173	}
174
175	if pusher.ids.app_id.as_str().len() > 64 {
176		return Err!(Request(InvalidParam("App ID length cannot be greater than 64 bytes.")));
177	}
178
179	if let PusherKind::Http(http) = &pusher.kind {
180		let url = &http.url;
181		let url = Url::parse(&http.url).map_err(|e| {
182			err!(Request(InvalidParam(warn!(%url, "HTTP pusher URL is not a valid URL: {e}"))))
183		})?;
184
185		self.check_http_pusher_url(&url)?;
186	}
187
188	let key = (sender, pushkey);
189	self.db.senderkey_pusher.put(key, Json(action));
190	self.db
191		.pushkey_deviceid
192		.insert(pushkey, sender_device);
193
194	self.forget_sent_badge(sender, pushkey);
195
196	Ok(())
197}
198
199#[implement(Service)]
200fn check_http_pusher_url(&self, url: &Url) -> Result {
201	if ["http", "https"]
202		.iter()
203		.all(|&scheme| !scheme.eq_ignore_ascii_case(url.scheme()))
204	{
205		return Err!(Request(InvalidParam(
206			warn!(%url, "HTTP pusher URL is not a valid HTTP/HTTPS URL")
207		)));
208	}
209
210	if self.services.client.proxy.resolver_alias(url) {
211		return Err!(Request(InvalidParam(
212			warn!(%url, "HTTP pusher URL is a forbidden proxy endpoint")
213		)));
214	}
215
216	if !self.services.client.valid_cidr_range_url(url) {
217		return Err!(Request(InvalidParam(
218			warn!(%url, "HTTP pusher URL is a forbidden remote address")
219		)));
220	}
221
222	Ok(())
223}
224
225#[implement(Service)]
226pub async fn delete_pusher(&self, sender: &UserId, pushkey: &str) {
227	let key = (sender, pushkey);
228	self.db.senderkey_pusher.del(key);
229	self.db.pushkey_deviceid.remove(pushkey);
230	self.clear_suppressed_pushkey(sender, pushkey);
231	self.forget_sent_badge(sender, pushkey);
232
233	self.services
234		.sending
235		.cleanup_events(None, Some(sender), Some(pushkey))
236		.await
237		.ok();
238}
239
240#[implement(Service)]
241pub async fn get_device_pushkeys(&self, sender: &UserId, device_id: &DeviceId) -> Vec<String> {
242	self.get_pushkeys(sender)
243		.map(ToOwned::to_owned)
244		.broad_filter_map(async |pushkey| {
245			self.get_pusher_device(&pushkey)
246				.await
247				.ok()
248				.as_ref()
249				.is_some_and(|pusher_device| pusher_device == device_id)
250				.then_some(pushkey)
251		})
252		.collect()
253		.await
254}
255
256#[implement(Service)]
257pub async fn get_pusher_device(&self, pushkey: &str) -> Result<OwnedDeviceId> {
258	self.db
259		.pushkey_deviceid
260		.get(pushkey)
261		.await
262		.deserialized()
263}
264
265#[implement(Service)]
266pub async fn get_pusher(&self, sender: &UserId, pushkey: &str) -> Result<Pusher> {
267	let senderkey = (sender, pushkey);
268	self.db
269		.senderkey_pusher
270		.qry(&senderkey)
271		.await
272		.deserialized()
273}
274
275#[implement(Service)]
276pub async fn get_pushers(&self, sender: &UserId) -> Vec<Pusher> {
277	let prefix = (sender, Interfix);
278	self.db
279		.senderkey_pusher
280		.stream_prefix(&prefix)
281		.ignore_err()
282		.map(|(_, pusher): (Ignore, Pusher)| pusher)
283		.collect()
284		.await
285}
286
287#[implement(Service)]
288pub fn get_pushkeys<'a>(&'a self, sender: &'a UserId) -> impl Stream<Item = &str> + Send + 'a {
289	let prefix = (sender, Interfix);
290	self.db
291		.senderkey_pusher
292		.keys_prefix(&prefix)
293		.ignore_err()
294		.map(|(_, pushkey): (Ignore, &str)| pushkey)
295}
296
297/// The push ruleset a user's notifications are evaluated against.
298///
299/// A user without stored rules gets the server default ruleset.
300#[implement(Service)]
301#[tracing::instrument(skip(self), level = "debug")]
302pub async fn ruleset(&self, user_id: &UserId) -> Ruleset {
303	self.services
304		.account_data
305		.get_global(user_id, GlobalAccountDataEventType::PushRules)
306		.await
307		.log_err(Level::TRACE)
308		.map_or_else(|_| Ruleset::server_default(user_id), |ev: PushRulesEvent| ev.content.global)
309}
310
311#[implement(Service)]
312#[tracing::instrument(level = "debug", skip_all)]
313pub fn get_notifications<'a>(
314	&'a self,
315	sender: &'a UserId,
316	from: Option<u64>,
317) -> impl Stream<Item = (u64, Notified)> + Send + 'a {
318	let from = from
319		.map(|from| from.saturating_sub(1))
320		.unwrap_or(u64::MAX);
321
322	self.db
323		.useridcount_notification
324		.rev_stream_from(&(sender, from))
325		.ignore_err()
326		.map(|item: ((&UserId, u64), _)| (item.0, item.1))
327		.ready_take_while(move |((user_id, _count), _)| sender == *user_id)
328		.map(|((_, count), notified)| (count, notified))
329}
330
331#[implement(Service)]
332#[tracing::instrument(level = "debug", skip_all)]
333pub async fn get_actions<'a>(
334	&self,
335	Evaluate {
336		user,
337		ruleset,
338		power_levels,
339		pdu,
340		room_id,
341		related_events,
342	}: Evaluate<'a, '_>,
343) -> &'a [Action] {
344	let user_display_name = self
345		.services
346		.profile
347		.displayname(user)
348		.unwrap_or_else(|_| user.localpart().to_owned());
349
350	let room_joined_count = self
351		.services
352		.state_cache
353		.room_joined_count(room_id)
354		.map_ok(TryInto::try_into)
355		.map_ok(|res| res.unwrap_or_else(|_| uint!(1)))
356		.unwrap_or_default();
357
358	let (room_joined_count, user_display_name) = join(room_joined_count, user_display_name).await;
359
360	let power_levels = power_levels.map(|power_levels| PushConditionPowerLevelsCtx {
361		users: power_levels.users.clone(),
362		users_default: power_levels.users_default,
363		notifications: power_levels.notifications.clone(),
364		rules: power_levels.rules.clone(),
365	});
366
367	let ctx = PushConditionRoomCtx::new(
368		room_id.to_owned(),
369		room_joined_count,
370		user.to_owned(),
371		user_display_name,
372	);
373
374	let ctx = match related_events {
375		| Some(related_events) => ctx.with_related_events(related_events.clone()),
376		| None => ctx,
377	};
378
379	let ctx = match power_levels {
380		| Some(pl) => ctx.with_power_levels(pl),
381		| None => ctx,
382	};
383
384	ruleset.get_actions(pdu, &ctx).await
385}
386
387/// Resolve the events an event relates to, for the MSC3664 push condition.
388///
389/// `None` while `msc3664_related_event_match` is disabled, which leaves every
390/// MSC3664 condition unable to match. The relations do not vary by recipient,
391/// so the append path resolves them once for the whole room; the push-notice
392/// path resolves them per notice, since it holds one event at a time.
393#[implement(Service)]
394#[tracing::instrument(level = "debug", skip_all)]
395pub async fn related_events<E: Event>(&self, event: &E) -> Option<Arc<RelatedEvents>> {
396	let config = &self.services.server.config;
397
398	if !config.msc3664_related_event_match {
399		return None;
400	}
401
402	let Ok(ExtractRelatesTo { relates_to }) = event.get_content() else {
403		return Some(NO_RELATED_EVENTS.clone());
404	};
405
406	let reply = relates_to
407		.in_reply_to
408		.map(|reply| (IN_REPLY_TO.to_owned(), reply.event_id));
409
410	let related = relates_to
411		.rel_type
412		.zip(relates_to.event_id)
413		.into_iter()
414		.chain(reply)
415		.stream()
416		.wide_filter_map(async |(rel_type, event_id)| {
417			let related = self
418				.services
419				.timeline
420				.get_pdu(&event_id)
421				.await
422				.ok()
423				// A relation crossing rooms is not one the sender could have made, and
424				// following it would expose an event the recipient may not be in a room for.
425				.filter(|related| related.room_id() == event.room_id())?;
426
427			let related: Raw<AnySyncTimelineEvent> = related.to_format();
428
429			Some((rel_type, FlattenedJson::from_raw(&related)))
430		})
431		.collect()
432		.await;
433
434	Some(Arc::new(related))
435}