tuwunel_service/migrations/
token_expiry.rs1use std::{pin::pin, sync::Arc};
2
3use futures::{TryFutureExt, TryStreamExt};
4use ruma::{DeviceId, UserId};
5use tuwunel_core::{Result, debug_warn, err, info, result::NotFound, warn};
6use tuwunel_database::{KeyVal, Map, deserialize_from_slice, serialize_key};
7
8use crate::Services;
9
10pub(super) const ADOPT_MARKER: &str = "adopt_foreign_token_expiry";
12
13pub(super) const RESTORE_MARKER: &str = "restore_foreign_token_expiry";
15
16const ORIGIN_COLUMN: &str = "userdeviceid_tokenexpires";
17
18type Device<'a> = (&'a UserId, &'a DeviceId);
19
20type TokenValue<'a> = (&'a UserId, &'a DeviceId, Option<u64>);
21
22#[derive(Default)]
28struct Tally {
29 applied: usize,
30 skipped: usize,
31 unreadable: usize,
32}
33
34#[derive(Clone, Copy)]
40enum Mode {
41 Probe,
43
44 Adopt,
46
47 Restore,
49}
50
51#[tracing::instrument(level = "debug", skip_all)]
64pub(super) async fn migrate_token_expiry(services: &Services) -> Result<bool> {
65 let Some(tokenexpires) = services.db.open_cf(ORIGIN_COLUMN)? else {
66 return Ok(true);
67 };
68
69 if services.oauth.get_server().is_err() {
70 let waiting = probe(services, &tokenexpires).await?;
71
72 if waiting {
73 info!("Migrated OAuth sessions keep their origin lifetime; no OIDC provider");
74 }
75
76 return Ok(!waiting);
77 }
78
79 let Tally { applied: adopted, skipped, .. } =
80 walk(services, &tokenexpires, Mode::Adopt).await?;
81
82 if adopted > 0 || skipped > 0 {
85 info!(%adopted, %skipped, "Adopted token expiry from a foreign database");
86 }
87
88 Ok(true)
89}
90
91#[tracing::instrument(level = "debug", skip_all)]
104pub(super) async fn restore_token_expiry(services: &Services) -> Result<bool> {
105 let Some(tokenexpires) = services.db.open_cf(ORIGIN_COLUMN)? else {
106 return Ok(true);
107 };
108
109 if services.oauth.get_server().is_ok() {
110 return Ok(false);
111 }
112
113 services.db["global"].remove(ADOPT_MARKER);
116
117 let Tally { applied: restored, skipped, .. } =
118 walk(services, &tokenexpires, Mode::Restore).await?;
119
120 if restored > 0 || skipped > 0 {
121 info!(
122 %restored,
123 %skipped,
124 "Restored the origin lifetime of adopted OAuth sessions; no OIDC provider"
125 );
126 }
127
128 Ok(true)
129}
130
131async fn probe(services: &Services, tokenexpires: &Arc<Map>) -> Result<bool> {
136 let (device_tokens, token_owners) = columns(services);
137 let mut adoptable = pin!(tokenexpires.raw_stream().try_filter_map(|row| {
138 apply_one(device_tokens, token_owners, row, Mode::Probe)
139 .map_ok(|adoptable| adoptable.then_some(()))
140 }));
141
142 adoptable
143 .try_next()
144 .map_ok(|hit| hit.is_some())
145 .await
146}
147
148fn columns(services: &Services) -> (&Arc<Map>, &Arc<Map>) {
149 (&services.db["userdeviceid_token"], &services.db["token_userdeviceid"])
150}
151
152async fn walk(services: &Services, tokenexpires: &Arc<Map>, mode: Mode) -> Result<Tally> {
161 let (device_tokens, token_owners) = columns(services);
162 let cork = services.db.cork_and_sync();
163
164 let tally = tokenexpires
165 .raw_stream()
166 .try_fold(Tally::default(), async |tally, row| {
167 Ok(tally.record(apply_one(device_tokens, token_owners, row, mode).await))
168 })
169 .await?;
170
171 drop(cork);
172
173 let unreadable = tally.unreadable;
174
175 unreadable
176 .eq(&0)
177 .then_some(tally)
178 .ok_or_else(|| err!(Database("{unreadable} token expiries could not be read")))
179}
180
181impl Tally {
182 fn record(mut self, result: Result<bool>) -> Self {
183 match result {
184 | Ok(true) => self.applied = self.applied.saturating_add(1),
185 | Ok(false) => self.skipped = self.skipped.saturating_add(1),
186 | Err(e) => {
187 warn!(error = %e, "a token expiry could not be read");
188 self.unreadable = self.unreadable.saturating_add(1);
189 },
190 }
191
192 self
193 }
194}
195
196async fn apply_one(
205 device_tokens: &Arc<Map>,
206 token_owners: &Arc<Map>,
207 (key, value): KeyVal<'_>,
208 mode: Mode,
209) -> Result<bool> {
210 let Ok((user_id, device_id)) = deserialize_from_slice::<Device<'_>>(key) else {
211 warn!("skipping a foreign token expiry whose device could not be read");
212 return Ok(false);
213 };
214
215 let Ok(expires) = deserialize_from_slice::<u64>(value) else {
216 warn!(%user_id, %device_id, "skipping a foreign token expiry that could not be read");
217 return Ok(false);
218 };
219
220 let Some(token_raw) = device_tokens
221 .qry(&(user_id, device_id))
222 .await
223 .optional()?
224 else {
225 debug_warn!(%user_id, %device_id, "skipping a device holding no access token here");
226 return Ok(false);
227 };
228
229 let Ok(token) = deserialize_from_slice::<&str>(&token_raw) else {
230 warn!(%user_id, %device_id, "skipping a device whose stored token could not be read");
231 return Ok(false);
232 };
233
234 let Some(stored) = token_owners.get(token).await.optional()? else {
235 debug_warn!(%user_id, %device_id, "skipping an access token this server does not hold");
236 return Ok(false);
237 };
238
239 let candidate = match mode {
240 | Mode::Probe | Mode::Adopt => adoptable(&stored),
241 | Mode::Restore => adopted(&stored, expires),
242 };
243
244 let Ok(candidate) = candidate else {
245 warn!(%user_id, %device_id, "skipping a token value that could not be read");
246 return Ok(false);
247 };
248
249 let Some((owner, device)) = candidate else {
250 debug_warn!(%user_id, %device_id, "skipping a token value this pass leaves as stored");
251 return Ok(false);
252 };
253
254 match mode {
255 | Mode::Probe => {},
256 | Mode::Adopt => {
257 token_owners.raw_put(token, (owner, device, Some(expires)));
258 info!(%user_id, %device_id, "adopted the origin expiry of an access token");
259 },
260 | Mode::Restore => {
261 token_owners.raw_put(token, (owner, device));
262 info!(%user_id, %device_id, "restored the origin lifetime of an access token");
263 },
264 }
265
266 Ok(true)
267}
268
269fn adoptable(stored: &[u8]) -> Result<Option<Device<'_>>> {
277 let (owner, device, carried): TokenValue<'_> = deserialize_from_slice(stored)?;
278
279 let adoptable = carried.is_none() && stored.len() == serialize_key((owner, device))?.len();
280
281 Ok(adoptable.then_some((owner, device)))
282}
283
284fn adopted(stored: &[u8], expires: u64) -> Result<Option<Device<'_>>> {
292 let (owner, device, carried): TokenValue<'_> = deserialize_from_slice(stored)?;
293
294 let adopted = carried == Some(expires);
295
296 Ok(adopted.then_some((owner, device)))
297}
298
299#[cfg(test)]
300mod tests {
301 use ruma::{device_id, user_id};
302 use tuwunel_database::{KeyBuf, deserialize_from_slice, serialize_key};
303
304 use super::{Device, TokenValue, adoptable, adopted};
305
306 const EXPIRES: u64 = 1_700_000_000;
307
308 fn owner_device() -> Device<'static> {
309 (user_id!("@alice:localhost"), device_id!("AAAAAAAAAA"))
310 }
311
312 fn foreign() -> KeyBuf {
314 serialize_key(owner_device()).expect("the foreign value serializes")
315 }
316
317 fn native(expires: Option<u64>) -> KeyBuf {
319 let (owner, device) = owner_device();
320
321 serialize_key((owner, device, expires)).expect("the native value serializes")
322 }
323
324 #[test]
325 fn foreign_value_reads_as_non_expiring() {
326 let (owner, device) = owner_device();
327 let foreign = foreign();
328
329 let (read_owner, read_device, carried): TokenValue<'_> =
330 deserialize_from_slice(&foreign).expect("the foreign value deserializes");
331
332 assert_eq!(read_owner, owner);
333 assert_eq!(read_device, device);
334 assert_eq!(carried, None, "a row without the tail must read as non-expiring");
335 }
336
337 #[test]
338 fn adopted_value_carries_its_expiry() {
339 let (owner, device) = owner_device();
340 let stamped = native(Some(EXPIRES));
341
342 let (read_owner, read_device, carried): TokenValue<'_> =
343 deserialize_from_slice(&stamped).expect("the stamped value deserializes");
344
345 assert_eq!(read_owner, owner);
346 assert_eq!(read_device, device);
347 assert_eq!(carried, Some(EXPIRES));
348 }
349
350 #[test]
351 fn a_past_expiry_survives_the_round_trip() {
352 let stamped = native(Some(1));
353
354 let (.., carried): TokenValue<'_> =
355 deserialize_from_slice(&stamped).expect("the stamped value deserializes");
356
357 assert_eq!(carried, Some(1), "an expiry already past is carried, not clamped");
358 }
359
360 #[test]
361 fn a_foreign_row_is_adoptable() {
362 let (owner, device) = owner_device();
363 let foreign = foreign();
364
365 let (read_owner, read_device) = adoptable(&foreign)
366 .expect("the foreign value is readable")
367 .expect("the foreign value is adoptable");
368
369 assert_eq!(read_owner, owner);
370 assert_eq!(read_device, device);
371 }
372
373 #[test]
375 fn a_row_this_server_wrote_is_refused() {
376 let native = native(None);
377 let candidate = adoptable(&native).expect("the native value is readable");
378
379 assert!(candidate.is_none());
380 }
381
382 #[test]
383 fn a_row_already_carrying_an_expiry_is_refused() {
384 let expiring = native(Some(EXPIRES));
385 let candidate = adoptable(&expiring).expect("the expiring value is readable");
386
387 assert!(candidate.is_none());
388 }
389
390 #[test]
391 fn an_adopted_row_is_restorable() {
392 let (owner, device) = owner_device();
393 let stamped = native(Some(EXPIRES));
394
395 let (read_owner, read_device) = adopted(&stamped, EXPIRES)
396 .expect("the stamped value is readable")
397 .expect("the stamped value is restorable");
398
399 assert_eq!(read_owner, owner);
400 assert_eq!(read_device, device);
401 }
402
403 #[test]
404 fn a_row_carrying_another_expiry_is_kept() {
405 let expiring = native(Some(EXPIRES.saturating_add(1)));
406 let candidate = adopted(&expiring, EXPIRES).expect("the expiring value is readable");
407
408 assert!(candidate.is_none());
409 }
410
411 #[test]
412 fn a_foreign_row_is_not_restorable() {
413 let foreign = foreign();
414 let candidate = adopted(&foreign, EXPIRES).expect("the foreign value is readable");
415
416 assert!(candidate.is_none());
417 }
418
419 #[test]
420 fn a_row_without_an_expiry_is_kept() {
421 let native = native(None);
422 let candidate = adopted(&native, EXPIRES).expect("the native value is readable");
423
424 assert!(candidate.is_none());
425 }
426}