use std::sync::Arc; use api_types::{Credentials, Me, OkResp, RootInfo, UserInfo}; use axum::Json; use axum::extract::State; use axum::http::{HeaderMap, StatusCode, header}; use axum::response::{IntoResponse, Response}; use serde::Deserialize; use crate::api::common::{display_name, hash_password, validate_account_name, validate_password}; use crate::auth::{self, clear_session_cookie, parse_session_cookie, session_cookie}; use crate::db::User; use crate::error::{ApiError, AppState}; /// GET /api/auth/me /// /// - No users at all → `200 {"first_boot": true}` /// - No/invalid session → `401` /// - Valid session → user info + visible roots pub async fn me( State(state): State>, headers: HeaderMap, ) -> Result, ApiError> { if state.db.user_count().await? == 0 { return Ok(Json(Me { first_boot: true, user: None, roots: Vec::new(), allow_writable_shares: false, })); } let user = require_user(&state, &headers).await?; Ok(Json(me_for(&state, &user).await?)) } /// The signed-in user behind the session cookie, or 401. async fn require_user(state: &AppState, headers: &HeaderMap) -> Result { let Some(token) = parse_session_cookie(headers) else { return Err(ApiError::new(StatusCode::UNAUTHORIZED, "not signed in")); }; state.db.session_user(&token).await?.ok_or_else(|| { ApiError::new( StatusCode::UNAUTHORIZED, "session expired, please sign in again", ) }) } /// Build the `/api/auth/me` payload for an authenticated user. async fn me_for(state: &AppState, user: &User) -> Result { let roots: Vec = state .db .user_roots(user.id) .await? .into_iter() .map(|r| RootInfo { id: r.id, name: display_name(&state.root, &r.path), path: r.path, mode: r.mode, }) .collect(); Ok(Me { first_boot: false, user: Some(UserInfo { id: user.id, name: user.name.clone(), is_admin: user.is_admin, single_click_open: user.single_click, }), roots, allow_writable_shares: state.db.allow_writable_shares().await?, }) } /// PUT /api/auth/me /// /// Update the signed-in user's profile settings. Returns the fresh `/me` /// payload so clients can apply the change immediately. #[derive(Deserialize)] pub(crate) struct ProfilePatch { pub single_click_open: bool, } pub async fn update_profile( State(state): State>, headers: HeaderMap, Json(body): Json, ) -> Result, ApiError> { let mut user = require_user(&state, &headers).await?; state .db .set_user_single_click(user.id, body.single_click_open) .await?; user.single_click = body.single_click_open; Ok(Json(me_for(&state, &user).await?)) } /// POST /api/auth/setup — create the first admin account. /// Only available while no users exist. pub async fn setup( State(state): State>, Json(body): Json, ) -> Result { let name = body.name.trim(); validate_account_name(name)?; validate_password(&body.password)?; if state.db.user_count().await? > 0 { return Err(ApiError::new( StatusCode::CONFLICT, "server is already set up", )); } let pass_hash = hash_password(&body.password).await?; let user = state.db.create_admin(name, &pass_hash).await?; let token = auth::random_token(); state.db.create_session(user.id, &token).await?; let mut res = Json(OkResp { ok: true }).into_response(); res.headers_mut().insert( header::SET_COOKIE, session_cookie(&token, state.https).parse().unwrap(), ); Ok(res) } /// POST /api/auth/login pub async fn login( State(state): State>, Json(body): Json, ) -> Result { let Some(user) = state.db.verify_password(&body.name, &body.password).await? else { return Err(ApiError::new( StatusCode::UNAUTHORIZED, "invalid name or password", )); }; let token = auth::random_token(); state.db.create_session(user.id, &token).await?; let mut res = Json(OkResp { ok: true }).into_response(); res.headers_mut().insert( header::SET_COOKIE, session_cookie(&token, state.https).parse().unwrap(), ); Ok(res) } /// POST /api/auth/logout pub async fn logout(State(state): State>, headers: HeaderMap) -> Response { if let Some(token) = parse_session_cookie(&headers) { let _ = state.db.delete_session(&token).await; } let mut res = Json(OkResp { ok: true }).into_response(); res.headers_mut().insert( header::SET_COOKIE, clear_session_cookie(state.https).parse().unwrap(), ); res }