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
45const K: usize = 10;
47
48#[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 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}