diff --git a/apps/desktop/vite.config.ts b/apps/desktop/vite.config.ts index 71b3e0811..9d86f534f 100644 --- a/apps/desktop/vite.config.ts +++ b/apps/desktop/vite.config.ts @@ -2,6 +2,7 @@ import { defineConfig } from "vite"; import vue from "@vitejs/plugin-vue"; import tailwindcss from "@tailwindcss/vite"; import path from "path"; +import { publicBasePathRedirectPlugin } from "./vitePublicBasePathRedirect"; const host = process.env.TAURI_DEV_HOST; const isTauri = !!host || !!process.env.TAURI_ENV_ARCH; @@ -63,7 +64,7 @@ const backendUrl = process.env.DBX_BACKEND_URL || "http://localhost:4224"; export default defineConfig(async () => ({ root: import.meta.dirname, base: viteBase, - plugins: [vue(), tailwindcss()], + plugins: [publicBasePathRedirectPlugin(publicBasePath), vue(), tailwindcss()], resolve: { alias: { "@": path.resolve(import.meta.dirname, "./src"), diff --git a/apps/desktop/vitePublicBasePathRedirect.ts b/apps/desktop/vitePublicBasePathRedirect.ts new file mode 100644 index 000000000..fab406b84 --- /dev/null +++ b/apps/desktop/vitePublicBasePathRedirect.ts @@ -0,0 +1,35 @@ +import type { Connect, Plugin } from "vite"; + +export function publicBasePathRedirectLocation(requestUrl: string | undefined, publicBasePath: string): string | null { + if (!requestUrl || !publicBasePath || publicBasePath === "/") return null; + + const queryStart = requestUrl.indexOf("?"); + const requestPath = queryStart >= 0 ? requestUrl.slice(0, queryStart) : requestUrl; + if (requestPath !== publicBasePath) return null; + + const query = queryStart >= 0 ? requestUrl.slice(queryStart) : ""; + return `${publicBasePath}/${query}`; +} + +export function publicBasePathRedirectMiddleware(publicBasePath: string): Connect.NextHandleFunction { + return (request, response, next) => { + const location = publicBasePathRedirectLocation(request.url, publicBasePath); + if (!location) { + next(); + return; + } + + response.statusCode = 308; + response.setHeader("Location", location); + response.end(); + }; +} + +export function publicBasePathRedirectPlugin(publicBasePath: string): Plugin { + return { + name: "dbx-public-base-path-redirect", + configureServer(server) { + server.middlewares.use(publicBasePathRedirectMiddleware(publicBasePath)); + }, + }; +} diff --git a/crates/dbx-web/src/main.rs b/crates/dbx-web/src/main.rs index ffd25ce8f..e98aeacd8 100644 --- a/crates/dbx-web/src/main.rs +++ b/crates/dbx-web/src/main.rs @@ -13,7 +13,9 @@ use argon2::password_hash::rand_core::OsRng; use argon2::password_hash::SaltString; use argon2::{Argon2, PasswordHasher}; use axum::extract::DefaultBodyLimit; +use axum::http::Uri; use axum::middleware; +use axum::response::Redirect; use axum::routing::{delete, get, post}; use axum::Router; use dbx_core::connection::AppState; @@ -101,6 +103,49 @@ fn normalize_public_base_path(value: Option) -> String { } } +fn add_public_base_path_redirect(app: Router, public_base_path: &str) -> Router +where + S: Clone + Send + Sync + 'static, +{ + if public_base_path == "/" { + return app; + } + + // Derive the target from the configured base path so single- and multi-segment prefixes both work. + let redirect_target = format!("{public_base_path}/"); + app.route( + public_base_path, + get(move |uri: Uri| { + let redirect_target = redirect_target.clone(); + async move { + let location = uri.query().map(|query| format!("{redirect_target}?{query}")).unwrap_or(redirect_target); + Redirect::permanent(&location) + } + }), + ) +} + +fn mount_public_base_path(mut app: Router, public_base_path: &str, static_dir: Option<&std::path::Path>) -> Router { + if let Some(static_dir) = static_dir { + use tower_http::services::{ServeDir, ServeFile}; + let index_path = static_dir.join("index.html"); + let serve_dir = ServeDir::new(static_dir).not_found_service(ServeFile::new(index_path)); + app = app.fallback_service(serve_dir); + } + + if public_base_path == "/" { + return app; + } + + app = Router::new().nest(public_base_path, app); + app = add_public_base_path_redirect(app, public_base_path); + if let Some(static_dir) = static_dir { + use tower_http::services::ServeFile; + app = app.route_service(&format!("{public_base_path}/"), ServeFile::new(static_dir.join("index.html"))); + } + app +} + #[cfg(feature = "mq-admin")] fn add_mq_routes(router: Router>) -> Router> { router @@ -903,25 +948,8 @@ async fn main() { .layer(CompressionLayer::new().compress_when(web_compression_predicate())) .layer(tower_http::trace::TraceLayer::new_for_http()); - // Static file serving - if let Ok(static_dir) = std::env::var("DBX_STATIC_DIR") { - use tower_http::services::{ServeDir, ServeFile}; - let index_path = format!("{}/index.html", static_dir); - let serve_dir = ServeDir::new(&static_dir).not_found_service(ServeFile::new(&index_path)); - app = app.fallback_service(serve_dir); - } - - if public_base_path != "/" { - app = Router::new().nest(&public_base_path, app); - // axum 的 nest 不匹配“子路径根目录”(带尾斜杠,如 /dbx/),导致子路径部署时首页 404。 - // 在 nest 外层显式把根目录挂到 index.html,浏览器地址栏保持 /dbx/ 不变, - // 相对资源与前端路径推断都依赖这个 URL 形态。见 issue #5518。 - if let Ok(static_dir) = std::env::var("DBX_STATIC_DIR") { - use tower_http::services::ServeFile; - let index_path = format!("{static_dir}/index.html"); - app = app.route_service(&format!("{public_base_path}/"), ServeFile::new(index_path)); - } - } + let static_dir = std::env::var_os("DBX_STATIC_DIR").map(std::path::PathBuf::from); + app = mount_public_base_path(app, &public_base_path, static_dir.as_deref()); // Bind address let port: u16 = std::env::var("DBX_PORT").ok().and_then(|p| p.parse().ok()).unwrap_or(4224); @@ -952,10 +980,15 @@ async fn main() { #[cfg(test)] mod tests { - use super::{normalize_public_base_path, web_agent_dir_from_env, web_compression_predicate, XLSX_CONTENT_TYPE}; + use super::{ + mount_public_base_path, normalize_public_base_path, web_agent_dir_from_env, web_compression_predicate, + XLSX_CONTENT_TYPE, + }; use axum::body::Body; use axum::http::header::CONTENT_TYPE; use axum::http::Response; + use axum::routing::get; + use axum::Router; use tower_http::compression::predicate::Predicate; fn compression_response(content_type: &str) -> Response { @@ -1005,4 +1038,99 @@ mod tests { std::path::PathBuf::from("/custom/agents") ); } + + #[tokio::test] + async fn public_base_path_routes_preserve_redirect_query_static_files_and_api() { + let client = + reqwest::Client::builder().redirect(reqwest::redirect::Policy::none()).build().expect("build test client"); + let static_dir = std::env::temp_dir().join(format!("dbx-web-public-base-{}", uuid::Uuid::new_v4())); + std::fs::create_dir_all(&static_dir).expect("create static directory"); + std::fs::write(static_dir.join("index.html"), "subpath index").expect("write index"); + std::fs::write(static_dir.join("app.js"), "subpath asset").expect("write asset"); + + for public_base_path in ["/dbx", "/xxxx/rsu"] { + let router = mount_public_base_path( + Router::new().route("/api/ping", get(|| async { "pong" })), + public_base_path, + Some(&static_dir), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind test listener"); + let address = listener.local_addr().expect("test listener address"); + let server = tokio::spawn(async move { + axum::serve(listener, router).await.expect("serve test router"); + }); + + let expected_target = format!("{public_base_path}/"); + let response = + client.get(format!("http://{address}{public_base_path}")).send().await.expect("GET bare base path"); + + assert_eq!(response.status(), reqwest::StatusCode::PERMANENT_REDIRECT); + assert_eq!( + response.headers().get(reqwest::header::LOCATION).and_then(|value| value.to_str().ok()), + Some(expected_target.as_str()) + ); + + let response = client + .get(format!("http://{address}{public_base_path}?next=%2Fworkspace&theme=dark")) + .send() + .await + .expect("GET bare base path with query"); + let expected_target_with_query = format!("{public_base_path}/?next=%2Fworkspace&theme=dark"); + assert_eq!(response.status(), reqwest::StatusCode::PERMANENT_REDIRECT); + assert_eq!( + response.headers().get(reqwest::header::LOCATION).and_then(|value| value.to_str().ok()), + Some(expected_target_with_query.as_str()) + ); + + let response = client + .get(format!("http://{address}{public_base_path}/?next=%2Fworkspace")) + .send() + .await + .expect("GET trailing slash base path"); + assert_eq!(response.status(), reqwest::StatusCode::OK); + assert_eq!(response.text().await.expect("read index response"), "subpath index"); + + let response = client + .get(format!("http://{address}{public_base_path}/app.js")) + .send() + .await + .expect("GET static asset"); + assert_eq!(response.status(), reqwest::StatusCode::OK); + assert_eq!(response.text().await.expect("read asset response"), "subpath asset"); + + let response = + client.get(format!("http://{address}{public_base_path}/api/ping")).send().await.expect("GET API route"); + assert_eq!(response.status(), reqwest::StatusCode::OK); + assert_eq!(response.text().await.expect("read API response"), "pong"); + + server.abort(); + } + + std::fs::remove_dir_all(static_dir).expect("remove static directory"); + } + + #[tokio::test] + async fn root_public_base_path_preserves_static_files_and_api() { + let static_dir = std::env::temp_dir().join(format!("dbx-web-root-base-{}", uuid::Uuid::new_v4())); + std::fs::create_dir_all(&static_dir).expect("create static directory"); + std::fs::write(static_dir.join("index.html"), "root index").expect("write index"); + std::fs::write(static_dir.join("app.js"), "root asset").expect("write asset"); + let router = + mount_public_base_path(Router::new().route("/api/ping", get(|| async { "pong" })), "/", Some(&static_dir)); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind test listener"); + let address = listener.local_addr().expect("test listener address"); + let server = tokio::spawn(async move { + axum::serve(listener, router).await.expect("serve test router"); + }); + let client = reqwest::Client::new(); + + for (request_path, expected_body) in [("/", "root index"), ("/app.js", "root asset"), ("/api/ping", "pong")] { + let response = client.get(format!("http://{address}{request_path}")).send().await.expect("GET root route"); + assert_eq!(response.status(), reqwest::StatusCode::OK); + assert_eq!(response.text().await.expect("read root response"), expected_body); + } + + server.abort(); + std::fs::remove_dir_all(static_dir).expect("remove static directory"); + } } diff --git a/packages/app-tests/vitePublicBasePathRedirect.test.ts b/packages/app-tests/vitePublicBasePathRedirect.test.ts new file mode 100644 index 000000000..c89bf6199 --- /dev/null +++ b/packages/app-tests/vitePublicBasePathRedirect.test.ts @@ -0,0 +1,48 @@ +import { expect, test, vi } from "vitest"; +import { publicBasePathRedirectMiddleware } from "../../apps/desktop/vitePublicBasePathRedirect"; + +type RedirectMiddleware = ReturnType; + +function runMiddleware(requestUrl: string, publicBasePath = "/dbx") { + const middleware = publicBasePathRedirectMiddleware(publicBasePath); + const request = { url: requestUrl } as Parameters[0]; + const response = { + statusCode: 200, + setHeader: vi.fn(), + end: vi.fn(), + }; + const next = vi.fn(); + + middleware(request, response as unknown as Parameters[1], next); + return { next, response }; +} + +test("redirects the exact bare public base path and preserves its query", () => { + for (const [requestUrl, expectedLocation] of [ + ["/dbx", "/dbx/"], + ["/dbx?next=%2Fworkspace&theme=dark", "/dbx/?next=%2Fworkspace&theme=dark"], + ]) { + const { next, response } = runMiddleware(requestUrl); + + expect(response.statusCode).toBe(308); + expect(response.setHeader).toHaveBeenCalledWith("Location", expectedLocation); + expect(response.end).toHaveBeenCalledOnce(); + expect(next).not.toHaveBeenCalled(); + } +}); + +test("passes root, trailing slash, static assets, API routes, and root deployments through", () => { + for (const requestUrl of ["/", "/dbx/", "/dbx/favicon.png", "/dbx/api/probe?value=1"]) { + const { next, response } = runMiddleware(requestUrl); + + expect(next).toHaveBeenCalledOnce(); + expect(response.statusCode).toBe(200); + expect(response.setHeader).not.toHaveBeenCalled(); + expect(response.end).not.toHaveBeenCalled(); + } + + const { next, response } = runMiddleware("/", "/"); + expect(next).toHaveBeenCalledOnce(); + expect(response.statusCode).toBe(200); + expect(response.end).not.toHaveBeenCalled(); +});