use std::time::Duration;
use argon2::password_hash::{Salt, SaltString};
use argon2::{Argon2, PasswordHash, PasswordHasher as _, PasswordVerifier as _};
use axum::extract::{FromRequestParts, Request, State};
use axum::http::header;
use axum::http::request::Parts;
use axum::middleware::Next;
use axum::response::{IntoResponse as _, Redirect, Response};
use axum_extra::extract::cookie::{Cookie, Key, SameSite, SignedCookieJar};
use base64::Engine as _;
use chrono::{DateTime, Utc};
use rand::Rng as _;
use secrecy::ExposeSecret as _;
use sha2::{Digest as _, Sha256};
use sqlx::SqlitePool;
use subtle::ConstantTimeEq as _;
use uuid::Uuid;
use crate::config::Config;
use crate::error::Error;
use crate::state::AppState;
pub type SessionRow = (Vec<u8>, String, String, DateTime<Utc>);
pub type SetPasswordTokenRow = (Vec<u8>, DateTime<Utc>, Option<DateTime<Utc>>);
fn build_session_cookie<'a>(
name: String,
value: String,
cfg: &Config,
max_age_seconds: i64,
) -> Cookie<'a> {
let mut cookie = Cookie::new(name, value);
cookie.set_path("/");
cookie.set_http_only(true);
cookie.set_secure(cfg.cookie_secure);
cookie.set_same_site(SameSite::Lax);
cookie.set_max_age(time::Duration::seconds(max_age_seconds));
cookie
}
#[must_use]
pub fn removal_cookie<'a>(cfg: &Config) -> Cookie<'a> {
let mut cookie = Cookie::new(cfg.cookie_name.clone(), String::new());
cookie.set_path("/");
cookie.set_http_only(true);
cookie.set_secure(cfg.cookie_secure);
cookie.set_same_site(SameSite::Lax);
cookie.set_max_age(time::Duration::seconds(0));
cookie
}
#[derive(Debug, Clone)]
pub struct CurrentUser {
pub user_id: Uuid,
pub legacy_name: String,
pub username: String,
}
impl FromRequestParts<AppState> for CurrentUser {
type Rejection = Error;
async fn from_request_parts(
parts: &mut Parts,
state: &AppState,
) -> Result<Self, Self::Rejection> {
let jar = SignedCookieJar::<Key>::from_headers(&parts.headers, state.cookie_key.clone());
let Some(cookie) = jar.get(&state.config.cookie_name) else {
return Err(Error::Unauthenticated);
};
let Some(session_id) = decode_session_id(cookie.value()) else {
return Err(Error::Unauthenticated);
};
let session_id_hash = Sha256::digest(session_id).to_vec();
let now = Utc::now();
let row: Option<SessionRow> = sqlx::query_as(
"SELECT users.user_id, users.legacy_name, users.username, sessions.expires_at \
FROM sessions \
JOIN users ON users.user_id = sessions.user_id \
WHERE sessions.session_id_hash = ?1",
)
.bind(session_id_hash.as_slice())
.fetch_optional(&state.db)
.await
.map_err(|err| {
tracing::error!("session lookup failed: {err}");
Error::Database
})?;
let Some((uid_bytes, legacy_name, username, expires_at)) = row else {
return Err(Error::Unauthenticated);
};
if expires_at <= now {
return Err(Error::Unauthenticated);
}
let user_id = uuid_from_bytes(&uid_bytes).ok_or(Error::Unauthenticated)?;
let stale_cutoff = now
.checked_sub_signed(chrono::Duration::seconds(60))
.unwrap_or(chrono::DateTime::<Utc>::MIN_UTC);
if let Err(err) = sqlx::query(
"UPDATE sessions SET last_seen_at = ?1 \
WHERE session_id_hash = ?2 AND last_seen_at < ?3",
)
.bind(now)
.bind(session_id_hash.as_slice())
.bind(stale_cutoff)
.execute(&state.db)
.await
{
tracing::warn!("failed to bump sessions.last_seen_at: {err}");
}
Ok(Self {
user_id,
legacy_name,
username,
})
}
}
pub async fn require_session(State(state): State<AppState>, req: Request, next: Next) -> Response {
let (mut parts, body) = req.into_parts();
match CurrentUser::from_request_parts(&mut parts, &state).await {
Ok(user) => {
parts.extensions.insert(user);
let req = Request::from_parts(parts, body);
next.run(req).await
}
Err(_) => {
let wants_html = parts.method == axum::http::Method::GET && prefers_html(&parts);
if wants_html {
let next_url = parts
.uri
.path_and_query()
.map_or_else(|| "/".to_owned(), |pq| pq.as_str().to_owned());
let encoded = percent_encode(&next_url).unwrap_or_else(|| "/".to_owned());
let target = format!("/login?next={encoded}");
Redirect::to(&target).into_response()
} else {
Error::Unauthenticated.into_response()
}
}
}
}
fn prefers_html(parts: &Parts) -> bool {
parts
.headers
.get(header::ACCEPT)
.and_then(|v| v.to_str().ok())
.is_some_and(|s| s.contains("text/html"))
}
fn percent_encode(input: &str) -> Option<String> {
let mut out = String::with_capacity(input.len());
for byte in input.bytes() {
if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.' | b'~' | b'/') {
out.push(char::from(byte));
} else {
use std::fmt::Write as _;
write!(out, "%{byte:02X}").ok()?;
}
}
Some(out)
}
#[derive(Debug)]
pub struct LslBearer;
impl FromRequestParts<AppState> for LslBearer {
type Rejection = Error;
async fn from_request_parts(
parts: &mut Parts,
state: &AppState,
) -> Result<Self, Self::Rejection> {
let Some(value) = parts.headers.get(header::AUTHORIZATION) else {
return Err(Error::Unauthenticated);
};
let bytes = value.as_bytes();
let prefix = b"Bearer ";
if bytes.len() <= prefix.len() || !bytes.starts_with(prefix) {
return Err(Error::Unauthenticated);
}
let presented = bytes.get(prefix.len()..).unwrap_or(&[]);
let expected = state
.config
.lsl_registration_bearer_token
.expose_secret()
.as_bytes();
let len_ok = presented.len() == expected.len();
let lhs = if len_ok { presented } else { expected };
let eq: bool = lhs.ct_eq(expected).into();
if len_ok && eq {
Ok(Self)
} else {
Err(Error::Unauthenticated)
}
}
}
pub fn hash_password(password: &str) -> Result<String, Error> {
let mut salt_bytes = [0_u8; Salt::RECOMMENDED_LENGTH];
rand::rng().fill_bytes(&mut salt_bytes);
let salt =
SaltString::encode_b64(&salt_bytes).map_err(|err| Error::PasswordHash(err.to_string()))?;
let argon = Argon2::default();
let hash = argon
.hash_password(password.as_bytes(), &salt)
.map_err(|err| Error::PasswordHash(err.to_string()))?;
Ok(hash.to_string())
}
pub fn verify_password(password: &str, stored_hash: &str) -> Result<bool, Error> {
let parsed =
PasswordHash::new(stored_hash).map_err(|err| Error::PasswordHash(err.to_string()))?;
match Argon2::default().verify_password(password.as_bytes(), &parsed) {
Ok(()) => Ok(true),
Err(argon2::password_hash::Error::Password) => Ok(false),
Err(err) => Err(Error::PasswordHash(err.to_string())),
}
}
#[must_use]
pub fn generate_token() -> (String, Vec<u8>) {
let mut raw = [0_u8; 32];
rand::rng().fill_bytes(&mut raw);
let encoded = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(raw);
let hash = Sha256::digest(raw).to_vec();
(encoded, hash)
}
#[must_use]
pub fn hash_token(raw_token: &str) -> Option<Vec<u8>> {
let bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(raw_token)
.ok()?;
if bytes.len() != 32 {
return None;
}
Some(Sha256::digest(&bytes).to_vec())
}
#[must_use]
pub fn generate_session_id() -> ([u8; 32], String) {
let mut raw = [0_u8; 32];
rand::rng().fill_bytes(&mut raw);
let encoded = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(raw);
(raw, encoded)
}
fn decode_session_id(value: &str) -> Option<[u8; 32]> {
let bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(value)
.ok()?;
if bytes.len() != 32 {
return None;
}
let mut out = [0_u8; 32];
out.copy_from_slice(&bytes);
Some(out)
}
#[must_use]
pub fn uuid_from_bytes(bytes: &[u8]) -> Option<Uuid> {
let arr: [u8; 16] = bytes.try_into().ok()?;
Some(Uuid::from_bytes(arr))
}
pub async fn create_session(
pool: &SqlitePool,
cfg: &Config,
user_id: Uuid,
client_ip: Option<String>,
) -> Result<Cookie<'static>, Error> {
let (raw, encoded) = generate_session_id();
let now = Utc::now();
let expires_at = now
.checked_add_signed(chrono::Duration::seconds(cfg.session_ttl_seconds))
.unwrap_or(chrono::DateTime::<Utc>::MAX_UTC);
let session_id_hash = Sha256::digest(raw).to_vec();
sqlx::query(
"INSERT INTO sessions \
(session_id_hash, user_id, expires_at, created_at, last_seen_at, client_ip) \
VALUES (?1, ?2, ?3, ?4, ?4, ?5)",
)
.bind(session_id_hash)
.bind(user_id.as_bytes().to_vec())
.bind(expires_at)
.bind(now)
.bind(client_ip)
.execute(pool)
.await
.map_err(|err| {
tracing::error!("failed to insert session row: {err}");
Error::Database
})?;
Ok(build_session_cookie(
cfg.cookie_name.clone(),
encoded,
cfg,
cfg.session_ttl_seconds,
))
}
pub async fn delete_session(pool: &SqlitePool, session_id: &[u8]) -> Result<(), Error> {
let session_id_hash = Sha256::digest(session_id).to_vec();
sqlx::query("DELETE FROM sessions WHERE session_id_hash = ?1")
.bind(session_id_hash)
.execute(pool)
.await
.map_err(|err| {
tracing::error!("failed to delete session row: {err}");
Error::Database
})?;
Ok(())
}
#[must_use]
pub fn session_id_from_jar(jar: &SignedCookieJar, cookie_name: &str) -> Option<[u8; 32]> {
let cookie = jar.get(cookie_name)?;
decode_session_id(cookie.value())
}
pub type UserRow = (Vec<u8>, String, String, Option<String>);
pub async fn lookup_user_by_identifier(
pool: &SqlitePool,
identifier: &str,
) -> Result<Option<UserRow>, Error> {
if let Ok(uuid) = Uuid::parse_str(identifier) {
let bytes = uuid.as_bytes().to_vec();
let row: Option<UserRow> = sqlx::query_as(
"SELECT user_id, legacy_name, username, password_hash FROM users WHERE user_id = ?1",
)
.bind(bytes)
.fetch_optional(pool)
.await
.map_err(|err| {
tracing::error!("user lookup (uuid) failed: {err}");
Error::Database
})?;
if row.is_some() {
return Ok(row);
}
}
let by_username: Option<UserRow> = sqlx::query_as(
"SELECT user_id, legacy_name, username, password_hash FROM users WHERE username = ?1",
)
.bind(identifier)
.fetch_optional(pool)
.await
.map_err(|err| {
tracing::error!("user lookup (username) failed: {err}");
Error::Database
})?;
if by_username.is_some() {
return Ok(by_username);
}
let by_legacy: Option<UserRow> = sqlx::query_as(
"SELECT user_id, legacy_name, username, password_hash FROM users WHERE legacy_name = ?1",
)
.bind(identifier)
.fetch_optional(pool)
.await
.map_err(|err| {
tracing::error!("user lookup (legacy_name) failed: {err}");
Error::Database
})?;
Ok(by_legacy)
}
pub async fn run_cleanup(pool: SqlitePool) {
let mut interval = tokio::time::interval(Duration::from_secs(300));
loop {
interval.tick().await;
let now = Utc::now();
if let Err(err) = sqlx::query("DELETE FROM sessions WHERE expires_at <= ?1")
.bind(now)
.execute(&pool)
.await
{
tracing::warn!("session cleanup failed: {err}");
}
if let Err(err) = sqlx::query(
"DELETE FROM set_password_tokens WHERE expires_at <= ?1 OR used_at IS NOT NULL",
)
.bind(now)
.execute(&pool)
.await
{
tracing::warn!("set-password token cleanup failed: {err}");
}
}
}