Skip to main content

sz_rust_core/
api_version.rs

1//! API 版本管理 — URL/Header/Query 多策略
2//!
3//! ## 设计目标
4//!
5//! 提供灵活的 API 版本协商机制,支持三种主流策略:
6//!
7//! 1. **URL 路径策略**:`/api/v1/users`、`/api/v2/users`
8//!    - 对齐 GitHub API、Twitter API 风格
9//!    - 版本号作为 URL 路径前缀
10//!
11//! 2. **Header 策略**:
12//!    - 自定义头 `X-API-Version: 2`
13//!    - 或 `Accept: application/vnd.api+json; version=2`
14//!    - 对齐 Stripe API、GitHub API(Accept header)风格
15//!
16//! 3. **Query 参数策略**:`/api/users?api_version=2`
17//!    - 对齐部分内部 API 风格
18//!    - 适合调试(浏览器直接访问)
19//!
20//! ## 中间件集成
21//!
22//! [`ApiVersionExtractor`] 是 axum 中间件,从请求中提取版本号并注入到请求扩展中。
23//! 后续 handler 可通过 [`Request::extensions`] 获取 [`ApiVersion`]。
24//!
25//! ## 路由分组
26//!
27//! [`VersionedRouter`] 提供按版本分组的路由构建器:
28//!
29//! ```ignore
30//! use sz_rust_core::api_version::{VersionedRouter, ApiVersion};
31//!
32//! let router = VersionedRouter::new()
33//!     .route("v1", "/users", axum::routing::get(get_users_v1))
34//!     .route("v2", "/users", axum::routing::get(get_users_v2))
35//!     .build();
36//! ```
37//!
38//! ## 默认版本与降级
39//!
40//! - 未指定版本时使用 `default_version`(通常为最新稳定版)
41//! - 不存在的版本返回 400 Bad Request(避免误用)
42
43use axum::extract::{Request, State};
44use axum::http::{HeaderMap, StatusCode, Uri};
45use axum::middleware::Next;
46use axum::response::{IntoResponse, Response};
47use std::collections::HashMap;
48use std::sync::Arc;
49
50// ============================================================================
51// API 版本号
52// ============================================================================
53
54/// API 版本号
55///
56/// 内部以 `u32` 存储,支持 `v1`、`v2` 等数字版本。
57/// 不支持语义化版本(如 `v1.2.3`),保持 API 版本协商简单。
58#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
59pub struct ApiVersion(u32);
60
61impl ApiVersion {
62    /// 创建新的 API 版本号
63    pub const fn new(version: u32) -> Self {
64        Self(version)
65    }
66
67    /// 获取版本号数值
68    pub fn as_u32(&self) -> u32 {
69        self.0
70    }
71
72    /// 从字符串解析版本号
73    ///
74    /// 支持格式:
75    /// - `"1"` / `"2"` → 纯数字
76    /// - `"v1"` / `"v2"` → v 前缀
77    /// - `"version=1"` → query 参数格式
78    pub fn parse(s: &str) -> Option<Self> {
79        let trimmed = s.trim();
80        // 处理 "version=1" 形式(先检测,避免被 trim_start_matches('v') 误删前缀 'v')
81        let num_str = if let Some(rest) = trimmed.strip_prefix("version=") {
82            rest.trim()
83        } else {
84            trimmed.trim_start_matches('v').trim()
85        };
86        num_str.parse::<u32>().ok().map(Self)
87    }
88
89    /// 转为 `vN` 字符串(如 `v1`、`v2`)
90    pub fn to_vstring(&self) -> String {
91        format!("v{}", self.0)
92    }
93}
94
95impl std::fmt::Display for ApiVersion {
96    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
97        write!(f, "v{}", self.0)
98    }
99}
100
101impl Default for ApiVersion {
102    fn default() -> Self {
103        Self::new(1)
104    }
105}
106
107// ============================================================================
108// 版本协商策略
109// ============================================================================
110
111/// 版本协商策略
112///
113/// 控制从请求中提取版本号的优先级和方式。
114#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
115pub enum VersionStrategy {
116    /// URL 路径策略:`/api/v1/users`
117    ///
118    /// 从 URL 路径的第一段提取版本号(如 `/v1/users` → `v1`)。
119    #[default]
120    UrlPath,
121    /// 自定义 Header 策略:`X-API-Version: 1`
122    Header,
123    /// Accept Header 策略:`Accept: application/vnd.api+json; version=1`
124    AcceptHeader,
125    /// Query 参数策略:`?api_version=1`
126    Query,
127}
128
129// ============================================================================
130// 版本协商器
131// ============================================================================
132
133/// 版本协商器
134///
135/// 从请求中提取 API 版本号,按配置的策略顺序尝试。
136#[derive(Debug, Clone)]
137pub struct VersionNegotiator {
138    /// 支持的版本列表(如 `[v1, v2, v3]`)
139    supported_versions: Vec<ApiVersion>,
140    /// 默认版本(未指定时使用)
141    default_version: ApiVersion,
142    /// 策略优先级顺序(前者优先)
143    strategies: Vec<VersionStrategy>,
144    /// URL 路径中版本前缀的识别前缀(如 `/api/v1/users` 中的 `api`)
145    url_prefix: Option<String>,
146    /// 版本 query 参数名(默认 `api_version`)
147    query_param_name: String,
148    /// 版本自定义 header 名(默认 `x-api-version`)
149    header_name: String,
150}
151
152impl Default for VersionNegotiator {
153    fn default() -> Self {
154        Self {
155            supported_versions: vec![ApiVersion::new(1)],
156            default_version: ApiVersion::new(1),
157            strategies: vec![
158                VersionStrategy::UrlPath,
159                VersionStrategy::Header,
160                VersionStrategy::AcceptHeader,
161                VersionStrategy::Query,
162            ],
163            url_prefix: Some("api".to_string()),
164            query_param_name: "api_version".to_string(),
165            header_name: "x-api-version".to_string(),
166        }
167    }
168}
169
170impl VersionNegotiator {
171    /// 创建新的版本协商器
172    pub fn new(default_version: ApiVersion) -> Self {
173        Self {
174            supported_versions: vec![default_version],
175            default_version,
176            ..Default::default()
177        }
178    }
179
180    /// 设置支持的版本列表
181    pub fn with_supported_versions(mut self, versions: Vec<ApiVersion>) -> Self {
182        self.supported_versions = versions;
183        self
184    }
185
186    /// 设置策略优先级顺序
187    pub fn with_strategies(mut self, strategies: Vec<VersionStrategy>) -> Self {
188        self.strategies = strategies;
189        self
190    }
191
192    /// 设置 URL 路径前缀(如 `api`、`v` 等)
193    pub fn with_url_prefix(mut self, prefix: impl Into<String>) -> Self {
194        self.url_prefix = Some(prefix.into());
195        self
196    }
197
198    /// 设置 query 参数名
199    pub fn with_query_param(mut self, name: impl Into<String>) -> Self {
200        self.query_param_name = name.into();
201        self
202    }
203
204    /// 设置自定义 header 名
205    pub fn with_header_name(mut self, name: impl Into<String>) -> Self {
206        self.header_name = name.into().to_lowercase();
207        self
208    }
209
210    /// 从请求中协商版本号
211    ///
212    /// 按策略顺序尝试,首个成功的版本号即为协商结果。
213    /// 若所有策略都未匹配,返回默认版本。
214    ///
215    /// # 返回
216    ///
217    /// - `Ok(ApiVersion)`:协商成功(可能是匹配的版本或默认版本)
218    /// - `Err(VersionError)`:客户端指定了不支持的版本
219    pub fn negotiate(&self, uri: &Uri, headers: &HeaderMap) -> Result<ApiVersion, VersionError> {
220        for strategy in &self.strategies {
221            let extracted = match strategy {
222                VersionStrategy::UrlPath => self.extract_from_url_path(uri),
223                VersionStrategy::Header => self.extract_from_header(headers),
224                VersionStrategy::AcceptHeader => self.extract_from_accept_header(headers),
225                VersionStrategy::Query => self.extract_from_query(uri),
226            };
227
228            if let Some(version) = extracted {
229                // 客户端指定了版本,但不在支持列表中
230                if !self.supported_versions.contains(&version) {
231                    return Err(VersionError::UnsupportedVersion(version));
232                }
233                return Ok(version);
234            }
235        }
236
237        // 所有策略都未匹配,使用默认版本
238        Ok(self.default_version)
239    }
240
241    /// 从 URL 路径提取版本号
242    ///
243    /// 形如 `/api/v1/users` → `v1`,要求 `url_prefix` 后紧跟版本段。
244    fn extract_from_url_path(&self, uri: &Uri) -> Option<ApiVersion> {
245        let path = uri.path();
246        let segments: Vec<&str> = path.trim_start_matches('/').split('/').collect();
247
248        // 查找 url_prefix 后的版本段
249        let start_idx = if let Some(ref prefix) = self.url_prefix {
250            segments.iter().position(|s| *s == prefix.as_str())? + 1
251        } else {
252            0
253        };
254
255        if start_idx >= segments.len() {
256            return None;
257        }
258
259        ApiVersion::parse(segments[start_idx])
260    }
261
262    /// 从自定义 header 提取版本号
263    fn extract_from_header(&self, headers: &HeaderMap) -> Option<ApiVersion> {
264        headers
265            .get(&self.header_name)
266            .and_then(|v| v.to_str().ok())
267            .and_then(ApiVersion::parse)
268    }
269
270    /// 从 Accept header 提取版本号
271    ///
272    /// 形如 `application/vnd.api+json; version=1`
273    fn extract_from_accept_header(&self, headers: &HeaderMap) -> Option<ApiVersion> {
274        let accept = headers.get(axum::http::header::ACCEPT)?.to_str().ok()?;
275        // 查找 `version=` 参数
276        for part in accept.split(';') {
277            let part = part.trim();
278            if let Some(rest) = part.strip_prefix("version=") {
279                return ApiVersion::parse(rest.trim_matches('"'));
280            }
281        }
282        None
283    }
284
285    /// 从 query 参数提取版本号
286    fn extract_from_query(&self, uri: &Uri) -> Option<ApiVersion> {
287        let query = uri.query()?;
288        for pair in query.split('&') {
289            let mut parts = pair.splitn(2, '=');
290            if parts.next()? == self.query_param_name {
291                return ApiVersion::parse(parts.next()?);
292            }
293        }
294        None
295    }
296}
297
298// ============================================================================
299// 版本错误
300// ============================================================================
301
302/// 版本协商错误
303#[derive(Debug, Clone, PartialEq, Eq)]
304pub enum VersionError {
305    /// 客户端指定了不支持的版本
306    UnsupportedVersion(ApiVersion),
307}
308
309impl std::fmt::Display for VersionError {
310    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
311        match self {
312            VersionError::UnsupportedVersion(v) => {
313                write!(f, "Unsupported API version: {}", v)
314            }
315        }
316    }
317}
318
319impl std::error::Error for VersionError {}
320
321impl IntoResponse for VersionError {
322    fn into_response(self) -> Response {
323        let body = match self {
324            VersionError::UnsupportedVersion(v) => {
325                format!(
326                    "{{\"code\":0,\"msg\":\"Unsupported API version: {}\",\"data\":{{}}}}",
327                    v
328                )
329            }
330        };
331        (
332            StatusCode::BAD_REQUEST,
333            [(
334                axum::http::header::CONTENT_TYPE,
335                "application/json; charset=utf-8",
336            )],
337            body,
338        )
339            .into_response()
340    }
341}
342
343// ============================================================================
344// 中间件
345// ============================================================================
346
347/// 版本协商中间件状态
348///
349/// 通过 [`State`] 注入到中间件,避免每次请求重建协商器。
350#[derive(Clone)]
351pub struct ApiVersionExtractor {
352    negotiator: Arc<VersionNegotiator>,
353}
354
355impl ApiVersionExtractor {
356    /// 创建新的版本提取器
357    pub fn new(negotiator: VersionNegotiator) -> Self {
358        Self {
359            negotiator: Arc::new(negotiator),
360        }
361    }
362
363    /// 获取协商器引用
364    pub fn negotiator(&self) -> &VersionNegotiator {
365        &self.negotiator
366    }
367}
368
369/// 版本协商中间件
370///
371/// 从请求中提取版本号并注入到请求扩展中。
372/// 后续 handler 可通过 `req.extensions().get::<ApiVersion>()` 获取。
373///
374/// # 错误处理
375///
376/// 若客户端指定了不支持的版本,直接返回 400 Bad Request。
377pub async fn version_negotiation_middleware(
378    State(extractor): State<ApiVersionExtractor>,
379    req: Request,
380    next: Next,
381) -> Response {
382    let (parts, body) = req.into_parts();
383    let uri = parts.uri.clone();
384    let headers = parts.headers.clone();
385
386    match extractor.negotiator.negotiate(&uri, &headers) {
387        Ok(version) => {
388            let mut req = Request::from_parts(parts, body);
389            req.extensions_mut().insert(version);
390            next.run(req).await
391        }
392        Err(err) => err.into_response(),
393    }
394}
395
396// ============================================================================
397// 版本化路由
398// ============================================================================
399
400/// 版本化路由构建器
401///
402/// 按 API 版本分组注册路由,自动添加版本前缀。
403///
404/// # 用法
405///
406/// ```ignore
407/// use sz_rust_core::api_version::{VersionedRouter, ApiVersion};
408///
409/// let router = VersionedRouter::new()
410///     .route("v1", "/users", axum::routing::get(get_users_v1))
411///     .route("v2", "/users", axum::routing::get(get_users_v2))
412///     .build();
413/// ```
414#[derive(Default)]
415pub struct VersionedRouter {
416    /// 版本 → (路径, MethodRouter) 列表
417    routes: HashMap<String, Vec<(String, axum::routing::MethodRouter)>>,
418    /// URL 前缀(如 `api`)
419    url_prefix: Option<String>,
420}
421
422impl VersionedRouter {
423    /// 创建新的版本化路由构建器
424    pub fn new() -> Self {
425        Self::default()
426    }
427
428    /// 设置 URL 前缀(如 `api`,最终路径为 `/api/v1/users`)
429    pub fn with_url_prefix(mut self, prefix: impl Into<String>) -> Self {
430        self.url_prefix = Some(prefix.into());
431        self
432    }
433
434    /// 注册版本化路由
435    ///
436    /// # 参数
437    ///
438    /// - `version`:版本字符串(如 `"v1"`、`"v2"`)
439    /// - `path`:路由路径(如 `"/users"`,不含版本前缀)
440    /// - `method_router`:方法路由器
441    pub fn route(
442        mut self,
443        version: impl Into<String>,
444        path: impl Into<String>,
445        method_router: axum::routing::MethodRouter,
446    ) -> Self {
447        self.routes
448            .entry(version.into())
449            .or_default()
450            .push((path.into(), method_router));
451        self
452    }
453
454    /// 构建最终的 axum Router
455    pub fn build(self) -> axum::Router {
456        let mut router = axum::Router::new();
457        let prefix = self.url_prefix.unwrap_or_default();
458
459        for (version, routes) in self.routes {
460            for (path, method_router) in routes {
461                let full_path = if prefix.is_empty() {
462                    format!("/{}/{}", version, path.trim_start_matches('/'))
463                } else {
464                    format!("/{}/{}/{}", prefix, version, path.trim_start_matches('/'))
465                };
466                router = router.route(&full_path, method_router);
467            }
468        }
469
470        router
471    }
472}
473
474// ============================================================================
475// 测试
476// ============================================================================
477
478#[cfg(test)]
479mod tests {
480    use super::*;
481    use axum::body::Body;
482    use axum::http::{HeaderValue, Method};
483    use http_body_util::BodyExt;
484    use tower::ServiceExt;
485
486    // --------------------------------------------------------------------
487    // ApiVersion
488    // --------------------------------------------------------------------
489
490    #[test]
491    fn test_api_version_new() {
492        let v = ApiVersion::new(2);
493        assert_eq!(v.as_u32(), 2);
494    }
495
496    #[test]
497    fn test_api_version_default() {
498        let v = ApiVersion::default();
499        assert_eq!(v.as_u32(), 1);
500    }
501
502    #[test]
503    fn test_api_version_parse_pure_number() {
504        assert_eq!(ApiVersion::parse("1"), Some(ApiVersion::new(1)));
505        assert_eq!(ApiVersion::parse("42"), Some(ApiVersion::new(42)));
506    }
507
508    #[test]
509    fn test_api_version_parse_with_v_prefix() {
510        assert_eq!(ApiVersion::parse("v1"), Some(ApiVersion::new(1)));
511        assert_eq!(ApiVersion::parse("v2"), Some(ApiVersion::new(2)));
512    }
513
514    #[test]
515    fn test_api_version_parse_with_spaces() {
516        assert_eq!(ApiVersion::parse("  v1  "), Some(ApiVersion::new(1)));
517    }
518
519    #[test]
520    fn test_api_version_parse_version_equals() {
521        assert_eq!(ApiVersion::parse("version=2"), Some(ApiVersion::new(2)));
522    }
523
524    #[test]
525    fn test_api_version_parse_invalid() {
526        assert_eq!(ApiVersion::parse("abc"), None);
527        assert_eq!(ApiVersion::parse(""), None);
528        assert_eq!(ApiVersion::parse("v"), None);
529        assert_eq!(ApiVersion::parse("vabc"), None);
530    }
531
532    #[test]
533    fn test_api_version_to_vstring() {
534        assert_eq!(ApiVersion::new(1).to_vstring(), "v1");
535        assert_eq!(ApiVersion::new(10).to_vstring(), "v10");
536    }
537
538    #[test]
539    fn test_api_version_display() {
540        assert_eq!(format!("{}", ApiVersion::new(1)), "v1");
541    }
542
543    #[test]
544    fn test_api_version_equality() {
545        assert_eq!(ApiVersion::new(1), ApiVersion::new(1));
546        assert_ne!(ApiVersion::new(1), ApiVersion::new(2));
547    }
548
549    #[test]
550    fn test_api_version_ordering() {
551        assert!(ApiVersion::new(1) < ApiVersion::new(2));
552        assert!(ApiVersion::new(3) > ApiVersion::new(2));
553    }
554
555    // --------------------------------------------------------------------
556    // VersionNegotiator
557    // --------------------------------------------------------------------
558
559    #[test]
560    fn test_negotiator_default() {
561        let n = VersionNegotiator::default();
562        assert_eq!(n.default_version, ApiVersion::new(1));
563        assert_eq!(n.supported_versions, vec![ApiVersion::new(1)]);
564    }
565
566    #[test]
567    fn test_negotiate_url_path_with_prefix() {
568        let n = VersionNegotiator::new(ApiVersion::new(1))
569            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)]);
570        let uri = Uri::from_static("/api/v2/users");
571        let headers = HeaderMap::new();
572
573        let version = n.negotiate(&uri, &headers).unwrap();
574        assert_eq!(version, ApiVersion::new(2));
575    }
576
577    #[test]
578    fn test_negotiate_url_path_without_prefix() {
579        // 没配置 url_prefix 时,从第一段提取版本
580        let n = VersionNegotiator::new(ApiVersion::new(1))
581            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)])
582            .with_url_prefix("");
583        let uri = Uri::from_static("/v1/users");
584        let headers = HeaderMap::new();
585
586        let version = n.negotiate(&uri, &headers).unwrap();
587        assert_eq!(version, ApiVersion::new(1));
588    }
589
590    #[test]
591    fn test_negotiate_header_custom() {
592        let n = VersionNegotiator::new(ApiVersion::new(1))
593            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)])
594            .with_strategies(vec![VersionStrategy::Header]);
595        let uri = Uri::from_static("/users");
596        let mut headers = HeaderMap::new();
597        headers.insert("x-api-version", HeaderValue::from_static("2"));
598
599        let version = n.negotiate(&uri, &headers).unwrap();
600        assert_eq!(version, ApiVersion::new(2));
601    }
602
603    #[test]
604    fn test_negotiate_header_v_prefix() {
605        let n = VersionNegotiator::new(ApiVersion::new(1))
606            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)])
607            .with_strategies(vec![VersionStrategy::Header]);
608        let uri = Uri::from_static("/users");
609        let mut headers = HeaderMap::new();
610        headers.insert("x-api-version", HeaderValue::from_static("v2"));
611
612        let version = n.negotiate(&uri, &headers).unwrap();
613        assert_eq!(version, ApiVersion::new(2));
614    }
615
616    #[test]
617    fn test_negotiate_accept_header() {
618        let n = VersionNegotiator::new(ApiVersion::new(1))
619            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)])
620            .with_strategies(vec![VersionStrategy::AcceptHeader]);
621        let uri = Uri::from_static("/users");
622        let mut headers = HeaderMap::new();
623        headers.insert(
624            axum::http::header::ACCEPT,
625            HeaderValue::from_static("application/vnd.api+json; version=2"),
626        );
627
628        let version = n.negotiate(&uri, &headers).unwrap();
629        assert_eq!(version, ApiVersion::new(2));
630    }
631
632    #[test]
633    fn test_negotiate_accept_header_quoted() {
634        let n = VersionNegotiator::new(ApiVersion::new(1))
635            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(3)])
636            .with_strategies(vec![VersionStrategy::AcceptHeader]);
637        let uri = Uri::from_static("/users");
638        let mut headers = HeaderMap::new();
639        headers.insert(
640            axum::http::header::ACCEPT,
641            HeaderValue::from_static("application/json; version=\"3\""),
642        );
643
644        let version = n.negotiate(&uri, &headers).unwrap();
645        assert_eq!(version, ApiVersion::new(3));
646    }
647
648    #[test]
649    fn test_negotiate_query_param() {
650        let n = VersionNegotiator::new(ApiVersion::new(1))
651            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)])
652            .with_strategies(vec![VersionStrategy::Query]);
653        let uri = Uri::from_static("/users?api_version=2");
654        let headers = HeaderMap::new();
655
656        let version = n.negotiate(&uri, &headers).unwrap();
657        assert_eq!(version, ApiVersion::new(2));
658    }
659
660    #[test]
661    fn test_negotiate_custom_query_param_name() {
662        let n = VersionNegotiator::new(ApiVersion::new(1))
663            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)])
664            .with_strategies(vec![VersionStrategy::Query])
665            .with_query_param("ver");
666        let uri = Uri::from_static("/users?ver=2");
667        let headers = HeaderMap::new();
668
669        let version = n.negotiate(&uri, &headers).unwrap();
670        assert_eq!(version, ApiVersion::new(2));
671    }
672
673    #[test]
674    fn test_negotiate_default_when_no_match() {
675        let n = VersionNegotiator::new(ApiVersion::new(2))
676            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)]);
677        let uri = Uri::from_static("/users");
678        let headers = HeaderMap::new();
679
680        let version = n.negotiate(&uri, &headers).unwrap();
681        assert_eq!(version, ApiVersion::new(2));
682    }
683
684    #[test]
685    fn test_negotiate_strategy_priority() {
686        // URL 路径优先于 Header
687        let n = VersionNegotiator::new(ApiVersion::new(1))
688            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)])
689            .with_strategies(vec![VersionStrategy::UrlPath, VersionStrategy::Header]);
690
691        let uri = Uri::from_static("/api/v1/users");
692        let mut headers = HeaderMap::new();
693        headers.insert("x-api-version", HeaderValue::from_static("2"));
694
695        let version = n.negotiate(&uri, &headers).unwrap();
696        assert_eq!(version, ApiVersion::new(1)); // URL 路径优先
697    }
698
699    #[test]
700    fn test_negotiate_unsupported_version() {
701        let n = VersionNegotiator::new(ApiVersion::new(1))
702            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)]);
703        let uri = Uri::from_static("/api/v3/users");
704        let headers = HeaderMap::new();
705
706        let result = n.negotiate(&uri, &headers);
707        assert_eq!(
708            result,
709            Err(VersionError::UnsupportedVersion(ApiVersion::new(3)))
710        );
711    }
712
713    // --------------------------------------------------------------------
714    // VersionError
715    // --------------------------------------------------------------------
716
717    #[test]
718    fn test_version_error_display() {
719        let err = VersionError::UnsupportedVersion(ApiVersion::new(3));
720        assert_eq!(err.to_string(), "Unsupported API version: v3");
721    }
722
723    #[tokio::test]
724    async fn test_version_error_into_response() {
725        let err = VersionError::UnsupportedVersion(ApiVersion::new(99));
726        let response = err.into_response();
727
728        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
729        assert_eq!(
730            response.headers().get("content-type").unwrap(),
731            "application/json; charset=utf-8"
732        );
733
734        let bytes = response.into_body().collect().await.unwrap().to_bytes();
735        let body = String::from_utf8(bytes.to_vec()).unwrap();
736        assert!(body.contains("Unsupported API version"));
737        assert!(body.contains("v99"));
738    }
739
740    // --------------------------------------------------------------------
741    // 中间件集成测试
742    // --------------------------------------------------------------------
743
744    #[tokio::test]
745    async fn test_middleware_injects_version() {
746        let negotiator = VersionNegotiator::new(ApiVersion::new(1))
747            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)]);
748        let extractor = ApiVersionExtractor::new(negotiator);
749
750        async fn handler(req: Request) -> String {
751            let version = req.extensions().get::<ApiVersion>().unwrap();
752            format!("version={}", version.as_u32())
753        }
754
755        let app = axum::Router::new()
756            .route("/api/{*path}", axum::routing::get(handler))
757            .layer(axum::middleware::from_fn_with_state(
758                extractor,
759                version_negotiation_middleware,
760            ));
761
762        let req = Request::builder()
763            .method(Method::GET)
764            .uri("/api/v2/users")
765            .body(Body::empty())
766            .unwrap();
767
768        let response = app.oneshot(req).await.unwrap();
769        let bytes = response.into_body().collect().await.unwrap().to_bytes();
770        let body = String::from_utf8(bytes.to_vec()).unwrap();
771        assert_eq!(body, "version=2");
772    }
773
774    #[tokio::test]
775    async fn test_middleware_unsupported_version_returns_400() {
776        let negotiator = VersionNegotiator::new(ApiVersion::new(1))
777            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)]);
778        let extractor = ApiVersionExtractor::new(negotiator);
779
780        async fn handler(_: Request) -> &'static str {
781            "should not reach"
782        }
783
784        let app = axum::Router::new()
785            .route("/api/{*path}", axum::routing::get(handler))
786            .layer(axum::middleware::from_fn_with_state(
787                extractor,
788                version_negotiation_middleware,
789            ));
790
791        let req = Request::builder()
792            .method(Method::GET)
793            .uri("/api/v99/users")
794            .body(Body::empty())
795            .unwrap();
796
797        let response = app.oneshot(req).await.unwrap();
798        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
799    }
800
801    #[tokio::test]
802    async fn test_middleware_default_version_when_unspecified() {
803        let negotiator = VersionNegotiator::new(ApiVersion::new(2))
804            .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)]);
805        let extractor = ApiVersionExtractor::new(negotiator);
806
807        async fn handler(req: Request) -> String {
808            let version = req.extensions().get::<ApiVersion>().unwrap();
809            format!("version={}", version.as_u32())
810        }
811
812        let app = axum::Router::new()
813            .route("/users", axum::routing::get(handler))
814            .layer(axum::middleware::from_fn_with_state(
815                extractor,
816                version_negotiation_middleware,
817            ));
818
819        // 未指定版本 → 使用默认 v2
820        let req = Request::builder()
821            .method(Method::GET)
822            .uri("/users")
823            .body(Body::empty())
824            .unwrap();
825
826        let response = app.oneshot(req).await.unwrap();
827        let bytes = response.into_body().collect().await.unwrap().to_bytes();
828        let body = String::from_utf8(bytes.to_vec()).unwrap();
829        assert_eq!(body, "version=2");
830    }
831
832    // --------------------------------------------------------------------
833    // VersionedRouter
834    // --------------------------------------------------------------------
835
836    #[tokio::test]
837    async fn test_versioned_router_routes_to_correct_version() {
838        async fn v1_handler() -> &'static str {
839            "v1 response"
840        }
841        async fn v2_handler() -> &'static str {
842            "v2 response"
843        }
844
845        let router = VersionedRouter::new()
846            .with_url_prefix("api")
847            .route("v1", "/users", axum::routing::get(v1_handler))
848            .route("v2", "/users", axum::routing::get(v2_handler))
849            .build();
850
851        // v1 路由
852        let req = Request::builder()
853            .method(Method::GET)
854            .uri("/api/v1/users")
855            .body(Body::empty())
856            .unwrap();
857        let response = router.clone().oneshot(req).await.unwrap();
858        let bytes = response.into_body().collect().await.unwrap().to_bytes();
859        assert_eq!(String::from_utf8_lossy(&bytes), "v1 response");
860
861        // v2 路由
862        let req = Request::builder()
863            .method(Method::GET)
864            .uri("/api/v2/users")
865            .body(Body::empty())
866            .unwrap();
867        let response = router.oneshot(req).await.unwrap();
868        let bytes = response.into_body().collect().await.unwrap().to_bytes();
869        assert_eq!(String::from_utf8_lossy(&bytes), "v2 response");
870    }
871
872    #[tokio::test]
873    async fn test_versioned_router_without_prefix() {
874        async fn handler() -> &'static str {
875            "ok"
876        }
877
878        let router = VersionedRouter::new()
879            .route("v1", "/posts", axum::routing::get(handler))
880            .build();
881
882        let req = Request::builder()
883            .method(Method::GET)
884            .uri("/v1/posts")
885            .body(Body::empty())
886            .unwrap();
887        let response = router.oneshot(req).await.unwrap();
888        assert_eq!(response.status(), StatusCode::OK);
889    }
890
891    #[tokio::test]
892    async fn test_versioned_router_unregistered_path_returns_404() {
893        async fn handler() -> &'static str {
894            "ok"
895        }
896
897        let router = VersionedRouter::new()
898            .with_url_prefix("api")
899            .route("v1", "/users", axum::routing::get(handler))
900            .build();
901
902        // v3 路径未注册 → 404
903        let req = Request::builder()
904            .method(Method::GET)
905            .uri("/api/v3/users")
906            .body(Body::empty())
907            .unwrap();
908        let response = router.oneshot(req).await.unwrap();
909        assert_eq!(response.status(), StatusCode::NOT_FOUND);
910    }
911}