use std::collections::BTreeMap;
use std::path::Path;
use std::sync::Arc;
use axum::body::Body;
use axum::extract::Request;
use axum::http::{header, StatusCode};
use axum::response::{IntoResponse, Response};
use futures::StreamExt;
use tracing::warn;
const HOP_BY_HOP: &[&str] = &[
"connection",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"te",
"trailer",
"transfer-encoding",
"upgrade",
"host",
"content-length",
];
fn is_hop_by_hop(name: &str) -> bool {
HOP_BY_HOP.contains(&name)
}
#[derive(Debug, Clone, Default)]
pub struct ProxyMap {
routes: Vec<(String, String)>,
}
impl ProxyMap {
pub fn new(entries: impl IntoIterator<Item = (String, String)>) -> Self {
let mut routes: Vec<(String, String)> = entries
.into_iter()
.map(|(prefix, base)| (normalize_prefix(&prefix), base))
.filter(|(prefix, _)| prefix != "/") .collect();
routes.sort_by(|a, b| b.0.len().cmp(&a.0.len()).then_with(|| a.0.cmp(&b.0)));
routes.dedup_by(|a, b| a.0 == b.0);
Self { routes }
}
pub fn from_json_file(path: &Path) -> anyhow::Result<Self> {
let body = std::fs::read_to_string(path)
.map_err(|e| anyhow::anyhow!("reading proxy map {}: {e}", path.display()))?;
let map: BTreeMap<String, String> = serde_json::from_str(&body)
.map_err(|e| anyhow::anyhow!("parsing proxy map {}: {e}", path.display()))?;
Ok(Self::new(map))
}
pub fn is_empty(&self) -> bool {
self.routes.is_empty()
}
pub fn routes(&self) -> &[(String, String)] {
&self.routes
}
pub fn match_base<'a>(&'a self, path: &str) -> Option<&'a str> {
self.routes.iter().find_map(|(prefix, base)| {
let is_match = path == prefix
|| (path.starts_with(prefix.as_str())
&& path.as_bytes().get(prefix.len()) == Some(&b'/'));
is_match.then_some(base.as_str())
})
}
}
fn normalize_prefix(raw: &str) -> String {
let trimmed = raw.trim();
let with_lead = if trimmed.starts_with('/') {
trimmed.to_string()
} else {
format!("/{trimmed}")
};
let no_trail = with_lead.trim_end_matches('/');
if no_trail.is_empty() {
"/".to_string()
} else {
no_trail.to_string()
}
}
#[derive(Clone)]
pub struct ProxyState {
map: Arc<ProxyMap>,
client: reqwest::Client,
}
impl ProxyState {
pub fn new(map: ProxyMap) -> Self {
Self {
map: Arc::new(map),
client: reqwest::Client::new(),
}
}
pub fn map(&self) -> &ProxyMap {
&self.map
}
pub async fn forward(&self, base: &str, req: Request) -> Response {
let (parts, body) = req.into_parts();
let path_and_query = parts
.uri
.path_and_query()
.map(|p| p.as_str())
.unwrap_or_else(|| parts.uri.path());
let target = format!("{}{}", base.trim_end_matches('/'), path_and_query);
let body_bytes = match collect_body(body.into_data_stream()).await {
Ok(b) => b,
Err(e) => {
warn!(error = %e, "proxy: failed to buffer request body");
return (StatusCode::BAD_GATEWAY, "request buffer failed").into_response();
}
};
let method = match reqwest::Method::from_bytes(parts.method.as_str().as_bytes()) {
Ok(m) => m,
Err(_) => return (StatusCode::BAD_GATEWAY, "bad method").into_response(),
};
let mut rb = self.client.request(method, &target);
for (k, v) in parts.headers.iter() {
if is_hop_by_hop(&k.as_str().to_ascii_lowercase()) {
continue;
}
rb = rb.header(k.as_str(), v.as_bytes());
}
if !body_bytes.is_empty() {
rb = rb.body(body_bytes);
}
let upstream = match rb.send().await {
Ok(r) => r,
Err(e) => {
warn!(error = %e, target = %target, "proxy: upstream request failed");
return (StatusCode::BAD_GATEWAY, format!("proxy upstream failed: {e}"))
.into_response();
}
};
let status = upstream.status();
let upstream_headers = upstream.headers().clone();
let upstream_body = match upstream.bytes().await {
Ok(b) => b,
Err(e) => {
warn!(error = %e, "proxy: failed to read upstream body");
return (StatusCode::BAD_GATEWAY, "upstream body read failed").into_response();
}
};
let mut builder = Response::builder().status(status.as_u16());
for (k, v) in upstream_headers.iter() {
if is_hop_by_hop(&k.as_str().to_ascii_lowercase()) {
continue;
}
if k == header::CONTENT_LENGTH {
continue;
}
builder = builder.header(k.as_str(), v.as_bytes());
}
builder
.body(Body::from(upstream_body))
.unwrap_or_else(|_| (StatusCode::BAD_GATEWAY, "proxy response build failed").into_response())
}
}
async fn collect_body(mut stream: axum::body::BodyDataStream) -> Result<Vec<u8>, axum::Error> {
let mut buf = Vec::new();
while let Some(chunk) = stream.next().await {
buf.extend_from_slice(&chunk?);
}
Ok(buf)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn normalizes_prefixes() {
assert_eq!(normalize_prefix("/auth"), "/auth");
assert_eq!(normalize_prefix("auth"), "/auth");
assert_eq!(normalize_prefix("/auth/"), "/auth");
assert_eq!(normalize_prefix(" /api/ "), "/api");
assert_eq!(normalize_prefix("/"), "/");
}
#[test]
fn matches_on_segment_boundary_only() {
let map = ProxyMap::new([("/auth".to_string(), "http://b:1".to_string())]);
assert_eq!(map.match_base("/auth"), Some("http://b:1"));
assert_eq!(map.match_base("/auth/magic-link/verify"), Some("http://b:1"));
assert_eq!(map.match_base("/authx"), None);
assert_eq!(map.match_base("/other"), None);
}
#[test]
fn longest_prefix_wins() {
let map = ProxyMap::new([
("/auth".to_string(), "http://broad:1".to_string()),
("/auth/admin".to_string(), "http://specific:2".to_string()),
]);
assert_eq!(map.match_base("/auth/admin/x"), Some("http://specific:2"));
assert_eq!(map.match_base("/auth/login"), Some("http://broad:1"));
}
#[test]
fn catch_all_root_is_dropped() {
let map = ProxyMap::new([("/".to_string(), "http://swallow:1".to_string())]);
assert!(map.is_empty(), "a '/' catch-all must not be installed");
assert_eq!(map.match_base("/anything"), None);
}
}