Skip to main content

tuwunel_service/sending/sender/dispatch/
appservice.rs

1use 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/// The appservice-injected recipient fields of a queued to-device event
34/// (MSC4203).
35///
36/// Parsed to scope MSC3202 one-time-key counts to the addressed devices.
37#[derive(Deserialize)]
38struct ToDeviceRecipient {
39	to_user_id: OwnedUserId,
40	to_device_id: OwnedDeviceId,
41}
42
43/// The wire pieces of one appservice transaction, accumulated per event.
44///
45/// `otk_users` and `otk_recipients` are the MSC3202 one-time-key scope: the
46/// appservice sender, the namespace-matched PDU senders, and the to-device
47/// recipients of the transaction.
48#[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
58/// One queued event rendered for the appservice transaction.
59///
60/// `None` is an event the appservice does not receive, or one that failed to
61/// load or decode; the transaction skips it.
62enum Part {
63	Pdu(Raw<AnyTimelineEvent>, Option<OwnedUserId>),
64	Edu(Raw<EphemeralData>),
65	ToDevice(Raw<AnyAppserviceToDeviceEvent>, Option<UserDevice>),
66	Changed(OwnedUserId),
67	None,
68}
69
70/// One device of one user.
71///
72/// The MSC3202 key-count maps are keyed by this pair.
73type UserDevice = (OwnedUserId, OwnedDeviceId);
74
75/// MSC3202 `device_one_time_keys_count`: unclaimed one-time-key counts per
76/// algorithm, keyed by user then device.
77///
78/// Matches the ruma request field type.
79type OtkCounts =
80	BTreeMap<OwnedUserId, BTreeMap<OwnedDeviceId, BTreeMap<OneTimeKeyAlgorithm, UInt>>>;
81
82/// MSC3202 `device_unused_fallback_key_types`: algorithms with an unused
83/// fallback key, keyed by user then device.
84///
85/// Matches the ruma request field type.
86type FallbackTypes = BTreeMap<OwnedUserId, BTreeMap<OwnedDeviceId, Vec<OneTimeKeyAlgorithm>>>;
87
88/// The MSC3202-interesting devices of one transaction: the appservice
89/// sender's plus matched PDU senders' devices and the to-device recipients.
90///
91/// Sorted and deduplicated before the per-device key lookups fan out.
92type 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		//TODO: appservice queue cleanup.
115		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/// MSC3202 one-time-key counts and unused fallback key types over every
281/// device of `users` plus the specific `recipients`.
282///
283/// Recomputed per build rather than snapshotted, so a retry ships fresh
284/// counts.
285#[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}