use std::sync::Arc;
use axum::extract::{Request, State};
use axum::http::{HeaderMap, Method, StatusCode};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use serde_json::json;
use tuitbot_core::auth::session;
use crate::state::AppState;
fn extract_session_cookie(headers: &HeaderMap) -> Option<String> {
headers
.get("cookie")
.and_then(|v| v.to_str().ok())
.and_then(|cookies| {
cookies.split(';').find_map(|c| {
let c = c.trim();
c.strip_prefix("tuitbot_session=").map(|v| v.to_string())
})
})
}
const AUTH_EXEMPT_PATHS: &[&str] = &[
"/health",
"/api/health",
"/settings/status",
"/api/settings/status",
"/settings/init",
"/api/settings/init",
"/settings/test-llm",
"/api/settings/test-llm",
"/ws",
"/api/ws",
"/auth/login",
"/api/auth/login",
"/auth/status",
"/api/auth/status",
"/connectors/google-drive/callback",
"/api/connectors/google-drive/callback",
"/media/file",
"/api/media/file",
"/onboarding/x-auth/start",
"/api/onboarding/x-auth/start",
"/onboarding/x-auth/callback",
"/api/onboarding/x-auth/callback",
"/onboarding/x-auth/status",
"/api/onboarding/x-auth/status",
"/onboarding/analyze-profile",
"/api/onboarding/analyze-profile",
];
pub async fn auth_middleware(
State(state): State<Arc<AppState>>,
headers: HeaderMap,
request: Request,
next: Next,
) -> Response {
let path = request.uri().path();
if AUTH_EXEMPT_PATHS.contains(&path) {
return next.run(request).await;
}
let bearer_ok = headers
.get("authorization")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.strip_prefix("Bearer "))
.is_some_and(|token| token == state.api_token);
if bearer_ok {
return next.run(request).await;
}
if let Some(session_token) = extract_session_cookie(&headers) {
match session::validate_session(&state.db, &session_token).await {
Ok(Some(sess)) => {
let method = request.method().clone();
if method == Method::POST
|| method == Method::PATCH
|| method == Method::DELETE
|| method == Method::PUT
{
let csrf_ok = headers
.get("x-csrf-token")
.and_then(|v| v.to_str().ok())
.is_some_and(|t| t == sess.csrf_token);
if !csrf_ok {
return (
StatusCode::FORBIDDEN,
axum::Json(json!({"error": "missing or invalid CSRF token"})),
)
.into_response();
}
}
return next.run(request).await;
}
Ok(None) => { }
Err(e) => {
tracing::error!(error = %e, "Session validation failed");
}
}
}
(
StatusCode::UNAUTHORIZED,
axum::Json(json!({"error": "unauthorized"})),
)
.into_response()
}