Skip to main content

tuwunel_service/sending/sender/select/
device_changes.rs

1use std::{
2	collections::{BTreeMap, BTreeSet},
3	sync::atomic::{AtomicU64, AtomicUsize, Ordering},
4};
5
6use futures::{StreamExt, future::join};
7use ruma::{
8	DeviceId, OwnedUserId, ServerName, UserId,
9	api::federation::transactions::edu::{DeviceListUpdateContent, Edu, SigningKeyUpdateContent},
10	device_id,
11};
12use tuwunel_core::{
13	implement,
14	itertools::Itertools,
15	smallstr::SmallString,
16	smallvec::SmallVec,
17	utils::{BoolExt, IterStream, ReadyExt, stream::WidebandExt},
18};
19
20use super::{Selected, edu_buf};
21use crate::{
22	sending::{EduBuf, Service},
23	users::{DeviceListChange, DeviceListRecord},
24};
25
26type Pending = BTreeMap<OwnedUserId, BTreeSet<u64>>;
27type Records = SmallVec<[DeviceListRecord; 1]>;
28
29enum Plan<D> {
30	Resync(u64),
31	Deltas {
32		signing_key_update: bool,
33		deltas: D,
34	},
35}
36
37#[derive(Debug, Eq, PartialEq)]
38struct Delta<'a> {
39	device_id: &'a DeviceId,
40	stream_id: u64,
41	prev_id: Option<u64>,
42	deleted: bool,
43}
44
45// Above ten distinct devices, a snapshot costs less than individual updates.
46const K: usize = 10;
47
48/// Select device-list deltas and signing-key updates for local users.
49///
50/// Legacy records and windows exceeding the device limit force a snapshot resync.
51#[implement(Service)]
52#[tracing::instrument(
53	name = "device_changes",
54	level = "trace",
55	skip(self, server_name, max_edu_count, events_len)
56)]
57pub(super) async fn select_edus_device_changes(
58	&self,
59	server_name: &ServerName,
60	since: (u64, u64),
61	max_edu_count: &AtomicU64,
62	events_len: &AtomicUsize,
63) -> Selected {
64	let pending = self
65		.services
66		.state_cache
67		.server_rooms(server_name)
68		.map(ToOwned::to_owned)
69		.fold(Pending::new(), async |pending, room_id| {
70			self.services
71				.users
72				.room_keys_changed(&room_id, since.0, Some(since.1))
73				.ready_filter_map(|(user_id, count)| {
74					if !self.services.globals.user_is_local(user_id) {
75						return None;
76					}
77
78					debug_assert!(count <= since.1, "exceeds upper-bound");
79					max_edu_count.fetch_max(count, Ordering::Relaxed);
80
81					Some((user_id.to_owned(), count))
82				})
83				.ready_fold(pending, |mut pending, (user_id, count)| {
84					pending.entry(user_id).or_default().insert(count);
85					pending
86				})
87				.await
88		})
89		.await;
90
91	pending
92		.into_iter()
93		.stream()
94		.fold(Selected::default(), async |selected, (user_id, counts)| {
95			self.select_user_devices(selected, &user_id, counts, server_name, events_len)
96				.await
97		})
98		.await
99}
100
101#[implement(Service)]
102async fn select_user_devices(
103	&self,
104	selected: Selected,
105	user_id: &UserId,
106	counts: BTreeSet<u64>,
107	server_name: &ServerName,
108	events_len: &AtomicUsize,
109) -> Selected {
110	let records: Records = counts
111		.into_iter()
112		.stream()
113		.wide_then(async |count| {
114			self.services
115				.users
116				.device_list_change(count)
117				.await
118				.unwrap_or(DeviceListRecord {
119					change: DeviceListChange::Resync,
120					stream_id: 0,
121				})
122		})
123		.collect()
124		.await;
125
126	let current_stream_id = self
127		.services
128		.users
129		.get_devicelist_version(user_id)
130		.await
131		.unwrap_or(0);
132
133	let (signing_key_update, deltas) = match plan_device_list_edus(&records, current_stream_id) {
134		| Plan::Deltas { signing_key_update, deltas } => (signing_key_update, deltas),
135		| Plan::Resync(stream_id) => {
136			return push(selected, device_list_edu(user_id, stream_id), events_len);
137		},
138	};
139
140	let signing = signing_key_update
141		.then_async(|| self.signing_key_edu(user_id, server_name))
142		.await
143		.flatten();
144
145	let selected = signing
146		.into_iter()
147		.fold(selected, |selected, edu| push(selected, edu, events_len));
148
149	deltas
150		.stream()
151		.fold(selected, async |selected, delta| {
152			let edu = self.device_delta_edu(user_id, delta).await;
153
154			// Current content and the maximum S keep ordered replay convergent.
155			push(selected, edu, events_len)
156		})
157		.await
158}
159
160fn plan_device_list_edus(
161	records: &[DeviceListRecord],
162	current_stream_id: u64,
163) -> Plan<impl Iterator<Item = Delta<'_>>> {
164	if records
165		.iter()
166		.any(|record| matches!(record.change, DeviceListChange::Resync))
167	{
168		return Plan::Resync(current_stream_id);
169	}
170
171	let latest: BTreeMap<_, _> = records
172		.iter()
173		.filter_map(record_delta)
174		.map(|delta| (delta.device_id, delta))
175		.collect();
176
177	if latest.len() > K {
178		return Plan::Resync(current_stream_id);
179	}
180
181	let anchor = records
182		.iter()
183		.filter_map(record_delta)
184		.map(|delta| delta.stream_id)
185		.min()
186		.unwrap_or(0)
187		.saturating_sub(1);
188
189	let deltas = latest
190		.into_values()
191		.sorted_unstable_by_key(|delta| delta.stream_id)
192		.scan(anchor, |anchor, delta| {
193			let prev_id = u64::ne(anchor, &0).then_some(*anchor);
194
195			*anchor = delta.stream_id;
196			Some(Delta { prev_id, ..delta })
197		});
198
199	let signing_key_update = records
200		.iter()
201		.any(|record| matches!(record.change, DeviceListChange::CrossSigning));
202
203	Plan::Deltas { signing_key_update, deltas }
204}
205
206fn record_delta(record: &DeviceListRecord) -> Option<Delta<'_>> {
207	let (device_id, deleted) = match &record.change {
208		| DeviceListChange::CrossSigning | DeviceListChange::Resync => return None,
209		| DeviceListChange::Device(device_id) => (device_id.as_ref(), false),
210		| DeviceListChange::Deleted(device_id) => (device_id.as_ref(), true),
211	};
212
213	Some(Delta {
214		device_id,
215		stream_id: record.stream_id,
216		prev_id: None,
217		deleted,
218	})
219}
220
221fn device_list_edu(user_id: &UserId, stream_id: u64) -> EduBuf {
222	edu_buf(&Edu::DeviceListUpdate(DeviceListUpdateContent {
223		user_id: user_id.to_owned(),
224		device_id: device_id!("placeholder").to_owned(),
225		device_display_name: Some("Placeholder".to_owned()),
226		stream_id: stream_id.try_into().unwrap_or_default(),
227		prev_id: Vec::new(),
228		deleted: None,
229		keys: None,
230	}))
231}
232
233fn push(mut selected: Selected, edu: EduBuf, events_len: &AtomicUsize) -> Selected {
234	selected.push(edu, events_len);
235	selected
236}
237
238#[implement(Service)]
239#[tracing::instrument(level = "trace", skip(self))]
240async fn signing_key_edu(&self, user_id: &UserId, server_name: &ServerName) -> Option<EduBuf> {
241	let allowed = |user: &UserId| user.server_name() == server_name;
242	let (master_key, self_signing_key) = join(
243		self.services
244			.users
245			.get_master_key(None, user_id, &allowed),
246		self.services
247			.users
248			.get_self_signing_key(None, user_id, &allowed),
249	)
250	.await;
251
252	let (master_key, self_signing_key) = (master_key.ok(), self_signing_key.ok());
253
254	(master_key.is_some() || self_signing_key.is_some()).then(|| {
255		edu_buf(&Edu::SigningKeyUpdate(SigningKeyUpdateContent {
256			user_id: user_id.to_owned(),
257			master_key,
258			self_signing_key,
259		}))
260	})
261}
262
263#[implement(Service)]
264#[tracing::instrument(level = "trace", skip(self))]
265async fn device_delta_edu(&self, user_id: &UserId, delta: Delta<'_>) -> EduBuf {
266	let (keys, device_display_name) = if delta.deleted {
267		(None, None)
268	} else {
269		let metadata = self
270			.server
271			.config
272			.allow_device_name_federation
273			.then_async(|| {
274				self.services
275					.users
276					.get_device_metadata(user_id, delta.device_id)
277			});
278
279		let (keys, metadata) = join(
280			self.services
281				.users
282				.get_device_keys(user_id, delta.device_id),
283			metadata,
284		)
285		.await;
286
287		let display_name = metadata
288			.and_then(Result::ok)
289			.and_then(|device| device.display_name)
290			.map(SmallString::into_string)
291			.or_else(|| Some(delta.device_id.as_str().into()));
292
293		(keys.ok(), display_name)
294	};
295
296	let prev_id = delta
297		.prev_id
298		.into_iter()
299		.map(|id| id.try_into().unwrap_or_default())
300		.collect();
301
302	edu_buf(&Edu::DeviceListUpdate(DeviceListUpdateContent {
303		user_id: user_id.to_owned(),
304		device_id: delta.device_id.to_owned(),
305		device_display_name,
306		stream_id: delta.stream_id.try_into().unwrap_or_default(),
307		prev_id,
308		deleted: delta.deleted.then_some(true),
309		keys,
310	}))
311}
312
313#[cfg(test)]
314mod tests {
315	use DeviceListChange::{CrossSigning, Deleted, Device, Resync};
316
317	use super::{DeviceListChange, DeviceListRecord, K, Plan, plan_device_list_edus};
318
319	#[test]
320	fn plans_dense_chains() {
321		let cases = [
322			(vec![(Device("X"), 5)], Some((false, vec![("X", 5, Some(4), false)]))),
323			(
324				vec![(Device("X"), 5), (Device("Y"), 6), (Device("X"), 7)],
325				Some((false, vec![("Y", 6, Some(4), false), ("X", 7, Some(6), false)])),
326			),
327			(
328				vec![(Device("X"), 5), (Deleted("X"), 6)],
329				Some((false, vec![("X", 6, Some(4), true)])),
330			),
331			(vec![(Resync, 5)], None),
332			(vec![(CrossSigning, 4)], Some((true, vec![]))),
333			(
334				vec![(CrossSigning, 4), (Device("X"), 5)],
335				Some((true, vec![("X", 5, Some(4), false)])),
336			),
337			(vec![(Device("X"), 1)], Some((false, vec![("X", 1, None, false)]))),
338		];
339
340		for (input, expected) in cases {
341			let records: Vec<_> = input
342				.into_iter()
343				.map(|(change, stream_id)| {
344					let change = match change {
345						| Resync => Resync,
346						| CrossSigning => CrossSigning,
347						| Device(id) => Device(id.into()),
348						| Deleted(id) => Deleted(id.into()),
349					};
350
351					DeviceListRecord { change, stream_id }
352				})
353				.collect();
354
355			let actual = match plan_device_list_edus(&records, 99) {
356				| Plan::Resync(stream_id) => {
357					assert_eq!(stream_id, 99);
358					None
359				},
360				| Plan::Deltas { signing_key_update, deltas } => {
361					let deltas = deltas
362						.map(|d| (d.device_id.as_str(), d.stream_id, d.prev_id, d.deleted))
363						.collect();
364
365					Some((signing_key_update, deltas))
366				},
367			};
368
369			assert_eq!(actual, expected, "{records:?}");
370		}
371	}
372
373	#[test]
374	fn too_many_devices_resync() {
375		let records: Vec<_> = (0..=K)
376			.map(|id| DeviceListRecord {
377				change: Device(id.to_string().into()),
378				stream_id: u64::try_from(id).unwrap() + 1,
379			})
380			.collect();
381
382		assert!(matches!(plan_device_list_edus(&records[..K], 99), Plan::Deltas { .. }));
383		assert!(matches!(plan_device_list_edus(&records, 99), Plan::Resync(99)));
384	}
385}