use std::collections::HashMap; use std::net::{IpAddr, Ipv6Addr, SocketAddr}; use std::sync::{LazyLock, Mutex}; use argon2::Argon2; use argon2::password_hash::{PasswordHasher, PasswordVerifier, phc::PasswordHash}; use axum::extract::{ConnectInfo, FromRequestParts}; use axum::http::request::Parts; use axum::http::{HeaderMap, header}; use rusqlite::{OptionalExtension, params}; use sha2::{Digest, Sha256}; use tokio::sync::Semaphore; use crate::{AppState, Error, now}; pub const SESSION_COOKIE: &str = "ot_session"; pub const SESSION_SECS: i64 = 30 * 86400; pub fn hash_password(password: &str) -> String { Argon2::default() .hash_password(password.as_bytes()) .expect("argon2 with default parameters") .to_string() } pub fn verify_password(password: &str, hash: &str) -> bool { PasswordHash::new(hash).is_ok_and(|h| { Argon2::default() .verify_password(password.as_bytes(), &h) .is_ok() }) } pub fn check_new_password(password: &str) -> Result<(), &'static str> { if password.chars().count() < 8 { return Err("password needs at least 8 characters"); } Ok(()) } /// A random secret for a session cookie or a device token, and the hash to store for it. pub fn new_secret() -> (String, Vec) { let mut bytes = [0u8; 32]; getrandom::fill(&mut bytes).expect("OS random number generator"); let secret = hex(&bytes); let hash = hash_secret(&secret); (secret, hash) } pub fn hex(bytes: &[u8]) -> String { bytes.iter().map(|b| format!("{b:02x}")).collect() } /// The secrets are 256 random bits, so a fast hash is enough. Argon2 is only for passwords. fn hash_secret(secret: &str) -> Vec { Sha256::digest(secret.as_bytes()).to_vec() } static DUMMY_HASH: LazyLock = LazyLock::new(|| hash_password("not a real password")); pub struct PasswordOk { pub id: i64, pub two_factor: bool, } /// Checks a username and password. An account without a password always fails. pub async fn check_password( state: &AppState, ip: IpAddr, username: &str, password: &str, ) -> Result { let row: Option<(i64, Option, bool)> = state .db() .query_row( "SELECT id, pw_hash, two_factor FROM users WHERE username = ?1", [username], |r| Ok((r.get(0)?, r.get(1)?, r.get(2)?)), ) .optional()?; let password = password.to_owned(); limited(state, ip, username, move || match row { Some((id, Some(hash), two_factor)) => { verify_password(&password, &hash).then_some(PasswordOk { id, two_factor }) } _ => { // Do the same work as for a real password, so timing does not reveal which accounts exist. verify_password(&password, &DUMMY_HASH); None } }) .await } /// Runs a password check for `name` from `ip` under the rate limits. None from `check` means a wrong password. pub async fn limited( state: &AppState, ip: IpAddr, name: &str, check: impl FnOnce() -> Option + Send + 'static, ) -> Result { let single = ip_group(ip); let name = name.to_lowercase(); let pair = format!("{single} {name}"); let any = format!("any {name}"); let mut keys = vec![ (pair.as_str(), MAX_FAILURES), (single.as_str(), MAX_IP_FAILURES), ]; // Guesses from many addresses lock the name for new addresses only, so a stranger cannot lock out the owner. if !state.limiter.known(&pair) { keys.push((&any, MAX_FAILURES)); } state.limiter.attempt(&keys)?; match argon(check).await? { Some(v) => { state.limiter.forgive(&pair, &[&single, &any]); Ok(v) } None => Err(Error::Unauthorized), } } /// The limits count an IPv6 /64 as one address, since one client usually holds the whole /64. pub fn ip_group(ip: IpAddr) -> String { match ip.to_canonical() { IpAddr::V4(v4) => v4.to_string(), IpAddr::V6(v6) => format!( "{}/64", Ipv6Addr::from(u128::from(v6) & !(u64::MAX as u128)) ), } } /// Each Argon2 run takes about 19 MiB. Without a bound, parallel requests could exhaust the memory. static ARGON: LazyLock = LazyLock::new(|| Semaphore::new(std::thread::available_parallelism().map_or(2, |n| n.get()))); /// Runs Argon2 work on the blocking pool, at most one run per CPU at a time. pub async fn argon( work: impl FnOnce() -> T + Send + 'static, ) -> Result { let _permit = ARGON .acquire() .await .map_err(|e| Error::Internal(e.to_string()))?; tokio::task::spawn_blocking(work) .await .map_err(|e| Error::Internal(e.to_string())) } pub async fn hash_password_async(password: String) -> Result { argon(move || hash_password(&password)).await } fn cookie(state: &AppState, value: &str, max_age: i64) -> String { let secure = if state.https() { "; Secure" } else { "" }; format!( "{SESSION_COOKIE}={value}; Path=/; Max-Age={max_age}; HttpOnly; SameSite=Strict{secure}" ) } /// Creates a session and returns its Set-Cookie value. pub fn create_session(state: &AppState, user_id: i64) -> Result { let (token, hash) = new_secret(); let now = now(); state.db().execute( "INSERT INTO sessions (token_hash, user_id, expires_at, signed_in_at) VALUES (?1, ?2, ?3, ?4)", params![hash, user_id, now + SESSION_SECS, now], )?; Ok(cookie(state, &token, SESSION_SECS)) } pub fn clear_session(state: &AppState) -> String { cookie(state, "", 0) } /// A credential changed. Sessions that existed before must not outlive it. pub fn end_other_sessions(state: &AppState, user: &User) -> Result<(), Error> { state.db().execute( "DELETE FROM sessions WHERE user_id = ?1 AND token_hash <> ?2", params![user.id, user.session_hash], )?; Ok(()) } /// Wrong passwords per name before a lockout: from one address, and from all addresses together. const MAX_FAILURES: u32 = 5; /// Wrong passwords per address across all names, so guessing one password for many users is limited too. const MAX_IP_FAILURES: u32 = 30; const LOCKOUT_SECS: i64 = 15 * 60; /// How long an address that passed a password check stays exempt from the lock across all addresses. const KNOWN_SECS: i64 = 30 * 86400; /// Attempts per key. An attempt counts before the check, so parallel requests cannot slip past the limit. #[derive(Default)] pub struct Limiter { counts: Mutex>, /// Address and name pairs that passed, with the time. known: Mutex>, } impl Limiter { /// Counts one attempt for every key, or none if any key is at its limit. pub fn attempt(&self, keys: &[(&str, u32)]) -> Result<(), Error> { let now = now(); let mut map = self.counts.lock().unwrap(); let locked = keys.iter().any(|(key, max)| { map.get(*key) .is_some_and(|&(n, since)| n >= *max && now - since < LOCKOUT_SECS) }); if locked { return Err(Error::TooManyRequests); } for (key, _) in keys { let entry = map.entry((*key).to_owned()).or_insert((0, now)); if now - entry.1 >= LOCKOUT_SECS { *entry = (0, now); } entry.0 += 1; } Ok(()) } /// The password was right: earlier failures for the pair are forgotten, and this attempt does not count for the other keys. pub fn forgive(&self, pair: &str, others: &[&str]) { let mut map = self.counts.lock().unwrap(); map.remove(pair); for key in others { if let Some(entry) = map.get_mut(*key) { entry.0 = entry.0.saturating_sub(1); } } self.known.lock().unwrap().insert(pair.to_owned(), now()); } pub fn known(&self, pair: &str) -> bool { let known = self.known.lock().unwrap(); known.get(pair).is_some_and(|&t| now() - t < KNOWN_SECS) } /// Attackers choose the names, so old entries must go. pub fn prune(&self) { let now = now(); self.counts .lock() .unwrap() .retain(|_, (_, since)| now - *since < LOCKOUT_SECS); self.known .lock() .unwrap() .retain(|_, t| now - *t < KNOWN_SECS); } } /// The client's address. With `--trusted-proxy`, the last `X-Forwarded-For` entry. pub struct ClientIp(pub IpAddr); /// The last `X-Forwarded-For` entry, across all header lines. The proxy appends the address it saw, /// so earlier entries may come from the client. Falls back to the peer, which then is the proxy. fn forwarded_ip(peer: IpAddr, headers: &HeaderMap) -> IpAddr { headers .get_all("x-forwarded-for") .iter() .filter_map(|v| v.to_str().ok()) .flat_map(|v| v.split(',')) .map(str::trim) .rfind(|v| !v.is_empty()) .and_then(|v| { v.parse::() .or_else(|_| v.parse::().map(|s| s.ip())) .ok() }) .unwrap_or(peer) } impl FromRequestParts for ClientIp { type Rejection = Error; async fn from_request_parts(parts: &mut Parts, state: &AppState) -> Result { let peer = parts .extensions .get::>() .map(|c| c.0.ip()) .ok_or_else(|| Error::Internal("no connect info".into()))?; Ok(ClientIp(match state.trusted_proxy { true => forwarded_ip(peer, &parts.headers), false => peer, })) } } /// A logged-in web user, from the session cookie. pub struct User { pub id: i64, pub username: String, pub is_admin: bool, pub session_hash: Vec, /// When this session proved a password or passkey. pub signed_in_at: i64, } /// How long after sign-in a session may change credentials. const RECENT_SECS: i64 = 10 * 60; impl User { /// Credential changes need a recent sign-in, so a stolen session cannot take over the account. pub fn check_recent(&self) -> Result<(), Error> { if now() - self.signed_in_at > RECENT_SECS { return Err(Error::BadRequest( "log out and log in again before you change this".into(), )); } Ok(()) } } impl FromRequestParts for User { type Rejection = Error; async fn from_request_parts(parts: &mut Parts, state: &AppState) -> Result { let token = parts .headers .get_all(header::COOKIE) .iter() .filter_map(|v| v.to_str().ok()) .flat_map(|v| v.split(';')) .find_map(|c| c.trim().strip_prefix(SESSION_COOKIE)?.strip_prefix('=')) .ok_or(Error::Unauthorized)?; let session_hash = hash_secret(token); let (id, username, is_admin, signed_in_at) = state .db() .query_row( "SELECT u.id, u.username, u.is_admin, s.signed_in_at FROM sessions s JOIN users u ON u.id = s.user_id WHERE s.token_hash = ?1 AND s.expires_at > ?2", params![session_hash, now()], |r| Ok((r.get(0)?, r.get(1)?, r.get(2)?, r.get(3)?)), ) .optional()? .ok_or(Error::Unauthorized)?; Ok(User { id, username, is_admin, session_hash, signed_in_at, }) } } /// A logged-in admin. pub struct Admin(pub User); impl FromRequestParts for Admin { type Rejection = Error; async fn from_request_parts(parts: &mut Parts, state: &AppState) -> Result { let user = User::from_request_parts(parts, state).await?; if !user.is_admin { return Err(Error::Forbidden); } Ok(Admin(user)) } } /// Who uploads points: a device with its `Authorization: Bearer` token, or the web UI with its session. pub struct Uploader { pub device_id: i64, } impl FromRequestParts for Uploader { type Rejection = Error; async fn from_request_parts(parts: &mut Parts, state: &AppState) -> Result { if parts.headers.contains_key(header::AUTHORIZATION) { let d = Device::from_request_parts(parts, state).await?; return Ok(Uploader { device_id: d.id }); } let user = User::from_request_parts(parts, state).await?; let db = state.db(); db.execute( "INSERT OR IGNORE INTO devices (user_id, name, created_at) VALUES (?1, 'Web', ?2)", params![user.id, now()], )?; let device_id = db.query_row( "SELECT id FROM devices WHERE user_id = ?1 AND token_hash IS NULL", [user.id], |r| r.get(0), )?; Ok(Uploader { device_id }) } } /// A device, from its `Authorization: Bearer` token. pub struct Device { pub id: i64, pub user_id: i64, } impl FromRequestParts for Device { type Rejection = Error; async fn from_request_parts(parts: &mut Parts, state: &AppState) -> Result { let token = parts .headers .get(header::AUTHORIZATION) .and_then(|v| v.to_str().ok()) .and_then(|v| v.strip_prefix("Bearer ")) .ok_or(Error::Unauthorized)?; let (id, user_id) = state .db() .query_row( "SELECT id, user_id FROM devices WHERE token_hash = ?1", [hash_secret(token)], |r| Ok((r.get(0)?, r.get(1)?)), ) .optional()? .ok_or(Error::Unauthorized)?; Ok(Device { id, user_id }) } } #[cfg(test)] mod tests { use super::*; #[test] fn limiter_locks_after_max_failures() { let l = Limiter::default(); let keys = |pair| [(pair, MAX_FAILURES), ("ip", MAX_IP_FAILURES)]; for _ in 0..MAX_FAILURES { assert!(l.attempt(&keys("ip a")).is_ok()); } assert!(matches!( l.attempt(&keys("ip a")), Err(Error::TooManyRequests) )); assert!(l.attempt(&keys("ip b")).is_ok()); l.forgive("ip a", &["ip"]); assert!(l.attempt(&keys("ip a")).is_ok()); for n in 0..MAX_IP_FAILURES { let pair = format!("ip {n}"); let _ = l.attempt(&[(&pair, MAX_FAILURES), ("ip", MAX_IP_FAILURES)]); } assert!(matches!( l.attempt(&keys("ip new")), Err(Error::TooManyRequests) )); } #[test] fn a_name_locks_across_addresses_but_not_for_known_ones() { let state = crate::test_state(); let rt = tokio::runtime::Runtime::new().unwrap(); let guess = |ip: &str, ok: bool| { rt.block_on(limited(&state, ip.parse().unwrap(), "Alice", move || { ok.then_some(()) })) }; assert!(guess("198.51.100.1", true).is_ok()); for n in 0..MAX_FAILURES { assert!(matches!( guess(&format!("2001:db8:{n}::1"), false), Err(Error::Unauthorized) )); } assert!(matches!( guess("2001:db8:99::1", true), Err(Error::TooManyRequests) )); assert!(guess("198.51.100.1", true).is_ok()); } #[test] fn ipv6_addresses_group_by_64() { let group = |s: &str| ip_group(s.parse().unwrap()); assert_eq!(group("2001:db8:1:2:aaaa::1"), group("2001:db8:1:2:bbbb::2")); assert_ne!(group("2001:db8:1:2::1"), group("2001:db8:1:3::1")); assert_eq!(group("::ffff:203.0.113.7"), "203.0.113.7"); } #[test] fn credential_changes_need_a_recent_sign_in() { let user = |signed_in_at| User { id: 1, username: "a".into(), is_admin: false, session_hash: vec![], signed_in_at, }; assert!(user(now() - 60).check_recent().is_ok()); assert!(user(now() - RECENT_SECS - 1).check_recent().is_err()); assert!(user(0).check_recent().is_err()); } #[test] fn forwarded_ip_takes_the_entry_the_proxy_wrote() { let proxy: IpAddr = "10.0.0.2".parse().unwrap(); let headers = |lines: &[&str]| { let mut h = HeaderMap::new(); for l in lines { h.append("x-forwarded-for", l.parse().unwrap()); } h }; let ip = |s: &str| s.parse::().unwrap(); // nginx appends the real address to a value the client sent. assert_eq!( forwarded_ip(proxy, &headers(&["6.6.6.6, 203.0.113.7"])), ip("203.0.113.7") ); assert_eq!( forwarded_ip(proxy, &headers(&["6.6.6.6", "203.0.113.7"])), ip("203.0.113.7") ); assert_eq!( forwarded_ip(proxy, &headers(&["[2001:db8::1]:4711"])), ip("2001:db8::1") ); assert_eq!( forwarded_ip(proxy, &headers(&["203.0.113.7,"])), ip("203.0.113.7") ); assert_eq!( forwarded_ip(proxy, &headers(&["203.0.113.7, nonsense"])), proxy ); assert_eq!(forwarded_ip(proxy, &HeaderMap::new()), proxy); } #[test] fn password_roundtrip() { let h = hash_password("correct horse"); assert!(verify_password("correct horse", &h)); assert!(!verify_password("wrong horse", &h)); } }