use axum::{
extract::Request,
http::{HeaderMap, HeaderValue, StatusCode},
middleware::Next,
response::{IntoResponse, Response},
};
use super::version::{ApiVersion, CURRENT_VERSION, SUPPORTED_VERSIONS};
#[derive(Debug, Clone, Copy)]
pub struct ResolvedVersion(pub ApiVersion);
pub async fn version_negotiation_middleware(
mut request: Request,
next: Next,
) -> Response {
let version = extract_version_from_path(request.uri().path())
.or_else(|| extract_version_from_header(request.headers()))
.or_else(|| extract_version_from_query(request.uri().query()))
.unwrap_or(CURRENT_VERSION);
request.extensions_mut().insert(ResolvedVersion(version));
let mut response = next.run(request).await;
if version.is_deprecated() {
if let Some(msg) = version.deprecation_message() {
response.headers_mut().insert(
"Deprecation",
HeaderValue::from_str(msg).unwrap_or_else(|_| HeaderValue::from_static("true")),
);
}
if let Some(date) = version.sunset_date() {
if let Ok(val) = HeaderValue::from_str(date) {
response.headers_mut().insert("Sunset", val);
}
}
}
response.headers_mut().insert(
"X-API-Version",
HeaderValue::from_str(version.as_str()).unwrap_or_else(|_| HeaderValue::from_static("v1")),
);
response
}
fn extract_version_from_path(path: &str) -> Option<ApiVersion> {
for version in SUPPORTED_VERSIONS {
let version_segment = format!("/{}/", version.as_str());
if path.contains(&version_segment) {
return Some(*version);
}
let version_end = format!("/{}", version.as_str());
if path.ends_with(&version_end) {
return Some(*version);
}
}
None
}
fn extract_version_from_header(headers: &HeaderMap) -> Option<ApiVersion> {
if let Some(value) = headers.get("X-API-Version") {
if let Ok(s) = value.to_str() {
if let Ok(v) = s.parse() {
return Some(v);
}
}
}
if let Some(accept) = headers.get("Accept") {
if let Ok(s) = accept.to_str() {
if let Some(version_part) = s.split(';').find(|p| p.trim().starts_with("version=")) {
let version_str = version_part.trim().trim_start_matches("version=");
if let Ok(v) = version_str.parse() {
return Some(v);
}
}
}
}
None
}
fn extract_version_from_query(query: Option<&str>) -> Option<ApiVersion> {
let query = query?;
for param in query.split('&') {
if let Some(value) = param.strip_prefix("version=") {
if let Ok(v) = value.parse() {
return Some(v);
}
}
}
None
}