Skip to main content

sz_rust_core/middleware/
builder.rs

1//! 中间件链构建器 — 统一管理 `MiddlewareChain` + 5 个 `Option<Config>`
2//!
3//! 提供 `MiddlewareBuilder` 用于:
4//! 1. 持有 `MiddlewareChain`(顺序定义,业务期望顺序,首元素最先执行)
5//! 2. 持有 5 个 `Option<Config>`:`Cors` / `Log` / `Auth` / `RateLimit` / `Trace`
6//! 3. 通过 `apply(self, router: Router) -> Router` 一次性应用所有中间件到 `axum::Router`
7//!
8//! ## 设计目标
9//!
10//! - **顺序保证**:按 `MiddlewareChain::service_builder_order()`(业务期望顺序的逆序)
11//!   调用 `Router::layer`,确保业务期望顺序与实际执行顺序一致(`Router::layer` 后注册先执行)
12//! - **配置可选**:每个中间件对应的 `Config` 都是 `Option`,链中包含但 Config 未设置时跳过
13//! - **链式 builder**:`with_xxx()` 方法支持链式配置
14//! - **PHP 对齐**:`php_global_builder()` 提供对齐 PHP `app/middleware.php` 全局中间件的默认配置
15//!
16//! ## 用法
17//!
18//! ```ignore
19//! use sz_rust_core::middleware::builder::{MiddlewareBuilder, php_global_builder};
20//! use sz_rust_core::middleware::cors::cors_layer;
21//! use axum::Router;
22//!
23//! // 1. 使用 PHP 全局默认(Trace + Cors)
24//! let builder = php_global_builder();
25//! let app: Router = Router::new()
26//!     .route("/", axum::routing::get(|| async { "ok" }))
27//!     .layer(cors_layer());
28//! // 注意:apply 会消耗 builder 并返回 Router
29//! // let app = builder.apply(app);
30//!
31//! // 2. 自定义完整链
32//! use sz_rust_core::middleware::auth::AuthConfig;
33//! use sz_rust_core::middleware::log::LogConfig;
34//! let builder = MiddlewareBuilder::default_builder()
35//!     .with_auth(AuthConfig::default())
36//!     .with_log(LogConfig::default());
37//! // let app = builder.apply(app);
38//! ```
39//!
40//! ## 与 `MiddlewareChain` 的关系
41//!
42//! `MiddlewareChain` 只负责「顺序定义」,不持有任何 Layer 实例。
43//! `MiddlewareBuilder` 在 `MiddlewareChain` 之上增加「Config 持有 + 应用到 Router」能力。
44//!
45//! ## `Router::layer` 语义
46//!
47//! `axum::Router::layer` 是「后注册先执行」(stack 反向):
48//! ```ignore
49//! let app = Router::new()
50//!     .route("/", get(handler))
51//!     .layer(A)  // A 后注册 → 先执行
52//!     .layer(B); // B 最后注册 → 最先执行
53//! // 执行顺序:B → A → handler
54//! ```
55//!
56//! 因此 `apply` 按 `chain.service_builder_order()`(业务期望顺序的逆序)遍历调用 `Router::layer`,
57//! 保证业务期望顺序(`chain.order()`)与实际执行顺序一致。
58
59use axum::Router;
60use tower_http::cors::CorsLayer;
61
62use super::auth::{auth_middleware, AuthConfig};
63use super::chain::MiddlewareChain;
64use super::cors;
65use super::log::{log_middleware_with_config, LogConfig};
66use super::order::MiddlewareKind;
67#[cfg(test)]
68use super::order::{DEFAULT_ORDER, PHP_GLOBAL_ORDER};
69use super::rate_limit::{rate_limit_middleware, RateLimitConfig};
70use super::trace::{trace_middleware, TraceConfig};
71
72/// 中间件链构建器
73///
74/// 持有 `MiddlewareChain`(顺序定义)+ 5 个 `Option<Config>`(各中间件配置),
75/// 通过 `apply()` 方法一次性应用到 `axum::Router`。
76///
77/// ## 字段说明
78///
79/// | 字段 | 类型 | 说明 |
80/// |------|------|------|
81/// | `chain` | `MiddlewareChain` | 中间件顺序定义(业务期望顺序,首元素最先执行) |
82/// | `cors` | `Option<CorsLayer>` | CORS Layer(基于 `tower-http::cors`,可直接应用) |
83/// | `log` | `Option<LogConfig>` | Log 中间件配置 |
84/// | `auth` | `Option<AuthConfig>` | Auth 中间件配置 |
85/// | `rate_limit` | `Option<RateLimitConfig>` | RateLimit 中间件配置 |
86/// | `trace` | `Option<TraceConfig>` | Trace 中间件配置 |
87#[derive(Debug, Clone)]
88pub struct MiddlewareBuilder {
89    chain: MiddlewareChain,
90    cors: Option<CorsLayer>,
91    log: Option<LogConfig>,
92    auth: Option<AuthConfig>,
93    rate_limit: Option<RateLimitConfig>,
94    trace: Option<TraceConfig>,
95}
96
97impl MiddlewareBuilder {
98    /// 创建空构建器(无中间件,无 Config)
99    pub fn new() -> Self {
100        Self {
101            chain: MiddlewareChain::new(),
102            cors: None,
103            log: None,
104            auth: None,
105            rate_limit: None,
106            trace: None,
107        }
108    }
109
110    /// 创建默认构建器(使用 `DEFAULT_ORDER`,但所有 Config 为 `None`)
111    ///
112    /// 调用方需通过 `with_xxx()` 方法显式设置 Config,否则 `apply()` 会跳过该中间件。
113    pub fn default_builder() -> Self {
114        Self {
115            chain: MiddlewareChain::default_chain(),
116            cors: None,
117            log: None,
118            auth: None,
119            rate_limit: None,
120            trace: None,
121        }
122    }
123
124    /// 创建 PHP 全局构建器(使用 `PHP_GLOBAL_ORDER`,对齐 `app/middleware.php`)
125    ///
126    /// 包含 `Trace` + `Cors` 两个中间件,对齐 PHP `app/middleware.php` 返回的全局中间件顺序。
127    /// 默认设置 `cors` 字段为 `Some(cors_layer())`,调用方可通过 `with_cors()` 覆盖。
128    pub fn php_global_builder() -> Self {
129        Self {
130            chain: MiddlewareChain::php_global(),
131            cors: Some(cors::cors_layer()),
132            log: None,
133            auth: None,
134            rate_limit: None,
135            trace: None,
136        }
137    }
138
139    /// 设置中间件链(替换现有链)
140    pub fn with_chain(mut self, chain: MiddlewareChain) -> Self {
141        self.chain = chain;
142        self
143    }
144
145    /// 设置 CORS Layer
146    pub fn with_cors(mut self, layer: CorsLayer) -> Self {
147        self.cors = Some(layer);
148        self
149    }
150
151    /// 设置 Log 配置
152    pub fn with_log(mut self, config: LogConfig) -> Self {
153        self.log = Some(config);
154        self
155    }
156
157    /// 设置 Auth 配置
158    pub fn with_auth(mut self, config: AuthConfig) -> Self {
159        self.auth = Some(config);
160        self
161    }
162
163    /// 设置 RateLimit 配置
164    pub fn with_rate_limit(mut self, config: RateLimitConfig) -> Self {
165        self.rate_limit = Some(config);
166        self
167    }
168
169    /// 设置 Trace 配置
170    pub fn with_trace(mut self, config: TraceConfig) -> Self {
171        self.trace = Some(config);
172        self
173    }
174
175    /// 从链中移除所有指定类型的中间件(同时清除对应 Config)
176    ///
177    /// 返回被移除的中间件数量(仅链中数量,不包括 Config)。
178    pub fn remove_kind(&mut self, kind: MiddlewareKind) -> usize {
179        let removed = self.chain.remove_kind(kind);
180        if removed > 0 {
181            match kind {
182                MiddlewareKind::Trace => self.trace = None,
183                MiddlewareKind::Cors => self.cors = None,
184                MiddlewareKind::Log => self.log = None,
185                MiddlewareKind::RateLimit => self.rate_limit = None,
186                MiddlewareKind::Auth => self.auth = None,
187            }
188        }
189        removed
190    }
191
192    /// 从链中移除指定类型及之后的所有中间件(含指定类型)
193    ///
194    /// 用于「公开 API 跳过 Auth 及之后中间件」场景。
195    /// 返回被移除的中间件数量;若 `kind` 不存在则不移除任何中间件,返回 0。
196    pub fn remove_from(&mut self, kind: MiddlewareKind) -> usize {
197        let removed_kinds: Vec<MiddlewareKind> = if let Some(pos) = self.chain.position(kind) {
198            self.chain.order()[pos..].to_vec()
199        } else {
200            return 0;
201        };
202        let removed = self.chain.remove_from(kind);
203        // 清除被移除中间件的 Config
204        for k in removed_kinds {
205            match k {
206                MiddlewareKind::Trace => self.trace = None,
207                MiddlewareKind::Cors => self.cors = None,
208                MiddlewareKind::Log => self.log = None,
209                MiddlewareKind::RateLimit => self.rate_limit = None,
210                MiddlewareKind::Auth => self.auth = None,
211            }
212        }
213        removed
214    }
215
216    /// 返回中间件链引用
217    pub fn chain(&self) -> &MiddlewareChain {
218        &self.chain
219    }
220
221    /// 返回 CORS Layer 引用
222    pub fn cors(&self) -> Option<&CorsLayer> {
223        self.cors.as_ref()
224    }
225
226    /// 返回 Log 配置引用
227    pub fn log(&self) -> Option<&LogConfig> {
228        self.log.as_ref()
229    }
230
231    /// 返回 Auth 配置引用
232    pub fn auth(&self) -> Option<&AuthConfig> {
233        self.auth.as_ref()
234    }
235
236    /// 返回 RateLimit 配置引用
237    pub fn rate_limit(&self) -> Option<&RateLimitConfig> {
238        self.rate_limit.as_ref()
239    }
240
241    /// 返回 Trace 配置引用
242    pub fn trace(&self) -> Option<&TraceConfig> {
243        self.trace.as_ref()
244    }
245
246    /// 判断指定中间件是否已启用(链中包含且 Config 已设置)
247    ///
248    /// 注意:`Cors` 的 Config 是 `CorsLayer`(必为 `Some` 才视为已启用)。
249    pub fn is_enabled(&self, kind: MiddlewareKind) -> bool {
250        if !self.chain.contains(kind) {
251            return false;
252        }
253        match kind {
254            MiddlewareKind::Trace => self.trace.is_some(),
255            MiddlewareKind::Cors => self.cors.is_some(),
256            MiddlewareKind::Log => self.log.is_some(),
257            MiddlewareKind::RateLimit => self.rate_limit.is_some(),
258            MiddlewareKind::Auth => self.auth.is_some(),
259        }
260    }
261
262    /// 应用所有中间件到 `axum::Router`
263    ///
264    /// 按 `chain.service_builder_order()`(业务期望顺序的逆序)遍历调用 `Router::layer`,
265    /// 保证业务期望顺序(`chain.order()`)与实际执行顺序一致(`Router::layer` 后注册先执行)。
266    ///
267    /// 链中包含但 Config 未设置的中间件会被跳过(不应用)。
268    ///
269    /// ## 消耗语义
270    ///
271    /// 此方法消耗 `self`(取出 Config 的所有权),返回应用了中间件的 `Router`。
272    pub fn apply(self, mut router: Router) -> Router {
273        let mut cors = self.cors;
274        let mut log = self.log;
275        let mut auth = self.auth;
276        let mut rate_limit = self.rate_limit;
277        let mut trace = self.trace;
278        for kind in self.chain.service_builder_order() {
279            router = match kind {
280                MiddlewareKind::Trace => {
281                    if let Some(cfg) = trace.take() {
282                        router.layer(axum::middleware::from_fn_with_state(cfg, trace_middleware))
283                    } else {
284                        router
285                    }
286                }
287                MiddlewareKind::Cors => {
288                    if let Some(layer) = cors.take() {
289                        router.layer(layer)
290                    } else {
291                        router
292                    }
293                }
294                MiddlewareKind::Log => {
295                    if let Some(cfg) = log.take() {
296                        router.layer(axum::middleware::from_fn_with_state(
297                            cfg,
298                            log_middleware_with_config,
299                        ))
300                    } else {
301                        router
302                    }
303                }
304                MiddlewareKind::RateLimit => {
305                    if let Some(cfg) = rate_limit.take() {
306                        router.layer(axum::middleware::from_fn_with_state(
307                            cfg,
308                            rate_limit_middleware,
309                        ))
310                    } else {
311                        router
312                    }
313                }
314                MiddlewareKind::Auth => {
315                    if let Some(cfg) = auth.take() {
316                        router.layer(axum::middleware::from_fn_with_state(cfg, auth_middleware))
317                    } else {
318                        router
319                    }
320                }
321            };
322        }
323        router
324    }
325}
326
327impl Default for MiddlewareBuilder {
328    fn default() -> Self {
329        Self::default_builder()
330    }
331}
332
333impl std::fmt::Display for MiddlewareBuilder {
334    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
335        write!(f, "MiddlewareBuilder(chain={}, ", self.chain)?;
336        write!(
337            f,
338            "cors={}, log={}, auth={}, rate_limit={}, trace={})",
339            self.cors.is_some(),
340            self.log.is_some(),
341            self.auth.is_some(),
342            self.rate_limit.is_some(),
343            self.trace.is_some()
344        )
345    }
346}
347
348/// 创建默认构建器(便捷函数,等价于 `MiddlewareBuilder::default_builder()`)
349pub fn default_builder() -> MiddlewareBuilder {
350    MiddlewareBuilder::default_builder()
351}
352
353/// 创建 PHP 全局构建器(便捷函数,等价于 `MiddlewareBuilder::php_global_builder()`)
354pub fn php_global_builder() -> MiddlewareBuilder {
355    MiddlewareBuilder::php_global_builder()
356}
357
358/// 创建带默认 CORS Layer 的 PHP 全局构建器(便捷函数)
359///
360/// 对齐 PHP `app/middleware.php` 全局中间件:
361/// - `SessionInit` → Rust `Trace`(Config 需调用方显式设置)
362/// - `AllowCrossDomain` → Rust `Cors`(已设置默认 `cors_layer()`)
363pub fn with_default_cors() -> MiddlewareBuilder {
364    MiddlewareBuilder::php_global_builder()
365}
366
367#[cfg(test)]
368mod tests {
369    use super::*;
370    use axum::body::Body;
371    use axum::http::Request;
372    use axum::http::StatusCode;
373    use http_body_util::BodyExt;
374    use std::sync::Arc;
375    use std::time::Duration;
376    use crate::orm::SlidingWindowRateLimiter;
377    use crate::orm::SzTracer;
378    use tower::ServiceExt;
379
380    // ====================================================================
381    // 辅助函数
382    // ====================================================================
383
384    async fn read_body(resp: axum::response::Response) -> String {
385        let bytes = resp.into_body().collect().await.unwrap().to_bytes();
386        String::from_utf8(bytes.to_vec()).unwrap()
387    }
388
389    fn make_request(method: &str, uri: &str) -> Request<Body> {
390        Request::builder()
391            .method(method)
392            .uri(uri)
393            .body(Body::empty())
394            .unwrap()
395    }
396
397    fn make_trace_config() -> TraceConfig {
398        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(SzTracer::new("test-service"));
399        TraceConfig::new(tracer)
400    }
401
402    fn make_rate_limit_config() -> RateLimitConfig {
403        let limiter: Arc<dyn RateLimiter + Send + Sync> =
404            Arc::new(SlidingWindowRateLimiter::new(1000, Duration::from_secs(60)));
405        RateLimitConfig::new(limiter)
406    }
407
408    // 引入 trait 以便 make_trace_config / make_rate_limit_config 编译
409    use crate::orm::RateLimiter;
410    use crate::orm::Tracer;
411
412    // ====================================================================
413    // 构造函数
414    // ====================================================================
415
416    #[test]
417    fn test_new_creates_empty_builder() {
418        let builder = MiddlewareBuilder::new();
419        assert!(builder.chain().is_empty());
420        assert_eq!(builder.chain().len(), 0);
421        assert!(builder.cors().is_none());
422        assert!(builder.log().is_none());
423        assert!(builder.auth().is_none());
424        assert!(builder.rate_limit().is_none());
425        assert!(builder.trace().is_none());
426    }
427
428    #[test]
429    fn test_default_builder_uses_default_order() {
430        let builder = MiddlewareBuilder::default_builder();
431        assert_eq!(builder.chain().order(), DEFAULT_ORDER);
432        assert_eq!(builder.chain().len(), 5);
433        // 默认所有 Config 为 None
434        assert!(builder.cors().is_none());
435        assert!(builder.log().is_none());
436        assert!(builder.auth().is_none());
437        assert!(builder.rate_limit().is_none());
438        assert!(builder.trace().is_none());
439    }
440
441    #[test]
442    fn test_default_trait_uses_default_builder() {
443        let builder = MiddlewareBuilder::default();
444        assert_eq!(builder.chain().order(), DEFAULT_ORDER);
445    }
446
447    #[test]
448    fn test_php_global_builder_uses_php_global_order() {
449        let builder = MiddlewareBuilder::php_global_builder();
450        assert_eq!(builder.chain().order(), PHP_GLOBAL_ORDER);
451        assert_eq!(builder.chain().len(), 2);
452        // 默认 cors 已设置
453        assert!(builder.cors().is_some());
454        // 其他 Config 为 None
455        assert!(builder.log().is_none());
456        assert!(builder.auth().is_none());
457        assert!(builder.rate_limit().is_none());
458        assert!(builder.trace().is_none());
459    }
460
461    // ====================================================================
462    // with_xxx 链式配置
463    // ====================================================================
464
465    #[test]
466    fn test_with_chain_replaces_chain() {
467        let custom_chain = MiddlewareChain::new()
468            .push(MiddlewareKind::Cors)
469            .push(MiddlewareKind::Log);
470        let builder = MiddlewareBuilder::new().with_chain(custom_chain);
471        assert_eq!(
472            builder.chain().order(),
473            [MiddlewareKind::Cors, MiddlewareKind::Log]
474        );
475    }
476
477    #[test]
478    fn test_with_cors_sets_layer() {
479        let builder = MiddlewareBuilder::new().with_cors(cors::cors_layer());
480        assert!(builder.cors().is_some());
481    }
482
483    #[test]
484    fn test_with_log_sets_config() {
485        let config = LogConfig::default().with_exclude_paths(vec!["/health".to_string()]);
486        let builder = MiddlewareBuilder::new().with_log(config);
487        assert!(builder.log().is_some());
488        assert_eq!(
489            builder.log().unwrap().exclude_paths,
490            vec!["/health".to_string()]
491        );
492    }
493
494    #[test]
495    fn test_with_auth_sets_config() {
496        let config = AuthConfig::default().with_secret("custom-secret");
497        let builder = MiddlewareBuilder::new().with_auth(config);
498        assert!(builder.auth().is_some());
499        assert_eq!(builder.auth().unwrap().secret, "custom-secret");
500    }
501
502    #[test]
503    fn test_with_rate_limit_sets_config() {
504        let config = make_rate_limit_config();
505        let builder = MiddlewareBuilder::new().with_rate_limit(config);
506        assert!(builder.rate_limit().is_some());
507    }
508
509    #[test]
510    fn test_with_trace_sets_config() {
511        let config = make_trace_config();
512        let builder = MiddlewareBuilder::new().with_trace(config);
513        assert!(builder.trace().is_some());
514    }
515
516    #[test]
517    fn test_chained_with_xxx_builders() {
518        let builder = MiddlewareBuilder::default_builder()
519            .with_cors(cors::cors_layer())
520            .with_log(LogConfig::default())
521            .with_auth(AuthConfig::default())
522            .with_rate_limit(make_rate_limit_config())
523            .with_trace(make_trace_config());
524        assert!(builder.cors().is_some());
525        assert!(builder.log().is_some());
526        assert!(builder.auth().is_some());
527        assert!(builder.rate_limit().is_some());
528        assert!(builder.trace().is_some());
529    }
530
531    // ====================================================================
532    // remove_kind / remove_from
533    // ====================================================================
534
535    #[test]
536    fn test_remove_kind_removes_from_chain_and_config() {
537        let mut builder = MiddlewareBuilder::default_builder()
538            .with_auth(AuthConfig::default())
539            .with_log(LogConfig::default());
540        assert!(builder.auth().is_some());
541        let removed = builder.remove_kind(MiddlewareKind::Auth);
542        assert_eq!(removed, 1);
543        assert!(builder.auth().is_none());
544        assert!(!builder.chain().contains(MiddlewareKind::Auth));
545    }
546
547    #[test]
548    fn test_remove_kind_not_present_returns_zero() {
549        let mut builder = MiddlewareBuilder::php_global_builder();
550        let removed = builder.remove_kind(MiddlewareKind::Auth);
551        assert_eq!(removed, 0);
552    }
553
554    #[test]
555    fn test_remove_from_removes_kind_and_after() {
556        let mut builder = MiddlewareBuilder::default_builder()
557            .with_rate_limit(make_rate_limit_config())
558            .with_auth(AuthConfig::default());
559        let removed = builder.remove_from(MiddlewareKind::RateLimit);
560        assert_eq!(removed, 2);
561        assert!(builder.rate_limit().is_none());
562        assert!(builder.auth().is_none());
563        assert!(!builder.chain().contains(MiddlewareKind::RateLimit));
564        assert!(!builder.chain().contains(MiddlewareKind::Auth));
565    }
566
567    // ====================================================================
568    // is_enabled 综合判断
569    // ====================================================================
570
571    #[test]
572    fn test_is_enabled_true_when_chain_and_config_present() {
573        let builder = MiddlewareBuilder::default_builder().with_auth(AuthConfig::default());
574        assert!(builder.is_enabled(MiddlewareKind::Auth));
575    }
576
577    #[test]
578    fn test_is_enabled_false_when_config_missing() {
579        let builder = MiddlewareBuilder::default_builder();
580        // Auth 在 DEFAULT_ORDER 中但 Config 未设置
581        assert!(!builder.is_enabled(MiddlewareKind::Auth));
582    }
583
584    #[test]
585    fn test_is_enabled_false_when_not_in_chain() {
586        let builder = MiddlewareBuilder::new().with_auth(AuthConfig::default());
587        // Config 设置但链中无 Auth
588        assert!(!builder.is_enabled(MiddlewareKind::Auth));
589    }
590
591    // ====================================================================
592    // apply 应用到 Router
593    // ====================================================================
594
595    #[test]
596    fn test_apply_empty_builder_returns_router_unchanged() {
597        let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
598        let builder = MiddlewareBuilder::new();
599        let app = builder.apply(router);
600        // 验证 Router 仍可正常使用(通过 oneshot 验证)
601        let rt = tokio::runtime::Runtime::new().unwrap();
602        rt.block_on(async {
603            let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
604            assert_eq!(resp.status(), StatusCode::OK);
605            let body = read_body(resp).await;
606            assert_eq!(body, "ok");
607        });
608    }
609
610    #[test]
611    fn test_apply_with_cors_only() {
612        let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
613        let builder = MiddlewareBuilder::new()
614            .with_chain(MiddlewareChain::new().push(MiddlewareKind::Cors))
615            .with_cors(cors::cors_layer());
616        let app = builder.apply(router);
617        let rt = tokio::runtime::Runtime::new().unwrap();
618        rt.block_on(async {
619            let req = Request::builder()
620                .method("GET")
621                .uri("/")
622                .header("origin", "https://example.com")
623                .body(Body::empty())
624                .unwrap();
625            let resp = app.oneshot(req).await.unwrap();
626            assert_eq!(resp.status(), StatusCode::OK);
627            // CORS 应设置 Access-Control-Allow-Origin
628            assert!(resp.headers().contains_key("access-control-allow-origin"));
629        });
630    }
631
632    #[test]
633    fn test_apply_skips_middlewares_without_config() {
634        let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
635        // 默认链包含 5 个中间件,但 Config 全为 None
636        let builder = MiddlewareBuilder::default_builder();
637        let app = builder.apply(router);
638        let rt = tokio::runtime::Runtime::new().unwrap();
639        rt.block_on(async {
640            let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
641            assert_eq!(resp.status(), StatusCode::OK);
642        });
643    }
644
645    #[test]
646    fn test_apply_with_all_configs_does_not_panic() {
647        let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
648        let builder = MiddlewareBuilder::default_builder()
649            .with_cors(cors::cors_layer())
650            .with_log(LogConfig::default())
651            .with_auth(AuthConfig::default())
652            .with_rate_limit(make_rate_limit_config())
653            .with_trace(make_trace_config());
654        let app = builder.apply(router);
655        let rt = tokio::runtime::Runtime::new().unwrap();
656        rt.block_on(async {
657            // Auth 未通过 → 401
658            let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
659            assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
660        });
661    }
662
663    #[test]
664    fn test_apply_preserves_business_order() {
665        // 业务期望顺序:Cors → Auth(Auth 在 Cors 之后执行)
666        // Router::layer 后注册先执行,因此 apply 应按 [Auth, Cors] 顺序注册
667        let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
668        let builder = MiddlewareBuilder::new()
669            .with_chain(
670                MiddlewareChain::new()
671                    .push(MiddlewareKind::Cors)
672                    .push(MiddlewareKind::Auth),
673            )
674            .with_cors(cors::cors_layer())
675            .with_auth(AuthConfig::default());
676        let app = builder.apply(router);
677        let rt = tokio::runtime::Runtime::new().unwrap();
678        rt.block_on(async {
679            // Auth 在最后注册 → 最先执行 → 401
680            let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
681            assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
682        });
683    }
684
685    // ====================================================================
686    // 便捷函数
687    // ====================================================================
688
689    #[test]
690    fn test_default_builder_helper() {
691        let builder = default_builder();
692        assert_eq!(builder.chain().order(), DEFAULT_ORDER);
693    }
694
695    #[test]
696    fn test_php_global_builder_helper() {
697        let builder = php_global_builder();
698        assert_eq!(builder.chain().order(), PHP_GLOBAL_ORDER);
699        assert!(builder.cors().is_some());
700    }
701
702    #[test]
703    fn test_with_default_cors_helper() {
704        let builder = with_default_cors();
705        assert_eq!(builder.chain().order(), PHP_GLOBAL_ORDER);
706        assert!(builder.cors().is_some());
707    }
708
709    // ====================================================================
710    // Display 格式化
711    // ====================================================================
712
713    #[test]
714    fn test_display_empty_builder() {
715        let builder = MiddlewareBuilder::new();
716        let s = format!("{builder}");
717        assert!(s.contains("MiddlewareBuilder"));
718        assert!(s.contains("chain=MiddlewareChain[]"));
719        assert!(s.contains("cors=false"));
720    }
721
722    #[test]
723    fn test_display_full_builder() {
724        let builder = MiddlewareBuilder::default_builder()
725            .with_cors(cors::cors_layer())
726            .with_log(LogConfig::default())
727            .with_auth(AuthConfig::default())
728            .with_rate_limit(make_rate_limit_config())
729            .with_trace(make_trace_config());
730        let s = format!("{builder}");
731        assert!(s.contains("cors=true"));
732        assert!(s.contains("log=true"));
733        assert!(s.contains("auth=true"));
734        assert!(s.contains("rate_limit=true"));
735        assert!(s.contains("trace=true"));
736    }
737
738    // ====================================================================
739    // Clone
740    // ====================================================================
741
742    #[test]
743    fn test_clone_preserves_state() {
744        let builder = MiddlewareBuilder::default_builder()
745            .with_cors(cors::cors_layer())
746            .with_log(LogConfig::default())
747            .with_auth(AuthConfig::default());
748        let cloned = builder.clone();
749        assert_eq!(builder.chain(), cloned.chain());
750        assert!(cloned.cors().is_some());
751        assert!(cloned.log().is_some());
752        assert!(cloned.auth().is_some());
753    }
754
755    // ====================================================================
756    // R5 PHP 行为对齐验证
757    // ====================================================================
758
759    #[test]
760    fn r5_1_php_global_order_matches_php_app_middleware() {
761        // PHP `app/middleware.php`:
762        //   \think\middleware\SessionInit::class,        // → Rust Trace
763        //   \think\middleware\AllowCrossDomain::class,  // → Rust Cors
764        let builder = php_global_builder();
765        assert_eq!(
766            builder.chain().order(),
767            [MiddlewareKind::Trace, MiddlewareKind::Cors]
768        );
769    }
770
771    #[test]
772    fn r5_2_php_global_builder_has_default_cors() {
773        // 对齐 PHP `AllowCrossDomain` 默认启用
774        let builder = php_global_builder();
775        assert!(builder.cors().is_some());
776    }
777
778    #[test]
779    fn r5_3_php_global_builder_trace_config_none_by_default() {
780        // PHP `SessionInit` 由框架自动配置,Rust 端需调用方显式设置 TraceConfig
781        // (因为 Tracer 实例需要服务名等业务参数)
782        let builder = php_global_builder();
783        assert!(builder.trace().is_none());
784    }
785
786    #[test]
787    fn r5_4_default_order_aligns_with_php_extension() {
788        // PHP 全局 + 业务中间件顺序:
789        //   SessionInit(Trace) → AllowCrossDomain(Cors) → [Log/RateLimit/Auth 业务追加]
790        let builder = default_builder();
791        assert_eq!(
792            builder.chain().order(),
793            [
794                MiddlewareKind::Trace,
795                MiddlewareKind::Cors,
796                MiddlewareKind::Log,
797                MiddlewareKind::RateLimit,
798                MiddlewareKind::Auth
799            ]
800        );
801        // PHP 全局顺序必须是默认顺序的前缀
802        assert!(builder.chain().order().starts_with(PHP_GLOBAL_ORDER));
803    }
804
805    #[test]
806    fn r5_5_php_public_api_skip_auth_via_remove_from() {
807        // 公开 API 跳过 RateLimit + Auth(对齐 PHP 公开路由不挂 Auth middleware)
808        let mut builder = default_builder().with_auth(AuthConfig::default());
809        let removed = builder.remove_from(MiddlewareKind::RateLimit);
810        assert_eq!(removed, 2);
811        assert!(!builder.is_enabled(MiddlewareKind::Auth));
812        assert!(!builder.is_enabled(MiddlewareKind::RateLimit));
813        // Trace/Cors/Log 仍保留
814        assert!(builder.chain().contains(MiddlewareKind::Trace));
815        assert!(builder.chain().contains(MiddlewareKind::Cors));
816        assert!(builder.chain().contains(MiddlewareKind::Log));
817    }
818
819    #[test]
820    fn r5_6_service_builder_order_reverses_for_router_layer() {
821        // 验证 apply 按 service_builder_order()(逆序)应用
822        let builder = default_builder();
823        let sb_order = builder.chain().service_builder_order();
824        // 业务期望:Trace, Cors, Log, RateLimit, Auth
825        // ServiceBuilder 注册顺序(逆序):Auth, RateLimit, Log, Cors, Trace
826        assert_eq!(
827            sb_order,
828            [
829                MiddlewareKind::Auth,
830                MiddlewareKind::RateLimit,
831                MiddlewareKind::Log,
832                MiddlewareKind::Cors,
833                MiddlewareKind::Trace,
834            ]
835        );
836    }
837
838    #[test]
839    fn r5_7_php_global_builder_skip_middlewares_without_config() {
840        // php_global_builder() 包含 Trace + Cors,但 TraceConfig 为 None
841        // apply 时应跳过 Trace,仅应用 Cors
842        let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
843        let app = php_global_builder().apply(router);
844        let rt = tokio::runtime::Runtime::new().unwrap();
845        rt.block_on(async {
846            let req = Request::builder()
847                .method("GET")
848                .uri("/")
849                .header("origin", "https://example.com")
850                .body(Body::empty())
851                .unwrap();
852            let resp = app.oneshot(req).await.unwrap();
853            // 应正常返回(Trace 跳过,Cors 设置)
854            assert_eq!(resp.status(), StatusCode::OK);
855            assert!(resp.headers().contains_key("access-control-allow-origin"));
856        });
857    }
858
859    #[test]
860    fn r5_8_is_enabled_aligns_with_php_middleware_registration() {
861        // PHP 端判断中间件是否「实际生效」:必须 (1) 在 middleware.php 中声明 + (2) 实例化成功
862        // Rust 端 `is_enabled` 对齐:(1) 链中包含 + (2) Config 已设置
863        let builder = MiddlewareBuilder::default_builder()
864            .with_cors(cors::cors_layer())
865            .with_auth(AuthConfig::default());
866        // Cors 已启用(链中包含 + Config 已设置)
867        assert!(builder.is_enabled(MiddlewareKind::Cors));
868        // Auth 已启用
869        assert!(builder.is_enabled(MiddlewareKind::Auth));
870        // Trace 未启用(Config 为 None)
871        assert!(!builder.is_enabled(MiddlewareKind::Trace));
872        // Log 未启用(Config 为 None)
873        assert!(!builder.is_enabled(MiddlewareKind::Log));
874        // RateLimit 未启用(Config 为 None)
875        assert!(!builder.is_enabled(MiddlewareKind::RateLimit));
876    }
877
878    // ====================================================================
879    // 集成测试(tokio::test)
880    // ====================================================================
881
882    #[tokio::test]
883    async fn integration_apply_returns_working_router() {
884        let router = Router::new().route("/health", axum::routing::get(|| async { "ok" }));
885        let app = MiddlewareBuilder::new()
886            .with_chain(MiddlewareChain::new().push(MiddlewareKind::Cors))
887            .with_cors(cors::cors_layer())
888            .apply(router);
889        let resp = app.oneshot(make_request("GET", "/health")).await.unwrap();
890        assert_eq!(resp.status(), StatusCode::OK);
891        let body = read_body(resp).await;
892        assert_eq!(body, "ok");
893    }
894
895    #[tokio::test]
896    async fn integration_cors_preflight_passes_through() {
897        let router = Router::new().route(
898            "/api",
899            axum::routing::get(|| async { "ok" }).post(|| async { "created" }),
900        );
901        let app = MiddlewareBuilder::new()
902            .with_chain(MiddlewareChain::new().push(MiddlewareKind::Cors))
903            .with_cors(cors::cors_layer())
904            .apply(router);
905        let req = Request::builder()
906            .method("OPTIONS")
907            .uri("/api")
908            .header("origin", "https://example.com")
909            .header("access-control-request-method", "POST")
910            .body(Body::empty())
911            .unwrap();
912        let resp = app.oneshot(req).await.unwrap();
913        // CORS 预检应返回 200 或 204(tower-http 行为)
914        assert!(resp.status().is_success());
915        assert!(resp.headers().contains_key("access-control-allow-methods"));
916    }
917
918    #[tokio::test]
919    async fn integration_auth_rejects_unauthenticated_request() {
920        let router = Router::new().route("/protected", axum::routing::get(|| async { "ok" }));
921        let app = MiddlewareBuilder::new()
922            .with_chain(MiddlewareChain::new().push(MiddlewareKind::Auth))
923            .with_auth(AuthConfig::default())
924            .apply(router);
925        let resp = app
926            .oneshot(make_request("GET", "/protected"))
927            .await
928            .unwrap();
929        assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
930        let body = read_body(resp).await;
931        assert!(body.contains("\"code\":-1"));
932    }
933
934    #[tokio::test]
935    async fn integration_log_does_not_block_request() {
936        let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
937        let app = MiddlewareBuilder::new()
938            .with_chain(MiddlewareChain::new().push(MiddlewareKind::Log))
939            .with_log(LogConfig::default())
940            .apply(router);
941        let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
942        assert_eq!(resp.status(), StatusCode::OK);
943        let body = read_body(resp).await;
944        assert_eq!(body, "ok");
945    }
946}