Skip to main content

tuwunel_service/oauth/
sessions.rs

1mod adopt;
2pub mod association;
3
4use std::{
5	iter::once,
6	sync::{Arc, Mutex},
7	time::SystemTime,
8};
9
10use futures::{FutureExt, Stream, StreamExt, TryFutureExt, TryStreamExt};
11use ruma::{OwnedUserId, UserId};
12use serde::{Deserialize, Serialize};
13use tuwunel_core::{
14	Err, Result, at, implement,
15	itertools::Itertools,
16	result::NotFound,
17	utils::{
18		MutexMap,
19		stream::{IterStream, ReadyExt, TryExpect},
20	},
21};
22use tuwunel_database::{Cbor, Database, Deserialized, Ignore, Map};
23use url::Url;
24
25pub use self::adopt::Counts;
26use super::{Provider, Providers, UserInfo, unique_id as session_unique_id};
27use crate::SelfServices;
28
29pub struct Sessions {
30	services: SelfServices,
31	association_pending: Mutex<association::Pending>,
32
33	/// Serializes probes and writes for each unique identity.
34	///
35	/// Transaction batches cannot conditionally claim an index key. Each
36	/// identity therefore has an independent critical section.
37	write_locks: MutexMap<String, ()>,
38
39	providers: Arc<Providers>,
40	db: Data,
41}
42
43struct Data {
44	oauthid_session: Arc<Map>,
45	oauthuniqid_oauthid: Arc<Map>,
46	userid_oauthid: Arc<Map>,
47	database: Arc<Database>,
48}
49
50/// Persistent state for one upstream OAuth authorization.
51///
52/// The record carries provider, redirect, PKCE, nonce, and token data across
53/// the authorization flow. Once linked, it also associates the provider
54/// identity with a Matrix user.
55#[derive(Clone, Debug, Default, Deserialize, Serialize)]
56pub struct Session {
57	/// Identity Provider ID (the `client_id` in the configuration) associated
58	/// with this session.
59	pub idp_id: Option<String>,
60
61	/// Session ID used as the index key for this session itself.
62	pub sess_id: Option<SessionId>,
63
64	/// Token type (bearer, mac, etc).
65	pub token_type: Option<String>,
66
67	/// Access token to the provider.
68	pub access_token: Option<String>,
69
70	/// OIDC ID token returned by the provider.
71	pub id_token: Option<String>,
72
73	/// Duration in seconds the access_token is valid for.
74	pub expires_in: Option<u64>,
75
76	/// Point in time that the access_token expires.
77	pub expires_at: Option<SystemTime>,
78
79	/// Token used to refresh the access_token.
80	pub refresh_token: Option<String>,
81
82	/// Duration in seconds the refresh_token is valid for
83	pub refresh_token_expires_in: Option<u64>,
84
85	/// Point in time that the refresh_token expires.
86	pub refresh_token_expires_at: Option<SystemTime>,
87
88	/// Access scope actually granted (if supported).
89	pub scope: Option<String>,
90
91	/// Redirect URL
92	pub redirect_url: Option<Url>,
93
94	/// Challenge preimage
95	pub code_verifier: Option<String>,
96
97	/// Random string passed exclusively in the grant session cookie.
98	pub cookie_nonce: Option<String>,
99
100	/// Random single-use string passed in the provider redirect.
101	pub query_nonce: Option<String>,
102
103	/// Point in time the authorization grant session expires.
104	pub authorize_expires_at: Option<SystemTime>,
105
106	/// Associated User Id registration.
107	pub user_id: Option<OwnedUserId>,
108
109	/// Last userinfo response persisted here.
110	pub user_info: Option<UserInfo>,
111}
112
113/// Session Identifier type.
114pub type SessionId = String;
115
116/// Number of characters generated for our code_verifier. The code_verifier is a
117/// random string which must be between 43 and 128 characters.
118pub const CODE_VERIFIER_LENGTH: usize = 64;
119
120/// Number of characters we will generate for the Session ID.
121pub const SESSION_ID_LENGTH: usize = 32;
122
123#[implement(Sessions)]
124pub(super) fn build(args: &crate::Args<'_>, providers: Arc<Providers>) -> Self {
125	Self {
126		services: args.services.clone(),
127		association_pending: Default::default(),
128		write_locks: MutexMap::new(),
129		providers,
130		db: Data {
131			oauthid_session: args.db["oauthid_session"].clone(),
132			oauthuniqid_oauthid: args.db["oauthuniqid_oauthid"].clone(),
133			userid_oauthid: args.db["userid_oauthid"].clone(),
134			database: args.db.clone(),
135		},
136	}
137}
138
139/// Delete database state for the session.
140///
141/// The canonical session and every association index that still refers to it
142/// commit together.
143#[implement(Sessions)]
144#[tracing::instrument(level = "debug", skip(self))]
145pub async fn delete(&self, sess_id: &str) {
146	let (session, unique_id, _write_guard) = loop {
147		let Ok(snapshot) = self.get(sess_id).await else {
148			return;
149		};
150
151		let provider = async {
152			let idp_id = snapshot.idp_id.as_deref()?;
153
154			self.providers.get(idp_id).map(Result::ok).await
155		}
156		.await;
157
158		let unique_id = provider
159			.as_ref()
160			.and_then(|provider| session_unique_id((provider, &snapshot)).ok());
161
162		let write_guard = match unique_id.as_deref() {
163			| Some(unique_id) => Some(self.write_locks.lock(unique_id).await),
164			| None => None,
165		};
166
167		let Ok(session) = self.get(sess_id).await else {
168			return;
169		};
170
171		if session.idp_id.as_deref() != snapshot.idp_id.as_deref() {
172			continue;
173		}
174
175		let current_unique_id = provider
176			.as_ref()
177			.and_then(|provider| session_unique_id((provider, &session)).ok());
178
179		if current_unique_id == unique_id {
180			break (session, unique_id, write_guard);
181		}
182	};
183
184	// Preserve a unique identity association updated to a newer session.
185	let unique_id = async {
186		let unique_id = unique_id.as_deref()?;
187		let assoc_id = self
188			.get_sess_id_by_unique_id(unique_id)
189			.map(Result::ok)
190			.await?;
191
192		(assoc_id == sess_id).then_some(unique_id)
193	}
194	.await;
195
196	let user_sessions = async {
197		let user_id = session.user_id.as_deref()?;
198		let sess_ids: Vec<_> = self
199			.get_sess_id_by_user(user_id)
200			.ready_filter_map(Result::ok)
201			.ready_filter(|assoc_id| assoc_id != sess_id)
202			.collect()
203			.await;
204
205		Some((user_id, sess_ids))
206	}
207	.await;
208
209	let mut txn = self.db.database.txn();
210
211	if let Some((user_id, sess_ids)) = user_sessions {
212		if !sess_ids.is_empty() {
213			txn.raw_put(&self.db.userid_oauthid, user_id, sess_ids);
214		} else {
215			txn.del_raw(&self.db.userid_oauthid, user_id);
216		}
217	}
218
219	if let Some(unique_id) = unique_id {
220		txn.del_raw(&self.db.oauthuniqid_oauthid, unique_id);
221	}
222
223	txn.del_raw(&self.db.oauthid_session, sess_id);
224	txn.execute();
225}
226
227/// Create or overwrite database state for the session.
228///
229/// The canonical session and its available identity and user indexes commit
230/// together.
231#[implement(Sessions)]
232#[tracing::instrument(level = "info", skip(self))]
233pub async fn put(&self, session: &Session) {
234	let unique_id = async {
235		let idp_id = session.idp_id.as_ref()?;
236		let provider = self.providers.get(idp_id).map(Result::ok).await?;
237
238		session_unique_id((&provider, session)).ok()
239	}
240	.await;
241
242	let _write_guard = match unique_id.as_deref() {
243		| Some(unique_id) => Some(self.write_locks.lock(unique_id).await),
244		| None => None,
245	};
246
247	self.put_locked(session, unique_id.as_deref())
248		.await;
249}
250
251/// Build and commit a session while exclusively claiming its identity key.
252///
253/// The callback's identity lookup and user selection remain ordered with bulk
254/// adoption until the canonical session and indexes have committed.
255#[implement(Sessions)]
256#[tracing::instrument(level = "debug", skip_all)]
257pub async fn commit_identity_session<T, F, Fut>(
258	&self,
259	unique_id: &str,
260	build: F,
261) -> Result<(Session, T, Option<SessionId>)>
262where
263	T: Send,
264	F: FnOnce(Option<OwnedUserId>) -> Fut + Send,
265	Fut: Future<Output = Result<(Session, T)>> + Send,
266{
267	let write_guard = self.write_locks.lock(unique_id).await;
268	let existing = self
269		.get_by_unique_id(unique_id)
270		.await
271		.optional()?;
272
273	let old_sess_id = existing
274		.as_ref()
275		.and_then(|session| session.sess_id.clone());
276
277	let old_user_id = existing.and_then(|session| session.user_id);
278	let (session, value) = build(old_user_id).await?;
279
280	self.put_locked(&session, Some(unique_id)).await;
281	drop(write_guard);
282
283	Ok((session, value, old_sess_id))
284}
285
286#[implement(Sessions)]
287async fn put_locked(&self, session: &Session, unique_id: Option<&str>) {
288	let sess_id = session
289		.sess_id
290		.as_deref()
291		.expect("Missing session.sess_id required for sessions.put()");
292
293	let user_sessions = async {
294		let user_id = session.user_id.as_deref()?;
295		let sess_ids: Vec<_> = self
296			.get_sess_id_by_user(user_id)
297			.ready_filter_map(Result::ok)
298			.chain(once(sess_id.to_owned()).stream())
299			.collect::<Vec<_>>()
300			.map(IntoIterator::into_iter)
301			.map(Itertools::sorted_unstable)
302			.map(Itertools::dedup)
303			.map(Iterator::collect)
304			.await;
305
306		Some((user_id, sess_ids))
307	}
308	.await;
309
310	let mut txn = self.db.database.txn();
311
312	txn.raw_put(&self.db.oauthid_session, sess_id, Cbor(session));
313
314	if let Some(unique_id) = unique_id {
315		txn.insert_raw(&self.db.oauthuniqid_oauthid, unique_id, sess_id);
316	}
317
318	if let Some((user_id, sess_ids)) = user_sessions {
319		txn.raw_put(&self.db.userid_oauthid, user_id, sess_ids);
320	}
321
322	txn.execute();
323}
324
325/// Fetch database state for a session from its associated `(iss,sub)`, in case
326/// `sess_id` is not known.
327#[implement(Sessions)]
328#[tracing::instrument(level = "debug", skip(self), ret(level = "debug"))]
329pub async fn get_by_unique_id(&self, unique_id: &str) -> Result<Session> {
330	self.get_sess_id_by_unique_id(unique_id)
331		.and_then(async |sess_id| self.get(&sess_id).await)
332		.await
333}
334
335/// Fetch database state for one or more sessions from its associated `user_id`,
336/// in case `sess_id` is not known.
337#[implement(Sessions)]
338#[tracing::instrument(level = "debug", skip(self))]
339pub fn get_by_user(&self, user_id: &UserId) -> impl Stream<Item = Result<Session>> + Send {
340	self.get_sess_id_by_user(user_id)
341		.and_then(async |sess_id| self.get(&sess_id).await)
342}
343
344/// Fetch database state for a session from its `sess_id`.
345#[implement(Sessions)]
346#[tracing::instrument(level = "debug", skip(self), ret(level = "debug"))]
347pub async fn get(&self, sess_id: &str) -> Result<Session> {
348	self.db
349		.oauthid_session
350		.get(sess_id)
351		.await
352		.deserialized::<Cbor<_>>()
353		.map(at!(0))
354}
355
356/// Resolve the `sess_id` associations with a `user_id`.
357#[implement(Sessions)]
358#[tracing::instrument(level = "debug", skip(self))]
359pub fn get_sess_id_by_user(&self, user_id: &UserId) -> impl Stream<Item = Result<String>> + Send {
360	self.db
361		.userid_oauthid
362		.get(user_id)
363		.map(Deserialized::deserialized)
364		.map_ok(Vec::into_iter)
365		.map_ok(IterStream::try_stream)
366		.try_flatten_stream()
367}
368
369/// Resolve the `sess_id` from an associated provider issuer and subject hash.
370#[implement(Sessions)]
371#[tracing::instrument(level = "debug", skip(self), ret(level = "debug"))]
372pub async fn get_sess_id_by_unique_id(&self, unique_id: &str) -> Result<String> {
373	self.db
374		.oauthuniqid_oauthid
375		.get(unique_id)
376		.await
377		.deserialized()
378}
379
380#[implement(Sessions)]
381pub fn users(&self) -> impl Stream<Item = OwnedUserId> + Send {
382	self.db
383		.userid_oauthid
384		.keys()
385		.expect_ok()
386		.map(UserId::to_owned)
387}
388
389#[implement(Sessions)]
390pub fn stream(&self) -> impl Stream<Item = Session> + Send {
391	self.db
392		.oauthid_session
393		.stream()
394		.expect_ok()
395		.map(|(_, session): (Ignore, Cbor<_>)| session.0)
396}
397
398#[implement(Sessions)]
399pub async fn provider(&self, session: &Session) -> Result<Provider> {
400	let Some(idp_id) = session.idp_id.as_deref() else {
401		return Err!(Request(NotFound("No provider for this session")));
402	};
403
404	self.providers.get(idp_id).await
405}