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.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 #[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 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 #[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 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 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 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}