Skip to main content

sz_rust_core/middleware/
trace.rs

1//! Trace 中间件 — 请求追踪 span(复用 sz-orm-tracing)
2//!
3//! sz-rust 自研中间件,对齐 PHP `think\middleware\SessionInit` 的「请求初始化」语义。
4//! PHP `SessionInit` 仅初始化会话,无追踪能力;sz-rust 的 Trace 中间件是自研增强,
5//! 提供 W3C TraceContext 传播 + Span 生命周期管理。
6//!
7//! 本模块在 [`crate::middleware::order::DEFAULT_ORDER`] 中位于第 1 位
8//! (**`Trace`** → `Cors` → `Log` → `RateLimit` → `Auth`),最先执行,
9//! 确保所有后续中间件和 handler 都能通过 `request.extensions()` 获取 Span。
10//!
11//! ## 行为
12//!
13//! 1. **排除路径检查**:如果请求路径在 `exclude_paths` 中,直接放行(不创建 Span)
14//! 2. **提取 traceparent**:从请求 headers 提取 W3C traceparent(如果存在)
15//!    - 存在 → 创建子 Span(继承 trace_id,parent_id = 提取的 span_id)
16//!    - 不存在 → 创建新 Span(新 trace_id + 新 span_id)
17//! 3. **注入 Span**:将 Span 注入到 request extensions
18//! 4. **调用 next**:传递请求给下游
19//! 5. **finish Span**:标记 Span 结束(写入 end_time)
20//! 6. **注入 traceparent**:将当前 Span 的 traceparent 注入到响应 headers
21//!
22//! ## W3C TraceContext
23//!
24//! traceparent 格式:`00-<trace_id>-<span_id>-<trace_flags>`
25//! - `trace_id`:32 字符 hex(16 字节)
26//! - `span_id`:16 字符 hex(8 字节)
27//! - `trace_flags`:2 字符 hex(1 字节,如 `01` 表示 sampled)
28//!
29//! sz-orm-tracing 的 `Tracer::inject` 生成 traceparent,`Tracer::extract` 解析 traceparent。
30//!
31//! ## PHP 对齐
32//!
33//! PHP `SessionInit` 仅初始化会话(`session_start()`),无追踪能力。
34//! sz-rust 的 Trace 中间件是自研增强,提供:
35//! - W3C TraceContext 标准传播(对齐 OpenTelemetry)
36//! - Span 生命周期管理(start_time/end_time/duration)
37//! - 跨服务追踪(通过 traceparent header 传递)
38//!
39//! ## 用法
40//!
41//! ```ignore
42//! use sz_rust_core::middleware::trace::{trace_middleware, TraceConfig};
43//! use sz_orm_tracing::SzTracer;
44//! use std::sync::Arc;
45//! use axum::Router;
46//!
47//! let tracer = Arc::new(SzTracer::new("my-service"));
48//! let config = TraceConfig::new(tracer);
49//! let app: Router = Router::new()
50//!     .route("/", axum::routing::get(|| async { "ok" }))
51//!     .layer(axum::middleware::from_fn_with_state(config, trace_middleware));
52//! ```
53
54use axum::extract::Request;
55use axum::http::{HeaderMap, HeaderValue};
56use axum::middleware::Next;
57use axum::response::Response;
58use std::collections::HashMap;
59use std::sync::Arc;
60
61use sz_orm_tracing::{Span, Tracer};
62
63/// Trace 中间件配置
64#[derive(Clone)]
65pub struct TraceConfig {
66    /// Tracer 实例(`Arc<dyn Tracer + Send + Sync>` 共享)
67    pub tracer: Arc<dyn Tracer + Send + Sync>,
68    /// 服务名(对齐 sz-orm-tracing 的 service_name)
69    pub service_name: String,
70    /// 排除路径(不创建 Span,复用 [`crate::middleware::auth::is_route_allowed`] 匹配)
71    pub exclude_paths: Vec<String>,
72}
73
74impl std::fmt::Debug for TraceConfig {
75    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
76        f.debug_struct("TraceConfig")
77            .field("service_name", &self.service_name)
78            .field("exclude_paths", &self.exclude_paths)
79            .finish_non_exhaustive()
80    }
81}
82
83impl TraceConfig {
84    /// 创建 TraceConfig
85    pub fn new(tracer: Arc<dyn Tracer + Send + Sync>) -> Self {
86        let service_name = "sz-rust".to_string();
87        Self {
88            tracer,
89            service_name,
90            exclude_paths: Vec::new(),
91        }
92    }
93
94    /// 设置服务名
95    pub fn with_service_name(mut self, name: impl Into<String>) -> Self {
96        self.service_name = name.into();
97        self
98    }
99
100    /// 设置排除路径
101    pub fn with_exclude_paths(mut self, paths: Vec<String>) -> Self {
102        self.exclude_paths = paths;
103        self
104    }
105
106    /// 判断路径是否被排除
107    pub fn is_excluded(&self, path: &str) -> bool {
108        crate::middleware::auth::is_route_allowed(path, &self.exclude_paths)
109    }
110}
111
112/// 从请求 headers 构建 HashMap(sz-orm-tracing 的 extract 需要 HashMap)
113fn headers_to_hashmap(headers: &HeaderMap) -> HashMap<String, String> {
114    let mut map = HashMap::new();
115    for (name, value) in headers.iter() {
116        if let Ok(v) = value.to_str() {
117            map.insert(name.as_str().to_lowercase(), v.to_string());
118        }
119    }
120    map
121}
122
123/// 从请求 headers 提取或创建 Span
124///
125/// 优先使用 W3C traceparent 提取(如果存在),否则创建新 Span。
126///
127/// ## 实现细节
128///
129/// `Tracer::start_span` 总是生成新的 `trace_id` + `span_id`。若 `extract`
130/// 返回 parent span,则通过 `Span` 的公共字段覆盖 `trace_id` 并设置
131/// `parent_id`,实现「子 span 继承 parent 的 trace_id」语义。
132/// 不使用 `Span::with_parent` 是因为它仅设置 `parent_id` 而不覆盖 `trace_id`。
133pub fn extract_or_create_span(headers: &HeaderMap, config: &TraceConfig) -> Span {
134    let headers_map = headers_to_hashmap(headers);
135
136    // start_span 生成新的 trace_id + span_id,service_name 来自 tracer
137    let mut span = config
138        .tracer
139        .start_span(&format!("{}:request", config.service_name));
140    // 覆盖 service_name 为 config 中配置的值(可能与 tracer 内部 service_name 不同)
141    span.service_name = config.service_name.clone();
142
143    // 如果提取到 parent span,覆盖 trace_id 并设置 parent_id(创建子 span)
144    if let Some(parent_span) = config.tracer.extract(&headers_map) {
145        span.trace_id = parent_span.trace_id.clone();
146        span.parent_id = Some(parent_span.span_id.clone());
147    }
148    span
149}
150
151/// 将 Span 的 traceparent 注入到响应 headers
152///
153/// 生成 W3C traceparent:`00-<trace_id>-<span_id>-01`
154pub fn inject_traceparent_to_response(response: &mut Response, span: &Span, config: &TraceConfig) {
155    let headers_map = config.tracer.inject(span);
156    let headers = response.headers_mut();
157    for (key, value) in headers_map {
158        // 将 String key 转换为 HeaderName(owned),避免 'static 生命周期约束
159        if let (Ok(name), Ok(header_value)) = (
160            axum::http::HeaderName::from_bytes(key.as_bytes()),
161            HeaderValue::from_str(&value),
162        ) {
163            headers.insert(name, header_value);
164        }
165    }
166}
167
168/// Trace 中间件主函数
169///
170/// ## 校验流程
171///
172/// 1. **排除路径检查**:如果请求路径在 `exclude_paths` 中,直接放行(不创建 Span)
173/// 2. **提取/创建 Span**:从请求 headers 提取 traceparent,或创建新 Span
174/// 3. **注入 Span**:将 Span 注入到 request extensions
175/// 4. **调用 next**:传递请求给下游
176/// 5. **end_span**:调用 `Tracer::end_span` 完成 span(写入 end_time + 存入 tracer 内部 buffer)
177/// 6. **注入 traceparent**:将 Span 的 traceparent 注入到响应 headers
178///
179/// ## end_span vs finish
180///
181/// 使用 `Tracer::end_span` 而非 `Span::finish` 是因为 `end_span` 内部会调用 `finish`
182/// 并将 span 存入 `SzTracer.spans`,便于后续通过 `tracer.get_spans()` 获取已完成 span
183/// 用于导出(如 OTLP exporter)。
184pub async fn trace_middleware(
185    axum::extract::State(config): axum::extract::State<TraceConfig>,
186    req: Request,
187    next: Next,
188) -> Response {
189    let path = req.uri().path().to_string();
190
191    // 1. 排除路径直接放行(不创建 Span)
192    if config.is_excluded(&path) {
193        return next.run(req).await;
194    }
195
196    // 2. 提取/创建 Span + 记录请求信息到 span tags(链式 builder)
197    let method = req.method().clone();
198    let uri = req.uri().path().to_string();
199    let mut span = extract_or_create_span(req.headers(), &config)
200        .with_tag("http.method", method.as_str())
201        .with_tag("http.uri", &uri)
202        .with_tag("http.path", &path);
203
204    // 3. 注入 Span 到 request extensions(下游 handler 可通过 extensions 获取)
205    let mut req = req;
206    req.extensions_mut().insert(span.clone());
207
208    // 4. 调用 next
209    let mut response = next.run(req).await;
210
211    // 5. 记录响应状态码到 span tags + end_span(finish + 存入 tracer 内部 buffer)
212    let status = response.status().as_u16();
213    span = span.with_tag("http.status_code", status.to_string());
214    // end_span 接收 owned Span,clone 一份用于后续 inject;end_span 内部会 finish
215    config.tracer.end_span(span.clone());
216
217    // 6. 注入 traceparent 到响应 headers(使用 clone 的 span,已 finish 但 traceparent 不变)
218    inject_traceparent_to_response(&mut response, &span, &config);
219
220    response
221}
222
223#[cfg(test)]
224mod tests {
225    use super::*;
226    use axum::body::Body;
227    use axum::Router;
228    use http_body_util::BodyExt;
229    use tower::ServiceExt;
230
231    // ====================================================================
232    // 辅助函数
233    // ====================================================================
234
235    async fn read_body(resp: Response) -> String {
236        let bytes = resp.into_body().collect().await.unwrap().to_bytes();
237        String::from_utf8(bytes.to_vec()).unwrap()
238    }
239
240    fn make_request(method: &str, uri: &str) -> Request {
241        Request::builder()
242            .method(method)
243            .uri(uri)
244            .body(Body::empty())
245            .unwrap()
246    }
247
248    fn make_request_with_traceparent(method: &str, uri: &str, traceparent: &str) -> Request {
249        Request::builder()
250            .method(method)
251            .uri(uri)
252            .header("traceparent", traceparent)
253            .body(Body::empty())
254            .unwrap()
255    }
256
257    /// 构建测试用 Router(使用 SzTracer)
258    fn build_app() -> Router {
259        let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
260        let config = TraceConfig::new(tracer).with_service_name("test-service");
261        Router::new()
262            .route(
263                "/api",
264                axum::routing::get(|| async { axum::http::StatusCode::OK }),
265            )
266            .layer(axum::middleware::from_fn_with_state(
267                config,
268                trace_middleware,
269            ))
270    }
271
272    // ====================================================================
273    // TraceConfig 单元测试
274    // ====================================================================
275
276    #[test]
277    fn test_trace_config_new() {
278        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
279        let config = TraceConfig::new(tracer);
280        assert_eq!(config.service_name, "sz-rust");
281        assert!(config.exclude_paths.is_empty());
282    }
283
284    #[test]
285    fn test_trace_config_with_service_name() {
286        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
287        let config = TraceConfig::new(tracer).with_service_name("my-service");
288        assert_eq!(config.service_name, "my-service");
289    }
290
291    #[test]
292    fn test_trace_config_with_exclude_paths() {
293        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
294        let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/health".to_string()]);
295        assert_eq!(config.exclude_paths, vec!["/health".to_string()]);
296    }
297
298    #[test]
299    fn test_trace_config_is_excluded_exact_match() {
300        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
301        let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/health".to_string()]);
302        assert!(config.is_excluded("/health"));
303        assert!(!config.is_excluded("/api"));
304    }
305
306    #[test]
307    fn test_trace_config_is_excluded_wildcard_match() {
308        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
309        let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/public/*".to_string()]);
310        assert!(config.is_excluded("/public/anything"));
311        assert!(!config.is_excluded("/api"));
312    }
313
314    #[test]
315    fn test_trace_config_is_excluded_empty_list() {
316        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
317        let config = TraceConfig::new(tracer);
318        assert!(!config.is_excluded("/any"));
319    }
320
321    #[test]
322    fn test_trace_config_clone() {
323        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
324        let config = TraceConfig::new(tracer).with_service_name("cloned-service");
325        let cloned = config.clone();
326        assert_eq!(config.service_name, cloned.service_name);
327    }
328
329    #[test]
330    fn test_trace_config_debug() {
331        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
332        let config = TraceConfig::new(tracer).with_service_name("debug-service");
333        let debug_str = format!("{:?}", config);
334        assert!(debug_str.contains("debug-service"));
335        assert!(debug_str.contains("TraceConfig"));
336    }
337
338    // ====================================================================
339    // headers_to_hashmap 单元测试
340    // ====================================================================
341
342    #[test]
343    fn test_headers_to_hashmap_empty() {
344        let headers = HeaderMap::new();
345        let map = headers_to_hashmap(&headers);
346        assert!(map.is_empty());
347    }
348
349    #[test]
350    fn test_headers_to_hashmap_single_header() {
351        let mut headers = HeaderMap::new();
352        headers.insert("x-custom", "value1".parse().unwrap());
353        let map = headers_to_hashmap(&headers);
354        assert_eq!(map.get("x-custom"), Some(&"value1".to_string()));
355    }
356
357    #[test]
358    fn test_headers_to_hashmap_multiple_headers() {
359        let mut headers = HeaderMap::new();
360        headers.insert("x-custom-1", "value1".parse().unwrap());
361        headers.insert("x-custom-2", "value2".parse().unwrap());
362        let map = headers_to_hashmap(&headers);
363        assert_eq!(map.len(), 2);
364        assert_eq!(map.get("x-custom-1"), Some(&"value1".to_string()));
365        assert_eq!(map.get("x-custom-2"), Some(&"value2".to_string()));
366    }
367
368    #[test]
369    fn test_headers_to_hashmap_lowercases_keys() {
370        let mut headers = HeaderMap::new();
371        headers.insert("X-Custom", "value".parse().unwrap());
372        let map = headers_to_hashmap(&headers);
373        // HeaderMap 已经将 name 存储为小写
374        assert_eq!(map.get("x-custom"), Some(&"value".to_string()));
375    }
376
377    #[test]
378    fn test_headers_to_hashmap_skips_invalid_ascii() {
379        let mut headers = HeaderMap::new();
380        // 插入一个包含非 ASCII 字符的 header value
381        let invalid_value = HeaderValue::from_bytes(b"\xff\xfe").unwrap();
382        headers.insert("x-invalid", invalid_value);
383        let map = headers_to_hashmap(&headers);
384        // to_str() 会失败,该 header 被跳过
385        assert!(!map.contains_key("x-invalid"));
386    }
387
388    // ====================================================================
389    // extract_or_create_span 单元测试
390    // ====================================================================
391
392    #[test]
393    fn test_extract_or_create_span_no_traceparent_creates_new_span() {
394        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
395        let config = TraceConfig::new(tracer).with_service_name("my-service");
396        let headers = HeaderMap::new();
397        let span = extract_or_create_span(&headers, &config);
398        // 新 span 应该有 trace_id 和 span_id
399        assert!(!span.trace_id().is_empty());
400        assert!(!span.span_id().is_empty());
401        // 新 span 没有 parent_id
402        assert!(span.parent_id().is_none());
403        // service_name 应该是 config 中设置的
404        assert_eq!(span.service_name(), "my-service");
405    }
406
407    #[test]
408    fn test_extract_or_create_span_with_traceparent_creates_child_span() {
409        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
410        let config = TraceConfig::new(tracer).with_service_name("my-service");
411
412        // 先创建一个 parent span,获取其 traceparent
413        let parent_span = config.tracer.start_span("parent");
414        let parent_trace_id = parent_span.trace_id().to_string();
415        let parent_span_id = parent_span.span_id().to_string();
416        let headers_map = config.tracer.inject(&parent_span);
417
418        // 构建包含 traceparent 的 HeaderMap
419        let mut headers = HeaderMap::new();
420        for (key, value) in &headers_map {
421            if let (Ok(name), Ok(header_value)) = (
422                axum::http::HeaderName::from_bytes(key.as_bytes()),
423                HeaderValue::from_str(value),
424            ) {
425                headers.insert(name, header_value);
426            }
427        }
428
429        let child_span = extract_or_create_span(&headers, &config);
430        // 子 span 应该继承 parent 的 trace_id
431        assert_eq!(child_span.trace_id(), parent_trace_id);
432        // 子 span 的 parent_id 应该是 parent 的 span_id
433        assert_eq!(child_span.parent_id(), Some(parent_span_id.as_str()));
434    }
435
436    #[test]
437    fn test_extract_or_create_span_with_invalid_traceparent_creates_new_span() {
438        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
439        let config = TraceConfig::new(tracer).with_service_name("my-service");
440
441        let mut headers = HeaderMap::new();
442        // 无效的 traceparent(格式错误)
443        headers.insert("traceparent", "invalid".parse().unwrap());
444
445        let span = extract_or_create_span(&headers, &config);
446        // 无效 traceparent 应回退到创建新 span
447        assert!(span.parent_id().is_none());
448    }
449
450    // ====================================================================
451    // inject_traceparent_to_response 单元测试
452    // ====================================================================
453
454    #[test]
455    fn test_inject_traceparent_to_response_adds_headers() {
456        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
457        let config = TraceConfig::new(tracer).with_service_name("my-service");
458        let span = config.tracer.start_span("test");
459
460        let mut response = Response::new(Body::from("body"));
461        inject_traceparent_to_response(&mut response, &span, &config);
462
463        // 应该注入 traceparent header
464        assert!(response.headers().contains_key("traceparent"));
465    }
466
467    #[test]
468    fn test_inject_traceparent_to_response_preserves_existing_headers() {
469        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
470        let config = TraceConfig::new(tracer).with_service_name("my-service");
471        let span = config.tracer.start_span("test");
472
473        let mut response = Response::builder()
474            .header("x-custom", "value")
475            .body(Body::from("body"))
476            .unwrap();
477        inject_traceparent_to_response(&mut response, &span, &config);
478
479        // 原有 header 应该保留
480        assert_eq!(
481            response
482                .headers()
483                .get("x-custom")
484                .unwrap()
485                .to_str()
486                .unwrap(),
487            "value"
488        );
489        // 新增的 traceparent 应该存在
490        assert!(response.headers().contains_key("traceparent"));
491    }
492
493    // ====================================================================
494    // trace_middleware 集成测试
495    // ====================================================================
496
497    #[tokio::test]
498    async fn test_trace_middleware_creates_span_for_request() {
499        let app = build_app();
500        let resp = app.oneshot(make_request("GET", "/api")).await.unwrap();
501        assert_eq!(resp.status(), axum::http::StatusCode::OK);
502        // 响应应该包含 traceparent header
503        assert!(resp.headers().contains_key("traceparent"));
504    }
505
506    #[tokio::test]
507    async fn test_trace_middleware_excluded_path_no_span() {
508        let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
509        let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/health".to_string()]);
510        let app = Router::new()
511            .route(
512                "/health",
513                axum::routing::get(|| async { axum::http::StatusCode::OK }),
514            )
515            .layer(axum::middleware::from_fn_with_state(
516                config,
517                trace_middleware,
518            ));
519
520        let resp = app.oneshot(make_request("GET", "/health")).await.unwrap();
521        assert_eq!(resp.status(), axum::http::StatusCode::OK);
522        // 排除路径不应创建 span,所以响应不应包含 traceparent
523        assert!(!resp.headers().contains_key("traceparent"));
524    }
525
526    #[tokio::test]
527    async fn test_trace_middleware_wildcard_exclude() {
528        let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
529        let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/public/*".to_string()]);
530        let app = Router::new()
531            .route(
532                "/public/asset",
533                axum::routing::get(|| async { axum::http::StatusCode::OK }),
534            )
535            .layer(axum::middleware::from_fn_with_state(
536                config,
537                trace_middleware,
538            ));
539
540        let resp = app
541            .oneshot(make_request("GET", "/public/asset"))
542            .await
543            .unwrap();
544        assert_eq!(resp.status(), axum::http::StatusCode::OK);
545        assert!(!resp.headers().contains_key("traceparent"));
546    }
547
548    #[tokio::test]
549    async fn test_trace_middleware_injects_span_into_extensions() {
550        let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
551        let config = TraceConfig::new(tracer);
552        let app = Router::new()
553            .route(
554                "/api",
555                axum::routing::get(|| async { axum::http::StatusCode::OK }),
556            )
557            .layer(axum::middleware::from_fn_with_state(
558                config,
559                trace_middleware,
560            ));
561
562        // 验证 span 被注入到 extensions(通过响应是否包含 traceparent 间接验证)
563        let resp = app.oneshot(make_request("GET", "/api")).await.unwrap();
564        assert_eq!(resp.status(), axum::http::StatusCode::OK);
565        assert!(resp.headers().contains_key("traceparent"));
566    }
567
568    #[tokio::test]
569    async fn test_trace_middleware_child_span_inherits_trace_id() {
570        let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
571        let config = TraceConfig::new(tracer);
572
573        // 先创建一个 parent span,获取其 traceparent
574        let parent_span = config.tracer.start_span("parent");
575        let parent_trace_id = parent_span.trace_id().to_string();
576        let headers_map = config.tracer.inject(&parent_span);
577        let traceparent = headers_map
578            .get("traceparent")
579            .expect("traceparent should be in injected headers");
580
581        let app = Router::new()
582            .route(
583                "/api",
584                axum::routing::get(|| async { axum::http::StatusCode::OK }),
585            )
586            .layer(axum::middleware::from_fn_with_state(
587                config.clone(),
588                trace_middleware,
589            ));
590
591        let resp = app
592            .oneshot(make_request_with_traceparent("GET", "/api", traceparent))
593            .await
594            .unwrap();
595        assert_eq!(resp.status(), axum::http::StatusCode::OK);
596
597        // 响应的 traceparent 应该包含与 parent 相同的 trace_id
598        let response_traceparent = resp.headers().get("traceparent").unwrap().to_str().unwrap();
599        // traceparent 格式:00-<trace_id>-<span_id>-01
600        let parts: Vec<&str> = response_traceparent.split('-').collect();
601        assert_eq!(parts.len(), 4);
602        // trace_id 应该继承自 parent
603        assert_eq!(parts[1], parent_trace_id);
604    }
605
606    #[tokio::test]
607    async fn test_trace_middleware_preserves_response_body() {
608        let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
609        let config = TraceConfig::new(tracer);
610        let app = Router::new()
611            .route("/body", axum::routing::get(|| async { "hello" }))
612            .layer(axum::middleware::from_fn_with_state(
613                config,
614                trace_middleware,
615            ));
616
617        let resp = app.oneshot(make_request("GET", "/body")).await.unwrap();
618        let body = read_body(resp).await;
619        assert_eq!(body, "hello");
620    }
621
622    #[tokio::test]
623    async fn test_trace_middleware_handles_post_request() {
624        let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
625        let config = TraceConfig::new(tracer);
626        let app = Router::new()
627            .route(
628                "/submit",
629                axum::routing::post(|| async { axum::http::StatusCode::CREATED }),
630            )
631            .layer(axum::middleware::from_fn_with_state(
632                config,
633                trace_middleware,
634            ));
635
636        let req = Request::builder()
637            .method("POST")
638            .uri("/submit")
639            .body(Body::empty())
640            .unwrap();
641        let resp = app.oneshot(req).await.unwrap();
642        assert_eq!(resp.status(), axum::http::StatusCode::CREATED);
643        assert!(resp.headers().contains_key("traceparent"));
644    }
645
646    #[tokio::test]
647    async fn test_trace_middleware_chains_with_other_middleware() {
648        async fn add_header_middleware(req: Request, next: Next) -> Response {
649            let mut resp = next.run(req).await;
650            resp.headers_mut()
651                .insert("X-Custom", "value".parse().unwrap());
652            resp
653        }
654
655        let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
656        let config = TraceConfig::new(tracer);
657        let app = Router::new()
658            .route("/", axum::routing::get(|| async { "ok" }))
659            .layer(axum::middleware::from_fn(add_header_middleware))
660            .layer(axum::middleware::from_fn_with_state(
661                config,
662                trace_middleware,
663            ));
664
665        let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
666        assert_eq!(resp.status(), axum::http::StatusCode::OK);
667        assert_eq!(
668            resp.headers().get("X-Custom").unwrap().to_str().unwrap(),
669            "value"
670        );
671        assert!(resp.headers().contains_key("traceparent"));
672    }
673
674    #[tokio::test]
675    async fn test_trace_middleware_different_requests_different_trace_ids() {
676        let app = build_app();
677        let resp1 = app
678            .clone()
679            .oneshot(make_request("GET", "/api"))
680            .await
681            .unwrap();
682        let resp2 = app.oneshot(make_request("GET", "/api")).await.unwrap();
683
684        let tp1 = resp1
685            .headers()
686            .get("traceparent")
687            .unwrap()
688            .to_str()
689            .unwrap();
690        let tp2 = resp2
691            .headers()
692            .get("traceparent")
693            .unwrap()
694            .to_str()
695            .unwrap();
696
697        // 两个请求应该有不同的 trace_id(除非有 parent traceparent)
698        let parts1: Vec<&str> = tp1.split('-').collect();
699        let parts2: Vec<&str> = tp2.split('-').collect();
700        assert_ne!(parts1[1], parts2[1]); // trace_id 不同
701    }
702
703    // ====================================================================
704    // PHP 行为对齐验证(R5 硬约束)
705    // ====================================================================
706
707    #[test]
708    fn test_php_session_init_no_tracing_capability() {
709        // 对齐 PHP `think\middleware\SessionInit` 的事实:
710        // PHP SessionInit 仅初始化会话(`session_start()`),无追踪能力
711        // sz-rust 的 Trace 中间件是自研增强,提供 W3C TraceContext 传播
712        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
713        let config = TraceConfig::new(tracer);
714        // 验证 sz-rust Trace 中间件的自研性质:默认 service_name 是 "sz-rust"(PHP 端无对应概念)
715        assert_eq!(config.service_name, "sz-rust");
716    }
717
718    #[test]
719    fn test_w3c_tracecontext_format_alignment() {
720        // 对齐 W3C TraceContext 标准(OpenTelemetry)
721        // traceparent 格式:00-<trace_id(32 hex)>-<span_id(16 hex)>-<trace_flags(2 hex)>
722        let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
723        let config = TraceConfig::new(tracer).with_service_name("test-service");
724        let span = config.tracer.start_span("test");
725        let headers_map = config.tracer.inject(&span);
726        let traceparent = headers_map
727            .get("traceparent")
728            .expect("traceparent should be present");
729
730        // 验证 W3C 格式
731        let parts: Vec<&str> = traceparent.split('-').collect();
732        assert_eq!(parts.len(), 4, "traceparent should have 4 parts");
733        assert_eq!(parts[0], "00", "version should be 00");
734        assert_eq!(parts[1].len(), 32, "trace_id should be 32 hex chars");
735        assert_eq!(parts[2].len(), 16, "span_id should be 16 hex chars");
736        assert_eq!(parts[3].len(), 2, "trace_flags should be 2 hex chars");
737    }
738
739    #[test]
740    fn test_trace_middleware_executes_first_in_order() {
741        // 对齐 DEFAULT_ORDER 中 Trace 位于第 1 位的约定
742        // Trace 必须最先执行,确保所有后续中间件都能通过 extensions 获取 Span
743        use crate::middleware::order::{MiddlewareKind, DEFAULT_ORDER};
744        assert_eq!(DEFAULT_ORDER.first(), Some(&MiddlewareKind::Trace));
745    }
746}