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
48type RelatedEvents = BTreeMap<String, FlattenedJson>;
50
51pub struct Evaluate<'a, 'b> {
53 pub user: &'b UserId,
55 pub ruleset: &'a Ruleset,
57 pub power_levels: Option<&'b RoomPowerLevels>,
59 pub pdu: &'b Raw<AnySyncTimelineEvent>,
61 pub room_id: &'b RoomId,
63 pub related_events: Option<&'b Arc<RelatedEvents>>,
65}
66
67const IN_REPLY_TO: &str = "m.in_reply_to";
69
70static 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#[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#[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 .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}