1use std::sync::{Arc, Mutex, MutexGuard};
2
3use axum::extract::{FromRequestParts, Request, State};
4use axum::http::header::{COOKIE, SET_COOKIE};
5use axum::http::request::Parts;
6use axum::http::{HeaderMap, HeaderValue};
7use axum::middleware::Next;
8use axum::response::Response;
9use cookie::{Cookie, CookieJar, Key, SameSite};
10use serde::de::DeserializeOwned;
11use serde::{Deserialize, Serialize};
12use serde_json::{Map, Value};
13
14use crate::crypto::random_token;
15use crate::{AppState, Error, Result};
16
17const OLD_INPUT: &str = "_old_input";
18const ERRORS: &str = "_errors";
19const ERROR_BAG: &str = "_error_bag";
21
22#[derive(Clone)]
28pub struct Session(Arc<Mutex<Inner>>);
29
30#[derive(Default)]
31struct Inner {
32 data: Map<String, Value>,
33 flashed: Map<String, Value>,
35 flash_next: Map<String, Value>,
37 token: String,
38 lifetime: Option<u64>,
40 rotate: bool,
43}
44
45#[derive(Serialize, Deserialize)]
49#[serde(untagged)]
50enum Stored {
51 Handle { sid: String },
52 Full(Payload),
53}
54
55#[derive(Serialize, Deserialize)]
56struct Payload {
57 #[serde(default)]
58 data: Map<String, Value>,
59 #[serde(default)]
60 flash: Map<String, Value>,
61 token: String,
62 expires: u64,
63 #[serde(default, skip_serializing_if = "Option::is_none")]
64 lifetime: Option<u64>,
65}
66
67const OLD_INPUT_LIMIT: usize = 2048;
70
71fn drop_passwords(value: &mut Value) {
72 match value {
73 Value::Object(map) => {
74 map.retain(|key, _| !key.to_ascii_lowercase().contains("password"));
75 map.values_mut().for_each(drop_passwords);
76 }
77 Value::Array(items) => items.iter_mut().for_each(drop_passwords),
78 _ => {}
79 }
80}
81
82fn old_input(input: Value) -> Value {
83 let Value::Object(mut map) = input else {
84 return Value::Object(Map::new());
85 };
86 map.retain(|key, _| !key.starts_with('_') && !key.to_ascii_lowercase().contains("password"));
87 for value in map.values_mut() {
89 drop_passwords(value);
90 }
91 let size = |map: &Map<String, Value>| serde_json::to_string(map).map_or(0, |s| s.len());
92 while size(&map) > OLD_INPUT_LIMIT {
93 let Some(largest) = map
94 .iter()
95 .max_by_key(|(_, v)| v.to_string().len())
96 .map(|(k, _)| k.clone())
97 else {
98 break;
99 };
100 map.remove(&largest);
101 }
102 Value::Object(map)
103}
104
105impl Session {
106 fn new(payload: Option<Payload>) -> Self {
107 let inner = match payload {
108 Some(p) => Inner {
109 data: p.data,
110 flashed: p.flash,
111 flash_next: Map::new(),
112 token: p.token,
113 lifetime: p.lifetime,
114 rotate: false,
115 },
116 None => Inner {
117 token: random_token(),
118 ..Inner::default()
119 },
120 };
121 Self(Arc::new(Mutex::new(inner)))
122 }
123
124 fn lock(&self) -> MutexGuard<'_, Inner> {
125 self.0
126 .lock()
127 .unwrap_or_else(|poisoned| poisoned.into_inner())
128 }
129
130 pub fn get<T: DeserializeOwned>(&self, key: &str) -> Option<T> {
132 let inner = self.lock();
133 let value = inner
134 .data
135 .get(key)
136 .or_else(|| inner.flash_next.get(key))
137 .or_else(|| inner.flashed.get(key))?;
138 serde_json::from_value(value.clone()).ok()
139 }
140
141 pub fn has(&self, key: &str) -> bool {
143 let inner = self.lock();
144 inner.data.contains_key(key)
145 || inner.flash_next.contains_key(key)
146 || inner.flashed.contains_key(key)
147 }
148
149 pub fn put(&self, key: &str, value: impl Serialize) -> Result {
151 let value = serde_json::to_value(value)?;
152 self.lock().data.insert(key.to_owned(), value);
153 Ok(())
154 }
155
156 pub fn remove(&self, key: &str) -> Option<Value> {
158 self.lock().data.remove(key)
159 }
160
161 pub fn push(&self, key: &str, value: impl Serialize) -> Result<usize> {
173 let value = serde_json::to_value(value)?;
174 let mut inner = self.lock();
175 let entry = inner
176 .data
177 .entry(key.to_owned())
178 .or_insert_with(|| Value::Array(Vec::new()));
179 if !entry.is_array() {
180 *entry = Value::Array(vec![entry.take()]);
181 }
182 let list = entry.as_array_mut().expect("made an array above");
183 list.push(value);
184 Ok(list.len())
185 }
186
187 pub fn increment(&self, key: &str, by: i64) -> Result<i64> {
190 let mut inner = self.lock();
191 let current = inner.data.get(key).and_then(Value::as_i64).unwrap_or(0);
192 let next = current.saturating_add(by);
193 inner.data.insert(key.to_owned(), Value::from(next));
194 Ok(next)
195 }
196
197 pub fn pull<T: DeserializeOwned>(&self, key: &str) -> Option<T> {
199 let value = self.lock().data.remove(key)?;
200 serde_json::from_value(value).ok()
201 }
202
203 pub fn set_lifetime(&self, lifetime: std::time::Duration) {
206 self.lock().lifetime = Some(lifetime.as_secs().div_ceil(60).max(1));
207 }
208
209 pub fn flash(&self, key: &str, value: impl Serialize) -> Result {
211 let value = serde_json::to_value(value)?;
212 self.lock().flash_next.insert(key.to_owned(), value);
213 Ok(())
214 }
215
216 pub fn reflash(&self) {
218 let mut inner = self.lock();
219 let flashed = inner.flashed.clone();
220 for (key, value) in flashed {
221 inner.flash_next.entry(key).or_insert(value);
222 }
223 }
224
225 pub fn keep(&self, keys: &[&str]) {
229 let mut inner = self.lock();
230 for key in keys {
231 if let Some(value) = inner.flashed.get(*key).cloned() {
232 inner.flash_next.entry((*key).to_owned()).or_insert(value);
233 }
234 }
235 }
236
237 pub fn flash_now(&self, key: &str, value: impl Serialize) -> Result {
241 let value = serde_json::to_value(value)?;
242 self.lock().flashed.insert(key.to_owned(), value);
243 Ok(())
244 }
245
246 pub fn flash_input(&self, input: &impl Serialize) -> Result {
251 self.flash(OLD_INPUT, old_input(serde_json::to_value(input)?))
252 }
253
254 pub fn old(&self, field: &str) -> Option<Value> {
258 let inner = self.lock();
259 let input = inner.flashed.get(OLD_INPUT)?;
260 input
261 .get(field)
262 .or_else(|| crate::validation::nested::lookup(input, field))
263 .cloned()
264 }
265
266 pub fn has_old_input(&self) -> bool {
270 self.lock().flashed.contains_key(OLD_INPUT)
271 }
272
273 pub fn flash_errors(&self, errors: &impl Serialize) -> Result {
275 self.flash(ERRORS, errors)
276 }
277
278 pub fn flash_errors_in(&self, bag: &str, errors: &impl Serialize) -> Result {
282 self.flash(ERRORS, errors)?;
283 self.flash(ERROR_BAG, bag)
284 }
285
286 pub fn errors(&self) -> Map<String, Value> {
289 let inner = self.lock();
290 match (inner.flashed.get(ERRORS), inner.flashed.get(ERROR_BAG)) {
291 (Some(Value::Object(errors)), None) => errors.clone(),
292 _ => Map::new(),
293 }
294 }
295
296 pub fn errors_in(&self, bag: &str) -> Map<String, Value> {
298 let inner = self.lock();
299 match (inner.flashed.get(ERRORS), inner.flashed.get(ERROR_BAG)) {
300 (Some(Value::Object(errors)), Some(Value::String(name))) if name == bag => {
301 errors.clone()
302 }
303 _ => Map::new(),
304 }
305 }
306
307 pub fn error_bag(&self) -> Option<String> {
309 match self.lock().flashed.get(ERROR_BAG) {
310 Some(Value::String(name)) => Some(name.clone()),
311 _ => None,
312 }
313 }
314
315 pub fn flashed(&self) -> Map<String, Value> {
317 let mut flashed = self.lock().flashed.clone();
318 flashed.remove(OLD_INPUT);
319 flashed.remove(ERRORS);
320 flashed.remove(ERROR_BAG);
321 flashed
322 }
323
324 pub fn token(&self) -> String {
326 self.lock().token.clone()
327 }
328
329 pub fn regenerate_token(&self) {
331 let mut inner = self.lock();
332 inner.token = random_token();
333 inner.rotate = true;
334 }
335
336 pub fn flush(&self) {
338 let mut inner = self.lock();
339 *inner = Inner {
340 token: random_token(),
341 rotate: true,
342 ..Inner::default()
343 };
344 }
345
346 fn take_rotate(&self) -> bool {
347 std::mem::take(&mut self.lock().rotate)
348 }
349
350 pub(crate) fn lifetime(&self) -> Option<u64> {
351 self.lock().lifetime
352 }
353
354 fn to_payload(&self, expires: u64) -> Payload {
355 let inner = self.lock();
356 Payload {
357 data: inner.data.clone(),
358 flash: inner.flash_next.clone(),
359 token: inner.token.clone(),
360 expires,
361 lifetime: inner.lifetime,
362 }
363 }
364}
365
366impl<S: Send + Sync> FromRequestParts<S> for Session {
367 type Rejection = Error;
368
369 async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self> {
370 parts
371 .extensions
372 .get::<Session>()
373 .cloned()
374 .ok_or_else(|| anyhow::anyhow!("the session middleware is not installed").into())
375 }
376}
377
378pub(crate) async fn middleware(
379 State(state): State<AppState>,
380 mut req: Request,
381 next: Next,
382) -> Response {
383 let config = &state.config;
384 let now = unix_now();
385 let database = config.session_driver == crate::SessionDriver::Database;
386 let (sid, payload) = match read_cookie(req.headers(), &config.session_cookie, &state.key) {
387 Some(Stored::Handle { sid }) if database => {
388 let payload = store::load(&state, &sid, now).await;
389 (payload.is_some().then_some(sid), payload)
390 }
391 Some(Stored::Full(payload)) if payload.expires > now => (None, Some(payload)),
392 _ => (None, None),
393 };
394 let loaded = payload.as_ref().map(store::fingerprint);
395 let session = Session::new(payload);
396 req.extensions_mut().insert(session.clone());
397
398 let mut res = next.run(req).await;
399
400 let lifetime = session
401 .lifetime()
402 .map_or(config.session_lifetime.as_secs(), |minutes| minutes * 60);
403 let payload = session.to_payload(now + lifetime);
404 let stored = if database {
405 let rotate = session.take_rotate();
406 let (sid, fresh) = match sid {
407 Some(old) if rotate => {
408 store::destroy(&state, &old).await;
409 (random_token(), true)
410 }
411 Some(sid) => (sid, false),
412 None => (random_token(), true),
413 };
414 let changed = fresh || loaded.as_deref() != Some(store::fingerprint(&payload).as_str());
416 if let Err(err) = store::save(&state, &sid, &payload, now, changed).await {
417 tracing::error!(error = ?err, "could not save the session");
418 return res;
419 }
420 store::maybe_prune(&state, now);
421 Stored::Handle { sid }
422 } else {
423 Stored::Full(payload)
424 };
425 let value = match serde_json::to_string(&stored) {
426 Ok(value) => value,
427 Err(err) => {
428 tracing::error!(error = %err, "could not serialize the session");
429 return res;
430 }
431 };
432 let cookie = Cookie::build((config.session_cookie.clone(), value))
433 .path("/")
434 .http_only(true)
435 .same_site(SameSite::Lax)
436 .secure(config.url.starts_with("https://"))
437 .max_age(cookie::time::Duration::seconds(lifetime as i64))
438 .build();
439 let mut jar = CookieJar::new();
440 jar.private_mut(&state.key).add(cookie);
441 for cookie in jar.delta() {
442 let encoded = cookie.encoded().to_string();
443 if encoded.len() > 4000 {
444 tracing::warn!(
445 bytes = encoded.len(),
446 "the session cookie is larger than browsers reliably store; \
447 SESSION_DRIVER=database keeps sessions on the server"
448 );
449 }
450 if let Ok(value) = HeaderValue::from_str(&encoded) {
451 res.headers_mut().append(SET_COOKIE, value);
452 }
453 }
454 res
455}
456
457mod store {
459 use super::Payload;
460 use crate::{AppState, Result};
461
462 fn key(sid: &str) -> String {
464 crate::webhook::sha256_hex(sid)
465 }
466
467 pub(super) fn fingerprint(payload: &Payload) -> String {
469 serde_json::to_string(&(
470 &payload.data,
471 &payload.flash,
472 &payload.token,
473 payload.lifetime,
474 ))
475 .unwrap_or_default()
476 }
477
478 pub(super) async fn load(state: &AppState, sid: &str, now: u64) -> Option<Payload> {
479 let key = key(sid);
480 if let Some(mirror) = &state.session_mirror {
481 return mirror
482 .lock()
483 .unwrap_or_else(|e| e.into_inner())
484 .get(&key)
485 .and_then(|json| serde_json::from_str::<Payload>(json).ok())
486 .filter(|p| p.expires > now);
487 }
488 let row: Option<String> =
489 crate::db::sql("SELECT payload FROM sessions WHERE id = ? AND expires_at > ?")
490 .bind(key)
491 .bind(now as i64)
492 .scalar_optional(&state.db)
493 .await
494 .map_err(|err| tracing::error!(error = ?err, "could not read a session"))
495 .ok()
496 .flatten();
497 row.and_then(|json| serde_json::from_str(&json).ok())
498 }
499
500 pub(super) async fn save(
503 state: &AppState,
504 sid: &str,
505 payload: &Payload,
506 now: u64,
507 changed: bool,
508 ) -> Result {
509 let key = key(sid);
510 let json = serde_json::to_string(payload).map_err(anyhow::Error::from)?;
511 if let Some(mirror) = &state.session_mirror {
512 mirror
513 .lock()
514 .unwrap_or_else(|e| e.into_inner())
515 .insert(key, json);
516 return Ok(());
517 }
518 let user_id = payload
519 .data
520 .get(crate::auth::AUTH_ID)
521 .and_then(serde_json::Value::as_i64);
522 if changed {
523 crate::db::sql(
524 "INSERT INTO sessions (id, user_id, payload, expires_at, last_activity) \
525 VALUES (?, ?, ?, ?, ?) ON CONFLICT (id) DO UPDATE SET user_id = excluded.user_id, \
526 payload = excluded.payload, expires_at = excluded.expires_at, \
527 last_activity = excluded.last_activity",
528 )
529 .bind(key)
530 .bind(user_id)
531 .bind(json)
532 .bind(payload.expires as i64)
533 .bind(now as i64)
534 .execute(&state.db)
535 .await?;
536 } else {
537 crate::db::sql(
538 "UPDATE sessions SET expires_at = ?, last_activity = ? \
539 WHERE id = ? AND last_activity < ?",
540 )
541 .bind(payload.expires as i64)
542 .bind(now as i64)
543 .bind(key)
544 .bind(now as i64 - 60)
545 .execute(&state.db)
546 .await?;
547 }
548 Ok(())
549 }
550
551 pub(super) async fn destroy(state: &AppState, sid: &str) {
552 let key = key(sid);
553 if let Some(mirror) = &state.session_mirror {
554 mirror
555 .lock()
556 .unwrap_or_else(|e| e.into_inner())
557 .remove(&key);
558 return;
559 }
560 if let Err(err) = crate::db::sql("DELETE FROM sessions WHERE id = ?")
561 .bind(key)
562 .execute(&state.db)
563 .await
564 {
565 tracing::error!(error = ?err, "could not delete a session");
566 }
567 }
568
569 pub(super) fn maybe_prune(state: &AppState, now: u64) {
571 if state.session_mirror.is_some() || rand::random_range(0..50) != 0 {
572 return;
573 }
574 let db = state.db.clone();
575 tokio::spawn(async move {
576 if let Err(err) = prune(&db, now).await {
577 tracing::warn!(error = ?err, "could not prune sessions");
578 }
579 });
580 }
581
582 pub(crate) async fn prune(db: &crate::db::Db, now: u64) -> Result<u64> {
583 Ok(crate::db::sql("DELETE FROM sessions WHERE expires_at <= ?")
584 .bind(now as i64)
585 .execute(db)
586 .await?)
587 }
588}
589
590impl Session {
591 pub async fn prune_expired(db: &crate::db::Db) -> Result<u64> {
595 store::prune(db, unix_now()).await
596 }
597}
598
599pub(crate) const MIGRATION: crate::db::Migration =
602 crate::db::framework_migration!("session", "00010101000210_create_sessions_table");
603
604fn read_cookie(headers: &HeaderMap, name: &str, key: &Key) -> Option<Stored> {
605 let mut jar = CookieJar::new();
606 for header in headers.get_all(COOKIE) {
607 let Ok(header) = header.to_str() else {
608 continue;
609 };
610 for cookie in Cookie::split_parse_encoded(header.to_owned()).flatten() {
611 if cookie.name() == name {
612 jar.add_original(cookie);
613 }
614 }
615 }
616 let cookie = jar.private(key).get(name)?;
617 serde_json::from_str(cookie.value()).ok()
618}
619
620fn unix_now() -> u64 {
621 crate::clock::unix_secs().max(0) as u64
622}
623
624pub(crate) fn from_cookie(state: &AppState, cookie_header: Option<&str>) -> Session {
627 let mut headers = HeaderMap::new();
628 if let Some(value) = cookie_header.and_then(|v| HeaderValue::from_str(v).ok()) {
629 headers.insert(COOKIE, value);
630 }
631 let now = unix_now();
632 let payload = match read_cookie(&headers, &state.config.session_cookie, &state.key) {
633 Some(Stored::Full(payload)) => Some(payload),
634 Some(Stored::Handle { sid }) => state.session_mirror.as_ref().and_then(|mirror| {
636 mirror
637 .lock()
638 .unwrap_or_else(|e| e.into_inner())
639 .get(&crate::webhook::sha256_hex(&sid))
640 .and_then(|json| serde_json::from_str::<Payload>(json).ok())
641 }),
642 None => None,
643 };
644 Session::new(payload.filter(|payload| payload.expires > now))
645}
646
647pub(crate) fn cookie_pair(state: &AppState, session: &Session) -> String {
649 let lifetime = session
650 .lifetime()
651 .map_or(state.config.session_lifetime.as_secs(), |minutes| {
652 minutes * 60
653 });
654 let value = serde_json::to_string(&Stored::Full(session.to_payload(unix_now() + lifetime)))
656 .expect("session payloads serialize");
657 let mut jar = CookieJar::new();
658 jar.private_mut(&state.key)
659 .add(Cookie::new(state.config.session_cookie.clone(), value));
660 let cookie = jar
661 .delta()
662 .next()
663 .expect("the cookie was just added")
664 .encoded()
665 .to_string();
666 cookie.split(';').next().unwrap_or_default().to_owned()
667}
668
669#[cfg(test)]
670mod tests {
671 use serde_json::json;
672
673 use super::*;
674
675 #[test]
676 fn old_input_keeps_objects_only_and_drops_passwords_and_the_largest_values() {
677 assert_eq!(old_input(json!(["a", "b"])), json!({}));
679 assert_eq!(old_input(json!("text")), json!({}));
680 let big = "x".repeat(OLD_INPUT_LIMIT);
681 let kept = old_input(json!({
682 "_token": "t",
683 "name": "Ann",
684 "users": [{ "email": "a@b.c", "Password": "secret" }],
685 "essay": big,
686 }));
687 assert_eq!(
688 kept,
689 json!({ "name": "Ann", "users": [{ "email": "a@b.c" }] })
690 );
691 }
692}