use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use arcature::prelude::*;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
const HANG_GUARD: Duration = Duration::from_secs(10);
struct RunningApp {
addr: SocketAddr,
shutdown: tokio::sync::oneshot::Sender<()>,
join: tokio::task::JoinHandle<()>,
}
impl RunningApp {
async fn start(app: Application<()>) -> RunningApp {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("bind ephemeral listener");
let addr = listener.local_addr().expect("read bound address");
let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel::<()>();
let join = tokio::spawn(async move {
app.serve_with_shutdown(listener, async {
let _ = shutdown_rx.await;
})
.await
.expect("application served without engine error");
});
RunningApp {
addr,
shutdown: shutdown_tx,
join,
}
}
fn addr(&self) -> SocketAddr {
self.addr
}
async fn stop(self) {
self.shutdown
.send(())
.expect("server task still alive to receive shutdown");
tokio::time::timeout(HANG_GUARD, self.join)
.await
.expect("server did not hang on shutdown")
.expect("server task did not panic");
}
}
async fn http_request(
addr: SocketAddr,
method: &str,
path: &str,
extra_headers: &[(&str, &str)],
) -> String {
let mut stream = tokio::time::timeout(HANG_GUARD, TcpStream::connect(addr))
.await
.expect("connect did not hang")
.expect("connect succeeds");
let mut request = format!("{method} {path} HTTP/1.1\r\nHost: localhost\r\n");
for (name, value) in extra_headers {
request.push_str(&format!("{name}: {value}\r\n"));
}
request.push_str("Connection: close\r\n\r\n");
stream
.write_all(request.as_bytes())
.await
.expect("write request");
let mut buffer = Vec::new();
tokio::time::timeout(HANG_GUARD, stream.read_to_end(&mut buffer))
.await
.expect("read did not hang")
.expect("read succeeds");
String::from_utf8_lossy(&buffer).into_owned()
}
fn status_line(response: &str) -> &str {
response.split("\r\n").next().unwrap_or(response)
}
fn body(response: &str) -> &str {
response
.split_once("\r\n\r\n")
.map(|(_, b)| b)
.unwrap_or("")
}
fn header<'a>(response: &'a str, name: &str) -> Option<&'a str> {
let lower = name.to_ascii_lowercase();
for line in response.split("\r\n") {
if let Some((k, v)) = line.split_once(": ")
&& k.to_ascii_lowercase() == lower
{
return Some(v.trim());
}
}
None
}
#[tokio::test]
async fn pre_routing_rewrite_hits_registered_target() {
let app = Application::new()
.routes(Routes::new().route("/new", get(|| async { "rewritten" })))
.proxy(|req| {
if req.uri().path() == "/old" {
ProxyAction::Rewrite {
uri: "/new".to_owned(),
}
} else {
ProxyAction::continue_default()
}
})
.build();
let server = RunningApp::start(app).await;
let response = http_request(server.addr(), "GET", "/old", &[]).await;
assert!(
response.starts_with("HTTP/1.1 200 OK"),
"expected 200 OK after rewrite, got: {}",
status_line(&response)
);
assert_eq!(body(&response), "rewritten");
let response = http_request(server.addr(), "GET", "/new", &[]).await;
assert!(
response.starts_with("HTTP/1.1 200 OK"),
"direct /new: {}",
status_line(&response)
);
server.stop().await;
}
#[tokio::test]
async fn pre_routing_rewrite_with_query_string() {
async fn echo_query(uri: Uri) -> String {
uri.query().unwrap_or("no query").to_owned()
}
let app = Application::new()
.routes(Routes::new().route("/search", get(echo_query)))
.proxy(|req| {
if req.uri().path() == "/find" {
ProxyAction::Rewrite {
uri: "/search?q=hello".to_owned(),
}
} else {
ProxyAction::continue_default()
}
})
.build();
let server = RunningApp::start(app).await;
let response = http_request(server.addr(), "GET", "/find", &[]).await;
assert!(
response.starts_with("HTTP/1.1 200 OK"),
"{}",
status_line(&response)
);
assert_eq!(body(&response), "q=hello");
server.stop().await;
}
#[tokio::test]
async fn proxy_redirect_302_found() {
let app = Application::new()
.routes(Routes::new().route("/here", get(|| async { "here" })))
.proxy(|_req| ProxyAction::Redirect {
location: "/here".to_owned(),
permanent: false,
})
.build();
let server = RunningApp::start(app).await;
let response = http_request(server.addr(), "GET", "/anything", &[]).await;
assert!(
response.starts_with("HTTP/1.1 302 Found"),
"expected 302, got: {}",
status_line(&response)
);
assert_eq!(header(&response, "location"), Some("/here"));
server.stop().await;
}
#[tokio::test]
async fn proxy_redirect_301_moved_permanently() {
let app = Application::new()
.routes(Routes::new().route("/here", get(|| async { "here" })))
.proxy(|_req| ProxyAction::Redirect {
location: "/here".to_owned(),
permanent: true,
})
.build();
let server = RunningApp::start(app).await;
let response = http_request(server.addr(), "GET", "/anything", &[]).await;
assert!(
response.starts_with("HTTP/1.1 301 Moved Permanently"),
"expected 301, got: {}",
status_line(&response)
);
assert_eq!(header(&response, "location"), Some("/here"));
server.stop().await;
}
#[tokio::test]
async fn proxy_short_circuit_with_status() {
let handler_called = Arc::new(AtomicUsize::new(0));
let counter = handler_called.clone();
let app = Application::new()
.routes(Routes::new().route(
"/",
get(move || {
let c = counter.clone();
async move {
c.fetch_add(1, Ordering::SeqCst);
"handler ran"
}
}),
))
.proxy(|_req| ProxyAction::ShortCircuit {
status: StatusCode::NOT_FOUND,
response: None,
})
.build();
let server = RunningApp::start(app).await;
let response = http_request(server.addr(), "GET", "/", &[]).await;
assert!(
response.starts_with("HTTP/1.1 404 Not Found"),
"expected 404, got: {}",
status_line(&response)
);
assert_eq!(
handler_called.load(Ordering::SeqCst),
0,
"handler must not run on short-circuit"
);
server.stop().await;
}
#[tokio::test]
async fn proxy_short_circuit_with_custom_response() {
let app = Application::new()
.routes(Routes::new().route("/", get(|| async { "handler" })))
.proxy(|_req| ProxyAction::ShortCircuit {
status: StatusCode::OK,
response: Some(("short-circuited").into_response()),
})
.build();
let server = RunningApp::start(app).await;
let response = http_request(server.addr(), "GET", "/", &[]).await;
assert!(
response.starts_with("HTTP/1.1 200 OK"),
"{}",
status_line(&response)
);
assert_eq!(body(&response), "short-circuited");
server.stop().await;
}
#[tokio::test]
async fn proxy_continue_with_set_headers() {
async fn echo_header(headers: HeaderMap) -> String {
headers
.get("x-proxy-marker")
.map(|v| v.to_str().unwrap_or_default().to_owned())
.unwrap_or_else(|| "no marker".to_owned())
}
let app = Application::new()
.routes(Routes::new().route("/", get(echo_header)))
.proxy(|_req| {
let mut headers = HeaderMap::new();
headers.insert("x-proxy-marker", "set-by-proxy".parse().unwrap());
ProxyAction::Continue {
set_headers: headers,
}
})
.build();
let server = RunningApp::start(app).await;
let response = http_request(server.addr(), "GET", "/", &[]).await;
assert!(
response.starts_with("HTTP/1.1 200 OK"),
"{}",
status_line(&response)
);
assert_eq!(body(&response), "set-by-proxy");
server.stop().await;
}
#[tokio::test]
async fn proxy_invalid_rewrite_uri_rejected_400() {
let app = Application::new()
.routes(Routes::new().route("/new", get(|| async { "ok" })))
.proxy(|_req| ProxyAction::Rewrite {
uri: "not a valid uri }}}".to_owned(),
})
.build();
let server = RunningApp::start(app).await;
let response = http_request(server.addr(), "GET", "/old", &[]).await;
assert!(
response.starts_with("HTTP/1.1 400 Bad Request"),
"expected 400 for invalid rewrite URI, got: {}",
status_line(&response)
);
server.stop().await;
}
#[tokio::test]
async fn proxy_rewrite_with_scheme_rejected_400() {
let app = Application::new()
.routes(Routes::new().route("/new", get(|| async { "ok" })))
.proxy(|_req| ProxyAction::Rewrite {
uri: "https://evil.example.com/new".to_owned(),
})
.build();
let server = RunningApp::start(app).await;
let response = http_request(server.addr(), "GET", "/old", &[]).await;
assert!(
response.starts_with("HTTP/1.1 400 Bad Request"),
"expected 400 for scheme-injected rewrite, got: {}",
status_line(&response)
);
server.stop().await;
}
#[tokio::test]
async fn proxy_redirect_crlf_injection_rejected_400() {
let app = Application::new()
.routes(Routes::new().route("/here", get(|| async { "here" })))
.proxy(|_req| ProxyAction::Redirect {
location: "/here\r\nSet-Cookie: stolen=1".to_owned(),
permanent: false,
})
.build();
let server = RunningApp::start(app).await;
let response = http_request(server.addr(), "GET", "/anything", &[]).await;
assert!(
response.starts_with("HTTP/1.1 400 Bad Request"),
"expected 400 for CRLF in redirect, got: {}",
status_line(&response)
);
assert!(
!response.contains("stolen"),
"CRLF-injected header must not appear in response"
);
server.stop().await;
}
#[tokio::test]
async fn proxy_set_headers_crlf_rejected_by_headervalue() {
let result = "value\r\nInjected: yes".parse::<arcature::axum::http::HeaderValue>();
assert!(
result.is_err(),
"HeaderValue must reject CRLF in its value — CRLF injection defense"
);
}
#[tokio::test]
async fn no_proxy_passes_through_unchanged() {
let app = Application::new()
.routes(Routes::new().route("/", get(|| async { "no proxy" })))
.build();
let server = RunningApp::start(app).await;
let response = http_request(server.addr(), "GET", "/", &[]).await;
assert!(
response.starts_with("HTTP/1.1 200 OK"),
"{}",
status_line(&response)
);
assert_eq!(body(&response), "no proxy");
server.stop().await;
}
#[tokio::test]
async fn proxy_continue_default_passes_through() {
let app = Application::new()
.routes(Routes::new().route("/", get(|| async { "continued" })))
.proxy(|_req| ProxyAction::continue_default())
.build();
let server = RunningApp::start(app).await;
let response = http_request(server.addr(), "GET", "/", &[]).await;
assert!(
response.starts_with("HTTP/1.1 200 OK"),
"{}",
status_line(&response)
);
assert_eq!(body(&response), "continued");
server.stop().await;
}