1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
59pub struct ApiVersion(u32);
60
61impl ApiVersion {
62 pub const fn new(version: u32) -> Self {
64 Self(version)
65 }
66
67 pub fn as_u32(&self) -> u32 {
69 self.0
70 }
71
72 pub fn parse(s: &str) -> Option<Self> {
79 let trimmed = s.trim();
80 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 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
115pub enum VersionStrategy {
116 #[default]
120 UrlPath,
121 Header,
123 AcceptHeader,
125 Query,
127}
128
129#[derive(Debug, Clone)]
137pub struct VersionNegotiator {
138 supported_versions: Vec<ApiVersion>,
140 default_version: ApiVersion,
142 strategies: Vec<VersionStrategy>,
144 url_prefix: Option<String>,
146 query_param_name: String,
148 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 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 pub fn with_supported_versions(mut self, versions: Vec<ApiVersion>) -> Self {
182 self.supported_versions = versions;
183 self
184 }
185
186 pub fn with_strategies(mut self, strategies: Vec<VersionStrategy>) -> Self {
188 self.strategies = strategies;
189 self
190 }
191
192 pub fn with_url_prefix(mut self, prefix: impl Into<String>) -> Self {
194 self.url_prefix = Some(prefix.into());
195 self
196 }
197
198 pub fn with_query_param(mut self, name: impl Into<String>) -> Self {
200 self.query_param_name = name.into();
201 self
202 }
203
204 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 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 if !self.supported_versions.contains(&version) {
231 return Err(VersionError::UnsupportedVersion(version));
232 }
233 return Ok(version);
234 }
235 }
236
237 Ok(self.default_version)
239 }
240
241 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 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 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 fn extract_from_accept_header(&self, headers: &HeaderMap) -> Option<ApiVersion> {
274 let accept = headers.get(axum::http::header::ACCEPT)?.to_str().ok()?;
275 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 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#[derive(Debug, Clone, PartialEq, Eq)]
304pub enum VersionError {
305 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#[derive(Clone)]
351pub struct ApiVersionExtractor {
352 negotiator: Arc<VersionNegotiator>,
353}
354
355impl ApiVersionExtractor {
356 pub fn new(negotiator: VersionNegotiator) -> Self {
358 Self {
359 negotiator: Arc::new(negotiator),
360 }
361 }
362
363 pub fn negotiator(&self) -> &VersionNegotiator {
365 &self.negotiator
366 }
367}
368
369pub 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#[derive(Default)]
415pub struct VersionedRouter {
416 routes: HashMap<String, Vec<(String, axum::routing::MethodRouter)>>,
418 url_prefix: Option<String>,
420}
421
422impl VersionedRouter {
423 pub fn new() -> Self {
425 Self::default()
426 }
427
428 pub fn with_url_prefix(mut self, prefix: impl Into<String>) -> Self {
430 self.url_prefix = Some(prefix.into());
431 self
432 }
433
434 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 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#[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 #[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 #[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 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 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)); }
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 #[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
735 .into_body()
736 .collect()
737 .await
738 .expect("响应体读取失败")
739 .to_bytes();
740 let body = String::from_utf8(bytes.to_vec()).unwrap();
741 assert!(body.contains("Unsupported API version"));
742 assert!(body.contains("v99"));
743 }
744
745 #[tokio::test]
750 async fn test_middleware_injects_version() {
751 let negotiator = VersionNegotiator::new(ApiVersion::new(1))
752 .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)]);
753 let extractor = ApiVersionExtractor::new(negotiator);
754
755 async fn handler(req: Request) -> String {
756 let version = req.extensions().get::<ApiVersion>().unwrap();
757 format!("version={}", version.as_u32())
758 }
759
760 let app = axum::Router::new()
761 .route("/api/{*path}", axum::routing::get(handler))
762 .layer(axum::middleware::from_fn_with_state(
763 extractor,
764 version_negotiation_middleware,
765 ));
766
767 let req = Request::builder()
768 .method(Method::GET)
769 .uri("/api/v2/users")
770 .body(Body::empty())
771 .expect("测试请求构造失败");
772
773 let response = app.oneshot(req).await.expect("测试请求执行失败");
774 let bytes = response
775 .into_body()
776 .collect()
777 .await
778 .expect("响应体读取失败")
779 .to_bytes();
780 let body = String::from_utf8(bytes.to_vec()).unwrap();
781 assert_eq!(body, "version=2");
782 }
783
784 #[tokio::test]
785 async fn test_middleware_unsupported_version_returns_400() {
786 let negotiator = VersionNegotiator::new(ApiVersion::new(1))
787 .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)]);
788 let extractor = ApiVersionExtractor::new(negotiator);
789
790 async fn handler(_: Request) -> &'static str {
791 "should not reach"
792 }
793
794 let app = axum::Router::new()
795 .route("/api/{*path}", axum::routing::get(handler))
796 .layer(axum::middleware::from_fn_with_state(
797 extractor,
798 version_negotiation_middleware,
799 ));
800
801 let req = Request::builder()
802 .method(Method::GET)
803 .uri("/api/v99/users")
804 .body(Body::empty())
805 .expect("测试请求构造失败");
806
807 let response = app.oneshot(req).await.expect("测试请求执行失败");
808 assert_eq!(response.status(), StatusCode::BAD_REQUEST);
809 }
810
811 #[tokio::test]
812 async fn test_middleware_default_version_when_unspecified() {
813 let negotiator = VersionNegotiator::new(ApiVersion::new(2))
814 .with_supported_versions(vec![ApiVersion::new(1), ApiVersion::new(2)]);
815 let extractor = ApiVersionExtractor::new(negotiator);
816
817 async fn handler(req: Request) -> String {
818 let version = req.extensions().get::<ApiVersion>().unwrap();
819 format!("version={}", version.as_u32())
820 }
821
822 let app = axum::Router::new()
823 .route("/users", axum::routing::get(handler))
824 .layer(axum::middleware::from_fn_with_state(
825 extractor,
826 version_negotiation_middleware,
827 ));
828
829 let req = Request::builder()
831 .method(Method::GET)
832 .uri("/users")
833 .body(Body::empty())
834 .expect("测试请求构造失败");
835
836 let response = app.oneshot(req).await.expect("测试请求执行失败");
837 let bytes = response
838 .into_body()
839 .collect()
840 .await
841 .expect("响应体读取失败")
842 .to_bytes();
843 let body = String::from_utf8(bytes.to_vec()).unwrap();
844 assert_eq!(body, "version=2");
845 }
846
847 #[tokio::test]
852 async fn test_versioned_router_routes_to_correct_version() {
853 async fn v1_handler() -> &'static str {
854 "v1 response"
855 }
856 async fn v2_handler() -> &'static str {
857 "v2 response"
858 }
859
860 let router = VersionedRouter::new()
861 .with_url_prefix("api")
862 .route("v1", "/users", axum::routing::get(v1_handler))
863 .route("v2", "/users", axum::routing::get(v2_handler))
864 .build();
865
866 let req = Request::builder()
868 .method(Method::GET)
869 .uri("/api/v1/users")
870 .body(Body::empty())
871 .expect("测试请求构造失败");
872 let response = router.clone().oneshot(req).await.expect("测试请求执行失败");
873 let bytes = response
874 .into_body()
875 .collect()
876 .await
877 .expect("响应体读取失败")
878 .to_bytes();
879 assert_eq!(String::from_utf8_lossy(&bytes), "v1 response");
880
881 let req = Request::builder()
883 .method(Method::GET)
884 .uri("/api/v2/users")
885 .body(Body::empty())
886 .expect("测试请求构造失败");
887 let response = router.oneshot(req).await.expect("测试请求执行失败");
888 let bytes = response
889 .into_body()
890 .collect()
891 .await
892 .expect("响应体读取失败")
893 .to_bytes();
894 assert_eq!(String::from_utf8_lossy(&bytes), "v2 response");
895 }
896
897 #[tokio::test]
898 async fn test_versioned_router_without_prefix() {
899 async fn handler() -> &'static str {
900 "ok"
901 }
902
903 let router = VersionedRouter::new()
904 .route("v1", "/posts", axum::routing::get(handler))
905 .build();
906
907 let req = Request::builder()
908 .method(Method::GET)
909 .uri("/v1/posts")
910 .body(Body::empty())
911 .expect("测试请求构造失败");
912 let response = router.oneshot(req).await.expect("测试请求执行失败");
913 assert_eq!(response.status(), StatusCode::OK);
914 }
915
916 #[tokio::test]
917 async fn test_versioned_router_unregistered_path_returns_404() {
918 async fn handler() -> &'static str {
919 "ok"
920 }
921
922 let router = VersionedRouter::new()
923 .with_url_prefix("api")
924 .route("v1", "/users", axum::routing::get(handler))
925 .build();
926
927 let req = Request::builder()
929 .method(Method::GET)
930 .uri("/api/v3/users")
931 .body(Body::empty())
932 .expect("测试请求构造失败");
933 let response = router.oneshot(req).await.expect("测试请求执行失败");
934 assert_eq!(response.status(), StatusCode::NOT_FOUND);
935 }
936}