use std::{
net::SocketAddr,
sync::{Arc, atomic::Ordering},
};
use anyhow::Context;
use axum::{
Router,
extract::{ConnectInfo, DefaultBodyLimit, Request},
http::StatusCode,
middleware::{self, Next},
response::{IntoResponse, Json},
routing::{any, delete, get, post},
};
use clap::Parser;
use tracing_subscriber::{EnvFilter, layer::SubscriberExt, util::SubscriberInitExt};
use aphrodite::{
config::{Cli, Command, MultiConfig, ProxyMode, SetupArgs},
proxy::{
self, handle_ccr_create, handle_ccr_delete, handle_ccr_list, handle_ccr_reload, handle_tool_relay, health_check,
},
retrieve, setup,
};
fn main() -> anyhow::Result<()> {
let args: Vec<String> = std::env::args().collect();
if args.iter().any(|a| a == "--version" || a == "-V") {
println!(
"aphrodite v{}",
option_env!("APHRODITE_VERSION").unwrap_or(env!("CARGO_PKG_VERSION"))
);
return Ok(());
}
if args.get(1).map(String::as_str) == Some("--help") || args.get(1).map(String::as_str) == Some("-h") {
Cli::parse();
return Ok(());
}
if args.get(1).map(String::as_str) == Some("setup") {
let cli = Cli::parse();
if let Some(Command::Setup { api_key, api_url, model, cache_port, token_port, no_launch, force }) = cli.command
{
let setup_args = SetupArgs { api_key, api_url, model, cache_port, token_port, no_launch, force };
match setup::run(&setup_args) {
Ok(()) => {
if setup_args.no_launch {
return Ok(());
}
println!("setup complete, starting proxy...");
},
Err(e) => {
eprintln!("setup failed: {e}");
std::process::exit(1);
},
}
}
}
let worker_threads = match std::env::var("APHRODITE_WORKER_THREADS") {
Ok(v) => match v.parse::<usize>() {
Ok(n) => n,
Err(_) => {
eprintln!("APHRODITE_WORKER_THREADS={v:?} is not a valid number; using the computed default");
let cpus = std::thread::available_parallelism().map(|n| n.get()).unwrap_or(8);
(cpus * 4).max(32)
},
},
Err(_) => {
let cpus = std::thread::available_parallelism().map(|n| n.get()).unwrap_or(8);
(cpus * 4).max(32)
},
};
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(worker_threads)
.enable_all()
.build()
.expect("tokio runtime");
runtime.block_on(run())
}
async fn run() -> anyhow::Result<()> {
let explicit_config_path = std::env::var("APHRODITE_CONFIG_PATH").ok();
let config_path = explicit_config_path.clone().unwrap_or_else(|| "aphrodite.toml".to_string());
let use_multi_config = std::path::Path::new(&config_path).exists();
let (config_path, use_multi_config) = if use_multi_config || explicit_config_path.is_some() {
(config_path, use_multi_config)
} else {
match dirs::home_dir().map(|h| h.join(".hermes").join("aphrodite").join("aphrodite.toml")) {
Some(p) if p.exists() => (p.to_string_lossy().into_owned(), true),
_ => (config_path, false),
}
};
let (multi_config, cli_fallback, log_compact) = if use_multi_config {
let config = MultiConfig::load(&config_path)?;
let log_compact = aphrodite::config::env_bool("APHRODITE_LOG_COMPACT");
(Some(config), None, log_compact)
} else {
let cli = Cli::parse();
let log_compact = cli.log_compact || aphrodite::config::env_bool("APHRODITE_LOG_COMPACT");
(None, Some(cli), log_compact)
};
let filter = EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("info"));
let subscriber = tracing_subscriber::registry().with(filter);
if log_compact {
subscriber
.with(tracing_subscriber::fmt::layer().compact().with_target(false).without_time())
.try_init()?;
} else {
subscriber.with(tracing_subscriber::fmt::layer()).try_init()?;
}
let compression = multi_config.as_ref().and_then(|c| c.compression.clone());
let proxies: Vec<(String, Cli)> = if let Some(config) = multi_config {
config
.proxies
.iter()
.map(|p| {
let cli = config.resolve(p)?;
let name = p.name.clone().unwrap_or_else(|| format!("{}", cli.listen));
Ok((name, cli))
})
.collect::<anyhow::Result<Vec<_>>>()?
} else {
let cli = cli_fallback.expect("cli_fallback set when use_multi_config is false");
if cli.api_key.trim().is_empty() {
anyhow::bail!(
"no API key configured - set APHRODITE_API_KEY env var or pass --api-key, or run `aphrodite setup`"
);
}
let name = format!("{}", cli.listen);
vec![(name, cli)]
};
tracing::info!(
"aphrodite v{} ({}{}) • {} • {}",
option_env!("APHRODITE_VERSION").unwrap_or("?"),
option_env!("APHRODITE_GIT_HASH").unwrap_or("?"),
option_env!("APHRODITE_PROFILE").map(|p| format!(", {p}")).unwrap_or_default(),
option_env!("APHRODITE_BUILD_DATE").unwrap_or("?"),
option_env!("APHRODITE_TARGET").unwrap_or("?"),
);
tracing::info!("starting {} proxy listener(s)", proxies.len());
let (shutdown_tx, shutdown_rx) = tokio::sync::watch::channel(false);
let mut bound = Vec::with_capacity(proxies.len());
for (name, mut cli) in proxies {
if let Some(ref db_path) = cli.ccr_db_path {
if !db_path.as_os_str().is_empty() && !db_path.is_absolute() {
if let Ok(exe_path) = std::env::current_exe() {
if let Some(exe_dir) = exe_path.parent() {
let old = db_path.display().to_string();
cli.ccr_db_path = Some(exe_dir.join(db_path));
tracing::info!(
"resolved relative ccr_db_path from {} to {}",
old,
cli.ccr_db_path.as_ref().unwrap().display()
);
}
}
}
if let Some(parent) = cli.ccr_db_path.as_ref().and_then(|p| p.parent()) {
std::fs::create_dir_all(parent)?;
}
}
let listener = tokio::net::TcpListener::bind(cli.listen)
.await
.with_context(|| format!("failed to bind listener \"{name}\" on {}", cli.listen))?;
let state = Arc::new(proxy::build_state(&cli, compression.as_ref()).await?);
bound.push((name, cli, listener, state));
}
let watch_path = {
let p = std::path::PathBuf::from(&config_path);
if p.is_relative() {
std::env::current_dir().unwrap_or_default().join(&p)
} else {
p
}
};
let watch_path_str = watch_path.to_string_lossy().to_string();
let states_for_watcher: Vec<_> = bound.iter().map(|(_, _, _, s)| s.clone()).collect();
tokio::spawn(async move {
use notify::{Event, EventKind, RecursiveMode, Watcher};
let (tx, mut rx) = tokio::sync::mpsc::channel(16);
let mut watcher = match notify::recommended_watcher(move |res: Result<Event, notify::Error>| {
if let Ok(event) = res {
let is_modify = matches!(event.kind, EventKind::Modify(_));
if is_modify && event.paths.iter().any(|p| p.to_string_lossy().contains("aphrodite.toml")) {
let _ = tx.try_send(());
}
}
}) {
Ok(w) => w,
Err(e) => {
tracing::warn!("failed to create config watcher: {e}");
return;
},
};
let watch_dir = std::path::Path::new(&watch_path).parent().unwrap_or(std::path::Path::new("."));
if let Err(e) = watcher.watch(watch_dir, RecursiveMode::NonRecursive) {
tracing::warn!("failed to start config watcher: {e}");
return;
}
tracing::info!(path = %watch_path_str, "config file watcher active");
loop {
if rx.recv().await.is_some() {
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
while rx.try_recv().is_ok() {}
match aphrodite::config::MultiConfig::load(&watch_path_str) {
Ok(config) => {
let thresholds = proxy::resolve_thresholds(config.compression.as_ref());
for state in &states_for_watcher {
state.cache_compress_threshold.store(thresholds.cache, Ordering::Relaxed);
state.token_compress_threshold.store(thresholds.token, Ordering::Relaxed);
state.inline_ccr_threshold.store(thresholds.inline, Ordering::Relaxed);
state
.code_multiplier_x100
.store((thresholds.code_multiplier * 100.0) as u64, Ordering::Relaxed);
}
tracing::info!(
path = %watch_path_str,
cache_threshold = thresholds.cache,
token_threshold = thresholds.token,
inline_threshold = thresholds.inline,
code_multiplier = thresholds.code_multiplier,
listeners = states_for_watcher.len(),
"config reloaded - compression thresholds applied to all live listeners"
);
},
Err(e) => {
tracing::warn!(error = %e, "failed to reload config on file change");
},
}
}
}
});
let mut handles = Vec::new();
for (name, cli, listener, state) in bound {
let rx = shutdown_rx.clone();
let handle = tokio::spawn(async move {
if let Err(e) = run_single(name, cli, listener, state, rx).await {
tracing::error!(%e, "proxy listener failed");
}
});
handles.push(handle);
}
drop(shutdown_rx);
shutdown_signal().await;
let _ = shutdown_tx.send(true);
tracing::info!("shutdown signal received, draining in-flight requests...");
let abort_handles: Vec<_> = handles.iter().map(|h| h.abort_handle()).collect();
let second_signal = async {
let _ = tokio::signal::ctrl_c().await;
tracing::warn!("second shutdown signal received, forcing immediate shutdown");
};
tokio::pin!(second_signal);
let drain_timeout = tokio::time::sleep(std::time::Duration::from_secs(5));
tokio::pin!(drain_timeout);
let drain_fut = futures::future::join_all(handles);
tokio::pin!(drain_fut);
tokio::select! {
_ = &mut drain_fut => {
tracing::info!("all proxy listeners completed gracefully");
}
_ = &mut drain_timeout => {
tracing::info!("drain timeout (5s) reached, aborting remaining tasks");
for h in &abort_handles {
h.abort();
}
}
_ = &mut second_signal => {
tracing::info!("force shutdown on second signal, aborting remaining tasks");
for h in &abort_handles {
h.abort();
}
}
}
Ok(())
}
async fn run_single(
name: String,
cli: Cli,
listener: tokio::net::TcpListener,
state: Arc<proxy::AppState>,
mut shutdown_rx: tokio::sync::watch::Receiver<bool>,
) -> anyhow::Result<()> {
let task_tracker = state.task_tracker.clone();
let mode_str = match cli.mode {
ProxyMode::Cache => "cache",
ProxyMode::Token => "token",
};
tracing::info!(
name = %name,
listen = %cli.listen,
mode = %mode_str,
api_url = %cli.api_url,
model = %cli.model,
tool_relay = cli.tool_relay,
"proxy starting"
);
if mgmt_token().is_none() {
tracing::warn!(
name = %name,
"APHRODITE_MGMT_TOKEN not set - management routes (/stats, /retrieve, /ccr/*, /reload, /tool/relay, ...) accept any loopback caller with no credential"
);
}
let restricted = Router::new()
.route("/health/upstream", get({
let s = state.clone();
move |ConnectInfo(addr): ConnectInfo<SocketAddr>| {
let s = s.clone();
async move {
if !addr.ip().is_loopback() {
return (StatusCode::FORBIDDEN, Json(serde_json::json!({
"error": "only loopback clients allowed"
}))).into_response();
}
const TTL: std::time::Duration = std::time::Duration::from_secs(60);
if let Some((ok, at)) = *s.upstream_health_cache.lock().unwrap_or_else(|e| e.into_inner()) {
if at.elapsed() < TTL {
return Json(serde_json::json!({"upstream": ok, "cached": true})).into_response();
}
}
let ok = s.client
.get(format!("{}/models", s.api_url.trim_end_matches('/')))
.header("Authorization", format!("Bearer {}", s.api_key.expose()))
.timeout(std::time::Duration::from_secs(5))
.send()
.await
.map(|r| r.status().is_success())
.unwrap_or(false);
if let Ok(mut cache) = s.upstream_health_cache.lock() {
*cache = Some((ok, std::time::Instant::now()));
}
Json(serde_json::json!({"upstream": ok, "cached": false})).into_response()
}
}
}))
.route("/version", get(|| async { env!("CARGO_PKG_VERSION") }))
.route("/stats", get({
let s = state.clone();
move || async move { Json(s.stats_json()) }
}))
.route("/stats/db", get({
let s = state.clone();
move || async move {
let mode = match s.mode {
ProxyMode::Cache => "cache",
ProxyMode::Token => "token",
};
match &s.ccr {
Some(ccr) => match ccr.stats_db() {
Some(stats) => Json(stats).into_response(),
None => (
StatusCode::OK,
Json(serde_json::json!({
"error": "stats_db not available for this backend",
"mode": mode,
})),
)
.into_response(),
},
None => (
StatusCode::OK,
Json(serde_json::json!({
"error": "CCR not enabled",
"mode": mode,
})),
)
.into_response(),
}
}
}))
.route("/metrics", get({
let s = state.clone();
move || async move {
let stats = s.stats_json();
let mut out = String::new();
let mode_str = stats["mode"].as_str().unwrap_or("unknown");
out.push_str(&format!("aphrodite_requests_total{{mode=\"{}\"}} {}\n",
mode_str, stats["requests"]["total"]));
out.push_str(&format!("aphrodite_requests_compressed_total{{mode=\"{}\"}} {}\n",
mode_str, stats["requests"]["compressed"]));
out.push_str(&format!("aphrodite_tokens_saved_total {}\n", stats["tokens_saved"]));
out.push_str(&format!("aphrodite_ccr_hits_total {}\n", stats["ccr"]["hits"]));
out.push_str(&format!("aphrodite_ccr_misses_total {}\n", stats["ccr"]["misses"]));
out.push_str(&format!("aphrodite_ccr_created_total {}\n", stats["ccr"]["created"]));
out.push_str(&format!("aphrodite_tool_relay_calls_total {}\n", stats["tool_relay_calls"]));
if let Some(cache) = stats["cache"].as_object() {
out.push_str(&format!("aphrodite_cache_hits_total {}\n", cache["hits"]));
out.push_str(&format!("aphrodite_cache_misses_total {}\n", cache["misses"]));
}
if let Some(buckets) = stats["latency_buckets_us"].as_array() {
let mut total_count: u64 = 0;
let le_labels = ["0.001", "0.01", "0.1", "1.0", "+Inf"];
for (i, v) in buckets.iter().enumerate() {
let le = le_labels.get(i).copied().unwrap_or("+Inf");
let count = v.as_u64().unwrap_or(0);
total_count += count;
out.push_str(&format!("aphrodite_latency_seconds_bucket{{le=\"{}\"}} {}\n", le, total_count));
}
out.push_str(&format!("aphrodite_latency_seconds_count {}\n", total_count));
if let Some(total_us) = stats["total_latency_micros"].as_u64() {
out.push_str(&format!("aphrodite_latency_seconds_sum {:.6}\n", total_us as f64 / 1_000_000.0));
}
}
if let Some(ratio) = stats["compression_ratio_ema"].as_f64() {
out.push_str(&format!("aphrodite_compression_ratio_ema {:.2}\n", ratio));
}
if let Some(icc) = stats["inline_ccr"].as_object() {
if let Some(h) = icc["hits"].as_u64() { out.push_str(&format!("aphrodite_inline_ccr_hits_total {h}\n")); }
if let Some(m) = icc["misses"].as_u64() { out.push_str(&format!("aphrodite_inline_ccr_misses_total {m}\n")); }
}
if let Some(tr) = stats["tool_relay"].as_object() {
if let Some(s) = tr["success"].as_u64() { out.push_str(&format!("aphrodite_tool_relay_success_total {s}\n")); }
if let Some(f) = tr["failure"].as_u64() { out.push_str(&format!("aphrodite_tool_relay_failure_total {f}\n")); }
}
if let Some(n) = stats["notify"].as_object() {
if let Some(s) = n["success"].as_u64() { out.push_str(&format!("aphrodite_notify_success_total {s}\n")); }
if let Some(f) = n["failure"].as_u64() { out.push_str(&format!("aphrodite_notify_failure_total {f}\n")); }
}
if let Some(ue) = stats["upstream_errors"].as_object() {
if let Some(c) = ue["4xx"].as_u64() { out.push_str(&format!("aphrodite_upstream_errors_total{{code=\"4xx\"}} {c}\n")); }
if let Some(c) = ue["5xx"].as_u64() { out.push_str(&format!("aphrodite_upstream_errors_total{{code=\"5xx\"}} {c}\n")); }
if let Some(t) = ue["timeouts"].as_u64() { out.push_str(&format!("aphrodite_upstream_timeouts_total {t}\n")); }
if let Some(c) = ue["connect_errors"].as_u64() { out.push_str(&format!("aphrodite_upstream_connect_errors_total {c}\n")); }
if let Some(c) = ue["sse_stream_errors"].as_u64() { out.push_str(&format!("aphrodite_sse_stream_errors_total {c}\n")); }
}
if let Some(cs) = stats["ccr_store"].as_object() {
if let Some(e) = cs["entries"].as_u64() { out.push_str(&format!("aphrodite_ccr_store_entries {e}\n")); }
if let Some(b) = cs["bytes_approx"].as_u64() { out.push_str(&format!("aphrodite_ccr_store_bytes {b}\n")); }
}
if let Some(bb) = stats["body_bytes"].as_object() {
if let Some(r) = bb["request"].as_u64() { out.push_str(&format!("aphrodite_request_body_bytes_total {r}\n")); }
if let Some(r) = bb["response"].as_u64() { out.push_str(&format!("aphrodite_response_body_bytes_total {r}\n")); }
}
if let Some(ul) = stats["upstream_latency_micros"].as_u64() {
out.push_str(&format!("aphrodite_upstream_latency_seconds_total {:.6}\n", ul as f64 / 1_000_000.0));
}
(StatusCode::OK, [(axum::http::header::CONTENT_TYPE, "text/plain; version=0.0.4")], out)
}
}))
.route("/history", get({
let s = state.clone();
move || async move {
Json(s.request_history.lock().map(|v| v.clone()).unwrap_or_default())
}
}))
.route("/retrieve", post(retrieve::handle_retrieve))
.route("/tool/relay", post(handle_tool_relay))
.route("/ccr/create", post(handle_ccr_create))
.route("/ccr/list", get(handle_ccr_list))
.route("/ccr/{hash}", delete(handle_ccr_delete))
.route("/reload", post(handle_ccr_reload))
.route("/favicon.ico", get(|| async { StatusCode::NOT_FOUND }))
.route("/robots.txt", get(|| async { "User-agent: *\nDisallow: /\n" }))
.route("/", get(|| async {
Json(serde_json::json!({
"proxy": "aphrodite",
"version": env!("CARGO_PKG_VERSION"),
"git_hash": option_env!("APHRODITE_GIT_HASH"),
"build_date": option_env!("APHRODITE_BUILD_DATE"),
"target": option_env!("APHRODITE_TARGET"),
"profile": option_env!("APHRODITE_PROFILE"),
}))
}))
.layer(DefaultBodyLimit::max(1024 * 1024))
.layer(middleware::from_fn(require_mgmt_token));
let catch_all = Router::new()
.route("/{*path}", any(proxy::proxy_handler))
.layer(DefaultBodyLimit::max(64 * 1024 * 1024));
let restricted = restricted
.merge(catch_all)
.layer(middleware::from_fn(loopback_only));
let app = Router::new()
.route("/health", get(health_check))
.merge(restricted)
.with_state(state);
tracing::info!(addr = %listener.local_addr()?, "listening");
let shutdown_fut = async move {
let _ = shutdown_rx.changed().await;
};
let serve_result = axum::serve(listener, app.into_make_service_with_connect_info::<SocketAddr>())
.with_graceful_shutdown(shutdown_fut)
.await;
task_tracker.close();
task_tracker.wait().await;
tracing::debug!("all background tasks completed");
serve_result?;
Ok(())
}
const ALLOWED_LOOPBACK_HOSTS: &[&str] = &["localhost", "127.0.0.1", "[::1]", "::1"];
fn host_header_to_hostname(host: &str) -> String {
if let Some(rest) = host.strip_prefix('[') {
rest.split(']').next().map(|h| format!("[{h}]")).unwrap_or_default()
} else {
host.split(':').next().unwrap_or("").to_string()
}
}
fn check_loopback_request(addr: SocketAddr, host_header: &str) -> Result<(), &'static str> {
if !addr.ip().is_loopback() {
return Err("only loopback clients allowed");
}
let hostname = host_header_to_hostname(host_header);
if !ALLOWED_LOOPBACK_HOSTS.contains(&hostname.as_str()) {
return Err("Host header does not name a loopback address");
}
Ok(())
}
async fn loopback_only(
ConnectInfo(addr): ConnectInfo<SocketAddr>,
request: Request,
next: Next,
) -> Result<impl IntoResponse, (StatusCode, Json<serde_json::Value>)> {
let host = request
.headers()
.get(axum::http::header::HOST)
.and_then(|v| v.to_str().ok())
.unwrap_or("");
if let Err(msg) = check_loopback_request(addr, host) {
return Err((StatusCode::FORBIDDEN, Json(serde_json::json!({"error": msg}))));
}
Ok(next.run(request).await)
}
fn mgmt_token() -> Option<String> {
std::env::var("APHRODITE_MGMT_TOKEN").ok().filter(|s| !s.is_empty())
}
fn check_bearer_token(configured: Option<&str>, auth_header: &str) -> Result<(), &'static str> {
let Some(token) = configured else {
return Ok(());
};
if auth_header.strip_prefix("Bearer ") == Some(token) {
Ok(())
} else {
Err("missing or invalid Authorization bearer token")
}
}
async fn require_mgmt_token(
request: Request,
next: Next,
) -> Result<impl IntoResponse, (StatusCode, Json<serde_json::Value>)> {
if request.uri().path() == "/metrics" {
return Ok(next.run(request).await);
}
let token = mgmt_token();
let auth = request
.headers()
.get(axum::http::header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.unwrap_or("");
if let Err(msg) = check_bearer_token(token.as_deref(), auth) {
return Err((StatusCode::UNAUTHORIZED, Json(serde_json::json!({"error": msg}))));
}
Ok(next.run(request).await)
}
async fn shutdown_signal() {
let ctrl_c = async {
let _ = tokio::signal::ctrl_c().await;
};
#[cfg(unix)]
let terminate = async {
if let Ok(mut s) = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate()) {
s.recv().await;
}
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! { _ = ctrl_c => {}, _ = terminate => {} }
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_host_header_to_hostname_strips_port() {
assert_eq!(host_header_to_hostname("127.0.0.1:9797"), "127.0.0.1");
assert_eq!(host_header_to_hostname("localhost:9798"), "localhost");
}
#[test]
fn test_host_header_to_hostname_no_port() {
assert_eq!(host_header_to_hostname("127.0.0.1"), "127.0.0.1");
assert_eq!(host_header_to_hostname("localhost"), "localhost");
}
#[test]
fn test_host_header_to_hostname_ipv6_bracketed() {
assert_eq!(host_header_to_hostname("[::1]:9797"), "[::1]");
assert_eq!(host_header_to_hostname("[::1]"), "[::1]");
}
#[test]
fn test_allowed_loopback_hosts_accepts_real_clients() {
for h in ["127.0.0.1", "localhost", "[::1]"] {
assert!(ALLOWED_LOOPBACK_HOSTS.contains(&h), "{h} must be an allowed loopback host");
}
}
#[test]
fn test_allowed_loopback_hosts_rejects_dns_rebinding_hostname() {
let hostname = host_header_to_hostname("attacker.example:9797");
assert!(!ALLOWED_LOOPBACK_HOSTS.contains(&hostname.as_str()));
}
fn loopback_v4(port: u16) -> SocketAddr {
use std::net::{Ipv4Addr, SocketAddrV4};
SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, port))
}
fn lan_v4(port: u16) -> SocketAddr {
use std::net::{Ipv4Addr, SocketAddrV4};
SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::new(192, 168, 1, 50), port))
}
#[test]
fn test_check_loopback_request_allows_real_client_host_headers() {
assert!(check_loopback_request(loopback_v4(9797), "127.0.0.1:9797").is_ok());
assert!(check_loopback_request(loopback_v4(9798), "127.0.0.1:9798").is_ok());
assert!(check_loopback_request(loopback_v4(9797), "localhost:9797").is_ok());
assert!(
check_loopback_request(loopback_v4(9797), "127.0.0.1").is_ok(),
"no-port Host must also pass"
);
}
#[test]
fn test_check_loopback_request_rejects_missing_host_header() {
assert!(check_loopback_request(loopback_v4(9797), "").is_err());
}
#[test]
fn test_check_loopback_request_allows_ipv6_loopback() {
use std::net::{Ipv6Addr, SocketAddrV6};
let addr = SocketAddr::V6(SocketAddrV6::new(Ipv6Addr::LOCALHOST, 9797, 0, 0));
assert!(check_loopback_request(addr, "[::1]:9797").is_ok());
}
#[test]
fn test_check_loopback_request_rejects_dns_rebinding() {
let result = check_loopback_request(loopback_v4(9797), "attacker.example:9797");
assert!(result.is_err());
}
#[test]
fn test_check_loopback_request_rejects_non_loopback_peer_regardless_of_host() {
let result = check_loopback_request(lan_v4(9797), "127.0.0.1:9797");
assert!(result.is_err());
}
#[test]
fn test_check_loopback_request_rejects_non_loopback_peer_with_no_host() {
let result = check_loopback_request(lan_v4(9797), "");
assert!(result.is_err());
}
#[test]
fn test_check_bearer_token_passes_when_unconfigured() {
assert!(check_bearer_token(None, "").is_ok());
assert!(check_bearer_token(None, "Bearer whatever").is_ok());
}
#[test]
fn test_check_bearer_token_accepts_matching_token() {
assert!(check_bearer_token(Some("secret123"), "Bearer secret123").is_ok());
}
#[test]
fn test_check_bearer_token_rejects_missing_header() {
assert!(check_bearer_token(Some("secret123"), "").is_err());
}
#[test]
fn test_check_bearer_token_rejects_wrong_token() {
assert!(check_bearer_token(Some("secret123"), "Bearer wrong").is_err());
}
#[test]
fn test_check_bearer_token_rejects_missing_bearer_prefix() {
assert!(check_bearer_token(Some("secret123"), "secret123").is_err());
}
}