tuwunel_service/sending/sender/dispatch/
appservice.rs1use std::{
2 collections::{BTreeMap, BTreeSet},
3 str::from_utf8,
4};
5
6use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
7use futures::{StreamExt, future::join};
8use ruma::{
9 OneTimeKeyAlgorithm, OwnedDeviceId, OwnedUserId, UInt, UserId,
10 api::appservice::event::push_events::v1::{
11 AnyAppserviceToDeviceEvent, DeviceLists, EphemeralData, Request as PushEventsRequest,
12 },
13 events::AnyTimelineEvent,
14 serde::Raw,
15};
16use serde::Deserialize;
17use tuwunel_core::{
18 Event, debug_warn, err, implement,
19 itertools::Itertools,
20 smallvec::SmallVec,
21 utils::{
22 IterStream, ReadyExt, calculate_hash,
23 stream::{BroadbandExt, WidebandExt},
24 },
25};
26
27use super::SendingResult;
28use crate::{
29 appservice::RegistrationInfo,
30 sending::{Destination, SendingEvent, Service, TAG_PREFIX_LEN},
31};
32
33#[derive(Deserialize)]
38struct ToDeviceRecipient {
39 to_user_id: OwnedUserId,
40 to_device_id: OwnedDeviceId,
41}
42
43#[derive(Default)]
49struct Parts {
50 pdus: Vec<Raw<AnyTimelineEvent>>,
51 edus: Vec<Raw<EphemeralData>>,
52 to_device: Vec<Raw<AnyAppserviceToDeviceEvent>>,
53 changed: Vec<OwnedUserId>,
54 otk_users: BTreeSet<OwnedUserId>,
55 otk_recipients: Devices,
56}
57
58enum Part {
63 Pdu(Raw<AnyTimelineEvent>, Option<OwnedUserId>),
64 Edu(Raw<EphemeralData>),
65 ToDevice(Raw<AnyAppserviceToDeviceEvent>, Option<UserDevice>),
66 Changed(OwnedUserId),
67 None,
68}
69
70type UserDevice = (OwnedUserId, OwnedDeviceId);
74
75type OtkCounts =
80 BTreeMap<OwnedUserId, BTreeMap<OwnedDeviceId, BTreeMap<OneTimeKeyAlgorithm, UInt>>>;
81
82type FallbackTypes = BTreeMap<OwnedUserId, BTreeMap<OwnedDeviceId, Vec<OneTimeKeyAlgorithm>>>;
87
88type Devices = SmallVec<[UserDevice; 1]>;
93
94#[implement(Service)]
95#[tracing::instrument(
96 name = "appservice",
97 level = "debug",
98 skip(self, events),
99 fields(
100 events = %events.len(),
101 ),
102)]
103pub(super) async fn send_events_dest_appservice(
104 &self,
105 id: String,
106 events: Vec<SendingEvent>,
107) -> SendingResult {
108 let Some(info) = self
109 .services
110 .appservice
111 .get_registration_info(&id)
112 .await
113 else {
114 return Err((
116 Destination::Appservice(id.clone()),
117 err!(Database(debug_warn!(?id, "Missing appservice registration"))),
118 ));
119 };
120
121 let msc3202 = info.registration.msc3202_transaction_extensions;
122 let parts = Parts {
123 otk_users: msc3202
124 .then(|| info.sender.clone())
125 .into_iter()
126 .collect(),
127 ..Default::default()
128 };
129
130 let Parts {
131 pdus,
132 edus,
133 to_device,
134 changed,
135 otk_users,
136 otk_recipients,
137 } = events
138 .iter()
139 .stream()
140 .wide_then(async |event| self.txn_part(&info, msc3202, event).await)
141 .ready_fold(parts, Parts::merge)
142 .await;
143
144 let txn_hash = calculate_hash(events.iter().filter_map(|e| match e {
145 | SendingEvent::Edu(b)
146 | SendingEvent::ToDevice(b)
147 | SendingEvent::DeviceListChanged(b) => Some(b.as_ref()),
148 | SendingEvent::Pdu(b) => Some(b.as_ref()),
149 | SendingEvent::BadgeRefresh | SendingEvent::Flush => None,
150 }));
151
152 let (device_lists, device_one_time_keys_count, device_unused_fallback_key_types) = if msc3202
153 {
154 let changed = changed
155 .into_iter()
156 .sorted_unstable()
157 .dedup()
158 .collect();
159
160 let (counts, fallbacks) = self
161 .msc3202_key_counts(otk_users, otk_recipients)
162 .await;
163
164 (DeviceLists { changed, left: Vec::new() }, counts, fallbacks)
165 } else {
166 (DeviceLists::new(), OtkCounts::new(), FallbackTypes::new())
167 };
168
169 if pdus.is_empty()
170 && edus.is_empty()
171 && to_device.is_empty()
172 && device_lists.is_empty()
173 && device_one_time_keys_count.is_empty()
174 && device_unused_fallback_key_types.is_empty()
175 {
176 return Ok(Destination::Appservice(id));
177 }
178
179 let request = PushEventsRequest {
180 txn_id: URL_SAFE_NO_PAD.encode(txn_hash).into(),
181 events: pdus,
182 ephemeral: edus,
183 to_device,
184 device_lists,
185 device_one_time_keys_count,
186 device_unused_fallback_key_types,
187 };
188
189 match self
190 .services
191 .appservice
192 .send_request(info.registration, request)
193 .await
194 {
195 | Ok(_) => Ok(Destination::Appservice(id)),
196 | Err(e) => Err((Destination::Appservice(id), e)),
197 }
198}
199
200#[implement(Service)]
201async fn txn_part(&self, info: &RegistrationInfo, msc3202: bool, event: &SendingEvent) -> Part {
202 match event {
203 | SendingEvent::Pdu(pdu_id) => {
204 let Ok(pdu) = self
205 .services
206 .timeline
207 .get_pdu_from_id(pdu_id)
208 .await
209 else {
210 return Part::None;
211 };
212
213 let sender =
214 (msc3202 && info.is_user_match(pdu.sender())).then(|| pdu.sender().to_owned());
215
216 Part::Pdu(pdu.to_format(), sender)
217 },
218 | SendingEvent::Edu(edu) => {
219 if !info.registration.receive_ephemeral {
220 return Part::None;
221 }
222
223 serde_json::from_slice::<EphemeralData>(edu)
224 .and_then(|edu| Raw::new(&edu))
225 .map_or(Part::None, Part::Edu)
226 },
227 | SendingEvent::ToDevice(buf) => {
228 let Some(bytes) = buf.get(TAG_PREFIX_LEN..) else {
229 debug_warn!("skipping malformed queued to-device event");
230 return Part::None;
231 };
232
233 let Ok(raw) = serde_json::from_slice(bytes) else {
234 debug_warn!("skipping malformed queued to-device event");
235 return Part::None;
236 };
237
238 let recipient = msc3202
239 .then(|| serde_json::from_slice(bytes).ok())
240 .flatten()
241 .map(|ToDeviceRecipient { to_user_id, to_device_id }| (to_user_id, to_device_id));
242
243 Part::ToDevice(raw, recipient)
244 },
245 | SendingEvent::DeviceListChanged(buf) => {
246 if msc3202
247 && let Some(bytes) = buf.get(TAG_PREFIX_LEN..)
248 && let Ok(user) = from_utf8(bytes)
249 && let Ok(user_id) = UserId::parse(user)
250 {
251 return Part::Changed(user_id);
252 }
253
254 Part::None
255 },
256 | SendingEvent::BadgeRefresh | SendingEvent::Flush => Part::None,
257 }
258}
259
260impl Parts {
261 fn merge(mut self, part: Part) -> Self {
262 match part {
263 | Part::Pdu(pdu, sender) => {
264 self.pdus.push(pdu);
265 self.otk_users.extend(sender);
266 },
267 | Part::Edu(edu) => self.edus.push(edu),
268 | Part::ToDevice(raw, recipient) => {
269 self.to_device.push(raw);
270 self.otk_recipients.extend(recipient);
271 },
272 | Part::Changed(user_id) => self.changed.push(user_id),
273 | Part::None => {},
274 }
275
276 self
277 }
278}
279
280#[implement(Service)]
286#[tracing::instrument(
287 level = "debug",
288 skip_all,
289 fields(
290 users = %users.len(),
291 recipients = %recipients.len(),
292 ),
293)]
294async fn msc3202_key_counts(
295 &self,
296 users: BTreeSet<OwnedUserId>,
297 recipients: Devices,
298) -> (OtkCounts, FallbackTypes) {
299 let devices: Devices = users
300 .iter()
301 .stream()
302 .flat_map(|user_id| {
303 self.services
304 .users
305 .all_device_ids(user_id)
306 .map(move |device_id| (user_id.to_owned(), device_id.to_owned()))
307 })
308 .chain(recipients.into_iter().stream())
309 .collect()
310 .await;
311
312 dedup(devices)
313 .into_iter()
314 .stream()
315 .broad_then(async |(user_id, device_id): UserDevice| {
316 let counts = self
317 .services
318 .users
319 .count_one_time_keys(&user_id, &device_id);
320
321 let fallbacks = self
322 .services
323 .users
324 .unused_fallback_key_algorithms(&user_id, &device_id)
325 .collect();
326
327 let (counts, fallbacks) = join(counts, fallbacks).await;
328
329 (user_id, device_id, counts, fallbacks)
330 })
331 .ready_fold(
332 (OtkCounts::new(), FallbackTypes::new()),
333 |(mut counts, mut fallbacks), (user_id, device_id, otk, fallback)| {
334 counts
335 .entry(user_id.clone())
336 .or_default()
337 .insert(device_id.clone(), otk);
338
339 fallbacks
340 .entry(user_id)
341 .or_default()
342 .insert(device_id, fallback);
343
344 (counts, fallbacks)
345 },
346 )
347 .await
348}
349
350fn dedup(mut devices: Devices) -> Devices {
351 devices.sort_unstable();
352 devices.dedup();
353
354 devices
355}