use std::collections::BTreeMap;
use std::future::Future;
use std::net::Ipv4Addr;
use std::pin::Pin;
use std::sync::Arc;
pub use server::Broker;
use crate::log::Logger;
use crate::log::fields;
pub fn host_allowed(host: &str, allow: &[String]) -> bool {
let candidate = host.trim().to_lowercase();
if candidate.is_empty() {
return false;
}
for rule in allow {
if rule == "*" {
return true;
}
if let Some(apex) = rule.strip_prefix("*.") {
let suffix = format!(".{}", apex.to_lowercase());
if candidate.len() > suffix.len() && candidate.ends_with(&suffix) {
return true;
}
} else if candidate == rule.to_lowercase() {
return true;
}
}
false
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ConnectTarget {
pub host: String,
pub port: u16,
}
pub fn parse_connect(request_line: &str) -> Option<ConnectTarget> {
let parts: Vec<&str> = request_line.split_whitespace().collect();
if parts.len() < 2 || !parts[0].eq_ignore_ascii_case("CONNECT") {
return None;
}
let authority = parts[1];
let colon = authority.rfind(':')?;
if colon == 0 || colon == authority.len() - 1 {
return None;
}
let host = &authority[..colon];
let port: i64 = authority[colon + 1..].parse().ok()?;
if !(1..=65_535).contains(&port) {
return None;
}
if host.contains('/') || host.contains('[') {
return None;
}
Some(ConnectTarget {
host: host.to_owned(),
#[expect(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
port: port as u16,
})
}
pub const ALLOWED_UPSTREAM_PORTS: [u16; 1] = [443];
#[derive(Debug, Clone)]
pub struct ProviderRoute {
pub prefix: String,
pub upstream: String,
pub nonce: String,
pub credential: String,
}
const HOP_BY_HOP: [&str; 9] = [
"connection",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"te",
"trailer",
"transfer-encoding",
"upgrade",
"host",
];
fn same_secret(a: &str, b: &str) -> bool {
if a.len() != b.len() {
return false;
}
let differences = a
.bytes()
.zip(b.bytes())
.fold(0_u8, |accumulated, (left, right)| {
accumulated | (left ^ right)
});
differences == 0
}
pub fn is_private_v4(address: &str) -> bool {
let Ok(address) = address.parse::<Ipv4Addr>() else {
return true;
};
let [first, second, third, _] = address.octets();
address.is_loopback()
|| address.is_private()
|| address.is_link_local()
|| address.is_multicast()
|| address.is_documentation()
|| first == 0
|| (first == 100 && (64..=127).contains(&second))
|| (first == 192 && second == 0 && third == 0)
|| (first == 192 && second == 88 && third == 99)
|| (first == 198 && (18..=19).contains(&second))
|| first >= 240
}
pub fn v6_groups(address: &str) -> Option<Vec<u32>> {
let bare = address
.to_lowercase()
.split('%')
.next()
.unwrap_or("")
.trim()
.to_owned();
if bare.is_empty() {
return None;
}
let mut text = bare.clone();
if let Some((head, quad)) = bare.rsplit_once(':') {
let octets: Vec<&str> = quad.split('.').collect();
if octets.len() == 4
&& octets
.iter()
.all(|part| part.parse::<i64>().is_ok_and(|n| (0..=255).contains(&n)))
{
let numbers: Vec<i64> = octets
.iter()
.map(|part| part.parse().expect("checked above"))
.collect();
let high = (numbers[0] << 8) | numbers[1];
let low = (numbers[2] << 8) | numbers[3];
text = format!("{head}:{high:x}:{low:x}");
}
}
let halves: Vec<&str> = text.split("::").collect();
if halves.len() > 2 {
return None;
}
let read = |part: &str| -> Option<Vec<u32>> {
if part.is_empty() {
return Some(Vec::new());
}
let mut groups = Vec::new();
for piece in part.split(':') {
let valid = (1..=4).contains(&piece.len())
&& piece.bytes().all(|byte| byte.is_ascii_hexdigit());
if !valid {
return None;
}
groups.push(u32::from_str_radix(piece, 16).expect("hex digits"));
}
Some(groups)
};
let head = read(halves[0])?;
if halves.len() == 1 {
return (head.len() == 8).then_some(head);
}
let tail = read(halves[1])?;
let gap = 8_usize.checked_sub(head.len() + tail.len())?;
if gap < 1 {
return None;
}
let mut groups = head;
groups.extend(std::iter::repeat_n(0, gap));
groups.extend(tail);
Some(groups)
}
#[expect(clippy::many_single_char_names)]
pub fn is_private_v6(address: &str) -> bool {
let Some(groups) = v6_groups(address) else {
return true;
};
if groups.len() != 8 {
return true;
}
let [a, b, c, d, e, f, g, h] = groups[..] else {
return true;
};
let leading = a | b | c | d | e;
let carries_v4 = (leading == 0 && f == 0xffff)
|| (leading == 0 && f == 0 && !(g == 0 && (h == 0 || h == 1)))
|| (a == 0x0064 && b == 0xff9b);
if carries_v4 {
let v4 = format!("{}.{}.{}.{}", g >> 8, g & 0xff, h >> 8, h & 0xff);
return is_private_v4(&v4);
}
if leading == 0 && f == 0 && g == 0 && (h == 0 || h == 1) {
return true; }
if (a & 0xffc0) == 0xfe80 {
return true; }
if (a & 0xfe00) == 0xfc00 {
return true; }
if (a & 0xff00) == 0xff00 {
return true; }
false
}
pub fn is_private_address(address: &str) -> bool {
if address.contains(':') {
is_private_v6(address)
} else {
is_private_v4(address)
}
}
pub type Resolve = Arc<dyn Fn(String) -> ResolveFuture + Send + Sync>;
pub async fn public_address(
host: &str,
allow_internal: bool,
resolve: Option<&Resolve>,
) -> Option<String> {
let literal = host.split('.').count() == 4 && host.parse::<std::net::Ipv4Addr>().is_ok()
|| host.contains(':');
if literal {
if !is_private_address(host) {
return Some(host.to_owned());
}
return allow_internal.then(|| host.to_owned());
}
let addresses = match resolve {
Some(resolve) => resolve(host.to_owned()).await,
None => default_resolve(host).await,
};
addresses
.into_iter()
.find(|address| !is_private_address(address))
}
async fn default_resolve(name: &str) -> Vec<String> {
match tokio::net::lookup_host((name, 0_u16)).await {
Ok(addrs) => addrs.map(|addr| addr.ip().to_string()).collect(),
Err(_) => Vec::new(),
}
}
pub(crate) struct ProviderState {
pub routes: Vec<ProviderRoute>,
pub allow: Vec<String>,
pub allow_internal: bool,
pub log: Logger,
pub resolve: Option<Resolve>,
pub provider_port: u16,
pub client: reqwest::Client,
}
fn route_for<'a>(state: &'a ProviderState, path: &str) -> Option<&'a ProviderRoute> {
let mut routes = state.routes.iter().collect::<Vec<_>>();
routes.sort_by_key(|route| std::cmp::Reverse(route.prefix.len()));
routes
.into_iter()
.find(|candidate| path.starts_with(&candidate.prefix))
}
pub(crate) async fn serve_provider_request(
state: Arc<ProviderState>,
request: http::Request<axum::body::Body>,
) -> axum::response::Response {
use axum::response::IntoResponse;
let path = request.uri().path().to_owned();
let query = request
.uri()
.query()
.map(|query| format!("?{query}"))
.unwrap_or_default();
let Some(route) = route_for(&state, &path) else {
return (
axum::http::StatusCode::NOT_FOUND,
"not a provider this broker serves\n",
)
.into_response();
};
let offered = request
.headers()
.get(http::header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.unwrap_or("")
.to_owned();
let expected = format!("Bearer {}", route.nonce);
if !same_secret(&offered, &expected) {
state.log.warn(
"a provider call arrived without this session's key",
&BTreeMap::new(),
);
return (axum::http::StatusCode::UNAUTHORIZED, "not this session\n").into_response();
}
let rest = path[route.prefix.len()..].to_owned();
let trimmed_upstream = route.upstream.strip_suffix('/').unwrap_or(&route.upstream);
let target = format!("{trimmed_upstream}{rest}{query}");
let hostname = url_host(&target);
let internal = match hostname {
Some(hostname) => public_address(&hostname, state.allow_internal, None)
.await
.is_none(),
None => true,
};
if internal {
state.log.warn(
"a provider is configured at a host-internal address",
&fields([("provider", route.prefix.as_str().into())]),
);
return (
axum::http::StatusCode::BAD_GATEWAY,
"the provider is not at a reachable address\n",
)
.into_response();
}
let forwarded = build_forwarded(&state, request, route, &target);
match forwarded.send().await {
Ok(answered) => stream_answer(answered),
Err(error) => {
state.log.warn(
"the provider could not be reached",
&fields([("detail", error.to_string().into())]),
);
(
axum::http::StatusCode::BAD_GATEWAY,
"the provider could not be reached\n",
)
.into_response()
}
}
}
fn build_forwarded(
state: &ProviderState,
request: http::Request<axum::body::Body>,
route: &ProviderRoute,
target: &str,
) -> reqwest::RequestBuilder {
let mut forwarded = state.client.request(
reqwest::Method::from_bytes(request.method().as_str().as_bytes())
.unwrap_or(reqwest::Method::GET),
target,
);
for (name, value) in request.headers() {
if name == http::header::AUTHORIZATION {
continue;
}
if !HOP_BY_HOP.contains(&name.as_str().to_lowercase().as_str()) {
forwarded = forwarded.header(name.as_str(), value.to_str().unwrap_or(""));
}
}
forwarded = forwarded.header(
http::header::AUTHORIZATION,
format!("Bearer {}", route.credential),
);
let body = request.into_body();
forwarded.body(reqwest::Body::wrap_stream(body.into_data_stream()))
}
fn stream_answer(answered: reqwest::Response) -> axum::response::Response {
use axum::response::IntoResponse;
use futures_util::TryStreamExt;
let mut builder = axum::http::Response::builder().status(answered.status().as_u16());
for (name, value) in answered.headers() {
if !HOP_BY_HOP.contains(&name.as_str().to_lowercase().as_str())
&& let Ok(value) = value.to_str()
{
builder = builder.header(name.as_str(), value);
}
}
let stream = answered
.bytes_stream()
.map_err(|error| std::io::Error::other(error.to_string()));
builder
.body(axum::body::Body::from_stream(stream))
.unwrap_or_else(|error| {
(axum::http::StatusCode::BAD_GATEWAY, error.to_string()).into_response()
})
}
pub(crate) type ResolveFuture = Pin<Box<dyn Future<Output = Vec<String>> + Send>>;
fn url_host(target: &str) -> Option<String> {
let rest = target
.strip_prefix("http://")
.or_else(|| target.strip_prefix("https://"))?;
let authority = rest.split(['/', '?']).next()?;
let host = authority
.rsplit_once(':')
.map_or(authority, |(host, _)| host);
Some(host.to_owned())
}
pub(crate) mod server;
#[cfg(test)]
mod tests;