use axum::http::HeaderValue; use axum::{ extract::DefaultBodyLimit, middleware, response::Json, routing::{delete, get, patch, post}, Router, }; use std::sync::Arc; use tower_http::{ cors::{AllowOrigin, CorsLayer}, trace::{DefaultMakeSpan, DefaultOnFailure, DefaultOnResponse, TraceLayer}, }; use utoipa::OpenApi; use crate::{auth::AuthenticatedUser, openapi::ApiDoc, state::AppState}; pub mod auth; pub mod correspondents; pub mod documents; pub mod folders; pub mod health; pub mod profile; pub mod tags; pub mod webdav; pub fn create_router(state: AppState) -> Router<()> { let cors = if let Some(origins) = state.config.cors_allowed_origin.as_ref() { let headers: Vec = origins .split(',') .filter_map(|value| { let trimmed = value.trim(); (!trimmed.is_empty()).then(|| { trimmed .parse::() .expect("invalid CORS allowed origin") }) }) .collect(); let allow_origin = AllowOrigin::list(headers); CorsLayer::new() .allow_origin(allow_origin) .allow_methods(tower_http::cors::AllowMethods::mirror_request()) .allow_headers(tower_http::cors::AllowHeaders::mirror_request()) .allow_credentials(true) } else { CorsLayer::new() .allow_origin(AllowOrigin::mirror_request()) .allow_methods(tower_http::cors::AllowMethods::mirror_request()) .allow_headers(tower_http::cors::AllowHeaders::mirror_request()) .allow_credentials(true) }; let auth_routes = Router::new() .route("/signup", post(auth::signup)) .route("/login", post(auth::login)) .route("/refresh", post(auth::refresh)) .route("/logout", post(auth::logout)) .route("/select-tenant", post(auth::select_tenant)) .route("/tenants", get(auth::list_tenants)) .route("/me", get(auth::me)); let documents_routes = Router::new() .route("/check", get(documents::check_document)) .route( "/", get(documents::list_documents).post(documents::upload_document), ) .route("/bulk/move", post(documents::bulk_move_documents)) .route("/bulk/tags", post(documents::bulk_update_tags)) .route( "/bulk/correspondents", post(documents::bulk_assign_correspondents), ) .route( "/bulk/reanalyze", post(documents::reanalyze_selected_documents), ) .route( "/:id", get(documents::get_document) .delete(documents::delete_document) .patch(documents::update_document), ) .route( "/:id/assets", get(documents::list_document_assets).post(documents::request_document_assets), ) .route("/:id/folder", patch(documents::move_document)) .route("/:id/versions", get(documents::list_document_versions)) .route( "/:id/versions/:version_id", get(documents::get_document_version), ) .route("/:id/restore", post(documents::restore_document)) .route("/:id/tags", post(documents::assign_tags)) .route("/:id/tags/:tag_id", delete(documents::remove_tag)) .route( "/:id/correspondents", post(documents::assign_correspondents), ) .route( "/:id/correspondents/:correspondent_id", delete(documents::remove_correspondent), ); let download_routes = Router::new().route("/download/:token", get(documents::download_with_token)); let folders_routes = Router::new() .route("/", post(folders::create_folder)) .route("/path", post(folders::ensure_folder_path)) .route("/:id", get(folders::get_folder)) .route("/:id", delete(folders::delete_folder)) .route("/:id", patch(folders::update_folder)) .route("/:id/contents", get(folders::list_folder_contents)); let tags_routes = Router::new() .route("/", get(tags::list_tags).post(tags::create_tag)) .route("/:id", patch(tags::update_tag).delete(tags::delete_tag)); let correspondents_routes = Router::new() .route( "/", get(correspondents::list_correspondents).post(correspondents::create_correspondent), ) .route( "/:id", patch(correspondents::update_correspondent) .delete(correspondents::delete_correspondent), ); let profile_routes = Router::new() .route( "/webdav-tokens", get(profile::list_webdav_tokens).post(profile::create_webdav_token), ) .route("/webdav-tokens/:id", delete(profile::delete_webdav_token)); let protected_state = state.clone(); let assets_routes = Router::new().route("/:asset_id", get(documents::get_document_asset)); let protected_routes = Router::new() .nest("/api/documents", documents_routes) .nest("/api/folders", folders_routes) .nest("/api/tags", tags_routes) .nest("/api/correspondents", correspondents_routes) .nest("/api/profile", profile_routes) .nest("/api/assets", assets_routes) .layer(middleware::from_extractor_with_state::(protected_state)); let openapi_arc = Arc::new(ApiDoc::openapi()); let docs_route = Router::new().route( "/api/docs/openapi.json", get({ let spec = openapi_arc.clone(); move || { let spec = spec.clone(); async move { Json((*spec).clone()) } } }), ); Router::new() .merge(download_routes) .merge(protected_routes) .merge(docs_route) .nest("/api/auth", auth_routes) .route("/api/health", get(health::health_check)) .with_state(state) .layer(cors) .layer(DefaultBodyLimit::max(1024 * 1024 * 512)) .layer( TraceLayer::new_for_http() .make_span_with(DefaultMakeSpan::new().level(tracing::Level::INFO)) .on_response(DefaultOnResponse::new().level(tracing::Level::INFO)) .on_failure(DefaultOnFailure::new().level(tracing::Level::ERROR)), ) }