1632 lines
54 KiB
Rust
1632 lines
54 KiB
Rust
mod helpers;
|
|
|
|
use std::path::Path;
|
|
use std::process::Command;
|
|
use std::sync::Arc;
|
|
use std::time::Duration;
|
|
|
|
use arc_swap::ArcSwap;
|
|
use axum::routing::{get, post};
|
|
use axum::Router;
|
|
|
|
use reverse_proxy::config::dynamic_config::{
|
|
BodyConfig, DynamicConfig, RateLimitConfig, SiteConfig,
|
|
};
|
|
use reverse_proxy::proxy::body_limit::DEFAULT_BODY_LIMIT_BYTES;
|
|
use reverse_proxy::proxy::router_with_body_limit;
|
|
use tower::ServiceExt;
|
|
|
|
#[tokio::test]
|
|
async fn test_upstream_spawn_and_connect() {
|
|
let upstream = helpers::http_test_helper::TestUpstream::spawn_ok().await;
|
|
let client = reqwest::Client::new();
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/", upstream.addr.port()))
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(resp.status(), reqwest::StatusCode::OK);
|
|
let _ = upstream.shutdown_tx.send(());
|
|
}
|
|
|
|
#[test]
|
|
fn test_self_signed_cert_generation() {
|
|
let cert = helpers::tls_test_helper::generate_self_signed_cert(&["test.local"]);
|
|
assert!(!cert.cert_pem.is_empty());
|
|
assert!(!cert.key_pem.is_empty());
|
|
assert!(cert.cert_pem.contains("BEGIN CERTIFICATE"));
|
|
assert!(cert.key_pem.contains("BEGIN"));
|
|
}
|
|
|
|
#[test]
|
|
fn test_config_fixtures() {
|
|
let static_config = reverse_proxy::config::test_fixtures::test_static_config();
|
|
assert!(!static_config.listeners.is_empty());
|
|
|
|
let dynamic_config = reverse_proxy::config::test_fixtures::test_dynamic_config();
|
|
assert!(!dynamic_config.sites.is_empty());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_health_check_local_port_returns_200() {
|
|
let (addr, handle) = reverse_proxy::health::start_health_check_listener(0, None, None)
|
|
.await
|
|
.unwrap();
|
|
|
|
let client = reqwest::Client::new();
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/health", addr.port()))
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(resp.status(), reqwest::StatusCode::OK);
|
|
let body = resp.text().await.unwrap();
|
|
assert!(body.is_empty());
|
|
|
|
handle.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_health_check_local_port_binds_localhost() {
|
|
let (addr, handle) = reverse_proxy::health::start_health_check_listener(0, None, None)
|
|
.await
|
|
.unwrap();
|
|
|
|
assert!(addr.ip().is_loopback());
|
|
assert_eq!(addr.ip().to_string(), "127.0.0.1");
|
|
|
|
handle.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_health_check_binds_random_port_when_zero() {
|
|
let result = reverse_proxy::health::start_health_check_listener(0, None, None).await;
|
|
assert!(result.is_ok());
|
|
let (addr, handle) = result.unwrap();
|
|
assert_ne!(addr.port(), 0);
|
|
handle.abort();
|
|
}
|
|
|
|
fn make_rate_limit_app(
|
|
limiter: Arc<reverse_proxy::rate_limit::RateLimiter>,
|
|
) -> axum::extract::connect_info::IntoMakeServiceWithConnectInfo<Router, std::net::SocketAddr> {
|
|
Router::new()
|
|
.route("/", get(|| async { "ok" }))
|
|
.layer(axum::middleware::from_fn_with_state(
|
|
limiter,
|
|
reverse_proxy::rate_limit::rate_limit_middleware,
|
|
))
|
|
.into_make_service_with_connect_info::<std::net::SocketAddr>()
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_rate_limit_allows_within_burst() {
|
|
let mut config = reverse_proxy::config::test_fixtures::test_dynamic_config();
|
|
config.rate_limit = reverse_proxy::config::RateLimitConfig {
|
|
requests_per_second: 10,
|
|
burst: 5,
|
|
};
|
|
let config_arc = Arc::new(ArcSwap::from_pointee(config));
|
|
let limiter = Arc::new(reverse_proxy::rate_limit::RateLimiter::new(config_arc));
|
|
|
|
let app = make_rate_limit_app(limiter);
|
|
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
|
let addr = listener.local_addr().unwrap();
|
|
|
|
tokio::spawn(async { axum::serve(listener, app).await.unwrap() });
|
|
|
|
let client = reqwest::Client::new();
|
|
for _ in 0..5 {
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/", addr.port()))
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(resp.status(), reqwest::StatusCode::OK);
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_rate_limit_rejects_above_burst() {
|
|
let mut config = reverse_proxy::config::test_fixtures::test_dynamic_config();
|
|
config.rate_limit = reverse_proxy::config::RateLimitConfig {
|
|
requests_per_second: 10,
|
|
burst: 2,
|
|
};
|
|
let config_arc = Arc::new(ArcSwap::from_pointee(config));
|
|
let limiter = Arc::new(reverse_proxy::rate_limit::RateLimiter::new(config_arc));
|
|
|
|
let app = make_rate_limit_app(limiter);
|
|
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
|
let addr = listener.local_addr().unwrap();
|
|
|
|
tokio::spawn(async { axum::serve(listener, app).await.unwrap() });
|
|
|
|
let client = reqwest::Client::new();
|
|
for _ in 0..2 {
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/", addr.port()))
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(resp.status(), reqwest::StatusCode::OK);
|
|
}
|
|
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/", addr.port()))
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(resp.status(), reqwest::StatusCode::TOO_MANY_REQUESTS);
|
|
let body = resp.text().await.unwrap();
|
|
assert_eq!(body, "Too Many Requests");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_rate_limit_429_response_body() {
|
|
let mut config = reverse_proxy::config::test_fixtures::test_dynamic_config();
|
|
config.rate_limit = reverse_proxy::config::RateLimitConfig {
|
|
requests_per_second: 10,
|
|
burst: 1,
|
|
};
|
|
let config_arc = Arc::new(ArcSwap::from_pointee(config));
|
|
let limiter = Arc::new(reverse_proxy::rate_limit::RateLimiter::new(config_arc));
|
|
|
|
let app = make_rate_limit_app(limiter);
|
|
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
|
let addr = listener.local_addr().unwrap();
|
|
|
|
tokio::spawn(async { axum::serve(listener, app).await.unwrap() });
|
|
|
|
let client = reqwest::Client::new();
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/", addr.port()))
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(resp.status(), reqwest::StatusCode::OK);
|
|
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/", addr.port()))
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(resp.status(), reqwest::StatusCode::TOO_MANY_REQUESTS);
|
|
let body = resp.text().await.unwrap();
|
|
assert_eq!(body, "Too Many Requests");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_rate_limit_per_ip_independent() {
|
|
let mut config = reverse_proxy::config::test_fixtures::test_dynamic_config();
|
|
config.rate_limit = reverse_proxy::config::RateLimitConfig {
|
|
requests_per_second: 10,
|
|
burst: 1,
|
|
};
|
|
let config_arc = Arc::new(ArcSwap::from_pointee(config));
|
|
let limiter = Arc::new(reverse_proxy::rate_limit::RateLimiter::new(config_arc));
|
|
|
|
let app = make_rate_limit_app(limiter);
|
|
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
|
let addr = listener.local_addr().unwrap();
|
|
|
|
tokio::spawn(async { axum::serve(listener, app).await.unwrap() });
|
|
|
|
let client = reqwest::Client::new();
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/", addr.port()))
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(resp.status(), reqwest::StatusCode::OK);
|
|
|
|
let resp2 = client
|
|
.get(format!("http://127.0.0.1:{}/", addr.port()))
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(resp2.status(), reqwest::StatusCode::TOO_MANY_REQUESTS);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_rate_limit_without_connect_info_rejected_with_429() {
|
|
let mut config = reverse_proxy::config::test_fixtures::test_dynamic_config();
|
|
config.rate_limit = reverse_proxy::config::RateLimitConfig {
|
|
requests_per_second: 10,
|
|
burst: 20,
|
|
};
|
|
let config_arc = Arc::new(ArcSwap::from_pointee(config));
|
|
let limiter = Arc::new(reverse_proxy::rate_limit::RateLimiter::new(config_arc));
|
|
|
|
let app = Router::new().route("/", get(|| async { "ok" })).layer(
|
|
axum::middleware::from_fn_with_state(
|
|
limiter,
|
|
reverse_proxy::rate_limit::rate_limit_middleware,
|
|
),
|
|
);
|
|
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
|
let addr = listener.local_addr().unwrap();
|
|
|
|
tokio::spawn(async { axum::serve(listener, app).await.unwrap() });
|
|
|
|
let client = reqwest::Client::new();
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/", addr.port()))
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(resp.status(), reqwest::StatusCode::TOO_MANY_REQUESTS);
|
|
let body = resp.text().await.unwrap();
|
|
assert_eq!(body, "Too Many Requests");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_rate_limit_xff_header_ignored_same_bucket() {
|
|
let mut config = reverse_proxy::config::test_fixtures::test_dynamic_config();
|
|
config.rate_limit = reverse_proxy::config::RateLimitConfig {
|
|
requests_per_second: 10,
|
|
burst: 2,
|
|
};
|
|
let config_arc = Arc::new(ArcSwap::from_pointee(config));
|
|
let limiter = Arc::new(reverse_proxy::rate_limit::RateLimiter::new(config_arc));
|
|
|
|
let app = make_rate_limit_app(limiter);
|
|
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
|
let addr = listener.local_addr().unwrap();
|
|
|
|
tokio::spawn(async { axum::serve(listener, app).await.unwrap() });
|
|
|
|
let client = reqwest::Client::new();
|
|
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/", addr.port()))
|
|
.header("X-Forwarded-For", "10.0.0.1")
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(resp.status(), reqwest::StatusCode::OK);
|
|
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/", addr.port()))
|
|
.header("X-Forwarded-For", "10.0.0.2")
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(resp.status(), reqwest::StatusCode::OK);
|
|
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/", addr.port()))
|
|
.header("X-Forwarded-For", "10.0.0.3")
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(resp.status(), reqwest::StatusCode::TOO_MANY_REQUESTS);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_rate_limit_eviction_task() {
|
|
let mut config = reverse_proxy::config::test_fixtures::test_dynamic_config();
|
|
config.rate_limit = reverse_proxy::config::RateLimitConfig {
|
|
requests_per_second: 10,
|
|
burst: 20,
|
|
};
|
|
let config_arc = Arc::new(ArcSwap::from_pointee(config));
|
|
let limiter = Arc::new(reverse_proxy::rate_limit::RateLimiter::new(config_arc));
|
|
|
|
limiter.check_and_consume(std::net::IpAddr::from([192, 168, 1, 1]));
|
|
|
|
let shutdown = Arc::new(reverse_proxy::shutdown::GracefulShutdown::new(30));
|
|
let handle = reverse_proxy::rate_limit::start_eviction_task(
|
|
limiter.clone(),
|
|
Duration::from_millis(50),
|
|
Duration::from_millis(100),
|
|
shutdown.subscribe(),
|
|
);
|
|
|
|
tokio::time::sleep(Duration::from_millis(200)).await;
|
|
|
|
assert!(!limiter.contains_ip(std::net::IpAddr::from([192, 168, 1, 1])));
|
|
|
|
handle.abort();
|
|
}
|
|
|
|
fn make_redirect_listener_config(
|
|
bind_addr: &str,
|
|
http_port: u16,
|
|
https_port: u16,
|
|
) -> reverse_proxy::config::static_config::ListenerConfig {
|
|
reverse_proxy::config::static_config::ListenerConfig {
|
|
bind_addr: bind_addr.to_string(),
|
|
http_port,
|
|
https_port,
|
|
tls: reverse_proxy::config::static_config::TlsConfig {
|
|
mode: "manual".to_string(),
|
|
acme_domains: vec![],
|
|
acme_cache_dir: String::new(),
|
|
acme_directory: "production".to_string(),
|
|
acme_contact: String::new(),
|
|
cert_path: String::new(),
|
|
key_path: String::new(),
|
|
},
|
|
sites: vec![],
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_http_redirect_returns_301_with_location() {
|
|
let config = make_redirect_listener_config("127.0.0.1", 0, 443);
|
|
let (addr, handle) = reverse_proxy::tls::redirect::start_http_redirect_listener(&config)
|
|
.await
|
|
.unwrap();
|
|
|
|
let client = reqwest::Client::builder()
|
|
.redirect(reqwest::redirect::Policy::none())
|
|
.build()
|
|
.unwrap();
|
|
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/some/path", addr.port()))
|
|
.header("Host", "example.com")
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(resp.status(), reqwest::StatusCode::MOVED_PERMANENTLY);
|
|
let location = resp.headers().get("location").unwrap().to_str().unwrap();
|
|
assert_eq!(location, "https://example.com/some/path");
|
|
|
|
handle.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_http_redirect_port_443_omitted_from_url() {
|
|
let config = make_redirect_listener_config("127.0.0.1", 0, 443);
|
|
let (addr, handle) = reverse_proxy::tls::redirect::start_http_redirect_listener(&config)
|
|
.await
|
|
.unwrap();
|
|
|
|
let client = reqwest::Client::builder()
|
|
.redirect(reqwest::redirect::Policy::none())
|
|
.build()
|
|
.unwrap();
|
|
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/", addr.port()))
|
|
.header("Host", "example.com")
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
|
|
let location = resp.headers().get("location").unwrap().to_str().unwrap();
|
|
assert_eq!(location, "https://example.com/");
|
|
|
|
handle.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_http_redirect_non_443_port_included_in_url() {
|
|
let config = make_redirect_listener_config("127.0.0.1", 0, 8443);
|
|
let (addr, handle) = reverse_proxy::tls::redirect::start_http_redirect_listener(&config)
|
|
.await
|
|
.unwrap();
|
|
|
|
let client = reqwest::Client::builder()
|
|
.redirect(reqwest::redirect::Policy::none())
|
|
.build()
|
|
.unwrap();
|
|
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/", addr.port()))
|
|
.header("Host", "example.com")
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
|
|
let location = resp.headers().get("location").unwrap().to_str().unwrap();
|
|
assert_eq!(location, "https://example.com:8443/");
|
|
|
|
handle.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_http_redirect_empty_host_returns_400() {
|
|
let config = make_redirect_listener_config("127.0.0.1", 0, 443);
|
|
let (addr, handle) = reverse_proxy::tls::redirect::start_http_redirect_listener(&config)
|
|
.await
|
|
.unwrap();
|
|
|
|
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
|
let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
|
|
stream
|
|
.write_all(b"GET / HTTP/1.1\r\nHost: \r\nConnection: close\r\n\r\n")
|
|
.await
|
|
.unwrap();
|
|
|
|
let mut response = vec![0u8; 4096];
|
|
let n = tokio::time::timeout(
|
|
std::time::Duration::from_secs(5),
|
|
stream.read(&mut response),
|
|
)
|
|
.await
|
|
.unwrap()
|
|
.unwrap();
|
|
let response_str = String::from_utf8_lossy(&response[..n]);
|
|
assert!(
|
|
response_str.contains(" 400 "),
|
|
"expected 400 status, got: {response_str}"
|
|
);
|
|
|
|
handle.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_http_redirect_no_host_header_returns_400() {
|
|
let config = make_redirect_listener_config("127.0.0.1", 0, 443);
|
|
let (addr, handle) = reverse_proxy::tls::redirect::start_http_redirect_listener(&config)
|
|
.await
|
|
.unwrap();
|
|
|
|
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
|
let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
|
|
stream
|
|
.write_all(b"GET / HTTP/1.0\r\nConnection: close\r\n\r\n")
|
|
.await
|
|
.unwrap();
|
|
|
|
let mut response = vec![0u8; 4096];
|
|
let n = tokio::time::timeout(
|
|
std::time::Duration::from_secs(5),
|
|
stream.read(&mut response),
|
|
)
|
|
.await
|
|
.unwrap()
|
|
.unwrap();
|
|
let response_str = String::from_utf8_lossy(&response[..n]);
|
|
assert!(
|
|
response_str.contains(" 400 "),
|
|
"expected 400 status, got: {response_str}"
|
|
);
|
|
|
|
handle.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_http_redirect_strips_host_port() {
|
|
let config = make_redirect_listener_config("127.0.0.1", 0, 443);
|
|
let (addr, handle) = reverse_proxy::tls::redirect::start_http_redirect_listener(&config)
|
|
.await
|
|
.unwrap();
|
|
|
|
let client = reqwest::Client::builder()
|
|
.redirect(reqwest::redirect::Policy::none())
|
|
.build()
|
|
.unwrap();
|
|
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/path", addr.port()))
|
|
.header("Host", "example.com:8080")
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
|
|
let location = resp.headers().get("location").unwrap().to_str().unwrap();
|
|
assert_eq!(location, "https://example.com/path");
|
|
|
|
handle.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_http_redirect_preserves_query_string() {
|
|
let config = make_redirect_listener_config("127.0.0.1", 0, 443);
|
|
let (addr, handle) = reverse_proxy::tls::redirect::start_http_redirect_listener(&config)
|
|
.await
|
|
.unwrap();
|
|
|
|
let client = reqwest::Client::builder()
|
|
.redirect(reqwest::redirect::Policy::none())
|
|
.build()
|
|
.unwrap();
|
|
|
|
let resp = client
|
|
.get(format!(
|
|
"http://127.0.0.1:{}/search?q=test&page=1",
|
|
addr.port()
|
|
))
|
|
.header("Host", "git.alk.dev")
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
|
|
let location = resp.headers().get("location").unwrap().to_str().unwrap();
|
|
assert_eq!(location, "https://git.alk.dev/search?q=test&page=1");
|
|
|
|
handle.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_http_redirect_acme_challenge_returns_404() {
|
|
let config = make_redirect_listener_config("127.0.0.1", 0, 443);
|
|
let (addr, handle) = reverse_proxy::tls::redirect::start_http_redirect_listener(&config)
|
|
.await
|
|
.unwrap();
|
|
|
|
let client = reqwest::Client::new();
|
|
let resp = client
|
|
.get(format!(
|
|
"http://127.0.0.1:{}/.well-known/acme-challenge/abc123",
|
|
addr.port()
|
|
))
|
|
.header("Host", "example.com")
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(resp.status(), reqwest::StatusCode::NOT_FOUND);
|
|
|
|
handle.abort();
|
|
}
|
|
|
|
async fn spawn_echoing_upstream() -> helpers::http_test_helper::TestUpstream {
|
|
helpers::http_test_helper::TestUpstream::spawn(|| {
|
|
Router::new().route(
|
|
"/",
|
|
get(|req: axum::extract::Request| async move {
|
|
let host = req
|
|
.headers()
|
|
.get("host")
|
|
.map(|v| v.to_str().unwrap().to_string())
|
|
.unwrap_or_default();
|
|
let proto = req
|
|
.headers()
|
|
.get("x-forwarded-proto")
|
|
.map(|v| v.to_str().unwrap().to_string())
|
|
.unwrap_or_default();
|
|
format!("host={}|proto={}", host, proto)
|
|
}),
|
|
)
|
|
})
|
|
.await
|
|
}
|
|
|
|
fn make_site_config(upstream_addr: &str) -> SiteConfig {
|
|
SiteConfig {
|
|
host: "test.local".to_string(),
|
|
upstream: upstream_addr.to_string(),
|
|
upstream_scheme: "http".to_string(),
|
|
upstream_connect_timeout_secs: 5,
|
|
upstream_request_timeout_secs: 60,
|
|
}
|
|
}
|
|
|
|
fn make_dynamic_config_with_site(upstream_addr: &str) -> DynamicConfig {
|
|
DynamicConfig::from_sites(
|
|
vec![make_site_config(upstream_addr)],
|
|
RateLimitConfig {
|
|
requests_per_second: 100,
|
|
burst: 100,
|
|
},
|
|
BodyConfig {
|
|
limit_bytes: 104857600,
|
|
},
|
|
)
|
|
}
|
|
|
|
fn make_https_test_proxy_state(upstream_addr: &str) -> Arc<reverse_proxy::proxy::ProxyState> {
|
|
Arc::new(reverse_proxy::proxy::ProxyState {
|
|
config: Arc::new(ArcSwap::from_pointee(make_dynamic_config_with_site(
|
|
upstream_addr,
|
|
))),
|
|
http_client: reverse_proxy::proxy::create_http_client(),
|
|
https_client: reverse_proxy::proxy::create_https_client(),
|
|
})
|
|
}
|
|
|
|
// Regression test: Gitea's self-check reported "Current URL doesn't match the
|
|
// URL seen by Gitea" because HTTP/2 requests carried the upstream authority
|
|
// (127.0.0.1:3000) as the Host header instead of the client's :authority
|
|
// (test.local). HTTP/2 requests have no literal Host header, so the proxy
|
|
// must reconstruct it from the request URI's authority. Without the fix, the
|
|
// hyper-util client (set_host=true) inserts the upstream authority.
|
|
#[tokio::test]
|
|
async fn test_proxy_forwards_client_authority_as_host() {
|
|
let upstream = spawn_echoing_upstream().await;
|
|
let upstream_addr = format!("127.0.0.1:{}", upstream.addr.port());
|
|
let proxy_state = make_https_test_proxy_state(&upstream_addr);
|
|
let config_arc = Arc::new(ArcSwap::from_pointee(make_dynamic_config_with_site(
|
|
&upstream_addr,
|
|
)));
|
|
let rate_limiter =
|
|
Arc::new(reverse_proxy::rate_limit::RateLimiter::new(config_arc.clone()));
|
|
let router = reverse_proxy::proxy::build_router(proxy_state, config_arc, rate_limiter);
|
|
|
|
// Simulates an HTTP/2 request translated by hyper: the :authority
|
|
// pseudo-header becomes the URI authority and no Host header is present.
|
|
let mut req = axum::http::Request::builder()
|
|
.method("GET")
|
|
.uri("http://test.local/")
|
|
.header("x-forwarded-proto", "https")
|
|
.body(axum::body::Body::empty())
|
|
.unwrap();
|
|
req.extensions_mut().insert(axum::extract::ConnectInfo(
|
|
std::net::SocketAddr::from(([127, 0, 0, 1], 54321)),
|
|
));
|
|
let resp = router.oneshot(req).await.unwrap();
|
|
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
|
|
let body = String::from_utf8(body.to_vec()).unwrap();
|
|
assert_eq!(body, "host=test.local|proto=https");
|
|
|
|
let _ = upstream.shutdown_tx.send(());
|
|
}
|
|
|
|
// HTTP/1.1 requests carry a real Host header and an origin-form request
|
|
// target — the Host header must be routed on and forwarded unchanged.
|
|
#[tokio::test]
|
|
async fn test_proxy_preserves_client_host_header_for_http1() {
|
|
let upstream = spawn_echoing_upstream().await;
|
|
let upstream_addr = format!("127.0.0.1:{}", upstream.addr.port());
|
|
let proxy_state = make_https_test_proxy_state(&upstream_addr);
|
|
let config_arc = Arc::new(ArcSwap::from_pointee(make_dynamic_config_with_site(
|
|
&upstream_addr,
|
|
)));
|
|
let rate_limiter =
|
|
Arc::new(reverse_proxy::rate_limit::RateLimiter::new(config_arc.clone()));
|
|
let router = reverse_proxy::proxy::build_router(proxy_state, config_arc, rate_limiter);
|
|
|
|
let mut req = axum::http::Request::builder()
|
|
.method("GET")
|
|
.uri("/")
|
|
.header("host", "test.local")
|
|
.header("x-forwarded-proto", "https")
|
|
.body(axum::body::Body::empty())
|
|
.unwrap();
|
|
req.extensions_mut().insert(axum::extract::ConnectInfo(
|
|
std::net::SocketAddr::from(([127, 0, 0, 1], 54322)),
|
|
));
|
|
let resp = router.oneshot(req).await.unwrap();
|
|
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
|
|
let body = String::from_utf8(body.to_vec()).unwrap();
|
|
assert_eq!(body, "host=test.local|proto=https");
|
|
|
|
let _ = upstream.shutdown_tx.send(());
|
|
}
|
|
|
|
fn write_valid_config(dir: &Path) -> std::path::PathBuf {
|
|
let config_path = dir.join("config.toml");
|
|
let config = r#"
|
|
health_check_port = 9900
|
|
admin_key_path = "/etc/reverse-proxy/admin-key"
|
|
|
|
[logging]
|
|
level = "info"
|
|
format = "text"
|
|
|
|
[[listeners]]
|
|
bind_addr = "127.0.0.1"
|
|
https_port = 443
|
|
|
|
[listeners.tls]
|
|
mode = "acme"
|
|
acme_domains = ["test.local"]
|
|
acme_cache_dir = "/tmp/acme-cache"
|
|
acme_contact = "mailto:admin@test.local"
|
|
|
|
[[listeners.sites]]
|
|
host = "test.local"
|
|
upstream = "127.0.0.1:8080"
|
|
|
|
[rate_limit]
|
|
requests_per_second = 10
|
|
burst = 20
|
|
|
|
[body]
|
|
limit_bytes = 104857600
|
|
"#;
|
|
std::fs::write(&config_path, config).unwrap();
|
|
config_path
|
|
}
|
|
|
|
fn write_invalid_config(dir: &Path) -> std::path::PathBuf {
|
|
let config_path = dir.join("config.toml");
|
|
let config = r#"
|
|
health_check_port = 9900
|
|
"#;
|
|
std::fs::write(&config_path, config).unwrap();
|
|
config_path
|
|
}
|
|
|
|
fn binary_path() -> std::path::PathBuf {
|
|
std::path::PathBuf::from(env!("CARGO_BIN_EXE_reverse-proxy"))
|
|
}
|
|
|
|
#[test]
|
|
fn test_validate_valid_config_exits_0() {
|
|
let dir = tempfile::tempdir().unwrap();
|
|
let config_path = write_valid_config(dir.path());
|
|
let output = Command::new(binary_path())
|
|
.arg("--config")
|
|
.arg(config_path.to_str().unwrap())
|
|
.arg("--validate")
|
|
.output()
|
|
.expect("failed to run binary");
|
|
assert_eq!(
|
|
output.status.code(),
|
|
Some(0),
|
|
"expected exit 0 with valid config, got {}: stderr={}",
|
|
output.status,
|
|
String::from_utf8_lossy(&output.stderr)
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn test_validate_invalid_config_exits_1() {
|
|
let dir = tempfile::tempdir().unwrap();
|
|
let config_path = write_invalid_config(dir.path());
|
|
let output = Command::new(binary_path())
|
|
.arg("--config")
|
|
.arg(config_path.to_str().unwrap())
|
|
.arg("--validate")
|
|
.output()
|
|
.expect("failed to run binary");
|
|
assert!(
|
|
output.status.code() == Some(1) || output.status.code() == Some(2),
|
|
"expected non-zero exit with invalid config, got {}: stderr={}",
|
|
output.status,
|
|
String::from_utf8_lossy(&output.stderr)
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn test_validate_missing_config_file_exits_1() {
|
|
let output = Command::new(binary_path())
|
|
.arg("--config")
|
|
.arg("/nonexistent/path/config.toml")
|
|
.arg("--validate")
|
|
.output()
|
|
.expect("failed to run binary");
|
|
assert_ne!(
|
|
output.status.code(),
|
|
Some(0),
|
|
"expected non-zero exit for missing config"
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn test_validate_wildcard_bind_via_cli_flag() {
|
|
let dir = tempfile::tempdir().unwrap();
|
|
let config_path = write_valid_config(dir.path());
|
|
let output = Command::new(binary_path())
|
|
.arg("--config")
|
|
.arg(config_path.to_str().unwrap())
|
|
.arg("--validate")
|
|
.arg("--allow-wildcard-bind")
|
|
.output()
|
|
.expect("failed to run binary");
|
|
assert_eq!(
|
|
output.status.code(),
|
|
Some(0),
|
|
"expected exit 0 with --allow-wildcard-bind, got {}: stderr={}",
|
|
output.status,
|
|
String::from_utf8_lossy(&output.stderr)
|
|
);
|
|
}
|
|
|
|
fn test_dynamic_config_with_limit(limit_bytes: u64) -> Arc<ArcSwap<DynamicConfig>> {
|
|
let sites = vec![SiteConfig {
|
|
host: "test.local".to_string(),
|
|
upstream: "127.0.0.1:8080".to_string(),
|
|
upstream_scheme: "http".to_string(),
|
|
upstream_connect_timeout_secs: 5,
|
|
upstream_request_timeout_secs: 60,
|
|
}];
|
|
let config = DynamicConfig::from_sites(
|
|
sites,
|
|
RateLimitConfig {
|
|
requests_per_second: 10,
|
|
burst: 20,
|
|
},
|
|
BodyConfig { limit_bytes },
|
|
);
|
|
Arc::new(ArcSwap::from_pointee(config))
|
|
}
|
|
|
|
async fn spawn_server_with_limit(limit_bytes: u64) -> helpers::http_test_helper::TestUpstream {
|
|
let config = test_dynamic_config_with_limit(limit_bytes);
|
|
helpers::http_test_helper::TestUpstream::spawn(|| {
|
|
let app = Router::new().route(
|
|
"/",
|
|
post(|body: axum::body::Body| async move {
|
|
let _ = body;
|
|
"ok"
|
|
}),
|
|
);
|
|
router_with_body_limit(app, config.clone())
|
|
})
|
|
.await
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_body_limit_rejects_oversized_request() {
|
|
let server = spawn_server_with_limit(100).await;
|
|
let client = reqwest::Client::new();
|
|
|
|
let large_body = vec![0u8; 200];
|
|
let resp = client
|
|
.post(format!("http://127.0.0.1:{}/", server.addr.port()))
|
|
.body(large_body)
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(resp.status(), reqwest::StatusCode::PAYLOAD_TOO_LARGE);
|
|
let body = resp.text().await.unwrap();
|
|
assert_eq!(body, "Payload Too Large");
|
|
|
|
let _ = server.shutdown_tx.send(());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_body_limit_allows_request_within_limit() {
|
|
let server = spawn_server_with_limit(100).await;
|
|
let client = reqwest::Client::new();
|
|
|
|
let small_body = vec![0u8; 50];
|
|
let resp = client
|
|
.post(format!("http://127.0.0.1:{}/", server.addr.port()))
|
|
.body(small_body)
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(resp.status(), reqwest::StatusCode::OK);
|
|
|
|
let _ = server.shutdown_tx.send(());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_body_limit_allows_request_at_exact_limit() {
|
|
let server = spawn_server_with_limit(100).await;
|
|
let client = reqwest::Client::new();
|
|
|
|
let exact_body = vec![0u8; 100];
|
|
let resp = client
|
|
.post(format!("http://127.0.0.1:{}/", server.addr.port()))
|
|
.body(exact_body)
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(resp.status(), reqwest::StatusCode::OK);
|
|
|
|
let _ = server.shutdown_tx.send(());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_body_limit_content_length_header_rejection() {
|
|
let server = spawn_server_with_limit(100).await;
|
|
let client = reqwest::Client::new();
|
|
|
|
let resp = client
|
|
.post(format!("http://127.0.0.1:{}/", server.addr.port()))
|
|
.header("content-length", "200")
|
|
.body(vec![0u8; 200])
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(resp.status(), reqwest::StatusCode::PAYLOAD_TOO_LARGE);
|
|
let body = resp.text().await.unwrap();
|
|
assert_eq!(body, "Payload Too Large");
|
|
|
|
let _ = server.shutdown_tx.send(());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_body_limit_default_is_100mb() {
|
|
assert_eq!(DEFAULT_BODY_LIMIT_BYTES, 104_857_600);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_body_limit_empty_body_request_succeeds() {
|
|
let server = spawn_server_with_limit(100).await;
|
|
let client = reqwest::Client::new();
|
|
|
|
let resp = client
|
|
.post(format!("http://127.0.0.1:{}/", server.addr.port()))
|
|
.body("")
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(resp.status(), reqwest::StatusCode::OK);
|
|
|
|
let _ = server.shutdown_tx.send(());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_graceful_shutdown_trigger() {
|
|
let shutdown = Arc::new(reverse_proxy::shutdown::GracefulShutdown::new(30));
|
|
|
|
assert!(!shutdown.is_shutdown_requested());
|
|
|
|
let mut rx = shutdown.subscribe();
|
|
assert!(!*rx.borrow_and_update());
|
|
|
|
shutdown.trigger_shutdown();
|
|
|
|
assert!(shutdown.is_shutdown_requested());
|
|
assert!(rx.has_changed().unwrap());
|
|
assert!(*rx.borrow_and_update());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_graceful_shutdown_custom_timeout() {
|
|
let shutdown = reverse_proxy::shutdown::GracefulShutdown::new(60);
|
|
assert_eq!(shutdown.shutdown_timeout(), Duration::from_secs(60));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_graceful_shutdown_subscribe_multiple_receivers() {
|
|
let shutdown = Arc::new(reverse_proxy::shutdown::GracefulShutdown::new(10));
|
|
|
|
let mut rx1 = shutdown.subscribe();
|
|
let mut rx2 = shutdown.subscribe();
|
|
|
|
assert!(!*rx1.borrow_and_update());
|
|
assert!(!*rx2.borrow_and_update());
|
|
|
|
shutdown.trigger_shutdown();
|
|
|
|
assert!(rx1.has_changed().unwrap());
|
|
assert!(rx2.has_changed().unwrap());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_sighup_config_reload_valid_config() {
|
|
let config_arc = Arc::new(ArcSwap::from_pointee(
|
|
reverse_proxy::config::test_fixtures::test_dynamic_config(),
|
|
));
|
|
let static_config = reverse_proxy::config::test_fixtures::test_static_config();
|
|
let reload_handle = Arc::new(reverse_proxy::config::ConfigReloadHandle::new(
|
|
config_arc.clone(),
|
|
static_config,
|
|
false,
|
|
));
|
|
|
|
let dir = tempfile::tempdir().unwrap();
|
|
let config_content = r#"
|
|
health_check_port = 9900
|
|
admin_key_path = "/tmp/test-admin-key"
|
|
|
|
[logging]
|
|
level = "info"
|
|
format = "text"
|
|
|
|
[rate_limit]
|
|
requests_per_second = 20
|
|
burst = 40
|
|
|
|
[body]
|
|
limit_bytes = 104857600
|
|
|
|
[[listeners]]
|
|
bind_addr = "127.0.0.1"
|
|
http_port = 80
|
|
https_port = 443
|
|
|
|
[listeners.tls]
|
|
mode = "acme"
|
|
acme_domains = ["test.local"]
|
|
acme_cache_dir = "/tmp/acme-cache"
|
|
acme_directory = "staging"
|
|
acme_contact = "mailto:admin@test.local"
|
|
|
|
[[listeners.sites]]
|
|
host = "test.local"
|
|
upstream = "127.0.0.1:8080"
|
|
"#;
|
|
let config_path = dir.path().join("config.toml");
|
|
tokio::fs::write(&config_path, config_content)
|
|
.await
|
|
.unwrap();
|
|
|
|
let config_path_str = config_path.to_str().unwrap().to_string();
|
|
reverse_proxy::shutdown::handle_sighup_reload(&reload_handle, &config_path_str).await;
|
|
|
|
let loaded = reload_handle.load();
|
|
assert_eq!(loaded.rate_limit.requests_per_second, 20);
|
|
assert_eq!(loaded.rate_limit.burst, 40);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_sighup_config_reload_invalid_config_keeps_old() {
|
|
let config_arc = Arc::new(ArcSwap::from_pointee(
|
|
reverse_proxy::config::test_fixtures::test_dynamic_config(),
|
|
));
|
|
let static_config = reverse_proxy::config::test_fixtures::test_static_config();
|
|
let reload_handle = Arc::new(reverse_proxy::config::ConfigReloadHandle::new(
|
|
config_arc.clone(),
|
|
static_config,
|
|
false,
|
|
));
|
|
|
|
let dir = tempfile::tempdir().unwrap();
|
|
let config_content = "invalid toml {{{";
|
|
let config_path = dir.path().join("config.toml");
|
|
tokio::fs::write(&config_path, config_content)
|
|
.await
|
|
.unwrap();
|
|
|
|
let config_path_str = config_path.to_str().unwrap().to_string();
|
|
let _ = reverse_proxy::shutdown::handle_sighup_reload(&reload_handle, &config_path_str).await;
|
|
|
|
let loaded = reload_handle.load();
|
|
assert_eq!(loaded.rate_limit.requests_per_second, 10);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_graceful_shutdown_with_health_check() {
|
|
let (addr, handle) = reverse_proxy::health::start_health_check_listener(0, None, None)
|
|
.await
|
|
.unwrap();
|
|
|
|
let client = reqwest::Client::new();
|
|
let resp = client
|
|
.get(format!("http://127.0.0.1:{}/health", addr.port()))
|
|
.send()
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(resp.status(), reqwest::StatusCode::OK);
|
|
|
|
let shutdown = Arc::new(reverse_proxy::shutdown::GracefulShutdown::new(5));
|
|
let rx = shutdown.subscribe();
|
|
|
|
assert!(!shutdown.is_shutdown_requested());
|
|
|
|
shutdown.trigger_shutdown();
|
|
assert!(shutdown.is_shutdown_requested());
|
|
assert!(rx.has_changed().unwrap());
|
|
|
|
handle.abort();
|
|
}
|
|
|
|
mod idle_timeout_tests {
|
|
use super::*;
|
|
use reverse_proxy::server::InFlightCounter;
|
|
use rustls::pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer};
|
|
use std::sync::Arc;
|
|
use std::time::Duration;
|
|
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
|
use tokio::net::TcpListener;
|
|
use tokio_rustls::TlsAcceptor;
|
|
|
|
fn make_test_tls_acceptor() -> TlsAcceptor {
|
|
let mut params = rcgen::CertificateParams::new(vec!["test.local".to_string()]).unwrap();
|
|
params.distinguished_name = rcgen::DistinguishedName::new();
|
|
params
|
|
.distinguished_name
|
|
.push(rcgen::DnType::CommonName, "test.local");
|
|
let key_pair = rcgen::KeyPair::generate().unwrap();
|
|
let cert = params.self_signed(&key_pair).unwrap();
|
|
let cert_der = cert.der().clone();
|
|
let key_der = key_pair.serialize_der();
|
|
let private_key = PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(key_der));
|
|
|
|
let server_config = reverse_proxy::tls::config::build_manual_server_config_from_certs(
|
|
vec![cert_der],
|
|
private_key,
|
|
)
|
|
.unwrap();
|
|
TlsAcceptor::from(Arc::new(server_config))
|
|
}
|
|
|
|
fn make_proxy_router_with_upstream(
|
|
upstream: String,
|
|
) -> (
|
|
Arc<reverse_proxy::proxy::ProxyState>,
|
|
Arc<ArcSwap<DynamicConfig>>,
|
|
Arc<reverse_proxy::rate_limit::RateLimiter>,
|
|
) {
|
|
let sites = vec![SiteConfig {
|
|
host: "test.local".to_string(),
|
|
upstream,
|
|
upstream_scheme: "http".to_string(),
|
|
upstream_connect_timeout_secs: 5,
|
|
upstream_request_timeout_secs: 60,
|
|
}];
|
|
let config = DynamicConfig::from_sites(
|
|
sites,
|
|
RateLimitConfig {
|
|
requests_per_second: 100,
|
|
burst: 100,
|
|
},
|
|
BodyConfig {
|
|
limit_bytes: 104857600,
|
|
},
|
|
);
|
|
let config_arc = Arc::new(ArcSwap::from_pointee(config));
|
|
let proxy_state = Arc::new(reverse_proxy::proxy::ProxyState {
|
|
config: config_arc.clone(),
|
|
http_client: reverse_proxy::proxy::create_http_client(),
|
|
https_client: reverse_proxy::proxy::create_https_client(),
|
|
});
|
|
let rate_limiter =
|
|
Arc::new(reverse_proxy::rate_limit::RateLimiter::new(config_arc.clone()));
|
|
(proxy_state, config_arc, rate_limiter)
|
|
}
|
|
|
|
async fn start_test_https_server_with_upstream(
|
|
idle_timeout: Duration,
|
|
upstream: String,
|
|
) -> (
|
|
std::net::SocketAddr,
|
|
Arc<InFlightCounter>,
|
|
tokio::task::JoinHandle<()>,
|
|
tokio::sync::watch::Sender<bool>,
|
|
) {
|
|
start_test_https_server_with_timeouts(idle_timeout, Duration::from_secs(10), upstream).await
|
|
}
|
|
|
|
async fn start_test_https_server_with_timeouts(
|
|
idle_timeout: Duration,
|
|
handshake_timeout: Duration,
|
|
upstream: String,
|
|
) -> (
|
|
std::net::SocketAddr,
|
|
Arc<InFlightCounter>,
|
|
tokio::task::JoinHandle<()>,
|
|
tokio::sync::watch::Sender<bool>,
|
|
) {
|
|
let conn_sem = reverse_proxy::server::ConnectionSemaphore::new(1024);
|
|
spawn_test_https_server(idle_timeout, handshake_timeout, upstream, conn_sem).await
|
|
}
|
|
|
|
async fn spawn_test_https_server(
|
|
idle_timeout: Duration,
|
|
handshake_timeout: Duration,
|
|
upstream: String,
|
|
conn_sem: Arc<reverse_proxy::server::ConnectionSemaphore>,
|
|
) -> (
|
|
std::net::SocketAddr,
|
|
Arc<InFlightCounter>,
|
|
tokio::task::JoinHandle<()>,
|
|
tokio::sync::watch::Sender<bool>,
|
|
) {
|
|
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
|
let addr = listener.local_addr().unwrap();
|
|
let tls_acceptor = make_test_tls_acceptor();
|
|
let (proxy_state, config_arc, rate_limiter) = make_proxy_router_with_upstream(upstream);
|
|
let router = reverse_proxy::proxy::build_router(proxy_state, config_arc, rate_limiter);
|
|
let in_flight = InFlightCounter::new();
|
|
let (shutdown_tx, shutdown_rx) = tokio::sync::watch::channel(false);
|
|
|
|
let in_flight_clone = in_flight.clone();
|
|
let handle = tokio::spawn(async move {
|
|
reverse_proxy::server::serve_https_listener(
|
|
listener,
|
|
tls_acceptor,
|
|
router,
|
|
shutdown_rx,
|
|
in_flight_clone,
|
|
idle_timeout,
|
|
handshake_timeout,
|
|
conn_sem,
|
|
)
|
|
.await;
|
|
});
|
|
|
|
(addr, in_flight, handle, shutdown_tx)
|
|
}
|
|
|
|
fn make_client_tls_config() -> Arc<rustls::ClientConfig> {
|
|
let mut roots = rustls::RootCertStore::empty();
|
|
let _ = roots.add(CertificateDer::from_slice(b"test".to_vec().as_slice()));
|
|
let config = rustls::ClientConfig::builder()
|
|
.dangerous()
|
|
.with_custom_certificate_verifier(Arc::new(NoVerifyVerifier))
|
|
.with_no_client_auth();
|
|
Arc::new(config)
|
|
}
|
|
|
|
#[derive(Debug)]
|
|
struct NoVerifyVerifier;
|
|
|
|
impl rustls::client::danger::ServerCertVerifier for NoVerifyVerifier {
|
|
fn verify_server_cert(
|
|
&self,
|
|
_end_entity: &CertificateDer<'_>,
|
|
_intermediates: &[CertificateDer<'_>],
|
|
_server_name: &rustls::pki_types::ServerName<'_>,
|
|
_ocsp_response: &[u8],
|
|
_now: rustls::pki_types::UnixTime,
|
|
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
|
|
Ok(rustls::client::danger::ServerCertVerified::assertion())
|
|
}
|
|
|
|
fn verify_tls12_signature(
|
|
&self,
|
|
_message: &[u8],
|
|
_cert: &CertificateDer<'_>,
|
|
_dss: &rustls::DigitallySignedStruct,
|
|
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
|
|
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
|
|
}
|
|
|
|
fn verify_tls13_signature(
|
|
&self,
|
|
_message: &[u8],
|
|
_cert: &CertificateDer<'_>,
|
|
_dss: &rustls::DigitallySignedStruct,
|
|
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
|
|
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
|
|
}
|
|
|
|
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
|
|
vec![
|
|
rustls::SignatureScheme::RSA_PKCS1_SHA256,
|
|
rustls::SignatureScheme::ECDSA_NISTP256_SHA256,
|
|
rustls::SignatureScheme::RSA_PSS_SHA256,
|
|
rustls::SignatureScheme::RSA_PKCS1_SHA384,
|
|
rustls::SignatureScheme::ECDSA_NISTP384_SHA384,
|
|
rustls::SignatureScheme::RSA_PSS_SHA384,
|
|
rustls::SignatureScheme::RSA_PKCS1_SHA512,
|
|
rustls::SignatureScheme::RSA_PSS_SHA512,
|
|
rustls::SignatureScheme::ED25519,
|
|
]
|
|
}
|
|
}
|
|
|
|
async fn connect_tls(addr: std::net::SocketAddr) -> tokio_rustls::client::TlsStream<tokio::net::TcpStream> {
|
|
let tcp_stream = tokio::net::TcpStream::connect(addr).await.unwrap();
|
|
let server_name = rustls::pki_types::ServerName::try_from("test.local").unwrap();
|
|
let connector = tokio_rustls::TlsConnector::from(make_client_tls_config());
|
|
connector.connect(server_name, tcp_stream).await.unwrap()
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn idle_http1_connection_closed_after_timeout() {
|
|
let idle_timeout = Duration::from_millis(500);
|
|
let upstream = helpers::http_test_helper::TestUpstream::spawn_ok().await;
|
|
let upstream_addr = format!("127.0.0.1:{}", upstream.addr.port());
|
|
let (addr, in_flight, _handle, _shutdown_tx) =
|
|
start_test_https_server_with_upstream(idle_timeout, upstream_addr).await;
|
|
|
|
let mut tls_stream = connect_tls(addr).await;
|
|
|
|
tls_stream
|
|
.write_all(b"GET / HTTP/1.1\r\nHost: test.local\r\nConnection: keep-alive\r\n\r\n")
|
|
.await
|
|
.unwrap();
|
|
|
|
let mut buf = vec![0u8; 4096];
|
|
let n = tokio::time::timeout(Duration::from_secs(2), tls_stream.read(&mut buf))
|
|
.await
|
|
.expect("timeout waiting for first response")
|
|
.expect("read error");
|
|
let response = String::from_utf8_lossy(&buf[..n]);
|
|
assert!(
|
|
response.starts_with("HTTP/1.1 200 OK"),
|
|
"expected a real 200 round-trip from the upstream, got: {response}"
|
|
);
|
|
|
|
tokio::time::sleep(Duration::from_millis(100)).await;
|
|
assert_eq!(
|
|
in_flight.count(),
|
|
1,
|
|
"connection should still be in-flight right after response"
|
|
);
|
|
|
|
let read_result = tokio::time::timeout(
|
|
idle_timeout + Duration::from_secs(2),
|
|
tls_stream.read(&mut buf),
|
|
)
|
|
.await
|
|
.expect("timeout waiting for idle close");
|
|
|
|
let closed = match read_result {
|
|
Ok(0) => true,
|
|
Ok(_) => false,
|
|
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => true,
|
|
Err(e) => panic!("unexpected read error: {e:?}"),
|
|
};
|
|
assert!(
|
|
closed,
|
|
"connection should be closed by server after idle timeout"
|
|
);
|
|
|
|
tokio::time::sleep(Duration::from_millis(100)).await;
|
|
assert_eq!(
|
|
in_flight.count(),
|
|
0,
|
|
"in-flight count should be 0 after connection closed"
|
|
);
|
|
|
|
let _ = upstream.shutdown_tx.send(());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn active_http1_connection_not_closed_within_timeout() {
|
|
let idle_timeout = Duration::from_secs(2);
|
|
let upstream = helpers::http_test_helper::TestUpstream::spawn_ok().await;
|
|
let upstream_addr = format!("127.0.0.1:{}", upstream.addr.port());
|
|
let (addr, in_flight, _handle, _shutdown_tx) =
|
|
start_test_https_server_with_upstream(idle_timeout, upstream_addr).await;
|
|
|
|
let mut tls_stream = connect_tls(addr).await;
|
|
|
|
tls_stream
|
|
.write_all(b"GET / HTTP/1.1\r\nHost: test.local\r\nConnection: keep-alive\r\n\r\n")
|
|
.await
|
|
.unwrap();
|
|
|
|
let mut buf = vec![0u8; 4096];
|
|
let n = tokio::time::timeout(Duration::from_secs(2), tls_stream.read(&mut buf))
|
|
.await
|
|
.expect("timeout waiting for first response")
|
|
.expect("read error");
|
|
let response = String::from_utf8_lossy(&buf[..n]);
|
|
assert!(
|
|
response.starts_with("HTTP/1.1 200 OK"),
|
|
"expected a real 200 round-trip from the upstream, got: {response}"
|
|
);
|
|
|
|
let read_result = tokio::time::timeout(
|
|
idle_timeout / 2,
|
|
tls_stream.read(&mut buf),
|
|
)
|
|
.await;
|
|
|
|
assert!(
|
|
read_result.is_err(),
|
|
"connection should NOT be closed within idle timeout"
|
|
);
|
|
|
|
assert_eq!(
|
|
in_flight.count(),
|
|
1,
|
|
"connection should still be in-flight (active)"
|
|
);
|
|
|
|
let _ = upstream.shutdown_tx.send(());
|
|
}
|
|
|
|
// Reproducer for review #009 C2: a streaming response body that outlasts
|
|
// the idle timeout must NOT be killed mid-stream by the watchdog. The
|
|
// watchdog should only fire when no request is in flight AND no response
|
|
// body is still streaming. With the current code (commit 4ab8c51),
|
|
// `in_flight` is decremented when the handler returns the Response, not
|
|
// when the body finishes streaming, so this test is expected to FAIL.
|
|
#[tokio::test]
|
|
async fn streaming_body_not_killed_by_idle_watchdog() {
|
|
// Stream one chunk every 200ms for 3s total (15 chunks). The stream
|
|
// outlasts the 500ms idle timeout by 6x.
|
|
let chunk_interval = Duration::from_millis(200);
|
|
let num_chunks: usize = 15;
|
|
let idle_timeout = Duration::from_millis(500);
|
|
|
|
let upstream =
|
|
helpers::http_test_helper::TestUpstream::spawn_slow_stream(chunk_interval, num_chunks)
|
|
.await;
|
|
let upstream_addr = format!("127.0.0.1:{}", upstream.addr.port());
|
|
|
|
let (addr, _in_flight, _handle, _shutdown_tx) =
|
|
start_test_https_server_with_upstream(idle_timeout, upstream_addr).await;
|
|
|
|
let mut tls_stream = connect_tls(addr).await;
|
|
|
|
tls_stream
|
|
.write_all(b"GET / HTTP/1.1\r\nHost: test.local\r\nConnection: keep-alive\r\n\r\n")
|
|
.await
|
|
.unwrap();
|
|
|
|
// Read the response headers + as much body as the server sends before
|
|
// closing the connection (or before a generous test deadline).
|
|
let mut buf = Vec::new();
|
|
let deadline = tokio::time::Instant::now() + Duration::from_secs(15);
|
|
let mut tmp = [0u8; 4096];
|
|
loop {
|
|
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
|
|
if remaining.is_zero() {
|
|
break;
|
|
}
|
|
match tokio::time::timeout(remaining, tls_stream.read(&mut tmp)).await {
|
|
Ok(Ok(0)) => break,
|
|
Ok(Ok(n)) => buf.extend_from_slice(&tmp[..n]),
|
|
Ok(Err(e)) if e.kind() == std::io::ErrorKind::UnexpectedEof => break,
|
|
Ok(Err(e)) => panic!("unexpected read error: {e:?}"),
|
|
Err(_) => break,
|
|
}
|
|
}
|
|
|
|
let response = String::from_utf8_lossy(&buf);
|
|
|
|
// The response must contain every chunk the upstream emitted, in
|
|
// order. If the watchdog killed the connection mid-stream, later
|
|
// chunks will be missing. We check for each chunk as a separate
|
|
// substring (in order) rather than a single contiguous string
|
|
// because HTTP chunked transfer-encoding interleaves frame
|
|
// delimiters ("\r\n3\r\n") between data chunks.
|
|
let mut search_from = 0;
|
|
for i in 0..num_chunks {
|
|
let needle = format!("<{i}>");
|
|
match response[search_from..].find(&needle) {
|
|
Some(pos) => search_from += pos + needle.len(),
|
|
None => {
|
|
panic!(
|
|
"streaming body was truncated — watchdog killed the connection mid-stream.\n\
|
|
missing chunk {i:?} ({needle:?}) after byte {search_from}.\n\
|
|
expected {num_chunks} chunks (<0>..<{n}>, n={n}).\n\
|
|
got: {response:?}",
|
|
n = num_chunks - 1,
|
|
);
|
|
}
|
|
}
|
|
}
|
|
|
|
let _ = upstream.shutdown_tx.send(());
|
|
}
|
|
|
|
// Reproducer for review #010 C3: a client that opens a TCP connection but
|
|
// never sends a TLS ClientHello must not hold its FD + semaphore permit
|
|
// indefinitely. The server must close the stalled handshake after
|
|
// tls_handshake_timeout and release the permit.
|
|
#[tokio::test]
|
|
async fn stalled_tls_handshake_closed_after_timeout_and_permit_released() {
|
|
let handshake_timeout = Duration::from_millis(300);
|
|
let upstream = helpers::http_test_helper::TestUpstream::spawn_ok().await;
|
|
let upstream_addr = format!("127.0.0.1:{}", upstream.addr.port());
|
|
let (addr, in_flight, _handle, _shutdown_tx) =
|
|
start_test_https_server_with_timeouts(
|
|
Duration::from_secs(60),
|
|
handshake_timeout,
|
|
upstream_addr,
|
|
)
|
|
.await;
|
|
|
|
let stalled = tokio::net::TcpStream::connect(addr).await.unwrap();
|
|
// Hold the connection open without sending anything (stalled handshake).
|
|
|
|
// Wait past the handshake timeout.
|
|
tokio::time::sleep(handshake_timeout + Duration::from_millis(500)).await;
|
|
|
|
// The server must have dropped the stalled connection: reading from it
|
|
// should yield EOF (or an error), not hang.
|
|
let mut stalled = stalled;
|
|
let mut buf = [0u8; 16];
|
|
let read_result =
|
|
tokio::time::timeout(Duration::from_secs(2), stalled.read(&mut buf)).await;
|
|
let closed = match read_result {
|
|
Err(_) => panic!("stalled handshake was not closed by the server"),
|
|
Ok(Ok(0)) => true,
|
|
Ok(Ok(_)) => false,
|
|
Ok(Err(e)) if e.kind() == std::io::ErrorKind::UnexpectedEof => true,
|
|
Ok(Err(_)) => true,
|
|
};
|
|
assert!(
|
|
closed,
|
|
"server should close stalled TLS handshake after timeout"
|
|
);
|
|
|
|
// The semaphore permit and in-flight guard must have been released.
|
|
tokio::time::sleep(Duration::from_millis(100)).await;
|
|
assert_eq!(
|
|
in_flight.count(),
|
|
0,
|
|
"in-flight count should return to 0 after stalled handshake closed"
|
|
);
|
|
|
|
let _ = upstream.shutdown_tx.send(());
|
|
}
|
|
|
|
// A real TLS handshake must complete comfortably within the timeout —
|
|
// the timeout must only kill stalled handshakes, not legitimate clients.
|
|
#[tokio::test]
|
|
async fn real_handshake_completes_within_handshake_timeout() {
|
|
let handshake_timeout = Duration::from_millis(300);
|
|
let upstream = helpers::http_test_helper::TestUpstream::spawn_ok().await;
|
|
let upstream_addr = format!("127.0.0.1:{}", upstream.addr.port());
|
|
let (addr, in_flight, _handle, _shutdown_tx) =
|
|
start_test_https_server_with_timeouts(
|
|
Duration::from_secs(60),
|
|
handshake_timeout,
|
|
upstream_addr,
|
|
)
|
|
.await;
|
|
|
|
let mut tls_stream = connect_tls(addr).await;
|
|
|
|
tls_stream
|
|
.write_all(b"GET / HTTP/1.1\r\nHost: test.local\r\nConnection: keep-alive\r\n\r\n")
|
|
.await
|
|
.unwrap();
|
|
|
|
let mut buf = vec![0u8; 4096];
|
|
let n = tokio::time::timeout(Duration::from_secs(2), tls_stream.read(&mut buf))
|
|
.await
|
|
.expect("timeout waiting for response")
|
|
.expect("read error");
|
|
let response = String::from_utf8_lossy(&buf[..n]);
|
|
assert!(
|
|
response.starts_with("HTTP/1.1 200 OK"),
|
|
"expected a real 200 round-trip, got: {response}"
|
|
);
|
|
assert_eq!(in_flight.count(), 1, "connection should remain in-flight");
|
|
|
|
let _ = upstream.shutdown_tx.send(());
|
|
}
|
|
|
|
// Reproducer for review #010 C4: the connection semaphore must be shared
|
|
// across listeners. Two listeners with a shared cap of 2 must admit only
|
|
// 2 concurrent connections total — the third connection, hitting either
|
|
// listener, must wait until a slot frees.
|
|
#[tokio::test]
|
|
async fn connection_semaphore_is_shared_across_listeners() {
|
|
let conn_sem = reverse_proxy::server::ConnectionSemaphore::new(2);
|
|
|
|
let upstream = helpers::http_test_helper::TestUpstream::spawn_ok().await;
|
|
let upstream_addr = format!("127.0.0.1:{}", upstream.addr.port());
|
|
|
|
let (addr1, _in_flight1, _h1, _s1) = spawn_test_https_server(
|
|
Duration::from_secs(60),
|
|
Duration::from_secs(10),
|
|
upstream_addr.clone(),
|
|
conn_sem.clone(),
|
|
)
|
|
.await;
|
|
let (addr2, _in_flight2, _h2, _s2) = spawn_test_https_server(
|
|
Duration::from_secs(60),
|
|
Duration::from_secs(10),
|
|
upstream_addr,
|
|
conn_sem.clone(),
|
|
)
|
|
.await;
|
|
|
|
// Two live connections saturate the shared pool: one per listener.
|
|
let tls1 = connect_tls(addr1).await;
|
|
let tls2 = connect_tls(addr2).await;
|
|
|
|
// A third connection (racing against listener 1) must not complete its
|
|
// handshake: the accept loop parks it on the shared semaphore before
|
|
// TLS ever starts.
|
|
let (third_tx, third_rx) = tokio::sync::oneshot::channel();
|
|
let third_task = tokio::spawn(async move {
|
|
let stream = connect_tls(addr1).await;
|
|
let _ = third_tx.send(stream);
|
|
});
|
|
|
|
tokio::time::sleep(Duration::from_millis(500)).await;
|
|
assert!(
|
|
third_rx.is_empty(),
|
|
"third connection must be capped by the shared semaphore"
|
|
);
|
|
|
|
// Release one slot; the parked third connection must now get through.
|
|
drop(tls2);
|
|
|
|
let mut tls3 = tokio::time::timeout(Duration::from_secs(5), third_rx)
|
|
.await
|
|
.expect("third connection should complete after a permit is released")
|
|
.expect("third connection task should not panic");
|
|
|
|
tls3.write_all(b"GET / HTTP/1.1\r\nHost: test.local\r\nConnection: close\r\n\r\n")
|
|
.await
|
|
.unwrap();
|
|
|
|
let mut buf = vec![0u8; 4096];
|
|
let n = tokio::time::timeout(Duration::from_secs(5), tls3.read(&mut buf))
|
|
.await
|
|
.expect("timeout waiting for response")
|
|
.expect("read error");
|
|
let response = String::from_utf8_lossy(&buf[..n]);
|
|
assert!(
|
|
response.starts_with("HTTP/1.1 200 OK"),
|
|
"expected a real 200 round-trip on the capped connection, got: {response}"
|
|
);
|
|
|
|
third_task.abort();
|
|
drop(tls1);
|
|
drop(tls3);
|
|
let _ = upstream.shutdown_tx.send(());
|
|
}
|
|
}
|