Skip to main content

sz_rust_middleware_facade/
log.rs

1//! 日志系统 — 对齐 PHP `think-logger`
2//!
3//! ## 设计
4//!
5//! - 基于 `sz-orm-logger` 的 `StructuredLogger` 提供日志收集
6//! - 同时通过 `tracing` 宏输出(与 SZ-ORM-Tracing 协同,未来接入 OpenTelemetry)
7//! - 全局单例 `LogFacade`,通过 [`LogFacade::init()`] 初始化、[`LogFacade::instance()`] 获取
8//! - 支持多通道(file/console),对齐 PHP `config/log.php` 的 `channels` 配置
9//!
10//! ## PHP 对齐
11//!
12//! ```php
13//! // PHP think-logger
14//! Log::info('hello');
15//! Log::error('error occurred', ['exception' => $e]);
16//! Log::channel('file')->info('file log');
17//! ```
18//!
19//! ```rust,ignore
20//! // SZ-Rust 等价
21//! use sz_rust_core::log::LogFacade;
22//! LogFacade::instance().unwrap().info("hello");
23//! LogFacade::instance().unwrap().error("error occurred");
24//! ```
25
26use parking_lot::RwLock;
27use std::collections::HashMap;
28use std::sync::OnceLock;
29use sz_rust_infra_facade::config::{LogChannel, LogSection};
30
31// 重导出 sz-orm-logger 核心类型,方便上层直接使用
32pub use sz_rust_orm_facade::logger::{LogEntry, LogLevel, Logger, LoggerFactory, StructuredLogger};
33
34/// 全局日志 facade 单例
35static LOG_FACADE: OnceLock<LogFacade> = OnceLock::new();
36
37/// 日志 facade — 持有默认 `StructuredLogger` 和命名通道
38///
39/// 对齐 PHP `think\facade\Log`,提供全局日志访问点。
40pub struct LogFacade {
41    /// 默认通道名(对应 PHP `config/log.php` 的 `default`)
42    default_channel: String,
43    /// 默认 logger 实例
44    logger: StructuredLogger,
45    /// 命名通道集合(对应 PHP `channels`)
46    channels: RwLock<HashMap<String, StructuredLogger>>,
47}
48
49impl LogFacade {
50    /// 构造 LogFacade 实例(不注册到全局单例)
51    pub fn new(section: &LogSection) -> Self {
52        let default_channel = section.default.clone();
53        let default_log_level = section
54            .channels
55            .get(&default_channel)
56            .map(|c| parse_level(&c.level))
57            .unwrap_or(LogLevel::Info);
58        let logger = StructuredLogger::with_level(default_log_level);
59
60        let mut channels = HashMap::new();
61        for (name, channel_cfg) in &section.channels {
62            channels.insert(name.clone(), channel_to_logger(channel_cfg));
63        }
64
65        LogFacade {
66            default_channel,
67            logger,
68            channels: RwLock::new(channels),
69        }
70    }
71
72    /// 初始化全局日志 facade
73    ///
74    /// 重复调用返回已有实例(不覆盖)。
75    pub fn init(section: &LogSection) -> &'static LogFacade {
76        LOG_FACADE.get_or_init(|| LogFacade::new(section))
77    }
78
79    /// 获取全局日志 facade 实例
80    ///
81    /// 必须先调用 [`LogFacade::init()`] 初始化,否则返回 `None`。
82    pub fn instance() -> Option<&'static LogFacade> {
83        LOG_FACADE.get()
84    }
85
86    /// 获取默认通道名
87    pub fn default_channel(&self) -> &str {
88        &self.default_channel
89    }
90
91    /// 获取默认 logger 引用
92    pub fn logger(&self) -> &StructuredLogger {
93        &self.logger
94    }
95
96    /// 获取指定通道的 logger 引用
97    ///
98    /// 对齐 PHP `Log::channel('file')->info(...)`。
99    pub fn channel(&self, name: &str) -> Option<ChannelRef<'_>> {
100        if self.channels.read().contains_key(name) {
101            Some(ChannelRef {
102                facade: self,
103                name: name.to_string(),
104            })
105        } else {
106            None
107        }
108    }
109
110    /// 获取所有通道名
111    pub fn channel_names(&self) -> Vec<String> {
112        self.channels.read().keys().cloned().collect()
113    }
114
115    /// 记录日志(同时输出到 StructuredLogger 和 tracing)
116    pub fn log(&self, level: LogLevel, msg: &str) {
117        self.logger.log(level, msg);
118        match level {
119            LogLevel::Trace => tracing::trace!("{}", msg),
120            LogLevel::Debug => tracing::debug!("{}", msg),
121            LogLevel::Info => tracing::info!("{}", msg),
122            LogLevel::Warn => tracing::warn!("{}", msg),
123            LogLevel::Error => tracing::error!("{}", msg),
124        }
125    }
126
127    /// 记录 DEBUG 级别日志
128    pub fn debug(&self, msg: &str) {
129        self.log(LogLevel::Debug, msg);
130    }
131
132    /// 记录 INFO 级别日志
133    pub fn info(&self, msg: &str) {
134        self.log(LogLevel::Info, msg);
135    }
136
137    /// 记录 WARN 级别日志
138    pub fn warn(&self, msg: &str) {
139        self.log(LogLevel::Warn, msg);
140    }
141
142    /// 记录 ERROR 级别日志
143    pub fn error(&self, msg: &str) {
144        self.log(LogLevel::Error, msg);
145    }
146}
147
148impl std::fmt::Debug for LogFacade {
149    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
150        f.debug_struct("LogFacade")
151            .field("default_channel", &self.default_channel)
152            .field("channels", &self.channels.read().keys().collect::<Vec<_>>())
153            .finish()
154    }
155}
156
157/// 命名通道引用
158///
159/// 通过 [`LogFacade::channel()`] 获取,提供与默认 logger 相同的日志方法。
160pub struct ChannelRef<'a> {
161    facade: &'a LogFacade,
162    name: String,
163}
164
165impl<'a> ChannelRef<'a> {
166    /// 通道名
167    pub fn name(&self) -> &str {
168        &self.name
169    }
170
171    /// 记录日志到指定通道
172    pub fn log(&self, level: LogLevel, msg: &str) {
173        let guard = self.facade.channels.read();
174        if let Some(logger) = guard.get(&self.name) {
175            logger.log(level, msg);
176        }
177        match level {
178            LogLevel::Trace => tracing::trace!("[{}] {}", self.name, msg),
179            LogLevel::Debug => tracing::debug!("[{}] {}", self.name, msg),
180            LogLevel::Info => tracing::info!("[{}] {}", self.name, msg),
181            LogLevel::Warn => tracing::warn!("[{}] {}", self.name, msg),
182            LogLevel::Error => tracing::error!("[{}] {}", self.name, msg),
183        }
184    }
185
186    /// 记录 debug 级别日志
187    pub fn debug(&self, msg: &str) {
188        self.log(LogLevel::Debug, msg);
189    }
190
191    /// 记录 info 级别日志
192    pub fn info(&self, msg: &str) {
193        self.log(LogLevel::Info, msg);
194    }
195
196    /// 记录 warn 级别日志
197    pub fn warn(&self, msg: &str) {
198        self.log(LogLevel::Warn, msg);
199    }
200
201    /// 记录 error 级别日志
202    pub fn error(&self, msg: &str) {
203        self.log(LogLevel::Error, msg);
204    }
205}
206
207/// 从字符串解析日志级别
208///
209/// 支持大小写不敏感:`"DEBUG"` / `"debug"` / `"Debug"` 均解析为 `LogLevel::Debug`。
210/// 未知字符串默认为 `LogLevel::Info`。
211pub fn parse_level(s: &str) -> LogLevel {
212    match s.to_lowercase().as_str() {
213        "trace" => LogLevel::Trace,
214        "debug" => LogLevel::Debug,
215        "info" => LogLevel::Info,
216        "warn" | "warning" => LogLevel::Warn,
217        "error" => LogLevel::Error,
218        _ => LogLevel::Info,
219    }
220}
221
222/// 从 `LogChannel` 配置构造 `StructuredLogger`
223fn channel_to_logger(channel: &LogChannel) -> StructuredLogger {
224    StructuredLogger::with_level(parse_level(&channel.level))
225}
226
227// ============================================================================
228// 单元测试
229// ============================================================================
230
231#[cfg(test)]
232mod tests {
233    use super::*;
234    use sz_rust_infra_facade::config::{LogChannel, LogSection};
235
236    /// 构造测试用的 LogSection(含 file + console 两个通道)
237    fn make_log_section() -> LogSection {
238        let mut channels = HashMap::new();
239        channels.insert(
240            "file".to_string(),
241            LogChannel {
242                r#type: "file".to_string(),
243                path: "runtime/logs".to_string(),
244                level: "info".to_string(),
245                max_files: 30,
246                format: "%{time} [%{level}] %{message}".to_string(),
247            },
248        );
249        channels.insert(
250            "console".to_string(),
251            LogChannel {
252                r#type: "console".to_string(),
253                path: String::new(),
254                level: "debug".to_string(),
255                max_files: 0,
256                format: "%{time} [%{level}] %{message}".to_string(),
257            },
258        );
259        LogSection {
260            default: "file".to_string(),
261            channels,
262        }
263    }
264
265    /// 测试 parse_level 各种输入
266    #[test]
267    fn test_parse_level() {
268        assert_eq!(parse_level("debug"), LogLevel::Debug);
269        assert_eq!(parse_level("DEBUG"), LogLevel::Debug);
270        assert_eq!(parse_level("Debug"), LogLevel::Debug);
271        assert_eq!(parse_level("info"), LogLevel::Info);
272        assert_eq!(parse_level("INFO"), LogLevel::Info);
273        assert_eq!(parse_level("warn"), LogLevel::Warn);
274        assert_eq!(parse_level("warning"), LogLevel::Warn);
275        assert_eq!(parse_level("WARN"), LogLevel::Warn);
276        assert_eq!(parse_level("error"), LogLevel::Error);
277        assert_eq!(parse_level("ERROR"), LogLevel::Error);
278        // 未知字符串默认 Info
279        assert_eq!(parse_level("unknown"), LogLevel::Info);
280        assert_eq!(parse_level(""), LogLevel::Info);
281    }
282
283    /// 测试 LogFacade 构造和默认通道
284    #[test]
285    fn test_log_facade_new() {
286        let section = make_log_section();
287        let facade = LogFacade::new(&section);
288
289        assert_eq!(facade.default_channel(), "file");
290        let names = facade.channel_names();
291        assert_eq!(names.len(), 2);
292        assert!(names.contains(&"file".to_string()));
293        assert!(names.contains(&"console".to_string()));
294    }
295
296    /// 测试默认 logger 级别取自 default 通道
297    #[test]
298    fn test_default_logger_level() {
299        let section = make_log_section();
300        let facade = LogFacade::new(&section);
301
302        // file 通道 level=info,所以默认 logger 级别为 Info
303        assert_eq!(facade.logger().level(), LogLevel::Info);
304
305        // Debug 级别应被过滤
306        facade.debug("debug msg - should be filtered");
307        let entries = facade.logger().entries();
308        assert!(entries.iter().all(|e| e.level != LogLevel::Debug));
309    }
310
311    /// 测试日志记录到默认 logger
312    #[test]
313    fn test_log_to_default_logger() {
314        let section = make_log_section();
315        let facade = LogFacade::new(&section);
316
317        facade.info("test info message");
318        facade.warn("test warn message");
319        facade.error("test error message");
320
321        let entries = facade.logger().entries();
322        assert!(entries.iter().any(|e| e.message == "test info message"));
323        assert!(entries.iter().any(|e| e.message == "test warn message"));
324        assert!(entries.iter().any(|e| e.message == "test error message"));
325    }
326
327    /// 测试通过 ChannelRef 访问命名通道
328    #[test]
329    fn test_channel_access() {
330        let section = make_log_section();
331        let facade = LogFacade::new(&section);
332
333        // file 通道存在
334        let file_channel = facade.channel("file");
335        assert!(file_channel.is_some());
336        let file_channel = file_channel.unwrap();
337        assert_eq!(file_channel.name(), "file");
338
339        // console 通道存在
340        let console_channel = facade.channel("console");
341        assert!(console_channel.is_some());
342
343        // 不存在的通道返回 None
344        assert!(facade.channel("nonexistent").is_none());
345    }
346
347    /// 测试 console 通道(level=debug)能记录所有级别
348    #[test]
349    fn test_console_channel_debug_level() {
350        let section = make_log_section();
351        let facade = LogFacade::new(&section);
352
353        let console = facade.channel("console").unwrap();
354        console.debug("debug msg");
355        console.info("info msg");
356        console.warn("warn msg");
357        console.error("error msg");
358
359        // console 通道 level=debug,所有级别都应记录
360        let guard = facade.channels.read();
361        let console_logger = guard.get("console").unwrap();
362        let entries = console_logger.entries();
363        assert_eq!(entries.len(), 4);
364    }
365
366    /// 测试 LogFacade init 全局单例
367    #[test]
368    fn test_log_facade_init_singleton() {
369        let section = make_log_section();
370        let facade = LogFacade::init(&section);
371
372        // instance() 应返回同一实例
373        let facade2 = LogFacade::instance();
374        assert!(facade2.is_some());
375        assert!(std::ptr::eq(facade, facade2.unwrap()));
376
377        // 再次 init 应返回同一实例(不覆盖)
378        let section2 = make_log_section();
379        let facade3 = LogFacade::init(&section2);
380        assert!(std::ptr::eq(facade, facade3));
381    }
382
383    /// 测试从实际配置文件加载日志配置
384    #[test]
385    fn test_load_from_config_file() {
386        // 查找 config 目录
387        let config_dir = std::env::current_dir().ok().and_then(|d| {
388            let mut current = d.clone();
389            for _ in 0..5 {
390                if current.join("config").exists() {
391                    return Some(current.join("config"));
392                }
393                if let Some(parent) = current.parent() {
394                    current = parent.to_path_buf();
395                } else {
396                    break;
397                }
398            }
399            None
400        });
401
402        let Some(config_dir) = config_dir else {
403            eprintln!("跳过:未找到 config 目录");
404            return;
405        };
406
407        let log_path = config_dir.join("log.yml");
408        if !log_path.exists() {
409            eprintln!("跳过:未找到 log.yml");
410            return;
411        }
412
413        let content = std::fs::read_to_string(&log_path).unwrap();
414        let section: LogSection = serde_yaml::from_str(&content).unwrap();
415
416        // 验证默认通道为 file
417        assert_eq!(section.default, "file");
418
419        // 验证有 file 和 console 两个通道
420        assert!(section.channels.contains_key("file"));
421        assert!(section.channels.contains_key("console"));
422
423        // 验证 file 通道配置
424        let file_channel = section.channels.get("file").unwrap();
425        assert_eq!(file_channel.r#type, "file");
426        assert_eq!(file_channel.level, "info");
427        assert_eq!(file_channel.max_files, 30);
428
429        // 验证 console 通道配置
430        let console_channel = section.channels.get("console").unwrap();
431        assert_eq!(console_channel.r#type, "console");
432        assert_eq!(console_channel.level, "debug");
433    }
434
435    /// 测试 LogFacade::new 处理空 channels(默认通道不存在时用 Info 级别)
436    #[test]
437    fn test_log_facade_with_empty_channels() {
438        let section = LogSection::default();
439        let facade = LogFacade::new(&section);
440
441        // 默认通道为空,logger 级别应为 Info(fallback)
442        assert_eq!(facade.logger().level(), LogLevel::Info);
443        assert_eq!(facade.default_channel(), "");
444    }
445
446    /// 测试 LogFacade::Debug 输出
447    #[test]
448    fn test_log_facade_debug_format() {
449        let section = make_log_section();
450        let facade = LogFacade::new(&section);
451
452        let debug_str = format!("{:?}", facade);
453        assert!(debug_str.contains("LogFacade"));
454        assert!(debug_str.contains("file"));
455    }
456}
457// Log 中间件 — 请求/响应日志(对齐 PHP `think-logger`)
458//
459// sz-rust 自研中间件,PHP 端无全局 Log 中间件(PHP `app/middleware.php` 仅含
460// `SessionInit` + `AllowCrossDomain`)。本模块在 [`crate::order::DEFAULT_ORDER`]
461// 中位于第 3 位(`Trace` → `Cors` → **`Log`** → `RateLimit` → `Auth`)。
462//
463// ## 行为
464//
465// 1. **入口**:生成 `RequestId`(如果 extensions 中没有,则新生成),注入 extensions
466// 2. **记录起始时间**:`std::time::Instant::now()`
467// 3. **调用 `next.run(req)`**:传递请求给下游
468// 4. **出口**:根据响应状态码记录日志
469//    - 2xx/3xx → `tracing::info!`
470//    - 4xx → `tracing::warn!`(对齐 PHP `apart_level=['error','sql']` 的级别分离思想)
471//    - 5xx → `tracing::error!`
472//
473// ## 日志字段
474//
475// | 字段 | 来源 | 说明 |
476// |------|------|------|
477// | `request_id` | `generate_request_id()` | 全局唯一计数器 + 时间戳,16 字符 hex |
478// | `method` | `Request::method()` | HTTP 方法 |
479// | `uri` | `Request::uri().path()` | 请求路径(不含查询字符串) |
480// | `status` | `Response::status().as_u16()` | HTTP 状态码 |
481// | `duration_ms` | `Instant::elapsed()` | 请求耗时(毫秒) |
482//
483// ## PHP 对齐
484//
485// PHP 端无 Log 中间件,业务代码通过 `Log::info()` 等主动调用。
486// sz-rust 的 Log 中间件是自研增强,提供:
487// - 请求生命周期自动日志(无需业务代码手动调用)
488// - 请求 ID 追踪(贯穿整个请求链路)
489// - 响应状态码分级日志(4xx Warn / 5xx Error)
490//
491// 日志级别对齐 think-logger 的 4 级(debug/info/warn/error),
492// `apart_level` 思想对齐 PHP `config/log.php` 的 `['error','sql']` 独立文件配置。
493//
494// ## 用法
495//
496// ```ignore
497// use sz_rust_core::middleware::log::log_middleware;
498// use axum::Router;
499//
500// let app: Router = Router::new()
501//     .route("/", axum::routing::get(|| async { "ok" }))
502//     .layer(axum::middleware::from_fn(log_middleware));
503// ```
504
505use axum::extract::Request;
506use axum::middleware::Next;
507use axum::response::Response;
508use std::sync::atomic::{AtomicU64, Ordering};
509use std::time::Instant;
510
511/// 请求 ID(注入到 request extensions,供下游 handler 和日志使用)
512///
513/// 生成方式:全局 `AtomicU64` 计数器 + 当前时间戳,保证进程内唯一。
514/// 格式:16 字符 hex(`{timestamp_secs:08x}{counter:08x}`)。
515#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
516pub struct RequestId {
517    /// 时间戳部分(UNIX 秒)
518    timestamp_secs: u64,
519    /// 计数器部分(进程内递增)
520    counter: u64,
521}
522
523impl RequestId {
524    /// 返回 16 字符 hex 字符串
525    ///
526    /// 格式:`{timestamp_secs:08x}{counter:08x}`(对齐 W3C traceparent 的 16 字符 span_id 长度)。
527    pub fn to_hex(&self) -> String {
528        format!("{:08x}{:08x}", self.timestamp_secs, self.counter)
529    }
530
531    /// 返回时间戳部分
532    pub fn timestamp_secs(&self) -> u64 {
533        self.timestamp_secs
534    }
535
536    /// 返回计数器部分
537    pub fn counter(&self) -> u64 {
538        self.counter
539    }
540}
541
542impl std::fmt::Display for RequestId {
543    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
544        f.write_str(&self.to_hex())
545    }
546}
547
548/// 全局 request_id 计数器(进程内递增)
549static REQUEST_ID_COUNTER: AtomicU64 = AtomicU64::new(0);
550
551/// 生成新的 `RequestId`
552///
553/// 使用全局 `AtomicU64` 计数器 + 当前 UNIX 时间戳,保证进程内唯一。
554/// 多线程安全(`fetch_add` 是原子操作)。
555pub fn generate_request_id() -> RequestId {
556    let counter = REQUEST_ID_COUNTER.fetch_add(1, Ordering::Relaxed);
557    let timestamp_secs = std::time::SystemTime::now()
558        .duration_since(std::time::UNIX_EPOCH)
559        .map(|d| d.as_secs())
560        .unwrap_or(0);
561    RequestId {
562        timestamp_secs,
563        counter,
564    }
565}
566
567/// Log 中间件配置
568#[derive(Debug, Clone, Default)]
569pub struct LogConfig {
570    /// 排除路径(不记录日志,对齐 PHP 端白名单思想)
571    ///
572    /// 支持精确匹配(如 `/health`)和通配符匹配(如 `/health/*`)。
573    pub exclude_paths: Vec<String>,
574}
575
576impl LogConfig {
577    /// 创建带排除路径的配置
578    pub fn with_exclude_paths(mut self, paths: Vec<String>) -> Self {
579        self.exclude_paths = paths;
580        self
581    }
582
583    /// 判断路径是否被排除
584    ///
585    /// 支持精确匹配和 `*` 通配符匹配(复用 [`crate::auth::is_route_allowed`] 的逻辑)。
586    pub fn is_excluded(&self, path: &str) -> bool {
587        crate::auth::is_route_allowed(path, &self.exclude_paths)
588    }
589}
590
591/// 根据响应状态码返回日志级别
592///
593/// 对齐 PHP `config/log.php` 的 `apart_level=['error','sql']` 思想:
594/// - 2xx/3xx → `Info`(成功请求)
595/// - 4xx → `Warn`(客户端错误)
596/// - 5xx → `Error`(服务端错误)
597///
598/// 其他状态码(如 1xx)默认为 `Info`。
599pub fn log_level_for_status(status: u16) -> LogLevel {
600    match status {
601        400..=499 => LogLevel::Warn,
602        500..=599 => LogLevel::Error,
603        _ => LogLevel::Info,
604    }
605}
606
607/// 格式化请求日志消息
608///
609/// 输出格式:`request_id=<hex> method=<METHOD> uri=<path> status=<code> duration_ms=<ms>`
610///
611/// 此函数主要用于测试可验证的纯函数,中间件实际输出通过 `tracing` 宏的结构化字段实现。
612pub fn format_request_log(
613    method: &str,
614    uri: &str,
615    status: u16,
616    duration_ms: u64,
617    request_id: &RequestId,
618) -> String {
619    format!(
620        "request_id={} method={} uri={} status={} duration_ms={}",
621        request_id.to_hex(),
622        method,
623        uri,
624        status,
625        duration_ms
626    )
627}
628
629/// Log 中间件 — 请求/响应日志
630///
631/// ## 校验流程
632///
633/// 1. **提取请求信息**:method, uri(在 `req` 被消费之前)
634/// 2. **生成 RequestId**:如果 extensions 中没有,则新生成
635/// 3. **记录起始时间**:`Instant::now()`
636/// 4. **注入 RequestId**:插入 request extensions
637/// 5. **调用 `next.run(req)`**:传递请求给下游
638/// 6. **计算耗时**:`start.elapsed()`
639/// 7. **记录日志**:根据状态码选择级别,输出结构化日志
640///
641/// ## 排除路径
642///
643/// 如果请求路径在 [`LogConfig::exclude_paths`] 中,则不记录日志(但仍注入 RequestId)。
644///
645/// ## 用法
646///
647/// ```ignore
648/// use sz_rust_core::middleware::log::{log_middleware, LogConfig};
649/// use axum::Router;
650///
651/// let config = LogConfig::default();
652/// let app: Router = Router::new()
653///     .route("/", axum::routing::get(|| async { "ok" }))
654///     .layer(axum::middleware::from_fn_with_state(config, log_middleware_with_config));
655/// ```
656pub async fn log_middleware(req: Request, next: Next) -> Response {
657    log_middleware_inner(req, next, &LogConfig::default()).await
658}
659
660/// 带配置的 Log 中间件
661pub async fn log_middleware_with_config(
662    axum::extract::State(config): axum::extract::State<LogConfig>,
663    req: Request,
664    next: Next,
665) -> Response {
666    log_middleware_inner(req, next, &config).await
667}
668
669async fn log_middleware_inner(req: Request, next: Next, config: &LogConfig) -> Response {
670    // 1. 提取请求信息(在 req 被消费之前)
671    let method = req.method().clone();
672    let uri = req.uri().path().to_string();
673
674    // 2. 生成 RequestId(如果 extensions 中没有,则新生成)
675    let request_id = req
676        .extensions()
677        .get::<RequestId>()
678        .copied()
679        .unwrap_or_else(generate_request_id);
680
681    // 3. 记录起始时间
682    let start = Instant::now();
683
684    // 4. 注入 RequestId 到 extensions
685    let mut req = req;
686    req.extensions_mut().insert(request_id);
687
688    // 5. 调用 next
689    let response = next.run(req).await;
690
691    // 6. 计算耗时
692    let duration_ms = start.elapsed().as_millis() as u64;
693
694    // 7. 记录日志(排除路径不记录)
695    if !config.is_excluded(&uri) {
696        let status = response.status().as_u16();
697        let level = log_level_for_status(status);
698        let request_id_hex = request_id.to_hex();
699
700        match level {
701            LogLevel::Trace => tracing::trace!(
702                request_id = %request_id_hex,
703                method = %method,
704                uri = %uri,
705                status = status,
706                duration_ms = duration_ms,
707                "request completed"
708            ),
709            LogLevel::Debug => tracing::debug!(
710                request_id = %request_id_hex,
711                method = %method,
712                uri = %uri,
713                status = status,
714                duration_ms = duration_ms,
715                "request completed"
716            ),
717            LogLevel::Info => tracing::info!(
718                request_id = %request_id_hex,
719                method = %method,
720                uri = %uri,
721                status = status,
722                duration_ms = duration_ms,
723                "request completed"
724            ),
725            LogLevel::Warn => tracing::warn!(
726                request_id = %request_id_hex,
727                method = %method,
728                uri = %uri,
729                status = status,
730                duration_ms = duration_ms,
731                "request completed"
732            ),
733            LogLevel::Error => tracing::error!(
734                request_id = %request_id_hex,
735                method = %method,
736                uri = %uri,
737                status = status,
738                duration_ms = duration_ms,
739                "request completed"
740            ),
741        }
742    }
743
744    response
745}
746
747#[cfg(test)]
748mod middleware_tests {
749    use super::*;
750    use axum::body::Body;
751    use axum::http::StatusCode;
752    use axum::Router;
753    use http_body_util::BodyExt;
754    use tower::ServiceExt;
755
756    // ====================================================================
757    // 辅助函数
758    // ====================================================================
759
760    async fn read_body(resp: Response) -> String {
761        let bytes = resp.into_body().collect().await.unwrap().to_bytes();
762        String::from_utf8(bytes.to_vec()).unwrap()
763    }
764
765    fn make_request(method: &str, uri: &str) -> Request {
766        Request::builder()
767            .method(method)
768            .uri(uri)
769            .body(Body::empty())
770            .unwrap()
771    }
772
773    /// 构建测试用 Router
774    fn build_app() -> Router {
775        Router::new()
776            .route(
777                "/ok",
778                axum::routing::get(|| async { axum::http::StatusCode::OK }),
779            )
780            .route(
781                "/notfound",
782                axum::routing::get(|| async { axum::http::StatusCode::NOT_FOUND }),
783            )
784            .route(
785                "/error",
786                axum::routing::get(|| async { axum::http::StatusCode::INTERNAL_SERVER_ERROR }),
787            )
788            .route("/body", axum::routing::get(|| async { "hello" }))
789            .layer(axum::middleware::from_fn(log_middleware))
790    }
791
792    // ====================================================================
793    // RequestId 单元测试
794    // ====================================================================
795
796    #[test]
797    fn test_request_id_to_hex_is_16_chars() {
798        let id = RequestId {
799            timestamp_secs: 0x12345678,
800            counter: 0x9ABCDEF0,
801        };
802        let hex = id.to_hex();
803        assert_eq!(hex.len(), 16);
804        assert_eq!(hex, "123456789abcdef0");
805    }
806
807    #[test]
808    fn test_request_id_to_hex_zero() {
809        let id = RequestId {
810            timestamp_secs: 0,
811            counter: 0,
812        };
813        assert_eq!(id.to_hex(), "0000000000000000");
814    }
815
816    #[test]
817    fn test_request_id_to_hex_max() {
818        let id = RequestId {
819            timestamp_secs: u64::MAX,
820            counter: u64::MAX,
821        };
822        // u64::MAX = 0xFFFFFFFFFFFFFFFF,但 format!("{:08x}", u64::MAX) 会输出 16 字符
823        let hex = id.to_hex();
824        assert_eq!(hex.len(), 32); // 每部分 16 字符,总共 32 字符
825    }
826
827    #[test]
828    fn test_request_id_display_matches_to_hex() {
829        let id = RequestId {
830            timestamp_secs: 0x12345678,
831            counter: 0x9ABCDEF0,
832        };
833        assert_eq!(format!("{}", id), id.to_hex());
834    }
835
836    #[test]
837    fn test_request_id_accessors() {
838        let id = RequestId {
839            timestamp_secs: 100,
840            counter: 200,
841        };
842        assert_eq!(id.timestamp_secs(), 100);
843        assert_eq!(id.counter(), 200);
844    }
845
846    #[test]
847    fn test_request_id_equality() {
848        let id1 = RequestId {
849            timestamp_secs: 1,
850            counter: 2,
851        };
852        let id2 = RequestId {
853            timestamp_secs: 1,
854            counter: 2,
855        };
856        let id3 = RequestId {
857            timestamp_secs: 1,
858            counter: 3,
859        };
860        assert_eq!(id1, id2);
861        assert_ne!(id1, id3);
862    }
863
864    // ====================================================================
865    // generate_request_id 单元测试
866    // ====================================================================
867
868    #[test]
869    fn test_generate_request_id_returns_unique() {
870        let id1 = generate_request_id();
871        let id2 = generate_request_id();
872        // 计数器递增,保证唯一
873        assert_ne!(id1.counter(), id2.counter());
874        assert_eq!(id2.counter(), id1.counter() + 1);
875    }
876
877    #[test]
878    fn test_generate_request_id_hex_is_16_chars() {
879        let id = generate_request_id();
880        let hex = id.to_hex();
881        // 注意:如果 timestamp_secs 或 counter 超过 u32::MAX,hex 会超过 16 字符
882        // 但在正常情况下(timestamp < 2106 年,counter < 40 亿次),hex 是 16 字符
883        assert!(hex.len() >= 16);
884    }
885
886    // ====================================================================
887    // log_level_for_status 单元测试
888    // ====================================================================
889
890    #[test]
891    fn test_log_level_for_2xx_returns_info() {
892        assert_eq!(log_level_for_status(200), LogLevel::Info);
893        assert_eq!(log_level_for_status(201), LogLevel::Info);
894        assert_eq!(log_level_for_status(204), LogLevel::Info);
895    }
896
897    #[test]
898    fn test_log_level_for_3xx_returns_info() {
899        assert_eq!(log_level_for_status(301), LogLevel::Info);
900        assert_eq!(log_level_for_status(302), LogLevel::Info);
901        assert_eq!(log_level_for_status(304), LogLevel::Info);
902    }
903
904    #[test]
905    fn test_log_level_for_4xx_returns_warn() {
906        assert_eq!(log_level_for_status(400), LogLevel::Warn);
907        assert_eq!(log_level_for_status(401), LogLevel::Warn);
908        assert_eq!(log_level_for_status(403), LogLevel::Warn);
909        assert_eq!(log_level_for_status(404), LogLevel::Warn);
910        assert_eq!(log_level_for_status(422), LogLevel::Warn);
911        assert_eq!(log_level_for_status(499), LogLevel::Warn);
912    }
913
914    #[test]
915    fn test_log_level_for_5xx_returns_error() {
916        assert_eq!(log_level_for_status(500), LogLevel::Error);
917        assert_eq!(log_level_for_status(501), LogLevel::Error);
918        assert_eq!(log_level_for_status(502), LogLevel::Error);
919        assert_eq!(log_level_for_status(503), LogLevel::Error);
920        assert_eq!(log_level_for_status(599), LogLevel::Error);
921    }
922
923    #[test]
924    fn test_log_level_for_1xx_returns_info() {
925        // 1xx 信息响应默认为 Info
926        assert_eq!(log_level_for_status(100), LogLevel::Info);
927        assert_eq!(log_level_for_status(101), LogLevel::Info);
928    }
929
930    #[test]
931    fn test_log_level_for_boundary() {
932        // 边界测试:399 → Info,400 → Warn,499 → Warn,500 → Error,599 → Error,600 → Info
933        assert_eq!(log_level_for_status(399), LogLevel::Info);
934        assert_eq!(log_level_for_status(400), LogLevel::Warn);
935        assert_eq!(log_level_for_status(499), LogLevel::Warn);
936        assert_eq!(log_level_for_status(500), LogLevel::Error);
937        assert_eq!(log_level_for_status(599), LogLevel::Error);
938        assert_eq!(log_level_for_status(600), LogLevel::Info);
939    }
940
941    // ====================================================================
942    // format_request_log 单元测试
943    // ====================================================================
944
945    #[test]
946    fn test_format_request_log_basic() {
947        let request_id = RequestId {
948            timestamp_secs: 0x12345678,
949            counter: 0x9ABCDEF0,
950        };
951        let msg = format_request_log("GET", "/api/users", 200, 15, &request_id);
952        assert_eq!(
953            msg,
954            "request_id=123456789abcdef0 method=GET uri=/api/users status=200 duration_ms=15"
955        );
956    }
957
958    #[test]
959    fn test_format_request_log_post_method() {
960        let request_id = RequestId {
961            timestamp_secs: 0,
962            counter: 1,
963        };
964        let msg = format_request_log("POST", "/api/orders", 201, 42, &request_id);
965        assert_eq!(
966            msg,
967            "request_id=0000000000000001 method=POST uri=/api/orders status=201 duration_ms=42"
968        );
969    }
970
971    #[test]
972    fn test_format_request_log_error_status() {
973        let request_id = RequestId {
974            timestamp_secs: 0,
975            counter: 0,
976        };
977        let msg = format_request_log("GET", "/missing", 404, 5, &request_id);
978        assert_eq!(
979            msg,
980            "request_id=0000000000000000 method=GET uri=/missing status=404 duration_ms=5"
981        );
982    }
983
984    #[test]
985    fn test_format_request_log_with_query_string_in_uri() {
986        // uri 应该是原始 path(含查询字符串),由调用方决定是否截取
987        let request_id = RequestId {
988            timestamp_secs: 0,
989            counter: 0,
990        };
991        let msg = format_request_log("GET", "/api?foo=bar", 200, 1, &request_id);
992        assert!(msg.contains("uri=/api?foo=bar"));
993    }
994
995    // ====================================================================
996    // LogConfig 单元测试
997    // ====================================================================
998
999    #[test]
1000    fn test_log_config_default_empty_exclude_paths() {
1001        let config = LogConfig::default();
1002        assert!(config.exclude_paths.is_empty());
1003    }
1004
1005    #[test]
1006    fn test_log_config_with_exclude_paths() {
1007        let config = LogConfig::default().with_exclude_paths(vec!["/health".to_string()]);
1008        assert_eq!(config.exclude_paths, vec!["/health".to_string()]);
1009    }
1010
1011    #[test]
1012    fn test_log_config_is_excluded_exact_match() {
1013        let config = LogConfig::default().with_exclude_paths(vec!["/health".to_string()]);
1014        assert!(config.is_excluded("/health"));
1015        assert!(!config.is_excluded("/health/detail"));
1016        assert!(!config.is_excluded("/api"));
1017    }
1018
1019    #[test]
1020    fn test_log_config_is_excluded_wildcard_match() {
1021        let config = LogConfig::default().with_exclude_paths(vec!["/health/*".to_string()]);
1022        assert!(config.is_excluded("/health/check"));
1023        assert!(config.is_excluded("/health/deep/nested"));
1024        assert!(!config.is_excluded("/health"));
1025        assert!(!config.is_excluded("/api"));
1026    }
1027
1028    #[test]
1029    fn test_log_config_is_excluded_empty_list() {
1030        let config = LogConfig::default();
1031        assert!(!config.is_excluded("/any"));
1032    }
1033
1034    #[test]
1035    fn test_log_config_is_excluded_multiple_entries() {
1036        let config = LogConfig::default()
1037            .with_exclude_paths(vec!["/health".to_string(), "/metrics/*".to_string()]);
1038        assert!(config.is_excluded("/health"));
1039        assert!(config.is_excluded("/metrics/prometheus"));
1040        assert!(!config.is_excluded("/api"));
1041    }
1042
1043    // ====================================================================
1044    // log_middleware 集成测试
1045    // ====================================================================
1046
1047    #[tokio::test]
1048    async fn test_log_middleware_returns_response_unchanged() {
1049        // 验证中间件不修改响应体
1050        let app = build_app();
1051        let resp = app.oneshot(make_request("GET", "/body")).await.unwrap();
1052        let body = read_body(resp).await;
1053        assert_eq!(body, "hello");
1054    }
1055
1056    #[tokio::test]
1057    async fn test_log_middleware_returns_correct_status() {
1058        let app = build_app();
1059        let resp = app.oneshot(make_request("GET", "/ok")).await.unwrap();
1060        assert_eq!(resp.status(), StatusCode::OK);
1061    }
1062
1063    #[tokio::test]
1064    async fn test_log_middleware_injects_request_id() {
1065        // 验证 request_id 被注入 extensions
1066        let app = Router::new()
1067            .route(
1068                "/",
1069                axum::routing::get(|req: Request| async move {
1070                    let request_id = req.extensions().get::<RequestId>().unwrap();
1071                    format!("request_id:{}", request_id.to_hex())
1072                }),
1073            )
1074            .layer(axum::middleware::from_fn(log_middleware));
1075
1076        let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
1077        assert_eq!(resp.status(), StatusCode::OK);
1078        let body = read_body(resp).await;
1079        assert!(body.starts_with("request_id:"));
1080        // 验证 hex 长度至少 16 字符
1081        let hex = body.strip_prefix("request_id:").unwrap();
1082        assert!(hex.len() >= 16);
1083    }
1084
1085    #[tokio::test]
1086    async fn test_log_middleware_generates_unique_request_ids() {
1087        // 验证多个请求生成不同的 request_id
1088        let app = Router::new()
1089            .route(
1090                "/",
1091                axum::routing::get(|req: Request| async move {
1092                    let request_id = req.extensions().get::<RequestId>().unwrap();
1093                    request_id.to_hex()
1094                }),
1095            )
1096            .layer(axum::middleware::from_fn(log_middleware));
1097
1098        let resp1 = app.clone().oneshot(make_request("GET", "/")).await.unwrap();
1099        let hex1 = read_body(resp1).await;
1100
1101        let resp2 = app.oneshot(make_request("GET", "/")).await.unwrap();
1102        let hex2 = read_body(resp2).await;
1103
1104        assert_ne!(hex1, hex2);
1105    }
1106
1107    #[tokio::test]
1108    async fn test_log_middleware_preserves_existing_request_id() {
1109        // 验证已存在的 request_id 不被覆盖
1110        let existing_id = RequestId {
1111            timestamp_secs: 0xDEADBEEF,
1112            counter: 0x12345678,
1113        };
1114        let app = Router::new()
1115            .route(
1116                "/",
1117                axum::routing::get(|req: Request| async move {
1118                    let request_id = req.extensions().get::<RequestId>().unwrap();
1119                    request_id.to_hex()
1120                }),
1121            )
1122            .layer(axum::middleware::from_fn(log_middleware))
1123            .layer(
1124                tower::ServiceBuilder::new().layer(tower::layer::layer_fn(move |service| {
1125                    tower::util::MapRequest::new(service, move |mut req: Request| {
1126                        req.extensions_mut().insert(existing_id);
1127                        req
1128                    })
1129                })),
1130            );
1131
1132        let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
1133        let body = read_body(resp).await;
1134        assert_eq!(body, "deadbeef12345678");
1135    }
1136
1137    #[tokio::test]
1138    async fn test_log_middleware_records_2xx_status() {
1139        // 验证 2xx 响应正常处理(日志级别由 log_level_for_status 决定)
1140        let app = build_app();
1141        let resp = app.oneshot(make_request("GET", "/ok")).await.unwrap();
1142        assert_eq!(resp.status(), StatusCode::OK);
1143    }
1144
1145    #[tokio::test]
1146    async fn test_log_middleware_records_4xx_status() {
1147        let app = build_app();
1148        let resp = app.oneshot(make_request("GET", "/notfound")).await.unwrap();
1149        assert_eq!(resp.status(), StatusCode::NOT_FOUND);
1150    }
1151
1152    #[tokio::test]
1153    async fn test_log_middleware_records_5xx_status() {
1154        let app = build_app();
1155        let resp = app.oneshot(make_request("GET", "/error")).await.unwrap();
1156        assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR);
1157    }
1158
1159    #[tokio::test]
1160    async fn test_log_middleware_with_config_excludes_path() {
1161        // 验证排除路径不记录日志(但仍注入 request_id)
1162        let config = LogConfig::default().with_exclude_paths(vec!["/health".to_string()]);
1163        let app = Router::new()
1164            .route("/health", axum::routing::get(|| async { "healthy" }))
1165            .layer(axum::middleware::from_fn_with_state(
1166                config,
1167                log_middleware_with_config,
1168            ));
1169
1170        let resp = app.oneshot(make_request("GET", "/health")).await.unwrap();
1171        assert_eq!(resp.status(), StatusCode::OK);
1172        let body = read_body(resp).await;
1173        assert_eq!(body, "healthy");
1174    }
1175
1176    #[tokio::test]
1177    async fn test_log_middleware_with_config_wildcard_exclude() {
1178        // 验证通配符排除路径
1179        let config = LogConfig::default().with_exclude_paths(vec!["/metrics/*".to_string()]);
1180        let app = Router::new()
1181            .route(
1182                "/metrics/prometheus",
1183                axum::routing::get(|| async { "metrics" }),
1184            )
1185            .layer(axum::middleware::from_fn_with_state(
1186                config,
1187                log_middleware_with_config,
1188            ));
1189
1190        let resp = app
1191            .oneshot(make_request("GET", "/metrics/prometheus"))
1192            .await
1193            .unwrap();
1194        assert_eq!(resp.status(), StatusCode::OK);
1195    }
1196
1197    #[tokio::test]
1198    async fn test_log_middleware_preserves_method_and_uri() {
1199        // 验证 method 和 uri 被正确提取(通过日志消息格式验证)
1200        // 由于 tracing 宏输出在测试中难以捕获,这里验证中间件不破坏请求
1201        let app = build_app();
1202        let resp = app.oneshot(make_request("GET", "/ok")).await.unwrap();
1203        assert_eq!(resp.status(), StatusCode::OK);
1204    }
1205
1206    #[tokio::test]
1207    async fn test_log_middleware_duration_is_non_negative() {
1208        // 验证 duration_ms 是非负的(通过响应正常返回间接验证)
1209        let app = build_app();
1210        let start = std::time::Instant::now();
1211        let resp = app.oneshot(make_request("GET", "/ok")).await.unwrap();
1212        let elapsed = start.elapsed();
1213        assert!(resp.status().is_success());
1214        // 中间件内部记录的 duration_ms 应该 <= 测试外部的 elapsed
1215        assert!(elapsed.as_millis() < 5000); // 5 秒上限(防止死循环)
1216    }
1217
1218    #[tokio::test]
1219    async fn test_log_middleware_handles_post_request() {
1220        let app = Router::new()
1221            .route(
1222                "/submit",
1223                axum::routing::post(|| async { axum::http::StatusCode::CREATED }),
1224            )
1225            .layer(axum::middleware::from_fn(log_middleware));
1226
1227        let req = Request::builder()
1228            .method("POST")
1229            .uri("/submit")
1230            .body(Body::empty())
1231            .unwrap();
1232        let resp = app.oneshot(req).await.unwrap();
1233        assert_eq!(resp.status(), StatusCode::CREATED);
1234    }
1235
1236    #[tokio::test]
1237    async fn test_log_middleware_chains_with_other_middleware() {
1238        // 验证 Log 中间件与其他中间件链式调用
1239        async fn add_header_middleware(req: Request, next: Next) -> Response {
1240            let mut resp = next.run(req).await;
1241            resp.headers_mut()
1242                .insert("X-Custom", "value".parse().unwrap());
1243            resp
1244        }
1245
1246        let app = Router::new()
1247            .route("/", axum::routing::get(|| async { "ok" }))
1248            .layer(axum::middleware::from_fn(add_header_middleware))
1249            .layer(axum::middleware::from_fn(log_middleware));
1250
1251        let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
1252        assert_eq!(resp.status(), StatusCode::OK);
1253        assert_eq!(resp.headers().get("X-Custom").unwrap(), "value");
1254    }
1255
1256    // ====================================================================
1257    // PHP 行为对齐验证
1258    // ====================================================================
1259
1260    #[test]
1261    fn test_php_apart_level_alignment() {
1262        // 对齐 PHP `config/log.php` 的 `apart_level=['error','sql']` 思想:
1263        // 4xx → Warn(客户端错误,类似 PHP warning)
1264        // 5xx → Error(服务端错误,对齐 PHP error 独立文件)
1265        assert_eq!(log_level_for_status(200), LogLevel::Info);
1266        assert_eq!(log_level_for_status(404), LogLevel::Warn);
1267        assert_eq!(log_level_for_status(500), LogLevel::Error);
1268    }
1269
1270    #[test]
1271    fn test_php_think_logger_level_alignment() {
1272        // 对齐 PHP think-logger 的 4 级日志(debug/info/warn/error)
1273        // sz-rust 的 LogLevel 也是 4 级,一一对应
1274        let levels = [
1275            LogLevel::Debug,
1276            LogLevel::Info,
1277            LogLevel::Warn,
1278            LogLevel::Error,
1279        ];
1280        assert_eq!(levels.len(), 4);
1281    }
1282
1283    #[test]
1284    fn test_request_id_format_aligns_with_w3c_span_id_length() {
1285        // 对齐 W3C traceparent 的 span_id 长度(16 字符 hex)
1286        // 便于未来 Trace 中间件实现时与 trace_id 格式兼容
1287        let id = RequestId {
1288            timestamp_secs: 0x12345678,
1289            counter: 0x9ABCDEF0,
1290        };
1291        assert_eq!(id.to_hex().len(), 16);
1292    }
1293}