fix(backend): P0 security hardening pass

Address the four P0 security items from the project review:

* token_version validation (#1): the JWT middleware now rejects access tokens
  whose token_version claim is stale (e.g. issued before a password change).
  A short-TTL (30s) in-memory cache (TokenVersionCache) avoids a Mongo lookup
  per request; credential changes invalidate the cache immediately on the
  handling instance.

* Fail-fast config (#2): add APP_ENVIRONMENT (development|production). In
  production the server refuses to boot unless JWT_SECRET and ENCRYPTION_KEY
  are set to non-default values; development keeps the insecure defaults with
  a warning.

* Real client IP in audit logs (#4): new client_ip middleware resolves the
  originating IP (X-Forwarded-For > X-Real-IP > ConnectInfo socket) and
  exposes it via a ClientIp extractor. All five hardcoded "0.0.0.0" audit
  calls are replaced, and the missing PasswordChanged audit event is added to
  change_password. axum::serve now uses into_make_service_with_connect_info.

* Refresh token persistence (#5): refresh tokens are now stored hashed in
  MongoDB (RefreshTokenRepository) instead of an in-memory map lost on restart.
  Added /api/auth/refresh (with rotation + token_version check) and
  /api/auth/logout routes; register/login return a refresh_token; password
  change/recovery revoke all of a user's refresh tokens. JwtService now honors
  JwtConfig expiries instead of hardcoding 15min/30d, and the dead in-memory
  refresh store is removed.

Also: wire up DatabaseInitializer (was never called), fix the refresh_tokens
index to tokenHash + add an expiresAt TTL index, add sha2 dep.

Rate limiting (#3) is deferred per scope; the stub remains.

Verified: cargo fmt --check clean, cargo build/clippy --all-targets clean,
18 unit tests pass (9 new). Integration tests (tests/*) still need a live
server — fixing them is tracked as P1.

BREAKING CHANGE: AuthResponse now includes a refresh_token field.
This commit is contained in:
goose 2026-06-27 13:12:19 -03:00
parent 46695ae2a0
commit 7ba78a31fb
17 changed files with 950 additions and 128 deletions

View file

@ -1,8 +1,21 @@
RUST_LOG=info RUST_LOG=info
# Deployment environment.
# development (default) — allows insecure JWT_SECRET/ENCRYPTION_KEY defaults (with warnings).
# production — REFUSES to boot unless JWT_SECRET and ENCRYPTION_KEY are set to
# strong, non-default values.
APP_ENVIRONMENT=development
SERVER_HOST=0.0.0.0 SERVER_HOST=0.0.0.0
SERVER_PORT=8000 SERVER_PORT=8000
MONGODB_URI=mongodb://mongodb:27017 MONGODB_URI=mongodb://mongodb:27017
MONGODB_DATABASE=normogen MONGODB_DATABASE=normogen
JWT_SECRET=change-this-to-a-random-secret-key
# Generate with: openssl rand -base64 48
# MUST be changed and MUST NOT equal "secret" when APP_ENVIRONMENT=production.
JWT_SECRET=change-this-to-a-strong-random-secret-key
JWT_ACCESS_TOKEN_EXPIRY_MINUTES=15 JWT_ACCESS_TOKEN_EXPIRY_MINUTES=15
JWT_REFRESH_TOKEN_EXPIRY_DAYS=30 JWT_REFRESH_TOKEN_EXPIRY_DAYS=30
# MUST be set (and not the default) when APP_ENVIRONMENT=production.
ENCRYPTION_KEY=change-this-to-a-32-byte-key

1
backend/Cargo.lock generated
View file

@ -1327,6 +1327,7 @@ dependencies = [
"reqwest", "reqwest",
"serde", "serde",
"serde_json", "serde_json",
"sha2",
"slog", "slog",
"strum", "strum",
"strum_macros", "strum_macros",

View file

@ -22,6 +22,7 @@ pbkdf2 = { version = "0.12.2", features = ["simple"] }
password-hash = "0.5.0" password-hash = "0.5.0"
rand = "0.8.5" rand = "0.8.5"
base64 = "0.22.1" base64 = "0.22.1"
sha2 = "0.10"
thiserror = "1.0.69" thiserror = "1.0.69"
anyhow = "1.0.94" anyhow = "1.0.94"
tracing = "0.1.41" tracing = "0.1.41"

View file

@ -3,9 +3,6 @@ use anyhow::Result;
use chrono::{Duration, Utc}; use chrono::{Duration, Utc};
use jsonwebtoken::{decode, encode, DecodingKey, EncodingKey, Header, Validation}; use jsonwebtoken::{decode, encode, DecodingKey, EncodingKey, Header, Validation};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use crate::config::JwtConfig; use crate::config::JwtConfig;
@ -21,9 +18,11 @@ pub struct Claims {
} }
impl Claims { impl Claims {
pub fn new(user_id: String, email: String, token_version: i32) -> Self { /// Build access-token claims. `expiry` is the access-token lifetime and is
/// taken from `JwtConfig` by callers (no longer hardcoded here).
pub fn new(user_id: String, email: String, token_version: i32, expiry: Duration) -> Self {
let now = Utc::now(); let now = Utc::now();
let exp = now + Duration::minutes(15); // Access token expires in 15 minutes let exp = now + expiry;
Self { Self {
sub: user_id.clone(), sub: user_id.clone(),
@ -47,9 +46,10 @@ pub struct RefreshClaims {
} }
impl RefreshClaims { impl RefreshClaims {
pub fn new(user_id: String, token_version: i32) -> Self { /// Build refresh-token claims. `expiry` is taken from `JwtConfig` by callers.
pub fn new(user_id: String, token_version: i32, expiry: Duration) -> Self {
let now = Utc::now(); let now = Utc::now();
let exp = now + Duration::days(30); // Refresh token expires in 30 days let exp = now + expiry;
Self { Self {
sub: user_id.clone(), sub: user_id.clone(),
@ -61,12 +61,15 @@ impl RefreshClaims {
} }
} }
/// JWT Service for token generation and validation /// JWT Service for token generation and validation.
///
/// Pure crypto service: it signs/verifies tokens but does NOT track them. Token
/// persistence (refresh-token storage, revocation) is handled by
/// `RefreshTokenRepository` + the auth handlers; token-version staleness is
/// enforced in the JWT middleware via `TokenVersionCache`.
#[derive(Clone)] #[derive(Clone)]
pub struct JwtService { pub struct JwtService {
config: JwtConfig, config: JwtConfig,
// In-memory storage for refresh tokens (user_id -> set of tokens)
refresh_tokens: Arc<RwLock<HashMap<String, Vec<String>>>>,
encoding_key: EncodingKey, encoding_key: EncodingKey,
decoding_key: DecodingKey, decoding_key: DecodingKey,
} }
@ -78,27 +81,38 @@ impl JwtService {
Self { Self {
config, config,
refresh_tokens: Arc::new(RwLock::new(HashMap::new())),
encoding_key, encoding_key,
decoding_key, decoding_key,
} }
} }
/// Generate access and refresh tokens /// Generate access and refresh tokens. Both expiries are derived from the
/// configured `JwtConfig` so callers don't need to thread durations through.
pub fn generate_tokens(&self, claims: Claims) -> Result<(String, String)> { pub fn generate_tokens(&self, claims: Claims) -> Result<(String, String)> {
// Generate access token
let access_token = encode(&Header::default(), &claims, &self.encoding_key) let access_token = encode(&Header::default(), &claims, &self.encoding_key)
.map_err(|e| anyhow::anyhow!("Failed to encode access token: {}", e))?; .map_err(|e| anyhow::anyhow!("Failed to encode access token: {}", e))?;
// Generate refresh token let refresh_expiry = Duration::days(self.config.refresh_token_expiry_days);
let refresh_claims = RefreshClaims::new(claims.user_id.clone(), claims.token_version); let refresh_claims =
RefreshClaims::new(claims.user_id.clone(), claims.token_version, refresh_expiry);
let refresh_token = encode(&Header::default(), &refresh_claims, &self.encoding_key) let refresh_token = encode(&Header::default(), &refresh_claims, &self.encoding_key)
.map_err(|e| anyhow::anyhow!("Failed to encode refresh token: {}", e))?; .map_err(|e| anyhow::anyhow!("Failed to encode refresh token: {}", e))?;
Ok((access_token, refresh_token)) Ok((access_token, refresh_token))
} }
/// Validate access token /// The configured refresh-token lifetime, as a chrono `Duration`.
pub fn refresh_expiry(&self) -> Duration {
Duration::days(self.config.refresh_token_expiry_days)
}
/// The configured access-token lifetime, as a chrono `Duration`.
pub fn access_expiry(&self) -> Duration {
Duration::minutes(self.config.access_token_expiry_minutes)
}
/// Validate access token (signature + expiry only). Callers are responsible
/// for any `token_version` comparison against the user record.
pub fn validate_token(&self, token: &str) -> Result<Claims> { pub fn validate_token(&self, token: &str) -> Result<Claims> {
let token_data = decode::<Claims>(token, &self.decoding_key, &Validation::default()) let token_data = decode::<Claims>(token, &self.decoding_key, &Validation::default())
.map_err(|e| anyhow::anyhow!("Invalid token: {}", e))?; .map_err(|e| anyhow::anyhow!("Invalid token: {}", e))?;
@ -106,73 +120,59 @@ impl JwtService {
Ok(token_data.claims) Ok(token_data.claims)
} }
/// Validate refresh token /// Validate refresh token (signature + expiry only). Callers must still
/// confirm the token exists, is not revoked, and has not expired in the
/// `RefreshTokenRepository`.
pub fn validate_refresh_token(&self, token: &str) -> Result<RefreshClaims> { pub fn validate_refresh_token(&self, token: &str) -> Result<RefreshClaims> {
let token_data = decode::<RefreshClaims>(token, &self.decoding_key, &Validation::default()) let token_data = decode::<RefreshClaims>(token, &self.decoding_key, &Validation::default())
.map_err(|e| anyhow::anyhow!("Invalid refresh token: {}", e))?; .map_err(|e| anyhow::anyhow!("Invalid refresh token: {}", e))?;
Ok(token_data.claims) Ok(token_data.claims)
} }
}
/// Store refresh token for a user #[cfg(test)]
pub async fn store_refresh_token(&self, user_id: &str, token: &str) -> Result<()> { mod tests {
let mut tokens = self.refresh_tokens.write().await; use super::*;
tokens
.entry(user_id.to_string())
.or_insert_with(Vec::new)
.push(token.to_string());
// Keep only last 5 tokens per user fn test_config() -> JwtConfig {
if let Some(user_tokens) = tokens.get_mut(user_id) { JwtConfig {
user_tokens.sort(); secret: "test-secret".to_string(),
user_tokens.dedup(); access_token_expiry_minutes: 15,
if user_tokens.len() > 5 { refresh_token_expiry_days: 7,
*user_tokens = user_tokens.split_off(user_tokens.len() - 5);
} }
} }
Ok(()) #[test]
fn generate_then_validate_roundtrip() {
let svc = JwtService::new(test_config());
let claims = Claims::new(
"user-1".to_string(),
"u@example.com".to_string(),
0,
Duration::minutes(15),
);
let (access, refresh) = svc.generate_tokens(claims).unwrap();
let access_claims = svc.validate_token(&access).unwrap();
assert_eq!(access_claims.user_id, "user-1");
assert_eq!(access_claims.token_version, 0);
let refresh_claims = svc.validate_refresh_token(&refresh).unwrap();
assert_eq!(refresh_claims.user_id, "user-1");
} }
/// Verify if a refresh token is stored #[test]
pub async fn verify_refresh_token_stored(&self, user_id: &str, token: &str) -> Result<bool> { fn refresh_version_is_carried_through() {
let tokens = self.refresh_tokens.read().await; let svc = JwtService::new(test_config());
if let Some(user_tokens) = tokens.get(user_id) { let claims = Claims::new(
Ok(user_tokens.contains(&token.to_string())) "user-2".to_string(),
} else { "u2@example.com".to_string(),
Ok(false) 3,
} Duration::minutes(15),
} );
let (_access, refresh) = svc.generate_tokens(claims).unwrap();
/// Rotate refresh token (remove old, add new) let refresh_claims = svc.validate_refresh_token(&refresh).unwrap();
pub async fn rotate_refresh_token( assert_eq!(refresh_claims.token_version, 3);
&self,
user_id: &str,
old_token: &str,
new_token: &str,
) -> Result<()> {
// Remove old token
self.revoke_refresh_token(old_token).await?;
// Add new token
self.store_refresh_token(user_id, new_token).await?;
Ok(())
}
/// Revoke a specific refresh token
pub async fn revoke_refresh_token(&self, token: &str) -> Result<()> {
let mut tokens = self.refresh_tokens.write().await;
for user_tokens in tokens.values_mut() {
user_tokens.retain(|t| t != token);
}
Ok(())
}
/// Revoke all refresh tokens for a user
pub async fn revoke_all_user_tokens(&self, user_id: &str) -> Result<()> {
let mut tokens = self.refresh_tokens.write().await;
tokens.remove(user_id);
Ok(())
} }
} }

View file

@ -1,5 +1,7 @@
#![allow(dead_code)] #![allow(dead_code)]
pub mod jwt; pub mod jwt;
pub mod password; pub mod password;
pub mod token_version_cache;
pub use jwt::JwtService; pub use jwt::JwtService;
pub use token_version_cache::TokenVersionCache;

View file

@ -0,0 +1,98 @@
use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use anyhow::Result;
use mongodb::bson::oid::ObjectId;
use tokio::sync::RwLock;
use crate::db::MongoDb;
/// How long a cached token_version is considered fresh before re-querying Mongo.
const DEFAULT_TTL: Duration = Duration::from_secs(30);
/// Short-TTL in-memory cache of `{ user_id -> token_version }`.
///
/// This lets the JWT auth middleware reject access tokens whose `token_version`
/// claim is stale (e.g. issued before a password change) without doing a Mongo
/// lookup on every request. Stale entries are evicted on credential changes via
/// [`TokenVersionCache::invalidate`], so a version bump becomes visible within
/// at most `ttl` on other instances, and immediately on the instance that
/// processed the change.
#[derive(Clone)]
pub struct TokenVersionCache {
map: Arc<RwLock<HashMap<ObjectId, (i32, Instant)>>>,
ttl: Duration,
}
impl TokenVersionCache {
pub fn new(ttl: Duration) -> Self {
Self {
map: Arc::new(RwLock::new(HashMap::new())),
ttl,
}
}
/// Build a cache with the default 30s TTL.
pub fn with_default_ttl() -> Self {
Self::new(DEFAULT_TTL)
}
/// Return the user's current `token_version`, loading it from Mongo on a
/// cache miss or once the cached value is older than the TTL.
///
/// Returns `Ok(None)` when the user no longer exists (deleted), so callers
/// can reject the token as unauthorized.
pub async fn get_or_load(&self, user_id: &ObjectId, db: &MongoDb) -> Result<Option<i32>> {
// Fast path: serve from cache if fresh.
{
let cache = self.map.read().await;
if let Some((version, fetched_at)) = cache.get(user_id) {
if fetched_at.elapsed() < self.ttl {
return Ok(Some(*version));
}
}
}
// Slow path: query Mongo and refresh the cache.
let version = match db.find_user_by_id(user_id).await? {
Some(user) => Some(user.token_version),
None => None,
};
let mut cache = self.map.write().await;
match version {
Some(v) => {
cache.insert(*user_id, (v, Instant::now()));
}
None => {
// User was deleted: drop any stale entry so we don't serve it.
cache.remove(user_id);
}
}
Ok(version)
}
/// Drop the cached entry for a user. Call this whenever a user's
/// `token_version` is bumped (password change / recovery) so that the new
/// version becomes visible immediately instead of waiting for TTL expiry.
pub async fn invalidate(&self, user_id: &ObjectId) {
let mut cache = self.map.write().await;
cache.remove(user_id);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn invalidate_removes_entry() {
// We can't easily exercise get_or_load without Mongo, but invalidate on
// an empty cache is a no-op and must not panic.
let cache = TokenVersionCache::with_default_ttl();
cache.invalidate(&ObjectId::new()).await;
assert!(cache.map.read().await.is_empty());
}
}

View file

@ -3,6 +3,34 @@
use anyhow::Result; use anyhow::Result;
use std::sync::Arc; use std::sync::Arc;
use crate::auth::token_version_cache::TokenVersionCache;
/// Deployment environment. In `Production`, insecure config defaults are rejected
/// at boot so the service never starts with a known secret/encryption key.
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum Environment {
Development,
Production,
}
impl Environment {
/// Parse `APP_ENVIRONMENT` (case-insensitive). Defaults to `Development`.
pub fn from_env() -> Self {
match std::env::var("APP_ENVIRONMENT")
.unwrap_or_default()
.to_ascii_lowercase()
.as_str()
{
"prod" | "production" => Environment::Production,
_ => Environment::Development,
}
}
}
/// Insecure sentinel values that must never be used in production.
const INSECURE_JWT_SECRET: &str = "secret";
const INSECURE_ENCRYPTION_KEY: &str = "default_key_32_bytes_long!";
#[derive(Clone)] #[derive(Clone)]
pub struct AppState { pub struct AppState {
pub db: crate::db::MongoDb, pub db: crate::db::MongoDb,
@ -16,6 +44,15 @@ pub struct AppState {
/// Phase 2.8: Interaction checker service /// Phase 2.8: Interaction checker service
pub interaction_service: Option<Arc<crate::services::InteractionService>>, pub interaction_service: Option<Arc<crate::services::InteractionService>>,
/// P0: short-TTL cache of {user_id -> token_version} so the JWT middleware
/// can reject tokens whose version is stale (e.g. after a password change)
/// without a Mongo lookup on every request.
pub token_version_cache: Arc<TokenVersionCache>,
/// P0: persisted (hashed) refresh tokens, used for rotation, revocation,
/// and surviving process restarts.
pub refresh_token_repo: Option<Arc<crate::models::refresh_token::RefreshTokenRepository>>,
} }
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
@ -25,6 +62,7 @@ pub struct Config {
pub jwt: JwtConfig, pub jwt: JwtConfig,
pub encryption: EncryptionConfig, pub encryption: EncryptionConfig,
pub cors: CorsConfig, pub cors: CorsConfig,
pub environment: Environment,
} }
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
@ -68,6 +106,43 @@ pub struct CorsConfig {
impl Config { impl Config {
pub fn from_env() -> Result<Self> { pub fn from_env() -> Result<Self> {
let environment = Environment::from_env();
// JWT_SECRET — required to be a real secret in production.
let jwt_secret =
std::env::var("JWT_SECRET").unwrap_or_else(|_| INSECURE_JWT_SECRET.to_string());
if environment == Environment::Production {
if jwt_secret.is_empty() || jwt_secret == INSECURE_JWT_SECRET {
anyhow::bail!(
"JWT_SECRET must be set to a strong, unique value in production (APP_ENVIRONMENT=production). \
Refusing to boot with an insecure default."
);
}
tracing::info!("JWT_SECRET validated for production");
} else if jwt_secret == INSECURE_JWT_SECRET {
tracing::warn!(
"JWT_SECRET is using the insecure default \"{}\". This is allowed only in development.",
INSECURE_JWT_SECRET
);
}
// ENCRYPTION_KEY — required to be a real key in production.
let encryption_key =
std::env::var("ENCRYPTION_KEY").unwrap_or_else(|_| INSECURE_ENCRYPTION_KEY.to_string());
if environment == Environment::Production {
if encryption_key.is_empty() || encryption_key == INSECURE_ENCRYPTION_KEY {
anyhow::bail!(
"ENCRYPTION_KEY must be set to a strong, unique value in production (APP_ENVIRONMENT=production). \
Refusing to boot with an insecure default."
);
}
tracing::info!("ENCRYPTION_KEY validated for production");
} else if encryption_key == INSECURE_ENCRYPTION_KEY {
tracing::warn!(
"ENCRYPTION_KEY is using the insecure default. This is allowed only in development."
);
}
Ok(Config { Ok(Config {
server: ServerConfig { server: ServerConfig {
host: std::env::var("NORMOGEN_HOST").unwrap_or_else(|_| "0.0.0.0".to_string()), host: std::env::var("NORMOGEN_HOST").unwrap_or_else(|_| "0.0.0.0".to_string()),
@ -82,7 +157,7 @@ impl Config {
.unwrap_or_else(|_| "normogen".to_string()), .unwrap_or_else(|_| "normogen".to_string()),
}, },
jwt: JwtConfig { jwt: JwtConfig {
secret: std::env::var("JWT_SECRET").unwrap_or_else(|_| "secret".to_string()), secret: jwt_secret,
access_token_expiry_minutes: std::env::var("JWT_ACCESS_TOKEN_EXPIRY_MINUTES") access_token_expiry_minutes: std::env::var("JWT_ACCESS_TOKEN_EXPIRY_MINUTES")
.unwrap_or_else(|_| "15".to_string()) .unwrap_or_else(|_| "15".to_string())
.parse()?, .parse()?,
@ -91,8 +166,7 @@ impl Config {
.parse()?, .parse()?,
}, },
encryption: EncryptionConfig { encryption: EncryptionConfig {
key: std::env::var("ENCRYPTION_KEY") key: encryption_key,
.unwrap_or_else(|_| "default_key_32_bytes_long!".to_string()),
}, },
cors: CorsConfig { cors: CorsConfig {
allowed_origins: std::env::var("CORS_ALLOWED_ORIGINS") allowed_origins: std::env::var("CORS_ALLOWED_ORIGINS")
@ -101,6 +175,7 @@ impl Config {
.map(|s| s.to_string()) .map(|s| s.to_string())
.collect(), .collect(),
}, },
environment,
}) })
} }
} }

View file

@ -1,35 +1,31 @@
#![allow(dead_code)] #![allow(dead_code)]
use mongodb::{bson::doc, Client, Collection, IndexModel}; use mongodb::{bson::doc, options::IndexOptions, Collection, IndexModel};
use anyhow::Result; use anyhow::Result;
/// Creates required collections and indexes. Best-effort by design: index
/// creation failures are logged as warnings and do not abort startup, matching
/// the existing swallow-on-warning style. Safe to run repeatedly (idempotent).
pub struct DatabaseInitializer { pub struct DatabaseInitializer {
client: Client, db: mongodb::Database,
db_name: String,
} }
impl DatabaseInitializer { impl DatabaseInitializer {
pub fn new(client: Client, db_name: String) -> Self { pub fn new(db: mongodb::Database) -> Self {
Self { client, db_name } Self { db }
} }
pub async fn initialize(&self) -> Result<()> { pub async fn initialize(&self) -> Result<()> {
let db = self.client.database(&self.db_name);
println!("[MongoDB] Initializing database collections and indexes..."); println!("[MongoDB] Initializing database collections and indexes...");
// Create users collection and index // Create users collection and index
{ {
let collection: Collection<mongodb::bson::Document> = db.collection("users"); let collection: Collection<mongodb::bson::Document> = self.db.collection("users");
// Create email index using the builder pattern // Create email index using the builder pattern
let index = IndexModel::builder() let index = IndexModel::builder()
.keys(doc! { "email": 1 }) .keys(doc! { "email": 1 })
.options( .options(IndexOptions::builder().unique(true).build())
mongodb::options::IndexOptions::builder()
.unique(true)
.build(),
)
.build(); .build();
match collection.create_index(index, None).await { match collection.create_index(index, None).await {
@ -40,7 +36,7 @@ impl DatabaseInitializer {
// Create families collection and indexes // Create families collection and indexes
{ {
let collection: Collection<mongodb::bson::Document> = db.collection("families"); let collection: Collection<mongodb::bson::Document> = self.db.collection("families");
let index1 = IndexModel::builder().keys(doc! { "userId": 1 }).build(); let index1 = IndexModel::builder().keys(doc! { "userId": 1 }).build();
@ -62,7 +58,7 @@ impl DatabaseInitializer {
// Create profiles collection and index // Create profiles collection and index
{ {
let collection: Collection<mongodb::bson::Document> = db.collection("profiles"); let collection: Collection<mongodb::bson::Document> = self.db.collection("profiles");
let index = IndexModel::builder().keys(doc! { "familyId": 1 }).build(); let index = IndexModel::builder().keys(doc! { "familyId": 1 }).build();
@ -77,31 +73,35 @@ impl DatabaseInitializer {
// Create health_data collection // Create health_data collection
{ {
let _collection: Collection<mongodb::bson::Document> = db.collection("health_data"); let _collection: Collection<mongodb::bson::Document> =
self.db.collection("health_data");
println!("✓ Created health_data collection"); println!("✓ Created health_data collection");
} }
// Create lab_results collection // Create lab_results collection
{ {
let _collection: Collection<mongodb::bson::Document> = db.collection("lab_results"); let _collection: Collection<mongodb::bson::Document> =
self.db.collection("lab_results");
println!("✓ Created lab_results collection"); println!("✓ Created lab_results collection");
} }
// Create medications collection // Create medications collection
{ {
let _collection: Collection<mongodb::bson::Document> = db.collection("medications"); let _collection: Collection<mongodb::bson::Document> =
self.db.collection("medications");
println!("✓ Created medications collection"); println!("✓ Created medications collection");
} }
// Create appointments collection // Create appointments collection
{ {
let _collection: Collection<mongodb::bson::Document> = db.collection("appointments"); let _collection: Collection<mongodb::bson::Document> =
self.db.collection("appointments");
println!("✓ Created appointments collection"); println!("✓ Created appointments collection");
} }
// Create shares collection and index // Create shares collection and index
{ {
let collection: Collection<mongodb::bson::Document> = db.collection("shares"); let collection: Collection<mongodb::bson::Document> = self.db.collection("shares");
let index = IndexModel::builder().keys(doc! { "familyId": 1 }).build(); let index = IndexModel::builder().keys(doc! { "familyId": 1 }).build();
@ -111,23 +111,41 @@ impl DatabaseInitializer {
} }
} }
// Create refresh_tokens collection and index // Create refresh_tokens collection and indexes.
// - unique index on tokenHash (the only thing we ever look tokens up by)
// - TTL index on expiresAt so expired tokens are auto-deleted by Mongo
{ {
let collection: Collection<mongodb::bson::Document> = db.collection("refresh_tokens"); let collection: Collection<mongodb::bson::Document> =
self.db.collection("refresh_tokens");
let index = IndexModel::builder() let unique_index = IndexModel::builder()
.keys(doc! { "token": 1 }) .keys(doc! { "tokenHash": 1 })
.options(IndexOptions::builder().unique(true).build())
.build();
// TTL: documents are deleted once their `expiresAt` timestamp passes.
// With expireAfterSeconds=0, Mongo removes each doc exactly when the
// datetime stored in `expiresAt` is reached.
let ttl_index = IndexModel::builder()
.keys(doc! { "expiresAt": 1 })
.options( .options(
mongodb::options::IndexOptions::builder() IndexOptions::builder()
.unique(true) .expire_after(std::time::Duration::from_secs(0))
.build(), .build(),
) )
.build(); .build();
match collection.create_index(index, None).await { match collection.create_index(unique_index, None).await {
Ok(_) => println!("✓ Created index on refresh_tokens.token"), Ok(_) => println!("✓ Created unique index on refresh_tokens.tokenHash"),
Err(e) => println!( Err(e) => println!(
"Warning: Failed to create index on refresh_tokens.token: {}", "Warning: Failed to create index on refresh_tokens.tokenHash: {}",
e
),
}
match collection.create_index(ttl_index, None).await {
Ok(_) => println!("✓ Created TTL index on refresh_tokens.expiresAt"),
Err(e) => println!(
"Warning: Failed to create TTL index on refresh_tokens.expiresAt: {}",
e e
), ),
} }

View file

@ -18,6 +18,7 @@ pub mod init; // Database initialization module
mod mongodb_impl; mod mongodb_impl;
pub use init::DatabaseInitializer;
pub use mongodb_impl::MongoDb; pub use mongodb_impl::MongoDb;
pub async fn create_database() -> Result<Database> { pub async fn create_database() -> Result<Database> {

View file

@ -1,9 +1,11 @@
use axum::{extract::State, http::StatusCode, response::IntoResponse, Json}; use axum::{extract::Extension, extract::State, http::StatusCode, response::IntoResponse, Json};
use mongodb::bson::oid::ObjectId;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use validator::Validate; use validator::Validate;
use crate::{ use crate::{
auth::jwt::Claims, config::AppState, models::audit_log::AuditEventType, models::user::User, auth::jwt::Claims, config::AppState, middleware::ClientIp, models::audit_log::AuditEventType,
models::user::User,
}; };
#[derive(Debug, Deserialize, Validate)] #[derive(Debug, Deserialize, Validate)]
@ -21,6 +23,7 @@ pub struct RegisterRequest {
#[derive(Debug, Serialize)] #[derive(Debug, Serialize)]
pub struct AuthResponse { pub struct AuthResponse {
pub token: String, pub token: String,
pub refresh_token: String,
pub user_id: String, pub user_id: String,
pub email: String, pub email: String,
pub username: String, pub username: String,
@ -28,6 +31,7 @@ pub struct AuthResponse {
pub async fn register( pub async fn register(
State(state): State<AppState>, State(state): State<AppState>,
Extension(client_ip): Extension<ClientIp>,
Json(req): Json<RegisterRequest>, Json(req): Json<RegisterRequest>,
) -> impl IntoResponse { ) -> impl IntoResponse {
if let Err(errors) = req.validate() { if let Err(errors) = req.validate() {
@ -98,7 +102,7 @@ pub async fn register(
AuditEventType::LoginSuccess, // Using LoginSuccess as registration event AuditEventType::LoginSuccess, // Using LoginSuccess as registration event
Some(id), Some(id),
Some(req.email.clone()), Some(req.email.clone()),
"0.0.0.0".to_string(), client_ip.as_str().to_string(),
None, None,
None, None,
) )
@ -127,9 +131,14 @@ pub async fn register(
} }
}; };
// Generate JWT token // Generate JWT token pair and persist the refresh token.
let claims = Claims::new(user_id.to_string(), user.email.clone(), token_version); let claims = Claims::new(
let (token, _refresh_token) = match state.jwt_service.generate_tokens(claims) { user_id.to_string(),
user.email.clone(),
token_version,
state.jwt_service.access_expiry(),
);
let (token, refresh_token) = match state.jwt_service.generate_tokens(claims) {
Ok(t) => t, Ok(t) => t,
Err(e) => { Err(e) => {
tracing::error!("Failed to generate token: {}", e); tracing::error!("Failed to generate token: {}", e);
@ -142,9 +151,19 @@ pub async fn register(
.into_response(); .into_response();
} }
}; };
if let Some(ref repo) = state.refresh_token_repo {
let expires_at = refresh_expiry_bson(&state.jwt_service);
if let Err(e) = repo
.create(&user_id.to_string(), &refresh_token, expires_at)
.await
{
tracing::warn!("Failed to persist refresh token on register: {}", e);
}
}
let response = AuthResponse { let response = AuthResponse {
token, token,
refresh_token,
user_id: user_id.to_string(), user_id: user_id.to_string(),
email: user.email, email: user.email,
username: user.username, username: user.username,
@ -163,6 +182,7 @@ pub struct LoginRequest {
pub async fn login( pub async fn login(
State(state): State<AppState>, State(state): State<AppState>,
Extension(client_ip): Extension<ClientIp>,
Json(req): Json<LoginRequest>, Json(req): Json<LoginRequest>,
) -> impl IntoResponse { ) -> impl IntoResponse {
if let Err(errors) = req.validate() { if let Err(errors) = req.validate() {
@ -212,7 +232,7 @@ pub async fn login(
AuditEventType::LoginFailed, AuditEventType::LoginFailed,
None, None,
Some(req.email.clone()), Some(req.email.clone()),
"0.0.0.0".to_string(), // TODO: Extract real IP client_ip.as_str().to_string(),
None, None,
None, None,
) )
@ -272,7 +292,7 @@ pub async fn login(
AuditEventType::LoginFailed, AuditEventType::LoginFailed,
Some(user_id), Some(user_id),
Some(req.email.clone()), Some(req.email.clone()),
"0.0.0.0".to_string(), // TODO: Extract real IP client_ip.as_str().to_string(),
None, None,
None, None,
) )
@ -301,9 +321,14 @@ pub async fn login(
// Update last active timestamp (TODO: Implement in database layer) // Update last active timestamp (TODO: Implement in database layer)
// Generate JWT token // Generate JWT token pair and persist the refresh token.
let claims = Claims::new(user_id.to_string(), user.email.clone(), user.token_version); let claims = Claims::new(
let (token, _refresh_token) = match state.jwt_service.generate_tokens(claims) { user_id.to_string(),
user.email.clone(),
user.token_version,
state.jwt_service.access_expiry(),
);
let (token, refresh_token) = match state.jwt_service.generate_tokens(claims) {
Ok(t) => t, Ok(t) => t,
Err(e) => { Err(e) => {
tracing::error!("Failed to generate token: {}", e); tracing::error!("Failed to generate token: {}", e);
@ -316,6 +341,15 @@ pub async fn login(
.into_response(); .into_response();
} }
}; };
if let Some(ref repo) = state.refresh_token_repo {
let expires_at = refresh_expiry_bson(&state.jwt_service);
if let Err(e) = repo
.create(&user_id.to_string(), &refresh_token, expires_at)
.await
{
tracing::warn!("Failed to persist refresh token on login: {}", e);
}
}
// Log successful login (Phase 2.6) // Log successful login (Phase 2.6)
if let Some(ref audit) = state.audit_logger { if let Some(ref audit) = state.audit_logger {
@ -324,7 +358,7 @@ pub async fn login(
AuditEventType::LoginSuccess, AuditEventType::LoginSuccess,
Some(user_id), Some(user_id),
Some(req.email.clone()), Some(req.email.clone()),
"0.0.0.0".to_string(), // TODO: Extract real IP client_ip.as_str().to_string(),
None, None,
None, None,
) )
@ -333,6 +367,7 @@ pub async fn login(
let response = AuthResponse { let response = AuthResponse {
token, token,
refresh_token,
user_id: user_id.to_string(), user_id: user_id.to_string(),
email: user.email, email: user.email,
username: user.username, username: user.username,
@ -353,6 +388,7 @@ pub struct RecoverPasswordRequest {
pub async fn recover_password( pub async fn recover_password(
State(state): State<AppState>, State(state): State<AppState>,
Extension(client_ip): Extension<ClientIp>,
Json(req): Json<RecoverPasswordRequest>, Json(req): Json<RecoverPasswordRequest>,
) -> impl IntoResponse { ) -> impl IntoResponse {
if let Err(errors) = req.validate() { if let Err(errors) = req.validate() {
@ -432,6 +468,20 @@ pub async fn recover_password(
// Save updated user // Save updated user
match state.db.update_user(&user).await { match state.db.update_user(&user).await {
Ok(_) => { Ok(_) => {
// Password change bumps token_version, so existing access tokens are
// rejected; also revoke every refresh token so they can't be used to
// mint new access tokens.
if let Some(ref repo) = state.refresh_token_repo {
if let Some(ref uid) = user.id.as_ref().map(|oid| oid.to_string()) {
if let Err(e) = repo.revoke_all_by_user(uid).await {
tracing::warn!("Failed to revoke refresh tokens on recovery: {}", e);
}
}
}
if let Some(ref uid) = user.id {
state.token_version_cache.invalidate(uid).await;
}
// Log password recovery (Phase 2.6) // Log password recovery (Phase 2.6)
if let Some(ref audit) = state.audit_logger { if let Some(ref audit) = state.audit_logger {
let user_id_for_log = user.id; let user_id_for_log = user.id;
@ -440,7 +490,7 @@ pub async fn recover_password(
AuditEventType::PasswordRecovery, AuditEventType::PasswordRecovery,
user_id_for_log, user_id_for_log,
Some(req.email.clone()), Some(req.email.clone()),
"0.0.0.0".to_string(), client_ip.as_str().to_string(),
None, None,
None, None,
) )
@ -461,3 +511,212 @@ pub async fn recover_password(
} }
} }
} }
/// Compute the BSON expiry timestamp for a freshly issued refresh token, from
/// the configured refresh lifetime on the JWT service.
fn refresh_expiry_bson(jwt: &crate::auth::JwtService) -> mongodb::bson::DateTime {
let now = std::time::SystemTime::now();
let expiry = now + jwt.refresh_expiry().to_std().unwrap_or_default();
mongodb::bson::DateTime::from_system_time(expiry)
}
#[derive(Debug, Deserialize)]
pub struct RefreshRequest {
pub refresh_token: String,
}
#[derive(Debug, Serialize)]
pub struct RefreshResponse {
pub token: String,
pub refresh_token: String,
}
/// Exchange a valid refresh token for a new access/refresh pair (rotation).
///
/// Validation: signature + expiry (JWT), then the token must exist in Mongo,
/// not be revoked, and not have passed its stored `expiresAt`; finally the
/// user's current `token_version` must match the refresh token's claim (so a
/// password change invalidates refresh tokens too). The old refresh token is
/// revoked and a new one issued + stored.
pub async fn refresh(
State(state): State<AppState>,
Json(req): Json<RefreshRequest>,
) -> impl IntoResponse {
// 1. Signature + expiry check on the presented refresh token.
let refresh_claims = match state.jwt_service.validate_refresh_token(&req.refresh_token) {
Ok(c) => c,
Err(_) => {
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({ "error": "invalid refresh token" })),
)
.into_response()
}
};
// 2. Must exist and be active in the store.
let record = match state.refresh_token_repo.as_ref() {
Some(repo) => match repo.find_by_raw_token(&req.refresh_token).await {
Ok(Some(r)) => r,
Ok(None) => {
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({ "error": "invalid refresh token" })),
)
.into_response()
}
Err(e) => {
tracing::error!("Refresh token lookup failed: {}", e);
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({ "error": "database error" })),
)
.into_response();
}
},
None => {
tracing::error!("Refresh token repository not configured");
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({ "error": "server misconfiguration" })),
)
.into_response();
}
};
if record.revoked {
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({ "error": "refresh token revoked" })),
)
.into_response();
}
if record.expires_at <= mongodb::bson::DateTime::now() {
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({ "error": "refresh token expired" })),
)
.into_response();
}
// 3. Load the user and verify token_version matches.
let user_id = match ObjectId::parse_str(&refresh_claims.sub) {
Ok(oid) => oid,
Err(_) => {
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({ "error": "invalid refresh token" })),
)
.into_response()
}
};
let user = match state.db.find_user_by_id(&user_id).await {
Ok(Some(u)) => u,
Ok(None) => {
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({ "error": "invalid refresh token" })),
)
.into_response()
}
Err(e) => {
tracing::error!("Failed to load user during refresh: {}", e);
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({ "error": "database error" })),
)
.into_response();
}
};
if user.token_version != refresh_claims.token_version {
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({ "error": "credentials have changed; please log in again" })),
)
.into_response();
}
// 4. Rotate: revoke the old refresh token, issue a new pair, store the new
// refresh token.
if let Some(ref repo) = state.refresh_token_repo {
if let Some(old_id) = record.id {
if let Err(e) = repo.revoke(&old_id).await {
tracing::warn!("Failed to revoke old refresh token: {}", e);
}
}
}
let claims = Claims::new(
user_id.to_string(),
user.email.clone(),
user.token_version,
state.jwt_service.access_expiry(),
);
let (access_token, new_refresh_token) = match state.jwt_service.generate_tokens(claims) {
Ok(t) => t,
Err(e) => {
tracing::error!("Failed to generate token: {}", e);
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({ "error": "failed to generate token" })),
)
.into_response();
}
};
if let Some(ref repo) = state.refresh_token_repo {
let expires_at = refresh_expiry_bson(&state.jwt_service);
if let Err(e) = repo
.create(&user_id.to_string(), &new_refresh_token, expires_at)
.await
{
tracing::warn!("Failed to persist rotated refresh token: {}", e);
}
}
(
StatusCode::OK,
Json(RefreshResponse {
token: access_token,
refresh_token: new_refresh_token,
}),
)
.into_response()
}
/// Revoke a single refresh token (client-side logout). The access token remains
/// valid until its short expiry; this prevents the refresh token from minting
/// any further access tokens.
pub async fn logout(
State(state): State<AppState>,
Json(req): Json<RefreshRequest>,
) -> impl IntoResponse {
if let Some(ref repo) = state.refresh_token_repo {
match repo.find_by_raw_token(&req.refresh_token).await {
Ok(Some(record)) => {
if let Some(id) = record.id {
if let Err(e) = repo.revoke(&id).await {
tracing::warn!("Failed to revoke refresh token on logout: {}", e);
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({ "error": "database error" })),
)
.into_response();
}
}
}
Ok(None) => {
// Already gone or never existed — treat as success (idempotent).
}
Err(e) => {
tracing::error!("Refresh token lookup failed on logout: {}", e);
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({ "error": "database error" })),
)
.into_response();
}
}
}
StatusCode::NO_CONTENT.into_response()
}

View file

@ -9,7 +9,7 @@ pub mod shares;
pub mod users; pub mod users;
// Re-export commonly used handler functions // Re-export commonly used handler functions
pub use auth::{login, recover_password, register}; pub use auth::{login, logout, recover_password, refresh, register};
pub use health::{health_check, ready_check}; pub use health::{health_check, ready_check};
pub use health_stats::{ pub use health_stats::{
create_health_stat, delete_health_stat, get_health_stat, get_health_trends, list_health_stats, create_health_stat, delete_health_stat, get_health_stat, get_health_trends, list_health_stats,

View file

@ -3,7 +3,10 @@ use mongodb::bson::oid::ObjectId;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use validator::Validate; use validator::Validate;
use crate::{auth::jwt::Claims, config::AppState, models::user::User}; use crate::{
auth::jwt::Claims, config::AppState, middleware::ClientIp, models::audit_log::AuditEventType,
models::user::User,
};
#[derive(Debug, Serialize)] #[derive(Debug, Serialize)]
pub struct UserProfileResponse { pub struct UserProfileResponse {
@ -162,6 +165,7 @@ pub struct ChangePasswordRequest {
pub async fn change_password( pub async fn change_password(
State(state): State<AppState>, State(state): State<AppState>,
Extension(claims): Extension<Claims>, Extension(claims): Extension<Claims>,
Extension(client_ip): Extension<ClientIp>,
Json(req): Json<ChangePasswordRequest>, Json(req): Json<ChangePasswordRequest>,
) -> impl IntoResponse { ) -> impl IntoResponse {
if let Err(errors) = req.validate() { if let Err(errors) = req.validate() {
@ -175,7 +179,18 @@ pub async fn change_password(
.into_response(); .into_response();
} }
let user_id = ObjectId::parse_str(&claims.sub).unwrap(); // The middleware already validated this against a real user, but guard
// against a malformed subject anyway rather than panicking.
let user_id = match ObjectId::parse_str(&claims.sub) {
Ok(oid) => oid,
Err(_) => {
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({ "error": "invalid token" })),
)
.into_response()
}
};
let mut user = match state.db.find_user_by_id(&user_id).await { let mut user = match state.db.find_user_by_id(&user_id).await {
Ok(Some(u)) => u, Ok(Some(u)) => u,
@ -224,7 +239,7 @@ pub async fn change_password(
} }
} }
// Update password // Update password (this also bumps token_version in memory)
match user.update_password(req.new_password) { match user.update_password(req.new_password) {
Ok(_) => {} Ok(_) => {}
Err(e) => { Err(e) => {
@ -240,7 +255,33 @@ pub async fn change_password(
} }
match state.db.update_user(&user).await { match state.db.update_user(&user).await {
Ok(_) => (StatusCode::NO_CONTENT, ()).into_response(), Ok(_) => {
// Invalidate all sessions: the bumped token_version rejects existing
// access tokens, and we revoke every refresh token immediately so
// they can't mint new access tokens.
if let Some(ref repo) = state.refresh_token_repo {
if let Err(e) = repo.revoke_all_by_user(&user_id.to_string()).await {
tracing::warn!("Failed to revoke refresh tokens on password change: {}", e);
}
}
state.token_version_cache.invalidate(&user_id).await;
// Audit the credential change.
if let Some(ref audit) = state.audit_logger {
let _ = audit
.log_event(
AuditEventType::PasswordChanged,
Some(user_id),
Some(user.email.clone()),
client_ip.as_str().to_string(),
None,
None,
)
.await;
}
(StatusCode::NO_CONTENT, ()).into_response()
}
Err(e) => { Err(e) => {
tracing::error!("Failed to update user: {}", e); tracing::error!("Failed to update user: {}", e);
( (

View file

@ -70,6 +70,25 @@ async fn main() -> anyhow::Result<()> {
// Get the underlying MongoDB database for security services // Get the underlying MongoDB database for security services
let database = db.get_database(); let database = db.get_database();
// Ensure collections/indexes exist (best-effort; failures are logged, not
// fatal). Includes the refresh_tokens tokenHash-unique + expiresAt TTL
// indexes and the previously-never-run users.email unique index.
if let Err(e) = db::DatabaseInitializer::new(database.clone())
.initialize()
.await
{
tracing::warn!("Database index initialization skipped: {}", e);
}
// Short-TTL cache of token_version, so the JWT middleware can reject stale
// access tokens without a Mongo lookup per request.
let token_version_cache = std::sync::Arc::new(auth::TokenVersionCache::with_default_ttl());
// Persisted refresh tokens (hashed) for rotation/revocation.
let refresh_token_repo = std::sync::Arc::new(
models::refresh_token::RefreshTokenRepository::new(&database),
);
// Initialize security services (Phase 2.6) // Initialize security services (Phase 2.6)
let audit_logger = security::AuditLogger::new(&database); let audit_logger = security::AuditLogger::new(&database);
let session_manager = security::SessionManager::new(&database); let session_manager = security::SessionManager::new(&database);
@ -101,6 +120,8 @@ async fn main() -> anyhow::Result<()> {
health_stats_repo: Some(health_stats_repo), health_stats_repo: Some(health_stats_repo),
mongo_client: None, mongo_client: None,
interaction_service: Some(interaction_service), interaction_service: Some(interaction_service),
token_version_cache,
refresh_token_repo: Some(refresh_token_repo),
}; };
eprintln!("Building router with security middleware..."); eprintln!("Building router with security middleware...");
@ -114,6 +135,8 @@ async fn main() -> anyhow::Result<()> {
.route("/ready", get(handlers::ready_check)) .route("/ready", get(handlers::ready_check))
.route("/api/auth/register", post(handlers::register)) .route("/api/auth/register", post(handlers::register))
.route("/api/auth/login", post(handlers::login)) .route("/api/auth/login", post(handlers::login))
.route("/api/auth/refresh", post(handlers::refresh))
.route("/api/auth/logout", post(handlers::logout))
.route( .route(
"/api/auth/recover-password", "/api/auth/recover-password",
post(handlers::recover_password), post(handlers::recover_password),
@ -175,6 +198,11 @@ async fn main() -> anyhow::Result<()> {
.with_state(state) .with_state(state)
.layer( .layer(
ServiceBuilder::new() ServiceBuilder::new()
// Resolve the client IP once and stash it for handlers/middleware
// (audit logging, account-lockout forensics).
.layer(axum::middleware::from_fn(
middleware::client_ip_middleware,
))
// Add security headers first (applies to all responses) // Add security headers first (applies to all responses)
.layer(axum::middleware::from_fn( .layer(axum::middleware::from_fn(
middleware::security_headers_middleware middleware::security_headers_middleware
@ -193,7 +221,14 @@ async fn main() -> anyhow::Result<()> {
eprintln!("Server listening on {}", &addr); eprintln!("Server listening on {}", &addr);
tracing::info!("Server listening on {}", &addr); tracing::info!("Server listening on {}", &addr);
axum::serve(listener, app).await?; // into_make_service_with_connect_info makes ConnectInfo<SocketAddr> available
// to the client_ip_middleware so we can fall back to the peer address when no
// proxy header (X-Forwarded-For / X-Real-IP) is present.
axum::serve(
listener,
app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
)
.await?;
Ok(()) Ok(())
} }

View file

@ -8,6 +8,7 @@ use axum::{
middleware::Next, middleware::Next,
response::Response, response::Response,
}; };
use mongodb::bson::oid::ObjectId;
pub async fn jwt_auth_middleware( pub async fn jwt_auth_middleware(
State(state): State<AppState>, State(state): State<AppState>,
@ -29,12 +30,26 @@ pub async fn jwt_auth_middleware(
let token = &auth_header[7..]; // Remove "Bearer " prefix let token = &auth_header[7..]; // Remove "Bearer " prefix
// Verify token // Verify signature and expiry.
let claims = state let claims = state
.jwt_service .jwt_service
.validate_token(token) .validate_token(token)
.map_err(|_| StatusCode::UNAUTHORIZED)?; .map_err(|_| StatusCode::UNAUTHORIZED)?;
// Reject tokens whose version is stale (e.g. issued before a password
// change). The user's current version is cached briefly so we don't hit
// Mongo on every request; a missing user (deleted) also means 401.
let user_oid = ObjectId::parse_str(&claims.sub).map_err(|_| StatusCode::UNAUTHORIZED)?;
let current_version = state
.token_version_cache
.get_or_load(&user_oid, &state.db)
.await
.map_err(|_| StatusCode::UNAUTHORIZED)?;
match current_version {
Some(v) if v == claims.token_version => {}
_ => return Err(StatusCode::UNAUTHORIZED),
}
// Add claims to request extensions for handlers to use // Add claims to request extensions for handlers to use
req.extensions_mut().insert(claims); req.extensions_mut().insert(claims);

View file

@ -0,0 +1,137 @@
//! Resolves the originating client IP for each request and exposes it to
//! handlers via the [`ClientIp`] request extension.
//!
//! Resolution order (most-trusted first):
//! 1. `X-Forwarded-For` — first hop (the original client). Normogen is deployed
//! behind a reverse proxy on Solaria, so this is the primary source.
//! 2. `X-Real-IP` — common single-IP proxy header.
//! 3. The peer address from Axum's `ConnectInfo<SocketAddr>` (direct connection).
//! 4. `"0.0.0.0"` only if none of the above are present.
#![allow(dead_code)]
use axum::{
extract::ConnectInfo, extract::Request, http::HeaderMap, middleware::Next, response::Response,
};
use std::net::SocketAddr;
/// Header carrying the originating client chain (proxy-populated).
const X_FORWARDED_FOR: &str = "x-forwarded-for";
/// Header some proxies set to the single originating IP.
const X_REAL_IP: &str = "x-real-ip";
/// Fallback when no IP can be determined.
const UNKNOWN_IP: &str = "0.0.0.0";
/// Extractor carrying the resolved client IP, inserted by [`client_ip_middleware`].
#[derive(Debug, Clone)]
pub struct ClientIp(pub String);
impl ClientIp {
pub fn as_str(&self) -> &str {
&self.0
}
}
impl std::ops::Deref for ClientIp {
type Target = str;
fn deref(&self) -> &str {
&self.0
}
}
/// Resolve the client IP from request headers, falling back to the peer socket.
///
/// Pure and unit-testable: pass the headers and an optional `ConnectInfo`.
pub fn extract_client_ip(headers: &HeaderMap, conn: Option<&ConnectInfo<SocketAddr>>) -> String {
// X-Forwarded-For: "client, proxy1, proxy2" — the leftmost entry is the
// original client. Only validate it looks like an IP; take the first token.
if let Some(xff) = headers.get(X_FORWARDED_FOR).and_then(|h| h.to_str().ok()) {
if let Some(first) = xff.split(',').map(str::trim).find(|s| !s.is_empty()) {
return first.to_string();
}
}
// X-Real-IP: single IP.
if let Some(ip) = headers.get(X_REAL_IP).and_then(|h| h.to_str().ok()) {
let ip = ip.trim();
if !ip.is_empty() {
return ip.to_string();
}
}
// Direct peer address, when ConnectInfo is wired (see main.rs).
if let Some(info) = conn {
return info.0.ip().to_string();
}
UNKNOWN_IP.to_string()
}
/// Middleware that computes the client IP once and stashes it in request
/// extensions so handlers can pull it out with `Extension(ClientIp(..))`.
pub async fn client_ip_middleware(mut req: Request, next: Next) -> Response {
let conn = req.extensions().get::<ConnectInfo<SocketAddr>>();
let ip = extract_client_ip(req.headers(), conn);
req.extensions_mut().insert(ClientIp(ip));
next.run(req).await
}
#[cfg(test)]
mod tests {
use super::*;
use axum::extract::ConnectInfo;
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
fn headers() -> HeaderMap {
HeaderMap::new()
}
fn conn() -> ConnectInfo<SocketAddr> {
ConnectInfo(SocketAddr::new(
IpAddr::V4(Ipv4Addr::new(203, 0, 113, 7)),
12345,
))
}
#[test]
fn prefers_x_forwarded_for_first_hop() {
let mut h = headers();
h.insert(
X_FORWARDED_FOR,
"198.51.100.2, 10.0.0.1, 10.0.0.2".parse().unwrap(),
);
assert_eq!(extract_client_ip(&h, Some(&conn())), "198.51.100.2");
}
#[test]
fn falls_back_to_x_real_ip_when_no_xff() {
let mut h = headers();
h.insert(X_REAL_IP, "198.51.100.9".parse().unwrap());
assert_eq!(extract_client_ip(&h, Some(&conn())), "198.51.100.9");
}
#[test]
fn xff_takes_precedence_over_x_real_ip() {
let mut h = headers();
h.insert(X_FORWARDED_FOR, "198.51.100.2".parse().unwrap());
h.insert(X_REAL_IP, "10.0.0.99".parse().unwrap());
assert_eq!(extract_client_ip(&h, Some(&conn())), "198.51.100.2");
}
#[test]
fn uses_socket_when_no_proxy_headers() {
assert_eq!(extract_client_ip(&headers(), Some(&conn())), "203.0.113.7");
}
#[test]
fn unknown_when_nothing_available() {
assert_eq!(extract_client_ip(&headers(), None), UNKNOWN_IP);
}
#[test]
fn ignores_blank_xff_entries() {
let mut h = headers();
// Leading whitespace/empty entry should be skipped, not returned blank.
h.insert(X_FORWARDED_FOR, " , 10.0.0.5".parse().unwrap());
assert_eq!(extract_client_ip(&h, None), "10.0.0.5");
}
}

View file

@ -1,7 +1,11 @@
pub mod auth; pub mod auth;
pub mod client_ip;
pub mod rate_limit; pub mod rate_limit;
pub use auth::jwt_auth_middleware; pub use auth::jwt_auth_middleware;
pub use client_ip::{client_ip_middleware, ClientIp};
// `extract_client_ip` is `pub` in the client_ip module and used by its own
// tests; handlers consume the resolved IP via the `ClientIp` extractor instead.
pub use rate_limit::general_rate_limit_middleware; pub use rate_limit::general_rate_limit_middleware;
// Simple security headers middleware // Simple security headers middleware

View file

@ -1,18 +1,24 @@
#![allow(dead_code)] #![allow(dead_code)]
#![allow(unused_imports)] use anyhow::Result;
use mongodb::bson::{oid::ObjectId, DateTime}; use base64::Engine;
use mongodb::{
bson::{doc, oid::ObjectId, DateTime},
Collection,
};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
/// A persisted, rotated refresh token. Only the SHA-256 hash of the raw JWT
/// refresh token is ever stored — the plaintext lives only inside the signed
/// JWT handed to the client.
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RefreshToken { pub struct RefreshToken {
#[serde(rename = "_id", skip_serializing_if = "Option::is_none")] #[serde(rename = "_id", skip_serializing_if = "Option::is_none")]
pub id: Option<ObjectId>, pub id: Option<ObjectId>,
#[serde(rename = "tokenId")]
pub token_id: String,
#[serde(rename = "userId")]
pub user_id: String,
#[serde(rename = "tokenHash")] #[serde(rename = "tokenHash")]
pub token_hash: String, pub token_hash: String,
#[serde(rename = "userId")]
pub user_id: String,
#[serde(rename = "expiresAt")] #[serde(rename = "expiresAt")]
pub expires_at: DateTime, pub expires_at: DateTime,
#[serde(rename = "createdAt")] #[serde(rename = "createdAt")]
@ -22,3 +28,119 @@ pub struct RefreshToken {
#[serde(rename = "revokedAt")] #[serde(rename = "revokedAt")]
pub revoked_at: Option<DateTime>, pub revoked_at: Option<DateTime>,
} }
/// SHA-256 the raw refresh token and base64-encode the digest. Two equal
/// tokens produce equal hashes, so lookups work; the plaintext is never stored.
pub fn hash_token(token: &str) -> String {
let digest = Sha256::digest(token.as_bytes());
base64::engine::general_purpose::STANDARD.encode(digest)
}
#[derive(Clone)]
pub struct RefreshTokenRepository {
collection: Collection<RefreshToken>,
}
impl RefreshTokenRepository {
pub fn new(db: &mongodb::Database) -> Self {
Self {
collection: db.collection("refresh_tokens"),
}
}
/// Persist a refresh token record from its raw JWT form. Hashes the token
/// before storage.
pub async fn create(
&self,
user_id: &str,
raw_token: &str,
expires_at: DateTime,
) -> Result<ObjectId> {
let now = DateTime::now();
let record = RefreshToken {
id: None,
token_hash: hash_token(raw_token),
user_id: user_id.to_string(),
expires_at,
created_at: now,
revoked: false,
revoked_at: None,
};
let id = self
.collection
.insert_one(record, None)
.await?
.inserted_id
.as_object_id()
.ok_or_else(|| anyhow::anyhow!("Failed to get inserted refresh token id"))?;
Ok(id)
}
/// Look up a token record by its raw JWT form. Returns `None` if not found.
/// Callers must still check `revoked` and `expires_at`.
pub async fn find_by_raw_token(&self, raw_token: &str) -> Result<Option<RefreshToken>> {
let hash = hash_token(raw_token);
let record = self
.collection
.find_one(doc! { "tokenHash": hash }, None)
.await?;
Ok(record)
}
/// Mark a single token record (by id) as revoked.
pub async fn revoke(&self, id: &ObjectId) -> Result<()> {
self.collection
.update_one(
doc! { "_id": id },
doc! {
"$set": {
"revoked": true,
"revokedAt": DateTime::now()
}
},
None,
)
.await?;
Ok(())
}
/// Revoke every active refresh token for a user. Called on credential
/// changes (password change / recovery) and explicit "logout everywhere".
pub async fn revoke_all_by_user(&self, user_id: &str) -> Result<()> {
self.collection
.update_many(
doc! { "userId": user_id, "revoked": false },
doc! {
"$set": {
"revoked": true,
"revokedAt": DateTime::now()
}
},
None,
)
.await?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hash_token_is_deterministic() {
assert_eq!(hash_token("abc"), hash_token("abc"));
}
#[test]
fn hash_token_differs_for_different_input() {
assert_ne!(hash_token("abc"), hash_token("abd"));
}
#[test]
fn hash_token_is_base64_sha256_length() {
// SHA-256 = 32 bytes -> 44 base64 chars (standard alphabet, padded).
assert_eq!(hash_token("anything").len(), 44);
}
}