//! Authenticated-user extractor shared by all protected API routes. use std::sync::Arc; use axum::extract::FromRequestParts; use axum::http::request::Parts; use axum::http::StatusCode; use crate::auth::parse_session_cookie; use crate::db::{RootRow, User}; use crate::error::{ApiError, AppState}; pub struct AuthUser { #[allow(dead_code)] // used from milestone 3 onwards pub user: User, pub roots: Vec, } impl FromRequestParts for AuthUser where S: HasState + Send + Sync, { type Rejection = ApiError; async fn from_request_parts(parts: &mut Parts, state: &S) -> Result { let state = state.state(); let Some(token) = parse_session_cookie(&parts.headers) else { return Err(ApiError::new( StatusCode::UNAUTHORIZED, "not signed in", )); }; let Some(user) = state.db.session_user(&token).await else { return Err(ApiError::new( StatusCode::UNAUTHORIZED, "session expired, please sign in again", )); }; let roots = state.db.user_roots(user.id).await; Ok(AuthUser { user, roots }) } } /// Lets the extractor pull the concrete state type out of a generic `S`. pub trait HasState { fn state(&self) -> &AppState; } impl HasState for Arc { fn state(&self) -> &AppState { self } } impl HasState for AppState { fn state(&self) -> &AppState { self } }