common.rs
⎇
Raw
1//! Authenticated-user extractor shared by all protected API routes.
2
3use std::sync::Arc;
4
5use api_types::P_SHARE;
6use axum::extract::FromRequestParts;
7use axum::http::StatusCode;
8use axum::http::request::Parts;
9
10use crate::auth::parse_session_cookie;
11use crate::db::{RootRow, ShareRow, User};
12use crate::error::{ApiError, AppState};
13
14pub struct AuthUser {
15 pub user: User,
16 pub roots: Vec<RootRow>,
17 /// Present when authenticated via a public share token. The single entry
18 /// in `roots` is the shared item (its path is the share's `target`), so all
19 /// file operations are scoped to it.
20 pub share: Option<ShareRow>,
21}
22
23impl<S> FromRequestParts<S> for AuthUser
24where
25 S: HasState + Send + Sync,
26{
27 type Rejection = ApiError;
28
29 async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Self::Rejection> {
30 let state = state.state();
31
32 // 1. Public share token (`?share=<token>` or `X-Share-Token`). Checked
33 // first: a request that explicitly carries a share token is a share
34 // request, even if a session is also present (so a signed-in user
35 // viewing a share link sees the shared scope).
36 if let Some(share_token) = share_token_from_request(parts) {
37 return match state.db.share_by_token(&share_token).await {
38 Some(share) if !share.is_expired() => {
39 let roots = vec![RootRow {
40 id: share.id,
41 path: share.target.clone(),
42 mode: share.mode.clone(),
43 }];
44 let user = User {
45 id: share.creator_id,
46 name: "shared".to_string(),
47 is_admin: false,
48 active: true,
49 };
50 Ok(AuthUser {
51 user,
52 roots,
53 share: Some(share),
54 })
55 }
56 Some(_) => Err(ApiError::new(StatusCode::GONE, "this share has expired")),
57 None => Err(ApiError::new(StatusCode::NOT_FOUND, "share not found")),
58 };
59 }
60
61 // 2. Signed-in session.
62 if let Some(token) = parse_session_cookie(&parts.headers) {
63 return match state.db.session_user(&token).await {
64 Some(user) => {
65 let roots = state.db.user_roots(user.id).await;
66 Ok(AuthUser {
67 user,
68 roots,
69 share: None,
70 })
71 }
72 None => Err(ApiError::new(
73 StatusCode::UNAUTHORIZED,
74 "session expired, please sign in again",
75 )),
76 };
77 }
78
79 Err(ApiError::new(StatusCode::UNAUTHORIZED, "not signed in"))
80 }
81}
82
83/// Extract the share token from a request, if present: a `?share=<token>`
84/// query param or an `X-Share-Token` header.
85fn share_token_from_request(parts: &Parts) -> Option<String> {
86 if let Some(q) = parts.uri.query() {
87 for pair in q.split('&') {
88 if let Some((k, v)) = pair.split_once('=')
89 && k == P_SHARE
90 && !v.is_empty()
91 {
92 return Some(v.to_string());
93 }
94 }
95 }
96 parts
97 .headers
98 .get("x-share-token")
99 .and_then(|v| v.to_str().ok())
100 .filter(|s| !s.is_empty())
101 .map(|s| s.to_string())
102}
103
104/// Lets the extractor pull the concrete state type out of a generic `S`.
105pub trait HasState {
106 fn state(&self) -> &AppState;
107}
108
109impl HasState for Arc<AppState> {
110 fn state(&self) -> &AppState {
111 self
112 }
113}
114
115// ---------------------------------------------------------------------------
116// Shared validation / naming helpers
117// ---------------------------------------------------------------------------
118
119/// Display name for a root path: the file/folder name, or the server root's
120/// own name when the path is the whole root (".").
121pub(crate) fn display_name(server_root: &std::path::Path, rel: &str) -> String {
122 let name = if rel == "." {
123 server_root.file_name()
124 } else {
125 std::path::Path::new(rel)
126 .file_name()
127 .filter(|_| !std::path::Path::new(rel).as_os_str().is_empty())
128 };
129 name.map(|s| s.to_string_lossy().into_owned())
130 .unwrap_or_else(|| rel.to_string())
131}
132
133pub(crate) fn validate_name(name: &str) -> Result<(), ApiError> {
134 let n = name.trim();
135 if n.is_empty() || n.len() > 64 {
136 return Err(ApiError::new(
137 StatusCode::BAD_REQUEST,
138 "name must be 1–64 characters",
139 ));
140 }
141 Ok(())
142}
143
144pub(crate) fn validate_password(pw: &str) -> Result<(), ApiError> {
145 if pw.len() < 8 {
146 return Err(ApiError::new(
147 StatusCode::BAD_REQUEST,
148 "password must be at least 8 characters",
149 ));
150 }
151 Ok(())
152}
153
154/// Extractor for admin-only routes: a signed-in user who is an admin.
155pub struct AdminUser {
156 pub user: User,
157}
158
159impl<S> FromRequestParts<S> for AdminUser
160where
161 S: HasState + Send + Sync,
162{
163 type Rejection = ApiError;
164
165 async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Self::Rejection> {
166 let auth = AuthUser::from_request_parts(parts, state).await?;
167 if !auth.user.is_admin {
168 return Err(ApiError::new(StatusCode::FORBIDDEN, "admin only"));
169 }
170 Ok(AdminUser { user: auth.user })
171 }
172}
173