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#[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 match &pusher.kind {
126 | PusherKind::Http(http) =>
127 self.send_http_event_notice(user_id, pusher, http, tweaks, event)
128 .await,
129 | _ => 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 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#[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}