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 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#[derive(Clone, Debug, Default, Deserialize, Serialize)]
56pub struct Session {
57 pub idp_id: Option<String>,
60
61 pub sess_id: Option<SessionId>,
63
64 pub token_type: Option<String>,
66
67 pub access_token: Option<String>,
69
70 pub id_token: Option<String>,
72
73 pub expires_in: Option<u64>,
75
76 pub expires_at: Option<SystemTime>,
78
79 pub refresh_token: Option<String>,
81
82 pub refresh_token_expires_in: Option<u64>,
84
85 pub refresh_token_expires_at: Option<SystemTime>,
87
88 pub scope: Option<String>,
90
91 pub redirect_url: Option<Url>,
93
94 pub code_verifier: Option<String>,
96
97 pub cookie_nonce: Option<String>,
99
100 pub query_nonce: Option<String>,
102
103 pub authorize_expires_at: Option<SystemTime>,
105
106 pub user_id: Option<OwnedUserId>,
108
109 pub user_info: Option<UserInfo>,
111}
112
113pub type SessionId = String;
115
116pub const CODE_VERIFIER_LENGTH: usize = 64;
119
120pub 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#[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 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#[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#[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#[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#[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#[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#[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#[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}