//! 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, ShareRow, User}; use crate::error::{ApiError, AppState}; pub struct AuthUser { pub user: User, pub roots: Vec, /// Present when authenticated via a public share token. The single entry /// in `roots` is the shared item (its path is the share's `target`), so all /// file operations are scoped to it. #[allow(dead_code)] pub share: Option, } 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(); // 1. Public share token (`?share=` or `X-Share-Token`). Checked // first: a request that explicitly carries a share token is a share // request, even if a session is also present (so a signed-in user // viewing a share link sees the shared scope). if let Some(share_token) = share_token_from_request(parts) { return match state.db.share_by_token(&share_token).await { Some(share) if !share.is_expired() => { let roots = vec![RootRow { id: share.id, path: share.target.clone(), mode: share.mode.clone(), }]; let user = User { id: share.creator_id, name: "shared".to_string(), is_admin: false, active: true, }; Ok(AuthUser { user, roots, share: Some(share), }) } Some(_) => Err(ApiError::new(StatusCode::GONE, "this share has expired")), None => Err(ApiError::new(StatusCode::NOT_FOUND, "share not found")), }; } // 2. Signed-in session. if let Some(token) = parse_session_cookie(&parts.headers) { return match state.db.session_user(&token).await { Some(user) => { let roots = state.db.user_roots(user.id).await; Ok(AuthUser { user, roots, share: None, }) } None => Err(ApiError::new( StatusCode::UNAUTHORIZED, "session expired, please sign in again", )), }; } Err(ApiError::new(StatusCode::UNAUTHORIZED, "not signed in")) } } /// Extract the share token from a request, if present: a `?share=` /// query param or an `X-Share-Token` header. fn share_token_from_request(parts: &Parts) -> Option { if let Some(q) = parts.uri.query() { for pair in q.split('&') { if let Some((k, v)) = pair.split_once('=') { if k == "share" && !v.is_empty() { return Some(v.to_string()); } } } } parts .headers .get("x-share-token") .and_then(|v| v.to_str().ok()) .filter(|s| !s.is_empty()) .map(|s| s.to_string()) } /// 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 } } /// Extractor for admin-only routes: a signed-in user who is an admin. pub struct AdminUser { pub user: User, } impl FromRequestParts for AdminUser where S: HasState + Send + Sync, { type Rejection = ApiError; async fn from_request_parts(parts: &mut Parts, state: &S) -> Result { let auth = AuthUser::from_request_parts(parts, state).await?; if !auth.user.is_admin { return Err(ApiError::new(StatusCode::FORBIDDEN, "admin only")); } Ok(AdminUser { user: auth.user }) } }