Skip to main content

tuwunel_service/pusher/
send.rs

1use futures::{
2	FutureExt,
3	future::{join, join4},
4};
5use ruma::{
6	UInt, UserId,
7	api::{
8		client::push::{Pusher, PusherKind},
9		push_gateway::send_event_notification::v1::{
10			Device, Notification, NotificationCounts, NotificationPriority, Request,
11		},
12	},
13	events::TimelineEventType,
14	push::{Action, HighlightTweakValue, HttpPusherData, PushFormat, Ruleset, Tweak},
15};
16use serde_json::Value;
17use tuwunel_core::{Result, err, error, implement, matrix::Event, trace, utils::BoolExt};
18use url::Url;
19
20use super::Evaluate;
21
22#[implement(super::Service)]
23#[tracing::instrument(level = "debug", skip_all)]
24pub async fn send_push_notice<E>(
25	&self,
26	user_id: &UserId,
27	pusher: &Pusher,
28	ruleset: &Ruleset,
29	event: &E,
30) -> Result
31where
32	E: Event,
33{
34	let power_levels = self
35		.services
36		.state_accessor
37		.get_power_levels(event.room_id())
38		.map(Result::ok);
39
40	let (power_levels, related_events) = join(power_levels, self.related_events(event)).await;
41
42	let serialized = event.to_format();
43	let actions = self
44		.get_actions(Evaluate {
45			user: user_id,
46			ruleset,
47			power_levels: power_levels.as_ref(),
48			pdu: &serialized,
49			room_id: event.room_id(),
50			related_events: related_events.as_ref(),
51		})
52		.await;
53
54	let notify = actions.iter().any(Action::should_notify);
55	let tweak_count = actions
56		.iter()
57		.filter(|action| matches!(action, Action::SetTweak(_)))
58		.count();
59
60	trace!(
61		%user_id,
62		event_id = %event.event_id(),
63		actions = %actions.len(),
64		notify,
65		tweaks = tweak_count,
66		"Push notice decision",
67	);
68
69	if notify || self.services.config.push_everything {
70		let tweaks: Vec<Tweak> = actions
71			.iter()
72			.filter_map(|action| match action {
73				| Action::SetTweak(tweak) => Some(tweak.clone()),
74				| _ => None,
75			})
76			.collect();
77
78		self.send_notice(user_id, pusher, tweaks, event)
79			.await?;
80	}
81
82	Ok(())
83}
84
85/// Send an account-wide counts-only notification to a push gateway.
86///
87/// Enabled HTTP pushers emit the request, including an explicit zero. The
88/// delivery is skipped only when the gateway is known to hold the current
89/// total already; an unknown gateway is always sent to.
90#[implement(super::Service)]
91#[tracing::instrument(level = "debug", skip_all)]
92pub async fn send_badge_notice(&self, user_id: &UserId, pusher: &Pusher) -> Result {
93	let PusherKind::Http(http) = &pusher.kind else {
94		return Ok(());
95	};
96
97	if badge_count_disabled(http) {
98		return Ok(());
99	}
100
101	let unread = UInt::new(self.global_notification_count(user_id).await).unwrap_or(UInt::MAX);
102
103	if self.sent_badge(user_id, &pusher.ids.pushkey) == Some(unread) {
104		return Ok(());
105	}
106
107	let device = self.prepare_http_pusher(pusher, http)?;
108	let mut notify = Notification::new(vec![device]);
109	notify.counts = NotificationCounts::new_explicit(Some(unread), None);
110
111	self.send_http_notice(user_id, pusher, http, notify, Some(unread))
112		.await
113}
114
115#[implement(super::Service)]
116#[tracing::instrument(level = "debug", skip_all)]
117async fn send_notice<Pdu: Event>(
118	&self,
119	user_id: &UserId,
120	pusher: &Pusher,
121	tweaks: Vec<Tweak>,
122	event: &Pdu,
123) -> Result {
124	// TODO: email
125	match &pusher.kind {
126		| PusherKind::Http(http) =>
127			self.send_http_event_notice(user_id, pusher, http, tweaks, event)
128				.await,
129		// TODO: Handle email
130		//PusherKind::Email(_) => Ok(()),
131		| _ => Ok(()),
132	}
133}
134
135#[implement(super::Service)]
136async fn send_http_event_notice<Pdu: Event>(
137	&self,
138	user_id: &UserId,
139	pusher: &Pusher,
140	http: &HttpPusherData,
141	tweaks: Vec<Tweak>,
142	event: &Pdu,
143) -> Result {
144	let mut device = self.prepare_http_pusher(pusher, http)?;
145
146	// TODO (timo): can pusher/devices have conflicting formats
147	let event_id_only = http.format == Some(PushFormat::EventIdOnly);
148
149	if !event_id_only {
150		device.tweaks.clone_from(&tweaks);
151	}
152
153	let mut notify = Notification::new(vec![device]);
154
155	notify.event_id = Some(event.event_id().to_owned());
156	notify.room_id = Some(event.room_id().to_owned());
157
158	let unread = badge_count_disabled(http)
159		.is_false()
160		.then_async(async || {
161			UInt::new(self.global_notification_count(user_id).await).unwrap_or(UInt::MAX)
162		});
163
164	let unread = if !event_id_only {
165		if *event.kind() == TimelineEventType::RoomEncrypted
166			|| tweaks.iter().any(|t| {
167				matches!(t, Tweak::Highlight(HighlightTweakValue::Yes) | Tweak::Sound(_))
168			}) {
169			notify.prio = NotificationPriority::High;
170		} else {
171			notify.prio = NotificationPriority::Low;
172		}
173		notify.sender = Some(event.sender().to_owned());
174		notify.event_type = Some(event.kind().to_owned());
175		notify.content = serde_json::value::to_raw_value(event.content()).ok();
176
177		if *event.kind() == TimelineEventType::RoomMember {
178			notify.user_is_target = event.state_key() == Some(user_id.as_str());
179		}
180
181		let (display_name, room_name, room_alias, unread) = join4(
182			self.services.profile.displayname(event.sender()),
183			self.services
184				.state_accessor
185				.get_name(event.room_id()),
186			self.services
187				.state_accessor
188				.get_canonical_alias(event.room_id()),
189			unread,
190		)
191		.await;
192
193		notify.sender_display_name = display_name.ok();
194		notify.room_name = room_name.ok();
195		notify.room_alias = room_alias.ok();
196
197		unread
198	} else {
199		unread.await
200	};
201
202	if let Some(unread) = unread {
203		notify.counts = NotificationCounts::new_explicit(Some(unread), None);
204	}
205
206	self.send_http_notice(user_id, pusher, http, notify, unread)
207		.await
208}
209
210#[implement(super::Service)]
211fn prepare_http_pusher(&self, pusher: &Pusher, http: &HttpPusherData) -> Result<Device> {
212	let address = &http.url;
213	let url = Url::parse(address).map_err(|e| {
214		err!(Request(InvalidParam(
215			warn!(url = %address, error = %e, "HTTP pusher URL is not a valid URL")
216		)))
217	})?;
218
219	self.check_http_pusher_url(&url)?;
220
221	let mut device = Device::new(pusher.ids.app_id.clone(), pusher.ids.pushkey.clone());
222	device.data.data.clone_from(&http.data);
223	device.data.format.clone_from(&http.format);
224
225	Ok(device)
226}
227
228/// Deliver one notification to the pusher's gateway and honor its verdict.
229///
230/// A pushkey the gateway names in `rejected` has its pusher removed. `unread`
231/// names the counts value on the wire; it is recorded as delivered only after
232/// the gateway accepts, so a failed or rejected send leaves the next refresh
233/// unconditional.
234#[implement(super::Service)]
235#[tracing::instrument(level = "debug", skip_all)]
236async fn send_http_notice(
237	&self,
238	user_id: &UserId,
239	pusher: &Pusher,
240	http: &HttpPusherData,
241	notify: Notification,
242	unread: Option<UInt>,
243) -> Result {
244	let response = self
245		.send_request(&http.url, Request::new(notify))
246		.await?;
247
248	let pushkey = &pusher.ids.pushkey;
249
250	if response.rejected.contains(pushkey) {
251		error!(
252			url = %http.url,
253			%pushkey,
254			"Push gateway rejected the pushkey; removing the pusher. Push notifications \
255			 for this device stop until the client registers a new pusher.",
256		);
257
258		self.delete_pusher(user_id, pushkey).await;
259
260		return Ok(());
261	}
262
263	if let Some(unread) = unread {
264		self.record_sent_badge(user_id, pushkey, unread);
265	}
266
267	Ok(())
268}
269
270fn badge_count_disabled(http: &HttpPusherData) -> bool {
271	["org.matrix.msc4076.disable_badge_count", "disable_badge_count"]
272		.iter()
273		.any(|key| http.data.get(*key).and_then(Value::as_bool) == Some(true))
274}