use std::sync::Arc;
use axum::extract::{Request, State};
use axum::http::{StatusCode, header};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
#[derive(Debug, Clone)]
pub(crate) struct Rebinding {
hosts: Arc<[String]>,
origins: Arc<[String]>,
origin_on_writes: bool,
}
impl Rebinding {
pub(crate) fn new(hosts: Vec<String>, origins: Vec<String>) -> Self {
Self {
hosts: hosts.into(),
origins: origins.into(),
origin_on_writes: false,
}
}
#[cfg(feature = "dev")]
pub(crate) const fn origin_on_writes(mut self) -> Self {
self.origin_on_writes = true;
self
}
}
pub(crate) async fn refuse_rebinding(
State(guard): State<Rebinding>,
request: Request,
next: Next,
) -> Response {
let host = request
.headers()
.get(header::HOST)
.and_then(|h| h.to_str().ok())
.map(str::to_owned)
.or_else(|| request.uri().authority().map(|a| a.as_str().to_owned()));
if !host
.as_deref()
.is_some_and(|h| host_allowed(h, &guard.hosts))
{
return (
StatusCode::FORBIDDEN,
"Forbidden: Host header is not allowed",
)
.into_response();
}
let origin = request.headers().get(header::ORIGIN);
if let Some(origin) = origin
&& !origin
.to_str()
.is_ok_and(|o| origin_allowed(o, &guard.origins))
{
return (
StatusCode::FORBIDDEN,
"Forbidden: Origin header is not allowed",
)
.into_response();
}
if guard.origin_on_writes && origin.is_none() && !request.method().is_safe() {
return (
StatusCode::FORBIDDEN,
"Forbidden: a request that changes state must name its Origin",
)
.into_response();
}
next.run(request).await
}
fn authority(value: &str) -> Option<(String, Option<u16>)> {
if let Ok(ip) = value.trim().parse::<std::net::Ipv6Addr>() {
return Some((ip.to_string(), None));
}
let parsed = value.trim().parse::<axum::http::uri::Authority>().ok()?;
let host = parsed
.host()
.trim_start_matches('[')
.trim_end_matches(']')
.to_ascii_lowercase();
Some((host, parsed.port_u16()))
}
pub(crate) fn host_allowed(value: &str, allowed: &[String]) -> bool {
let Some((host, port)) = authority(value) else {
return false;
};
allowed
.iter()
.filter_map(|a| authority(a))
.any(|(h, p)| h == host && p.is_none_or(|p| Some(p) == port))
}
fn origin(value: &str) -> Option<(String, String, Option<u16>)> {
let value = value.trim();
if value == "null" {
return Some(("null".to_owned(), String::new(), None));
}
let uri = value.parse::<axum::http::Uri>().ok()?;
let scheme = uri.scheme_str()?.to_ascii_lowercase();
let (host, port) = authority(uri.authority()?.as_str())?;
let port = port.or(match scheme.as_str() {
"http" => Some(80),
"https" => Some(443),
_ => None,
});
Some((scheme, host, port))
}
pub(crate) fn origin_allowed(value: &str, allowed: &[String]) -> bool {
origin(value).is_some_and(|o| allowed.iter().filter_map(|a| origin(a)).any(|a| a == o))
}