Skip to main content

tuwunel_service/migrations/
token_expiry.rs

1use 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
10/// Stamped once every origin expiry has been adopted.
11pub(super) const ADOPT_MARKER: &str = "adopt_foreign_token_expiry";
12
13/// Stamped once every expiry adopted without a provider has been restored.
14pub(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/// Counts of what one walk of the origin column did.
23///
24/// A skipped row and an unreadable one are counted apart because only the
25/// second withholds the marker, and withholding it retries the whole pass on the
26/// next boot.
27#[derive(Default)]
28struct Tally {
29	applied: usize,
30	skipped: usize,
31	unreadable: usize,
32}
33
34/// What a walk does to the token an origin row names.
35///
36/// One walk serves every pass over the column, resolving each row through the
37/// same lookups and the same provenance guard before the mode decides what, if
38/// anything, is written.
39#[derive(Clone, Copy)]
40enum Mode {
41	/// Reports whether the row could be adopted, without writing.
42	Probe,
43
44	/// Stamps the origin expiry onto the token.
45	Adopt,
46
47	/// Puts a token carrying the stamped expiry back into the foreign shape.
48	Restore,
49}
50
51/// Adopts the access-token expiry a foreign database keeps in a column of its
52/// own, returning whether the pass is finished with that column.
53///
54/// That column is keyed by device while this server keeps the expiry alongside
55/// the owner in the shared token column, so a migrated session keeps
56/// authenticating while its expiry does not.
57///
58/// Every row names an OAuth session whose client can only recover from the
59/// expiry by refreshing against an OIDC provider here, so without one the pass
60/// waits, unfinished, rather than stranding sessions. With one it adopts once,
61/// since this server writes the same column afterward and a second pass would
62/// resurrect a replaced expiry.
63#[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	// A skipped row is usually a device this pass correctly left alone, so the
83	// summary counts them without implying a loss; a real loss logs per row.
84	if adopted > 0 || skipped > 0 {
85		info!(%adopted, %skipped, "Adopted token expiry from a foreign database");
86	}
87
88	Ok(true)
89}
90
91/// Restores the origin lifetime of a session an earlier release adopted with
92/// no provider to refresh it against, returning whether the pass is finished.
93///
94/// Before the adoption waited for a provider it stamped every origin expiry,
95/// each long past, so a session that has not authenticated since is refused and
96/// removed on its next request while its client, with nothing to refresh
97/// against, never signs out. The origin column still names every adopted device
98/// with the value stamped, so a token carrying exactly that value goes back to
99/// the foreign shape, and the adoption marker is cleared so the gated pass owns
100/// the column again from this boot on. With a provider an adopted expiry can be
101/// refreshed and stays, so the pass waits rather than finishing, and a provider
102/// removed later still gets the restore.
103#[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	// Handed back before the walk, so however it ends the restored rows are the
114	// adoption's to own.
115	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
131/// Whether the origin column holds a row the adoption could still carry.
132///
133/// The first adoptable row answers, so a column of rows this server can never
134/// adopt is walked whole while one with work waiting costs a single hit.
135async 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
152/// Walks the origin column once, applying the mode to every row.
153///
154/// One row at a time: each is read before it is rewritten, so a concurrent walk
155/// could let two rows naming one token both clear the provenance guard. A
156/// cursor error ends the walk rather than being counted, because the status is
157/// sticky and the iterator cannot advance past it. An unreadable row fails the
158/// walk too, since leaving the marker unstamped is what makes an engine failure
159/// recoverable: the pass is idempotent, so the next boot retries it whole.
160async 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
196/// Applies the mode to the token one origin row names, reporting whether the
197/// row was acted on.
198///
199/// A `false` return is a row with nothing to do: one this pass cannot make
200/// sense of, a device holding no token here, a token this server does not hold,
201/// or a stored value the mode leaves as it is. Only a failed lookup returns an
202/// error, because a row that will never decode would otherwise refuse every
203/// later boot as well.
204async 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
269/// Reads a stored token value, yielding its owner only when the row is one this
270/// pass may annotate.
271///
272/// A foreign row omits the trailing expiry field and reads as `None`. This
273/// server always writes that field, even empty, so every row it wrote encodes
274/// shorter than it is stored. The length comparison is what skips a token issued
275/// here after the import, whose foreign expiry no longer describes it.
276fn 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
284/// Reads a stored token value, yielding its owner only when the row carries the
285/// expiry the adoption stamped.
286///
287/// A token this server issued carries an expiry it computed itself, or none, so
288/// a stored value equal to the origin's is the adoption's own write. A foreign
289/// row, and a token a later login replaced, read as something else and are
290/// left alone.
291fn 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	/// The shape a foreign database writes: owner and device, no expiry field.
313	fn foreign() -> KeyBuf {
314		serialize_key(owner_device()).expect("the foreign value serializes")
315	}
316
317	/// The shape this server writes, whose expiry field is present either way.
318	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	// A token issued after the import must never be given a stale expiry.
374	#[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}