use std::net::SocketAddr;
use axum::http::{HeaderMap, HeaderName, Uri, header};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SelfAuthorities(Vec<String>);
impl SelfAuthorities {
pub fn new(listen: SocketAddr, magic_dns: Option<&str>) -> Self {
let mut authorities = vec![listen.to_string()];
if let Some(name) = magic_dns {
authorities.push(format!("{name}:{}", listen.port()));
}
Self(authorities)
}
fn contains(&self, authority: &str) -> bool {
self.0.iter().any(|a| a.eq_ignore_ascii_case(authority))
}
}
pub fn check_target(
headers: &HeaderMap,
uri: &Uri,
allowed: &SelfAuthorities,
) -> Result<(), &'static str> {
let uri_authority = uri.authority().map(|a| a.as_str());
let host = match single_header(headers, &header::HOST)? {
Some(host) => {
if uri_authority.is_some_and(|a| !a.eq_ignore_ascii_case(host)) {
return Err("request target and Host disagree");
}
host
}
None => uri_authority.ok_or("no Host")?,
};
if !allowed.contains(host) {
return Err("Host is not this listener");
}
if let Some(origin) = single_header(headers, &header::ORIGIN)?
&& !origin.eq_ignore_ascii_case(&format!("http://{host}"))
{
return Err("Origin is not this listener");
}
Ok(())
}
fn single_header<'a>(
headers: &'a HeaderMap,
name: &HeaderName,
) -> Result<Option<&'a str>, &'static str> {
let mut values = headers.get_all(name).iter();
let Some(first) = values.next() else {
return Ok(None);
};
if values.next().is_some() {
return Err("repeated Host or Origin header");
}
first
.to_str()
.map(Some)
.map_err(|_| "unreadable Host or Origin header")
}