Skip to main content

id_effect_rpc/
versioning.rs

1//! API version negotiation for RPC-shaped Axum routes.
2
3use axum::extract::Request;
4use axum::http::{HeaderMap, StatusCode};
5use axum::middleware::Next;
6use axum::response::{IntoResponse, Response};
7use std::sync::Arc;
8
9/// Parsed API version label (e.g. `"v1"`).
10#[derive(Clone, Debug, PartialEq, Eq, Hash)]
11pub struct ApiVersion(pub String);
12
13impl ApiVersion {
14  /// Construct a version label.
15  #[inline]
16  pub fn new(label: impl Into<String>) -> Self {
17    Self(label.into())
18  }
19}
20
21/// Configuration for version routing middleware.
22#[derive(Clone, Debug)]
23pub struct VersionConfig {
24  /// Default version when the client omits a preference.
25  pub default: ApiVersion,
26  /// Supported version labels (must include `default`).
27  pub supported: Arc<[ApiVersion]>,
28}
29
30impl VersionConfig {
31  /// Build config ensuring `default` is listed in `supported`.
32  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
60/// Header name for explicit version negotiation.
61pub const ACCEPT_VERSION: &str = "accept-version";
62
63/// Response header echoing the negotiated version.
64pub const API_VERSION_HEADER: &str = "api-version";
65
66/// Extract version from `Accept-Version` header.
67pub 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
75/// Extract version from path prefix `/vN/`.
76pub 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
89/// Insert `api-version` on successful responses.
90pub 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
96/// Middleware: attach resolved [`ApiVersion`] as request extension; 406 when unsupported.
97pub 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}