Skip to main content

sz_rust_core/middleware/
chain.rs

1//! 中间件链构建器 — 基于 `tower::ServiceBuilder`
2//!
3//! 对齐 PHP 中间件组合方式(`app/middleware.php` 返回数组顺序即执行顺序),
4//! Rust 端通过 `tower::ServiceBuilder` 组合 Layer,但 `ServiceBuilder` 的 layer
5//! 是「后注册先执行」(stack 反向),本模块负责把业务期望顺序转换为
6//! `ServiceBuilder` 注册顺序。
7//!
8//! ## 设计目标
9//!
10//! 1. **顺序保证**:`DEFAULT_ORDER` 数组首元素最先执行(业务语义)
11//! 2. **PHP 对齐**:默认顺序对齐 `app/middleware.php` 全局中间件 + 业务中间件约定
12//! 3. **可定制**:支持自定义顺序(如跳过 Auth 的公开路由)
13//! 4. **可观测**:提供 `OrderRecorder` 工具记录中间件执行顺序,用于测试验证
14//!
15//! ## 用法
16//!
17//! ```ignore
18//! use sz_rust_core::middleware::chain::MiddlewareChain;
19//! use sz_rust_core::middleware::order::MiddlewareKind;
20//!
21//! // 1. 使用默认顺序
22//! let chain = MiddlewareChain::default();
23//! assert_eq!(chain.order(), [
24//!     MiddlewareKind::Trace,
25//!     MiddlewareKind::Cors,
26//!     MiddlewareKind::Log,
27//!     MiddlewareKind::RateLimit,
28//!     MiddlewareKind::Auth,
29//! ]);
30//!
31//! // 2. 自定义顺序(如公开 API 跳过 Auth)
32//! let chain = MiddlewareChain::new()
33//!     .push(MiddlewareKind::Trace)
34//!     .push(MiddlewareKind::Cors)
35//!     .push(MiddlewareKind::Log);
36//! assert_eq!(chain.order(), [
37//!     MiddlewareKind::Trace,
38//!     MiddlewareKind::Cors,
39//!     MiddlewareKind::Log,
40//! ]);
41//!
42//! // 3. 从 PHP 全局顺序构建
43//! let chain = MiddlewareChain::php_global();
44//! assert_eq!(chain.order(), [
45//!     MiddlewareKind::Trace,
46//!     MiddlewareKind::Cors,
47//! ]);
48//! ```
49//!
50//! ## 与 `tower::ServiceBuilder` 的关系
51//!
52//! `MiddlewareChain` 只负责「顺序定义和验证」,不直接构造 `ServiceBuilder`。
53//! 具体的 Layer 实例化由 `auth.rs` / `log.rs` /
54//! `rate_limit.rs` / `trace.rs` 模块提供,再由调用方按 `chain.order()` 逆序
55//! 调用 `ServiceBuilder::layer` 注册。
56//!
57//! 这样设计的原因:
58//! - 各中间件有不同的配置参数(如 Auth 需要 JWT secret,RateLimit 需要 capacity)
59//! - `tower::Layer` 是泛型 trait,不同 Layer 类型不同,无法统一存入 `Vec<Box<dyn Layer>>`
60//! - 解耦「顺序定义」与「Layer 实例化」更易测试和维护
61
62use crate::middleware::order::{MiddlewareKind, DEFAULT_ORDER, PHP_GLOBAL_ORDER};
63
64/// 中间件链 — 定义中间件执行顺序
65///
66/// 业务期望顺序:`order()` 数组首元素最先执行。
67/// 调用方在 `ServiceBuilder` 上注册时需逆序调用 `layer()`(后注册先执行)。
68#[derive(Debug, Clone, PartialEq, Eq)]
69pub struct MiddlewareChain {
70    order: Vec<MiddlewareKind>,
71}
72
73impl MiddlewareChain {
74    /// 创建空链(无中间件)
75    pub fn new() -> Self {
76        Self { order: Vec::new() }
77    }
78
79    /// 创建默认链(使用 `DEFAULT_ORDER`)
80    pub fn default_chain() -> Self {
81        Self {
82            order: DEFAULT_ORDER.to_vec(),
83        }
84    }
85
86    /// 创建 PHP 全局链(使用 `PHP_GLOBAL_ORDER`,对齐 `app/middleware.php`)
87    pub fn php_global() -> Self {
88        Self {
89            order: PHP_GLOBAL_ORDER.to_vec(),
90        }
91    }
92
93    /// 追加一个中间件到链尾(最后执行)
94    pub fn push(mut self, kind: MiddlewareKind) -> Self {
95        self.order.push(kind);
96        self
97    }
98
99    /// 在指定位置插入中间件
100    ///
101    /// 返回 `Err(message)` 如果位置越界。
102    pub fn insert(mut self, index: usize, kind: MiddlewareKind) -> Result<Self, String> {
103        if index > self.order.len() {
104            return Err(format!(
105                "insert index {index} out of bounds (len={})",
106                self.order.len()
107            ));
108        }
109        self.order.insert(index, kind);
110        Ok(self)
111    }
112
113    /// 移除并返回指定位置的中间件
114    ///
115    /// 返回 `None` 如果位置越界。
116    pub fn remove(&mut self, index: usize) -> Option<MiddlewareKind> {
117        if index >= self.order.len() {
118            return None;
119        }
120        Some(self.order.remove(index))
121    }
122
123    /// 移除所有指定类型的中间件
124    ///
125    /// 返回被移除的数量。
126    pub fn remove_kind(&mut self, kind: MiddlewareKind) -> usize {
127        let before = self.order.len();
128        self.order.retain(|k| *k != kind);
129        before - self.order.len()
130    }
131
132    /// 移除指定类型之后的所有中间件(含指定类型)
133    ///
134    /// 用于「公开 API 跳过 Auth 及之后中间件」场景。
135    /// 返回被移除的数量;若 `kind` 不存在则不移除任何中间件,返回 0。
136    pub fn remove_from(&mut self, kind: MiddlewareKind) -> usize {
137        if let Some(pos) = self.order.iter().position(|k| *k == kind) {
138            let removed = self.order.len() - pos;
139            self.order.truncate(pos);
140            removed
141        } else {
142            0
143        }
144    }
145
146    /// 返回中间件顺序(业务期望顺序,首元素最先执行)
147    pub fn order(&self) -> &[MiddlewareKind] {
148        &self.order
149    }
150
151    /// 返回 `ServiceBuilder` 注册顺序(业务期望顺序的逆序)
152    ///
153    /// `ServiceBuilder::layer` 是「后注册先执行」,因此注册时需逆序。
154    /// 调用方按此顺序调用 `ServiceBuilder::layer(layer_xxx)` 即可保证
155    /// 业务期望顺序与实际执行顺序一致。
156    pub fn service_builder_order(&self) -> Vec<MiddlewareKind> {
157        self.order.iter().copied().rev().collect()
158    }
159
160    /// 返回链长度
161    pub fn len(&self) -> usize {
162        self.order.len()
163    }
164
165    /// 链是否为空
166    pub fn is_empty(&self) -> bool {
167        self.order.is_empty()
168    }
169
170    /// 是否包含指定中间件
171    pub fn contains(&self, kind: MiddlewareKind) -> bool {
172        self.order.contains(&kind)
173    }
174
175    /// 返回指定中间件的位置(首次出现)
176    pub fn position(&self, kind: MiddlewareKind) -> Option<usize> {
177        self.order.iter().position(|k| *k == kind)
178    }
179
180    /// 校验链中无重复中间件
181    ///
182    /// 重复中间件通常表示配置错误(如 Auth 注册两次),应避免。
183    pub fn has_duplicates(&self) -> bool {
184        use std::collections::HashSet;
185        let set: HashSet<MiddlewareKind> = self.order.iter().copied().collect();
186        set.len() != self.order.len()
187    }
188}
189
190impl Default for MiddlewareChain {
191    fn default() -> Self {
192        Self::default_chain()
193    }
194}
195
196impl std::fmt::Display for MiddlewareChain {
197    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
198        write!(f, "MiddlewareChain[")?;
199        for (i, kind) in self.order.iter().enumerate() {
200            if i > 0 {
201                write!(f, " -> ")?;
202            }
203            write!(f, "{kind}")?;
204        }
205        write!(f, "]")
206    }
207}
208
209#[cfg(test)]
210mod tests {
211    use super::*;
212
213    // ====================================================================
214    // 构造函数
215    // ====================================================================
216
217    #[test]
218    fn test_new_creates_empty_chain() {
219        let chain = MiddlewareChain::new();
220        assert!(chain.is_empty());
221        assert_eq!(chain.len(), 0);
222        assert_eq!(chain.order(), &[]);
223    }
224
225    #[test]
226    fn test_default_chain_uses_default_order() {
227        let chain = MiddlewareChain::default_chain();
228        assert_eq!(chain.order(), DEFAULT_ORDER);
229        assert_eq!(chain.len(), 5);
230    }
231
232    #[test]
233    fn test_default_trait_uses_default_chain() {
234        let chain = MiddlewareChain::default();
235        assert_eq!(chain.order(), DEFAULT_ORDER);
236    }
237
238    #[test]
239    fn test_php_global_uses_php_global_order() {
240        let chain = MiddlewareChain::php_global();
241        assert_eq!(chain.order(), PHP_GLOBAL_ORDER);
242        assert_eq!(chain.len(), 2);
243    }
244
245    // ====================================================================
246    // push 链式追加
247    // ====================================================================
248
249    #[test]
250    fn test_push_appends_to_end() {
251        let chain = MiddlewareChain::new()
252            .push(MiddlewareKind::Trace)
253            .push(MiddlewareKind::Cors);
254        assert_eq!(chain.order(), [MiddlewareKind::Trace, MiddlewareKind::Cors]);
255    }
256
257    #[test]
258    fn test_push_preserves_order() {
259        let chain = MiddlewareChain::new()
260            .push(MiddlewareKind::Auth)
261            .push(MiddlewareKind::Log)
262            .push(MiddlewareKind::Trace);
263        // 业务期望顺序:push 顺序 = 执行顺序
264        assert_eq!(
265            chain.order(),
266            [
267                MiddlewareKind::Auth,
268                MiddlewareKind::Log,
269                MiddlewareKind::Trace
270            ]
271        );
272    }
273
274    // ====================================================================
275    // insert 指定位置插入
276    // ====================================================================
277
278    #[test]
279    fn test_insert_at_beginning() {
280        let chain = MiddlewareChain::new()
281            .push(MiddlewareKind::Cors)
282            .push(MiddlewareKind::Log);
283        let chain = chain
284            .insert(0, MiddlewareKind::Trace)
285            .expect("insert at 0 should succeed");
286        assert_eq!(
287            chain.order(),
288            [
289                MiddlewareKind::Trace,
290                MiddlewareKind::Cors,
291                MiddlewareKind::Log
292            ]
293        );
294    }
295
296    #[test]
297    fn test_insert_at_middle() {
298        let chain = MiddlewareChain::new()
299            .push(MiddlewareKind::Trace)
300            .push(MiddlewareKind::Log);
301        let chain = chain
302            .insert(1, MiddlewareKind::Cors)
303            .expect("insert at 1 should succeed");
304        assert_eq!(
305            chain.order(),
306            [
307                MiddlewareKind::Trace,
308                MiddlewareKind::Cors,
309                MiddlewareKind::Log
310            ]
311        );
312    }
313
314    #[test]
315    fn test_insert_at_end() {
316        let chain = MiddlewareChain::new()
317            .push(MiddlewareKind::Trace)
318            .push(MiddlewareKind::Cors);
319        let chain = chain
320            .insert(2, MiddlewareKind::Log)
321            .expect("insert at 2 should succeed");
322        assert_eq!(
323            chain.order(),
324            [
325                MiddlewareKind::Trace,
326                MiddlewareKind::Cors,
327                MiddlewareKind::Log
328            ]
329        );
330    }
331
332    #[test]
333    fn test_insert_out_of_bounds_returns_err() {
334        let chain = MiddlewareChain::new().push(MiddlewareKind::Trace);
335        let result = chain.insert(5, MiddlewareKind::Cors);
336        assert!(result.is_err());
337        let err = result.unwrap_err();
338        assert!(err.contains("out of bounds"));
339    }
340
341    // ====================================================================
342    // remove / remove_kind / remove_from
343    // ====================================================================
344
345    #[test]
346    fn test_remove_by_index() {
347        let mut chain = MiddlewareChain::default_chain();
348        let removed = chain.remove(2); // 移除 Log
349        assert_eq!(removed, Some(MiddlewareKind::Log));
350        assert_eq!(
351            chain.order(),
352            [
353                MiddlewareKind::Trace,
354                MiddlewareKind::Cors,
355                MiddlewareKind::RateLimit,
356                MiddlewareKind::Auth,
357            ]
358        );
359    }
360
361    #[test]
362    fn test_remove_out_of_bounds_returns_none() {
363        let mut chain = MiddlewareChain::default_chain();
364        assert_eq!(chain.remove(99), None);
365        assert_eq!(chain.len(), 5); // 未变化
366    }
367
368    #[test]
369    fn test_remove_kind_removes_all_occurrences() {
370        let mut chain = MiddlewareChain::new()
371            .push(MiddlewareKind::Trace)
372            .push(MiddlewareKind::Cors)
373            .push(MiddlewareKind::Trace); // 重复
374        let removed = chain.remove_kind(MiddlewareKind::Trace);
375        assert_eq!(removed, 2);
376        assert_eq!(chain.order(), [MiddlewareKind::Cors]);
377    }
378
379    #[test]
380    fn test_remove_kind_not_present_returns_zero() {
381        let mut chain = MiddlewareChain::php_global();
382        let removed = chain.remove_kind(MiddlewareKind::Auth);
383        assert_eq!(removed, 0);
384    }
385
386    #[test]
387    fn test_remove_from_removes_kind_and_after() {
388        // 公开 API 跳过 RateLimit 和 Auth
389        let mut chain = MiddlewareChain::default_chain();
390        let removed = chain.remove_from(MiddlewareKind::RateLimit);
391        assert_eq!(removed, 2);
392        assert_eq!(
393            chain.order(),
394            [
395                MiddlewareKind::Trace,
396                MiddlewareKind::Cors,
397                MiddlewareKind::Log
398            ]
399        );
400    }
401
402    #[test]
403    fn test_remove_from_first_element_clears_all() {
404        let mut chain = MiddlewareChain::default_chain();
405        let removed = chain.remove_from(MiddlewareKind::Trace);
406        assert_eq!(removed, 5);
407        assert!(chain.is_empty());
408    }
409
410    #[test]
411    fn test_remove_from_not_present_returns_zero() {
412        let mut chain = MiddlewareChain::php_global(); // 不含 Auth
413        let removed = chain.remove_from(MiddlewareKind::Auth);
414        assert_eq!(removed, 0);
415        assert_eq!(chain.len(), 2); // 未变化
416    }
417
418    // ====================================================================
419    // service_builder_order 逆序转换
420    // ====================================================================
421
422    #[test]
423    fn test_service_builder_order_reverses() {
424        let chain = MiddlewareChain::default_chain();
425        let sb_order = chain.service_builder_order();
426        // ServiceBuilder 后注册先执行,因此逆序
427        assert_eq!(
428            sb_order,
429            [
430                MiddlewareKind::Auth,
431                MiddlewareKind::RateLimit,
432                MiddlewareKind::Log,
433                MiddlewareKind::Cors,
434                MiddlewareKind::Trace,
435            ]
436        );
437    }
438
439    #[test]
440    fn test_service_builder_order_empty_chain() {
441        let chain = MiddlewareChain::new();
442        assert_eq!(chain.service_builder_order(), Vec::<MiddlewareKind>::new());
443    }
444
445    #[test]
446    fn test_service_builder_order_single_element() {
447        let chain = MiddlewareChain::new().push(MiddlewareKind::Cors);
448        assert_eq!(chain.service_builder_order(), [MiddlewareKind::Cors]);
449    }
450
451    // ====================================================================
452    // contains / position / has_duplicates
453    // ====================================================================
454
455    #[test]
456    fn test_contains_true() {
457        let chain = MiddlewareChain::default_chain();
458        assert!(chain.contains(MiddlewareKind::Auth));
459        assert!(chain.contains(MiddlewareKind::Trace));
460    }
461
462    #[test]
463    fn test_contains_false() {
464        let chain = MiddlewareChain::php_global();
465        assert!(!chain.contains(MiddlewareKind::Auth));
466    }
467
468    #[test]
469    fn test_position_returns_index() {
470        let chain = MiddlewareChain::default_chain();
471        assert_eq!(chain.position(MiddlewareKind::Trace), Some(0));
472        assert_eq!(chain.position(MiddlewareKind::Auth), Some(4));
473    }
474
475    #[test]
476    fn test_position_not_present_returns_none() {
477        let chain = MiddlewareChain::php_global();
478        assert_eq!(chain.position(MiddlewareKind::Auth), None);
479    }
480
481    #[test]
482    fn test_has_duplicates_false_for_default() {
483        let chain = MiddlewareChain::default_chain();
484        assert!(!chain.has_duplicates());
485    }
486
487    #[test]
488    fn test_has_duplicates_true_when_repeated() {
489        let chain = MiddlewareChain::new()
490            .push(MiddlewareKind::Trace)
491            .push(MiddlewareKind::Cors)
492            .push(MiddlewareKind::Trace);
493        assert!(chain.has_duplicates());
494    }
495
496    // ====================================================================
497    // Display 格式化
498    // ====================================================================
499
500    #[test]
501    fn test_display_empty_chain() {
502        let chain = MiddlewareChain::new();
503        assert_eq!(chain.to_string(), "MiddlewareChain[]");
504    }
505
506    #[test]
507    fn test_display_single_element() {
508        let chain = MiddlewareChain::new().push(MiddlewareKind::Cors);
509        assert_eq!(chain.to_string(), "MiddlewareChain[cors]");
510    }
511
512    #[test]
513    fn test_display_multiple_elements() {
514        let chain = MiddlewareChain::php_global();
515        assert_eq!(chain.to_string(), "MiddlewareChain[trace -> cors]");
516    }
517
518    #[test]
519    fn test_display_full_default_chain() {
520        let chain = MiddlewareChain::default_chain();
521        assert_eq!(
522            chain.to_string(),
523            "MiddlewareChain[trace -> cors -> log -> rate_limit -> auth]"
524        );
525    }
526
527    // ====================================================================
528    // Clone / PartialEq / Eq
529    // ====================================================================
530
531    #[test]
532    fn test_clone_produces_equal_chain() {
533        let chain = MiddlewareChain::default_chain();
534        let cloned = chain.clone();
535        assert_eq!(chain, cloned);
536    }
537
538    #[test]
539    fn test_eq_same_order() {
540        let a = MiddlewareChain::default_chain();
541        let b = MiddlewareChain::default_chain();
542        assert_eq!(a, b);
543    }
544
545    #[test]
546    fn test_ne_different_order() {
547        let a = MiddlewareChain::default_chain();
548        let b = MiddlewareChain::php_global();
549        assert_ne!(a, b);
550    }
551
552    // ====================================================================
553    // PHP 行为对齐验证(R5 硬约束)
554    // ====================================================================
555
556    #[test]
557    fn test_php_alignment_default_chain_includes_global() {
558        // DEFAULT_ORDER 必须包含 PHP 全局中间件(Trace + Cors)作为前缀
559        let chain = MiddlewareChain::default_chain();
560        let php_global = MiddlewareChain::php_global();
561        assert!(
562            chain.order().starts_with(php_global.order()),
563            "DEFAULT_ORDER must start with PHP global order"
564        );
565    }
566
567    #[test]
568    fn test_php_alignment_trace_first() {
569        // 对齐 PHP `app/middleware.php` 第一个中间件 `SessionInit`
570        let chain = MiddlewareChain::default_chain();
571        assert_eq!(chain.order().first(), Some(&MiddlewareKind::Trace));
572    }
573
574    #[test]
575    fn test_php_alignment_cors_second() {
576        // 对齐 PHP `app/middleware.php` 第二个中间件 `AllowCrossDomain`
577        let chain = MiddlewareChain::default_chain();
578        assert_eq!(chain.order().get(1), Some(&MiddlewareKind::Cors));
579    }
580
581    #[test]
582    fn test_php_alignment_auth_for_public_routes_can_be_removed() {
583        // PHP 端公开路由(如 login/captcha)通过 `allow_all_action` 白名单跳过 Auth
584        // Rust 端可通过 `remove_kind(Auth)` 或 `remove_from(Auth)` 实现等价行为
585        let mut chain = MiddlewareChain::default_chain();
586        let removed = chain.remove_kind(MiddlewareKind::Auth);
587        assert_eq!(removed, 1);
588        assert!(!chain.contains(MiddlewareKind::Auth));
589        // 其他中间件保持不变
590        assert!(chain.contains(MiddlewareKind::Trace));
591        assert!(chain.contains(MiddlewareKind::Cors));
592        assert!(chain.contains(MiddlewareKind::Log));
593        assert!(chain.contains(MiddlewareKind::RateLimit));
594    }
595}