id_effect_rpc/
versioning.rs1use axum::extract::Request;
4use axum::http::{HeaderMap, StatusCode};
5use axum::middleware::Next;
6use axum::response::{IntoResponse, Response};
7use std::sync::Arc;
8
9#[derive(Clone, Debug, PartialEq, Eq, Hash)]
11pub struct ApiVersion(pub String);
12
13impl ApiVersion {
14 #[inline]
16 pub fn new(label: impl Into<String>) -> Self {
17 Self(label.into())
18 }
19}
20
21#[derive(Clone, Debug)]
23pub struct VersionConfig {
24 pub default: ApiVersion,
26 pub supported: Arc<[ApiVersion]>,
28}
29
30impl VersionConfig {
31 pub fn new(default: ApiVersion, supported: Vec<ApiVersion>) -> Self {
33 let mut supported = supported;
34 if !supported.iter().any(|v| v == &default) {
35 supported.push(default.clone());
36 }
37 Self {
38 default,
39 supported: supported.into(),
40 }
41 }
42
43 fn resolve(&self, headers: &HeaderMap, path: &str) -> Option<ApiVersion> {
44 if let Some(v) = extract_header_version(headers) {
45 if self.supported.iter().any(|s| s == &v) {
46 return Some(v);
47 }
48 return None;
49 }
50 if let Some(v) = extract_path_version(path) {
51 if self.supported.iter().any(|s| s == &v) {
52 return Some(v);
53 }
54 return None;
55 }
56 Some(self.default.clone())
57 }
58}
59
60pub const ACCEPT_VERSION: &str = "accept-version";
62
63pub const API_VERSION_HEADER: &str = "api-version";
65
66pub fn extract_header_version(headers: &HeaderMap) -> Option<ApiVersion> {
68 headers
69 .get(ACCEPT_VERSION)
70 .and_then(|v| v.to_str().ok())
71 .map(|s| ApiVersion(s.trim().to_string()))
72 .filter(|v| !v.0.is_empty())
73}
74
75pub fn extract_path_version(path: &str) -> Option<ApiVersion> {
77 let rest = path.strip_prefix('/')?;
78 let segment = rest.split('/').next()?;
79 if segment.starts_with('v')
80 && segment.len() > 1
81 && segment[1..].chars().all(|c| c.is_ascii_digit())
82 {
83 Some(ApiVersion(segment.to_string()))
84 } else {
85 None
86 }
87}
88
89pub fn set_response_version(headers: &mut HeaderMap, version: &ApiVersion) {
91 if let Ok(value) = version.0.parse() {
92 headers.insert(API_VERSION_HEADER, value);
93 }
94}
95
96pub async fn negotiate_api_version(
98 axum::extract::State(config): axum::extract::State<VersionConfig>,
99 mut request: Request,
100 next: Next,
101) -> Response {
102 let path = request.uri().path().to_string();
103 match config.resolve(request.headers(), &path) {
104 Some(version) => {
105 request.extensions_mut().insert(version);
106 next.run(request).await
107 }
108 None => (
109 StatusCode::NOT_ACCEPTABLE,
110 format!("unsupported API version; supported: {:?}", config.supported),
111 )
112 .into_response(),
113 }
114}
115
116#[cfg(test)]
117mod tests {
118 use super::*;
119 use axum::Router;
120 use axum::body::Body;
121 use axum::http::HeaderMap;
122 use axum::routing::get;
123 use tower::ServiceExt;
124
125 #[test]
126 fn extract_path_version_parses_v_prefix() {
127 assert_eq!(
128 extract_path_version("/v2/users"),
129 Some(ApiVersion("v2".into()))
130 );
131 assert_eq!(extract_path_version("/health"), None);
132 }
133
134 #[test]
135 fn extract_header_version_reads_accept_version() {
136 let mut headers = HeaderMap::new();
137 headers.insert(ACCEPT_VERSION, "v2".parse().unwrap());
138 assert_eq!(
139 extract_header_version(&headers),
140 Some(ApiVersion("v2".into()))
141 );
142 }
143
144 #[test]
145 fn version_config_includes_default_in_supported() {
146 let cfg = VersionConfig::new(ApiVersion("v1".into()), vec![ApiVersion("v2".into())]);
147 assert!(cfg.supported.iter().any(|v| v.0 == "v1"));
148 assert!(cfg.supported.iter().any(|v| v.0 == "v2"));
149 }
150
151 #[test]
152 fn set_response_version_inserts_header() {
153 let mut headers = HeaderMap::new();
154 set_response_version(&mut headers, &ApiVersion("v3".into()));
155 assert_eq!(headers.get(API_VERSION_HEADER).unwrap(), "v3");
156 }
157
158 #[tokio::test]
159 async fn middleware_returns_406_for_unsupported_header_version() {
160 let config = VersionConfig::new(ApiVersion("v1".into()), vec![ApiVersion("v1".into())]);
161 let app =
162 Router::new()
163 .route("/", get(|| async { "ok" }))
164 .layer(axum::middleware::from_fn_with_state(
165 config,
166 negotiate_api_version,
167 ));
168 let res = app
169 .oneshot(
170 Request::builder()
171 .uri("/")
172 .header(ACCEPT_VERSION, "v9")
173 .body(Body::empty())
174 .unwrap(),
175 )
176 .await
177 .unwrap();
178 assert_eq!(res.status(), StatusCode::NOT_ACCEPTABLE);
179 }
180
181 #[tokio::test]
182 async fn middleware_uses_path_version_when_present() {
183 let config = VersionConfig::new(
184 ApiVersion("v1".into()),
185 vec![ApiVersion("v1".into()), ApiVersion("v2".into())],
186 );
187 let app = Router::new()
188 .route(
189 "/v2/hello",
190 get(|ext: axum::Extension<ApiVersion>| async move { ext.0.0.clone() }),
191 )
192 .layer(axum::middleware::from_fn_with_state(
193 config,
194 negotiate_api_version,
195 ));
196 let res = app
197 .oneshot(
198 Request::builder()
199 .uri("/v2/hello")
200 .body(Body::empty())
201 .unwrap(),
202 )
203 .await
204 .unwrap();
205 assert_eq!(res.status(), StatusCode::OK);
206 }
207
208 #[tokio::test]
209 async fn middleware_defaults_when_unspecified() {
210 let config = VersionConfig::new(ApiVersion("v1".into()), vec![ApiVersion("v1".into())]);
211 let app = Router::new()
212 .route(
213 "/",
214 get(|ext: axum::Extension<ApiVersion>| async move { ext.0.0.clone() }),
215 )
216 .layer(axum::middleware::from_fn_with_state(
217 config,
218 negotiate_api_version,
219 ));
220
221 let res = app
222 .oneshot(Request::builder().uri("/").body(Body::empty()).unwrap())
223 .await
224 .unwrap();
225 assert_eq!(res.status(), StatusCode::OK);
226 }
227}