common.rs
⎇
Raw
1//! Authenticated-user extractor shared by all protected API routes.
2
3use std::sync::Arc;
4
5use axum::extract::FromRequestParts;
6use axum::http::request::Parts;
7use axum::http::StatusCode;
8
9use crate::auth::parse_session_cookie;
10use crate::db::{RootRow, User};
11use crate::error::{ApiError, AppState};
12
13pub struct AuthUser {
14 #[allow(dead_code)] // used from milestone 3 onwards
15 pub user: User,
16 pub roots: Vec<RootRow>,
17}
18
19impl<S> FromRequestParts<S> for AuthUser
20where
21 S: HasState + Send + Sync,
22{
23 type Rejection = ApiError;
24
25 async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Self::Rejection> {
26 let state = state.state();
27 let Some(token) = parse_session_cookie(&parts.headers) else {
28 return Err(ApiError::new(
29 StatusCode::UNAUTHORIZED,
30 "not signed in",
31 ));
32 };
33 let Some(user) = state.db.session_user(&token).await else {
34 return Err(ApiError::new(
35 StatusCode::UNAUTHORIZED,
36 "session expired, please sign in again",
37 ));
38 };
39 let roots = state.db.user_roots(user.id).await;
40 Ok(AuthUser { user, roots })
41 }
42}
43
44/// Lets the extractor pull the concrete state type out of a generic `S`.
45pub trait HasState {
46 fn state(&self) -> &AppState;
47}
48
49impl HasState for Arc<AppState> {
50 fn state(&self) -> &AppState {
51 self
52 }
53}
54
55impl HasState for AppState {
56 fn state(&self) -> &AppState {
57 self
58 }
59}
60