use super::*;
use crate::net::qos::{QosPolicy, QosRuntime};
use axum::{
Router,
body::{Body, to_bytes},
routing::{get, post},
};
use serde_json::json;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf};
use tower::ServiceExt;
#[test]
fn parses_overridden_http_bind_address() {
let overrides = NetworkEnvOverrides {
http_bind: Some("localhost:9000".to_owned()),
http_request_timeout_ms: Some(1500),
http_shutdown_timeout_ms: Some(2500),
http_max_body_bytes: Some(4096),
proxy: Some("https://proxy.internal:8443".to_owned()),
no_proxy: Some("localhost,.internal".to_owned()),
ssl_verify: Some(false),
..NetworkEnvOverrides::default()
};
let config = HttpConfig::from_overrides(&overrides).expect("config should parse");
assert_eq!(config.bind_address.to_string(), "localhost:9000");
assert_eq!(config.bind_address.port(), 9000);
assert_eq!(config.request_timeout, Duration::from_millis(1500));
assert_eq!(
config.graceful_shutdown_timeout,
Duration::from_millis(2500)
);
assert_eq!(config.max_request_body_bytes, 4096);
assert_eq!(
config.proxy.proxy,
Some("https://proxy.internal:8443".to_owned())
);
assert_eq!(config.proxy.no_proxy_rules, ["localhost", ".internal"]);
assert!(!config.proxy.ssl_verify);
}
#[test]
fn rejects_invalid_bind_addresses() {
let overrides = NetworkEnvOverrides {
http_bind: Some("localhost".to_owned()),
..NetworkEnvOverrides::default()
};
let error = HttpConfig::from_overrides(&overrides)
.expect_err("bind address must include host and port");
assert_eq!(
error,
HttpConfigError::InvalidBindAddress {
value: "localhost".to_owned()
}
);
}
#[test]
fn rejects_ephemeral_ports() {
let error = HttpBindAddress::parse("127.0.0.1:0").expect_err("port zero should fail");
assert_eq!(error, HttpConfigError::EphemeralPort);
}
#[test]
fn rejects_proxy_urls_without_supported_scheme_or_host() {
for proxy in [
"socks5://proxy.internal:1080",
"http://:8080",
"https://user@:443",
] {
let overrides = NetworkEnvOverrides {
proxy: Some(proxy.to_owned()),
..NetworkEnvOverrides::default()
};
let error = HttpConfig::from_overrides(&overrides).expect_err("invalid proxy should fail");
assert_eq!(error, HttpConfigError::InvalidProxyUrl);
}
}
#[test]
fn rejects_empty_no_proxy_entries() {
let overrides = NetworkEnvOverrides {
no_proxy: Some("localhost,,example.com".to_owned()),
..NetworkEnvOverrides::default()
};
let error =
HttpConfig::from_overrides(&overrides).expect_err("empty no-proxy entry should fail");
assert_eq!(error, HttpConfigError::EmptyNoProxyRule);
}
#[tokio::test]
async fn post_json_with_qos_reuses_outer_http_request_admission() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should load");
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("client should connect");
let mut buffer = [0; 512];
let _ = stream.read(&mut buffer).await.expect("request should read");
stream
.write_all(
b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 11\r\n\r\n{\"ok\":true}",
)
.await
.expect("response should write");
});
let config = HttpConfig::new(
HttpBindAddress::parse("127.0.0.1:8791").expect("bind should parse"),
Duration::from_secs(5),
Duration::from_secs(5),
1024,
HttpProxyConfig::new(None, Vec::new(), true).expect("proxy should build"),
)
.expect("config should build");
let qos = QosRuntime::default();
let policy = QosPolicy::new(8, 1, 8).expect("policy should build");
let url = format!("http://{addr}/worker");
let router = router_with_qos_request_admission(
Router::new().route(
"/nested-worker",
post({
let config = config.clone();
let qos = qos.clone();
let policy = policy.clone();
move || {
let config = config.clone();
let qos = qos.clone();
let policy = policy.clone();
let url = url.clone();
async move {
post_json_with_qos(&config, &qos, &policy, &url, &json!({"task": "ocr"}))
.await
.expect("nested raw worker call should not be double charged")
.to_string()
}
}
}),
),
qos.clone(),
policy,
);
let response = router
.oneshot(
axum::http::Request::builder()
.method("POST")
.uri("/nested-worker")
.body(Body::empty())
.expect("request should build"),
)
.await
.expect("router should respond");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(qos.diagnostics_snapshot().admitted_total, 1);
assert_eq!(qos.diagnostics_snapshot().rejected_total, 0);
server.await.expect("server task should finish");
}
#[tokio::test]
async fn qos_request_layer_rejects_excess_web_requests() {
let qos = QosRuntime::default();
let policy = QosPolicy::new(8, 1, 8).expect("policy should build");
let (request_started, request_started_waiter) = tokio::sync::oneshot::channel();
let request_started = Arc::new(std::sync::Mutex::new(Some(request_started)));
let router = router_with_qos_request_admission(
Router::new()
.route(
"/hold",
get(move || signal_pending_route(request_started.clone())),
)
.route("/ok", get(|| async { "ok" }))
.route("/api/ok", get(|| async { "ok" })),
qos.clone(),
policy,
);
let first = tokio::spawn(
router.clone().oneshot(
axum::http::Request::builder()
.uri("/hold")
.body(Body::empty())
.expect("request should build"),
),
);
tokio::time::timeout(Duration::from_secs(5), request_started_waiter)
.await
.expect("handler should start")
.expect("request should signal startup");
let response = router
.oneshot(
axum::http::Request::builder()
.uri("/api/ok")
.body(Body::empty())
.expect("request should build"),
)
.await
.expect("router should respond");
assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS);
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("body should be readable");
let error: serde_json::Value =
serde_json::from_slice(&body).expect("API QoS rejection should be JSON");
assert_eq!(error["error_kind"], "qos_rejected");
assert_eq!(qos.diagnostics_snapshot().rejected_total, 1);
first.abort();
}
#[tokio::test]
async fn qos_request_layer_does_not_double_charge_nested_outbound_request() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should load");
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("client should connect");
let mut buffer = [0; 512];
let _ = stream.read(&mut buffer).await.expect("request should read");
stream
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok")
.await
.expect("response should write");
});
let qos = QosRuntime::default();
let policy = QosPolicy::new(8, 1, 8).expect("policy should build");
let client = reqwest::Client::new();
let url = format!("http://{addr}/catalog");
let router = router_with_qos_request_admission(
Router::new().route(
"/nested",
get({
let client = client.clone();
let qos = qos.clone();
let policy = policy.clone();
move || {
let client = client.clone();
let qos = qos.clone();
let policy = policy.clone();
let url = url.clone();
async move {
let response = send_request_with_qos(&qos, &policy, client.get(url))
.await
.expect("nested outbound request should not be double charged");
response.text().await.expect("nested body should read")
}
}
}),
),
qos.clone(),
policy,
);
let response = router
.oneshot(
axum::http::Request::builder()
.uri("/nested")
.body(Body::empty())
.expect("request should build"),
)
.await
.expect("router should respond");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(qos.diagnostics_snapshot().admitted_total, 1);
assert_eq!(qos.diagnostics_snapshot().rejected_total, 0);
server.await.expect("server task should finish");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn server_timeout_with_qos_request_layer_counts_once() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("listener should bind");
let bind = listener
.local_addr()
.expect("listener should expose address")
.to_string();
let config = HttpConfig::new(
HttpBindAddress::parse(&bind).expect("bind should parse"),
Duration::from_millis(20),
Duration::from_secs(5),
1024,
HttpProxyConfig::new(None, Vec::new(), true).expect("proxy should build"),
)
.expect("config should build");
let qos = QosRuntime::default();
let policy = QosPolicy::new(8, 8, 8).expect("policy should build");
let router = router_with_qos_request_admission(
Router::new().route(
"/hold",
get(|| async { std::future::pending::<&'static str>().await }),
),
qos.clone(),
policy,
);
let (shutdown, shutdown_waiter) = tokio::sync::oneshot::channel();
let server = tokio::spawn(serve_listener(
listener,
router,
config,
Some(qos.clone()),
async {
let _ = shutdown_waiter.await;
},
));
let response = reqwest::Client::new()
.get(format!("http://{bind}/hold"))
.send()
.await
.expect("server should respond with timeout");
assert_eq!(response.status(), StatusCode::REQUEST_TIMEOUT);
assert_eq!(qos.diagnostics_snapshot().timed_out_total, 1);
let _ = shutdown.send(());
server
.await
.expect("server task should join")
.expect("server should stop");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn serve_router_enforces_graceful_shutdown_timeout() {
let bind = "127.0.0.1:8791";
let listener =
InMemoryRequestListener::new(b"GET /hold HTTP/1.1\r\nHost: localhost\r\n\r\n".to_vec());
let config = HttpConfig::new(
HttpBindAddress::parse(bind).expect("bind should parse"),
Duration::from_secs(30),
Duration::from_millis(10),
1024,
HttpProxyConfig::new(None, Vec::new(), true).expect("proxy should build"),
)
.expect("config should build");
let (request_started, request_started_waiter) = tokio::sync::oneshot::channel();
let request_started = Arc::new(std::sync::Mutex::new(Some(request_started)));
let router = Router::new().route(
"/hold",
get(move || signal_pending_route(request_started.clone())),
);
let (shutdown, shutdown_waiter) = tokio::sync::oneshot::channel();
let server = tokio::spawn(serve_listener(listener, router, config, None, async {
let _ = shutdown_waiter.await;
}));
tokio::time::timeout(Duration::from_secs(10), request_started_waiter)
.await
.expect("handler should start before shutdown")
.expect("request should signal startup");
let _ = shutdown.send(());
let error = server
.await
.expect("server task should join")
.expect_err("active request should exceed shutdown timeout");
assert!(matches!(error, HttpServeError::ShutdownTimeout));
}
async fn signal_pending_route(
request_started: Arc<std::sync::Mutex<Option<tokio::sync::oneshot::Sender<()>>>>,
) -> &'static str {
let sender = request_started
.lock()
.expect("request signal mutex should not be poisoned")
.take();
if let Some(sender) = sender {
let _ = sender.send(());
}
std::future::pending().await
}
struct InMemoryRequestListener {
request: Option<InMemoryRequestStream>,
address: std::net::SocketAddr,
}
impl InMemoryRequestListener {
fn new(request: Vec<u8>) -> Self {
Self {
request: Some(InMemoryRequestStream { request, offset: 0 }),
address: "127.0.0.1:8791"
.parse()
.expect("loopback test address should parse"),
}
}
}
impl axum::serve::Listener for InMemoryRequestListener {
type Io = InMemoryRequestStream;
type Addr = std::net::SocketAddr;
async fn accept(&mut self) -> (Self::Io, Self::Addr) {
if let Some(request) = self.request.take() {
return (request, self.address);
}
std::future::pending().await
}
fn local_addr(&self) -> std::io::Result<Self::Addr> {
Ok(self.address)
}
}
struct InMemoryRequestStream {
request: Vec<u8>,
offset: usize,
}
impl AsyncRead for InMemoryRequestStream {
fn poll_read(
mut self: std::pin::Pin<&mut Self>,
_context: &mut std::task::Context<'_>,
buffer: &mut ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
let remaining = &self.request[self.offset..];
if remaining.is_empty() {
return std::task::Poll::Pending;
}
let readable = remaining.len().min(buffer.remaining());
buffer.put_slice(&remaining[..readable]);
self.offset += readable;
std::task::Poll::Ready(Ok(()))
}
}
impl AsyncWrite for InMemoryRequestStream {
fn poll_write(
self: std::pin::Pin<&mut Self>,
_context: &mut std::task::Context<'_>,
buffer: &[u8],
) -> std::task::Poll<std::io::Result<usize>> {
std::task::Poll::Ready(Ok(buffer.len()))
}
fn poll_flush(
self: std::pin::Pin<&mut Self>,
_context: &mut std::task::Context<'_>,
) -> std::task::Poll<std::io::Result<()>> {
std::task::Poll::Ready(Ok(()))
}
fn poll_shutdown(
self: std::pin::Pin<&mut Self>,
_context: &mut std::task::Context<'_>,
) -> std::task::Poll<std::io::Result<()>> {
std::task::Poll::Ready(Ok(()))
}
}
#[tokio::test]
async fn serve_router_with_qos_rejects_excess_connections() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("listener should bind");
let bind = listener
.local_addr()
.expect("listener should expose address")
.to_string();
let config = HttpConfig::new(
HttpBindAddress::parse(&bind).expect("bind should parse"),
Duration::from_secs(5),
Duration::from_millis(100),
1024,
HttpProxyConfig::new(None, Vec::new(), true).expect("proxy should build"),
)
.expect("config should build");
let router = Router::new().route("/ok", get(|| async { "ok" }));
let qos = QosRuntime::default();
let policy = QosPolicy::new(1, 4, 4).expect("policy should build");
let (shutdown, shutdown_waiter) = tokio::sync::oneshot::channel();
let server_qos = qos.clone();
let listener = QosTcpListener::new(listener, server_qos, policy);
let server = tokio::spawn(serve_listener(
listener,
router,
config,
Some(qos.clone()),
async {
let _ = shutdown_waiter.await;
},
));
let first = connect_with_retry(&bind).await;
wait_for_connection_count(&qos, 1).await;
let second = connect_with_retry(&bind).await;
wait_for_peer_close(&second).await;
drop(first);
let _ = shutdown.send(());
server
.await
.expect("server task should join")
.expect("server should stop");
}
async fn connect_with_retry(bind: &str) -> tokio::net::TcpStream {
for _ in 0..50 {
if let Ok(stream) = tokio::net::TcpStream::connect(bind).await {
return stream;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
panic!("server did not accept connections on {bind}");
}
async fn wait_for_connection_count(qos: &QosRuntime, expected: usize) {
for _ in 0..50 {
if qos.snapshot().connections == expected {
return;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
panic!("connection count did not reach {expected}");
}
async fn wait_for_peer_close(stream: &tokio::net::TcpStream) {
let mut buffer = [0; 1024];
for _ in 0..50 {
stream.readable().await.expect("stream should be readable");
match stream.try_read(&mut buffer) {
Ok(0) => return,
Ok(_) => panic!("over-budget connection should close before serving data"),
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {}
Err(error) if error.kind() == std::io::ErrorKind::ConnectionReset => return,
Err(error) => panic!("response read failed: {error}"),
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
panic!("server did not close over-budget connection");
}