use axum::http::header::{HeaderMap, HeaderValue};
use axum::http::{Method, StatusCode};
use axum::response::{IntoResponse, Response};
use super::dto::{new_hex_id, ApiError, ErrorCode};
const REJECTED_TOKEN_PARAMS: [&str; 4] = ["token", "access_token", "auth", "bearer"];
#[derive(Debug, Clone)]
pub struct CorrelationId(pub String);
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct NormalizedOrigin {
pub scheme: String,
pub host: String,
pub port: u16,
}
pub fn normalize_origin(value: &str) -> Option<NormalizedOrigin> {
let value = value.trim();
if value.is_empty() || value.contains('*') {
return None;
}
let (scheme, rest) = value.split_once("://")?;
let scheme = scheme.to_ascii_lowercase();
let default_port = match scheme.as_str() {
"http" => 80,
"https" => 443,
_ => return None,
};
if rest.contains('/') || rest.contains('?') || rest.contains('#') || rest.is_empty() {
return None;
}
let (host, port) = if let Some(inner) = rest.strip_prefix('[') {
let (host, tail) = inner.split_once(']')?;
let port = match tail {
"" => default_port,
other => other.strip_prefix(':')?.parse().ok()?,
};
(host.to_ascii_lowercase(), port)
} else if let Some((host, port)) = rest.rsplit_once(':') {
(host.to_ascii_lowercase(), port.parse().ok()?)
} else {
(rest.to_ascii_lowercase(), default_port)
};
if host.is_empty() {
return None;
}
Some(NormalizedOrigin { scheme, host, port })
}
#[derive(Debug, Clone, Default)]
pub struct RemoteControlAuth {
pub token: Option<String>,
pub allowed_origins: Vec<NormalizedOrigin>,
}
impl RemoteControlAuth {
pub fn new(token: Option<String>, origins: &[String]) -> Result<Self, String> {
let mut allowed_origins = Vec::with_capacity(origins.len());
for origin in origins {
match normalize_origin(origin) {
Some(normalized) => allowed_origins.push(normalized),
None => {
return Err(format!(
"invalid allowed origin '{origin}': expected an exact \
http(s)://host[:port] value with no wildcard or path"
))
}
}
}
Ok(Self {
token: token.filter(|t| !t.is_empty()),
allowed_origins,
})
}
pub fn is_enforced(&self) -> bool {
self.token.is_some()
}
pub fn check_bearer(&self, headers: &HeaderMap, correlation_id: &str) -> Result<(), ApiError> {
let Some(expected) = self.token.as_deref() else {
return Ok(());
};
let provided = headers
.get(axum::http::header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.strip_prefix("Bearer "))
.unwrap_or("");
if provided.is_empty() || !constant_time_eq(provided.as_bytes(), expected.as_bytes()) {
return Err(ApiError::new(
ErrorCode::Unauthorized,
"missing or invalid Authorization: Bearer credentials",
correlation_id,
));
}
Ok(())
}
pub fn reject_out_of_band_credentials(
&self,
query: Option<&str>,
headers: &HeaderMap,
correlation_id: &str,
) -> Result<(), ApiError> {
if let Some(query) = query {
for pair in query.split('&') {
let name = pair.split('=').next().unwrap_or("").to_ascii_lowercase();
if REJECTED_TOKEN_PARAMS.contains(&name.as_str()) {
return Err(ApiError::new(
ErrorCode::Unauthorized,
"credentials in query parameters are not accepted; use Authorization: Bearer",
correlation_id,
));
}
}
}
if headers.contains_key("sec-websocket-protocol") {
return Err(ApiError::new(
ErrorCode::Unauthorized,
"credentials in Sec-WebSocket-Protocol are not accepted; use Authorization: Bearer",
correlation_id,
));
}
Ok(())
}
pub fn check_origin(
&self,
headers: &HeaderMap,
correlation_id: &str,
) -> Result<Option<String>, ApiError> {
let Some(origin_header) = headers
.get(axum::http::header::ORIGIN)
.and_then(|value| value.to_str().ok())
else {
return Ok(None);
};
let deny = || {
ApiError::new(
ErrorCode::Forbidden,
"origin is not allowed for /api/v2; configure an exact allowed origin",
correlation_id,
)
};
let origin = normalize_origin(origin_header).ok_or_else(deny)?;
if let Some(direct) = direct_origin(headers) {
if direct == origin {
return Ok(Some(origin_header.to_string()));
}
}
if self.allowed_origins.contains(&origin) {
return Ok(Some(origin_header.to_string()));
}
Err(deny())
}
}
fn direct_origin(headers: &HeaderMap) -> Option<NormalizedOrigin> {
let host = headers
.get(axum::http::header::HOST)
.and_then(|value| value.to_str().ok())?;
normalize_origin(&format!("http://{host}"))
}
fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
if a.len() != b.len() {
return false;
}
a.iter().zip(b).fold(0u8, |acc, (x, y)| acc | (x ^ y)) == 0
}
pub fn resolve_correlation_id(headers: &HeaderMap) -> Result<String, ApiError> {
match headers
.get("x-correlation-id")
.and_then(|value| value.to_str().ok())
{
None => Ok(new_hex_id()),
Some(value) if super::dto::is_valid_correlation_id(value) => Ok(value.to_string()),
Some(_) => Err(ApiError::new(
ErrorCode::ValidationFailed,
"correlation_id must be 1-64 characters matching [A-Za-z0-9._:-]",
&new_hex_id(),
)),
}
}
pub fn cors_headers(allowed_origin: Option<&str>) -> Vec<(axum::http::HeaderName, HeaderValue)> {
let Some(origin) = allowed_origin else {
return Vec::new();
};
let Ok(origin_value) = HeaderValue::from_str(origin) else {
return Vec::new();
};
vec![
(
axum::http::header::ACCESS_CONTROL_ALLOW_ORIGIN,
origin_value,
),
(axum::http::header::VARY, HeaderValue::from_static("Origin")),
(
axum::http::header::ACCESS_CONTROL_ALLOW_METHODS,
HeaderValue::from_static("GET, POST, OPTIONS"),
),
(
axum::http::header::ACCESS_CONTROL_ALLOW_HEADERS,
HeaderValue::from_static("Authorization, Content-Type, X-Correlation-Id"),
),
]
}
pub fn preflight_response(allowed_origin: Option<&str>) -> Response {
let mut response = StatusCode::NO_CONTENT.into_response();
for (name, value) in cors_headers(allowed_origin) {
response.headers_mut().insert(name, value);
}
response
}
pub fn is_preflight(method: &Method) -> bool {
method == Method::OPTIONS
}