Skip to main content

yggdrasil/
middleware.rs

1//! Axum 中间件与启动期纯函数。
2//!
3//! 从 `main.rs` 抽出的、可独立测试的服务端 HTTP 中间件(cache-control、admin
4//! 守卫)与压缩层构造逻辑。整体 server-only——WASM 构建不会编译本模块。
5//!
6//! 这些函数此前散落在 `main.rs`,既无法单独测试也使入口职责过载。迁移后
7//! `startup.rs` 的路由组装以全路径 `crate::middleware::xxx` 引用,语义不变。
8
9#![cfg(feature = "server")]
10
11/// 压缩算法配置。
12#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13pub(crate) struct CompressionAlgorithms {
14    gzip: bool,
15    brotli: bool,
16    deflate: bool,
17    zstd: bool,
18}
19
20impl CompressionAlgorithms {
21    fn all_enabled() -> Self {
22        Self {
23            gzip: true,
24            brotli: true,
25            deflate: true,
26            zstd: true,
27        }
28    }
29
30    fn is_empty(&self) -> bool {
31        !self.gzip && !self.brotli && !self.deflate && !self.zstd
32    }
33}
34
35/// 解析 COMPRESSION_ALGORITHMS 环境变量值。
36/// ""、"none"、"off" 返回 None;"all" 或未识别到任何算法时启用全部。
37pub(crate) fn parse_compression_algorithms(env: &str) -> Option<CompressionAlgorithms> {
38    let env = env.trim();
39    if env.is_empty() || env.eq_ignore_ascii_case("none") || env.eq_ignore_ascii_case("off") {
40        return None;
41    }
42
43    let mut all = false;
44    let mut gzip = false;
45    let mut brotli = false;
46    let mut deflate = false;
47    let mut zstd = false;
48
49    for part in env.split(',') {
50        match part.trim().to_lowercase().as_str() {
51            "all" => all = true,
52            "gzip" => gzip = true,
53            "brotli" | "br" => brotli = true,
54            "deflate" => deflate = true,
55            "zstd" => zstd = true,
56            other => tracing::warn!(
57                "Unknown compression algorithm in COMPRESSION_ALGORITHMS: '{}'",
58                other
59            ),
60        }
61    }
62
63    if all {
64        return Some(CompressionAlgorithms::all_enabled());
65    }
66
67    let algorithms = CompressionAlgorithms {
68        gzip,
69        brotli,
70        deflate,
71        zstd,
72    };
73    if algorithms.is_empty() {
74        return None;
75    }
76
77    Some(algorithms)
78}
79
80/// 根据 COMPRESSION_ALGORITHMS 环境变量构造 CompressionLayer。
81/// 默认(未设置时)关闭压缩;显式设为 "all" 启用全部算法,设为 "gzip,brotli,..."
82/// 按需选择;设为 ""、"none" 或 "off" 等价于不启用。
83///
84/// CompressionLayer 使用 tower-http 的 `DefaultPredicate`,开箱即用即:
85/// - 跳过 `image/*` content-type(WebP/PNG/JPEG/GIF 等已是压缩格式,再压浪费 CPU,
86///   唯一例外是 `image/svg+xml`,作为 XML 文本可被压缩);
87/// - 跳过 gRPC 与 `text/event-stream`(SSE);
88/// - 跳过小于 32 字节的响应。
89///
90/// 因此无需在此处对图片响应做额外的 content-type 过滤。另:图片实际挂在
91/// `static_routes`(无中间件),根本不经此层,详见 startup.rs 路由 merge 处。
92pub(crate) fn compression_layer_from_env() -> Option<tower_http::compression::CompressionLayer> {
93    use tower_http::compression::CompressionLayer;
94
95    let env = std::env::var("COMPRESSION_ALGORITHMS").unwrap_or_else(|_| "off".to_string());
96    let algorithms = parse_compression_algorithms(&env)?;
97
98    Some(
99        CompressionLayer::new()
100            .gzip(algorithms.gzip)
101            .br(algorithms.brotli)
102            .deflate(algorithms.deflate)
103            .zstd(algorithms.zstd),
104    )
105}
106
107/// 根据请求路径和方法决定公开页面的 Cache-Control 头。
108/// 返回 None 表示不添加缓存头(保留现有行为或避免覆盖)。
109pub(crate) fn cache_control_for_path(
110    path: &str,
111    method: &axum::http::Method,
112) -> Option<axum::http::HeaderValue> {
113    use axum::http::{HeaderValue, Method};
114
115    // 只对 GET/HEAD 请求添加缓存头
116    if *method != Method::GET && *method != Method::HEAD {
117        return None;
118    }
119
120    // 可撤回的笔记与受保护图片不缓存,包含未公开时的 404 响应。
121    if path == "/notes" || path.starts_with("/notes/") {
122        return Some(HeaderValue::from_static("no-store"));
123    }
124    if path.starts_with("/admin/notes")
125        || path == "/admin/notebooks"
126        || path.starts_with("/note-media/")
127    {
128        return Some(HeaderValue::from_static("private, no-store"));
129    }
130    if path.starts_with("/api") {
131        return None;
132    }
133
134    // 管理后台和认证页面:不缓存
135    if path.starts_with("/admin") || path == "/login" || path == "/register" {
136        return None;
137    }
138
139    // 固定路径的脚本、样式和 WASM 会随部署变化,允许存储但每次使用前须校验。
140    // Dioxus 已为真正带内容哈希的资源设置 immutable;中间件不会覆盖该响应头。
141    if path.starts_with("/_dioxus/")
142        || path.starts_with("/wasm/")
143        || path.ends_with(".wasm")
144        || path.ends_with(".js")
145        || path.ends_with(".css")
146    {
147        return Some(HeaderValue::from_static("public, no-cache"));
148    }
149
150    // 公开页面:5 分钟新鲜期,过期后 1 小时内可提供过期内容并后台重新验证
151    Some(HeaderValue::from_static(
152        "public, max-age=300, stale-while-revalidate=3600",
153    ))
154}
155
156/// Axum 中间件:为公开页面和静态资源附加 Cache-Control 头。
157pub(crate) async fn add_cache_control(
158    req: axum::extract::Request,
159    next: axum::middleware::Next,
160) -> axum::response::Response {
161    use axum::http::header;
162
163    let path = req.uri().path().to_string();
164    let cache_value = cache_control_for_path(&path, req.method());
165
166    let mut response = next.run(req).await;
167
168    if let Some(value) = cache_value {
169        // 仅当响应尚未设置 Cache-Control 时才添加,避免覆盖已有策略
170        response
171            .headers_mut()
172            .entry(header::CACHE_CONTROL)
173            .or_insert(value);
174    }
175
176    response
177}
178
179/// Axum 中间件:`/admin*` 的 SSR 层认证守卫。
180///
181/// 未登录访问后台时,服务端**直接 302 跳转 `/login`**,根本不进入 Dioxus
182/// SSR 渲染器。此前后台鉴权完全在客户端 WASM 完成(SSR 渲染骨架屏 → WASM
183/// 下载/编译 → hydrate → 异步 `get_current_user()` → 客户端 `navigator.push`),
184/// 整条链串行,未登录用户首屏要"空白好久"才跳登录。
185///
186/// - 只匹配 `/admin*`,其它路径(`/login`、公开页、`/api/*`)直接放行。
187/// - 复用 `get_user_by_token`:命中内存缓存 + 校验 `session_generation`
188///   (封禁/降级后旧 session 立即失效),与客户端鉴权同一套语义。
189/// - DB 错误时 **fail-open**(放行进入 SSR):避免数据库抖动把已登录的
190///   管理员也踢到登录页;客户端 `AdminLayout` 仍有兜底校验。
191/// - `/admin` 与 `/login` 本就不进 `cache_control_for_path` 缓存,302 不会被缓存。
192pub(crate) async fn admin_guard(
193    req: axum::extract::Request,
194    next: axum::middleware::Next,
195) -> axum::response::Response {
196    use crate::models::user::UserRole;
197    use axum::body::Body;
198    use axum::http::{header, StatusCode};
199    use axum::response::Response;
200
201    let path = req.uri().path().to_string();
202    if !path.starts_with("/admin") {
203        return next.run(req).await;
204    }
205
206    // 从 Cookie 头读 session token(与 export.rs / upload.rs 同款手法)。
207    let cookie = req
208        .headers()
209        .get("cookie")
210        .and_then(|h| h.to_str().ok())
211        .unwrap_or("");
212    let token = crate::auth::session::parse_session_token(cookie);
213
214    let is_admin = match token {
215        Some(t) => match crate::api::auth::get_user_by_token(t).await {
216            Ok(Some(user)) => user.role == UserRole::Admin,
217            // Err(DB 抖动)/ Ok(None)(token 无效):fail-open,交给客户端兜底。
218            _ => true,
219        },
220        // 无 token:明确未登录,拦截。
221        None => false,
222    };
223
224    if is_admin {
225        next.run(req).await
226    } else {
227        Response::builder()
228            .status(StatusCode::FOUND)
229            .header(header::LOCATION, "/login")
230            .body(Body::empty())
231            .expect("静态 302 重定向响应(合法 status + 固定 header + 空 body)必然构造成功")
232    }
233}
234
235/// Axum 中间件:把当前 SSR 全局世代号注入请求扩展,并对 GET 请求的响应附加
236/// `X-SSR-Generation` 头。这是为未来 Dioxus 支持自定义 SSR 缓存键预留的钩子;
237/// 目前主要提供可观测性,不会实际失效 SSR 缓存。
238pub(crate) async fn ssr_generation_middleware(
239    req: axum::extract::Request,
240    next: axum::middleware::Next,
241) -> axum::response::Response {
242    let generation = crate::ssr_cache::current_global_generation();
243    let is_get = req.method() == axum::http::Method::GET;
244    let (mut parts, body) = req.into_parts();
245    parts
246        .extensions
247        .insert(crate::ssr_cache::SsrGeneration(generation));
248    let mut response = next.run(axum::http::Request::from_parts(parts, body)).await;
249    if is_get {
250        response.headers_mut().insert(
251            axum::http::header::HeaderName::from_static("x-ssr-generation"),
252            axum::http::HeaderValue::from_str(&generation.to_string())
253                .unwrap_or_else(|_| axum::http::HeaderValue::from_static("0")),
254        );
255    }
256    response
257}
258
259/// Axum 中间件:为所有响应附加 Server / X-Yggdrasil-Version / X-Yggdrasil-Git / X-Yggdrasil-Hash 头。
260/// 数据源与启动日志 `log_build_info()` 同源(`crate::build_info::BUILD_INFO`)。
261/// 受 `EXPOSE_VERSION_HEADERS` 控制——该开关在 `startup.rs` 决定是否挂载本层。
262///
263/// 暴露 hash 头的原因:tag 精确 checkout 时 `git describe` 只返回纯 tag 名(如 `v0.10.1`),
264/// 不含 commit hash;单独暴露 `git_hash`(完整 40 位)以精确定位线上二进制对应的提交。
265pub(crate) async fn version_headers_middleware(
266    req: axum::extract::Request,
267    next: axum::middleware::Next,
268) -> axum::response::Response {
269    let mut response = next.run(req).await;
270    let h = response.headers_mut();
271    h.insert(
272        axum::http::header::SERVER,
273        axum::http::HeaderValue::from_str(&format!(
274            "yggdrasil/{}",
275            crate::build_info::BUILD_INFO.version
276        ))
277        .unwrap_or_else(|_| axum::http::HeaderValue::from_static("yggdrasil")),
278    );
279    h.insert(
280        axum::http::header::HeaderName::from_static("x-yggdrasil-version"),
281        axum::http::HeaderValue::from_static(crate::build_info::BUILD_INFO.version),
282    );
283    h.insert(
284        axum::http::header::HeaderName::from_static("x-yggdrasil-git"),
285        axum::http::HeaderValue::from_str(crate::build_info::BUILD_INFO.git_describe)
286            .unwrap_or_else(|_| axum::http::HeaderValue::from_static("unknown")),
287    );
288    h.insert(
289        axum::http::header::HeaderName::from_static("x-yggdrasil-hash"),
290        axum::http::HeaderValue::from_str(crate::build_info::BUILD_INFO.git_hash)
291            .unwrap_or_else(|_| axum::http::HeaderValue::from_static("unknown")),
292    );
293    response
294}
295
296#[cfg(test)]
297mod tests {
298    use super::{cache_control_for_path, parse_compression_algorithms, CompressionAlgorithms};
299    use axum::http::Method;
300
301    fn cache_value(path: &str, method: Method) -> Option<String> {
302        cache_control_for_path(path, &method).map(|v| v.to_str().unwrap().to_string())
303    }
304
305    #[test]
306    fn public_page_is_cached() {
307        assert_eq!(
308            cache_value("/", Method::GET),
309            Some("public, max-age=300, stale-while-revalidate=3600".to_string())
310        );
311        assert_eq!(
312            cache_value("/post/hello-world", Method::GET),
313            Some("public, max-age=300, stale-while-revalidate=3600".to_string())
314        );
315        assert_eq!(
316            cache_value("/tags/rust", Method::GET),
317            Some("public, max-age=300, stale-while-revalidate=3600".to_string())
318        );
319    }
320
321    #[test]
322    fn revocable_notes_never_use_browser_or_cdn_cache() {
323        for path in ["/notes", "/notes/example", "/notes/book/1"] {
324            for method in [Method::GET, Method::HEAD] {
325                assert_eq!(cache_value(path, method).as_deref(), Some("no-store"));
326            }
327        }
328        assert_eq!(
329            cache_value("/admin/notes/edit/1", Method::GET).as_deref(),
330            Some("private, no-store")
331        );
332        assert_eq!(
333            cache_value("/note-media/example", Method::GET).as_deref(),
334            Some("private, no-store")
335        );
336    }
337
338    #[test]
339    fn unversioned_assets_require_revalidation() {
340        for path in [
341            "/style.css",
342            "/highlight.css",
343            "/tiptap/editor.css",
344            "/xterm/terminal.css",
345            "/tiptap/editor.js",
346            "/codemirror/editor.js",
347            "/yggdrasil-core/yggdrasil-core.js",
348            "/mermaid/mermaid.js",
349            "/wasm/app.wasm",
350            "/wasm/app.js",
351            "/_dioxus/assets/main.js",
352        ] {
353            for method in [Method::GET, Method::HEAD] {
354                assert_eq!(
355                    cache_value(path, method).as_deref(),
356                    Some("public, no-cache"),
357                    "{path}"
358                );
359            }
360        }
361    }
362
363    #[test]
364    fn api_and_admin_and_auth_are_not_cached() {
365        assert_eq!(cache_value("/api/posts", Method::GET), None);
366        assert_eq!(cache_value("/admin", Method::GET), None);
367        assert_eq!(cache_value("/admin/posts", Method::GET), None);
368        assert_eq!(cache_value("/login", Method::GET), None);
369        assert_eq!(cache_value("/register", Method::GET), None);
370    }
371
372    #[test]
373    fn non_get_requests_are_not_cached() {
374        assert_eq!(cache_value("/", Method::POST), None);
375        assert_eq!(cache_value("/post/hello-world", Method::POST), None);
376        assert_eq!(cache_value("/style.css", Method::POST), None);
377    }
378
379    #[test]
380    fn head_requests_are_cached_like_get() {
381        assert_eq!(
382            cache_value("/", Method::HEAD),
383            Some("public, max-age=300, stale-while-revalidate=3600".to_string())
384        );
385    }
386
387    #[test]
388    fn compression_all_enables_everything() {
389        assert_eq!(
390            parse_compression_algorithms("all"),
391            Some(CompressionAlgorithms::all_enabled())
392        );
393    }
394
395    #[test]
396    fn compression_default_env_is_off() {
397        // 模拟未设置环境变量时的默认值
398        assert_eq!(parse_compression_algorithms("off"), None);
399    }
400
401    #[test]
402    fn compression_empty_none_off_disable() {
403        assert_eq!(parse_compression_algorithms(""), None);
404        assert_eq!(parse_compression_algorithms("none"), None);
405        assert_eq!(parse_compression_algorithms("NONE"), None);
406        assert_eq!(parse_compression_algorithms("off"), None);
407        assert_eq!(parse_compression_algorithms("OFF"), None);
408    }
409
410    #[test]
411    fn compression_single_algorithm() {
412        assert_eq!(
413            parse_compression_algorithms("gzip"),
414            Some(CompressionAlgorithms {
415                gzip: true,
416                brotli: false,
417                deflate: false,
418                zstd: false,
419            })
420        );
421        assert_eq!(
422            parse_compression_algorithms("br"),
423            Some(CompressionAlgorithms {
424                gzip: false,
425                brotli: true,
426                deflate: false,
427                zstd: false,
428            })
429        );
430    }
431
432    #[test]
433    fn compression_multiple_algorithms() {
434        assert_eq!(
435            parse_compression_algorithms("gzip, zstd"),
436            Some(CompressionAlgorithms {
437                gzip: true,
438                brotli: false,
439                deflate: false,
440                zstd: true,
441            })
442        );
443    }
444
445    #[test]
446    fn compression_case_insensitive_and_whitespace_tolerant() {
447        assert_eq!(
448            parse_compression_algorithms("GZIP, Brotli, Deflate, Zstd"),
449            Some(CompressionAlgorithms::all_enabled())
450        );
451        assert_eq!(
452            parse_compression_algorithms(" gzip , br , deflate , zstd "),
453            Some(CompressionAlgorithms::all_enabled())
454        );
455    }
456
457    #[test]
458    fn compression_unknown_algorithms_are_ignored() {
459        assert_eq!(
460            parse_compression_algorithms("gzip, unknown, lz4"),
461            Some(CompressionAlgorithms {
462                gzip: true,
463                brotli: false,
464                deflate: false,
465                zstd: false,
466            })
467        );
468    }
469}