use super::*;
use std::sync::Mutex;
use std::sync::atomic::{AtomicUsize, Ordering};
use http_body_util::BodyExt;
#[test]
fn ephemeral_bind_addr_is_loopback_only() {
assert_eq!(
ephemeral_bind_addr().ip(),
std::net::IpAddr::V4(std::net::Ipv4Addr::LOCALHOST)
);
}
#[test]
fn backoff_delays_double_up_to_a_cap() {
let mut backoff = Backoff::new();
assert_eq!(backoff.next_delay(), ACCEPT_BACKOFF_INITIAL);
assert_eq!(backoff.next_delay(), ACCEPT_BACKOFF_INITIAL * 2);
assert_eq!(backoff.next_delay(), ACCEPT_BACKOFF_INITIAL * 4);
let mut last = Duration::ZERO;
for _ in 0..20 {
last = backoff.next_delay();
}
assert_eq!(last, ACCEPT_BACKOFF_MAX);
}
#[test]
fn backoff_reset_returns_to_initial_delay() {
let mut backoff = Backoff::new();
backoff.next_delay();
backoff.next_delay();
backoff.reset();
assert_eq!(backoff.next_delay(), ACCEPT_BACKOFF_INITIAL);
}
struct FlakyListener {
inner: TcpListener,
remaining_failures: AtomicUsize,
attempts: Mutex<Vec<tokio::time::Instant>>,
}
impl TcpAccept for FlakyListener {
async fn accept(&self) -> std::io::Result<(TcpStream, SocketAddr)> {
self.attempts.lock().unwrap().push(tokio::time::Instant::now());
if self.remaining_failures.fetch_sub(1, Ordering::SeqCst) > 0 {
Err(std::io::Error::other("simulated accept error"))
} else {
TcpAccept::accept(&self.inner).await
}
}
}
#[tokio::test(start_paused = true)]
async fn accept_loop_backs_off_between_repeated_errors_instead_of_busy_spinning() {
let inner = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let addr = inner.local_addr().unwrap();
let flaky = FlakyListener {
inner,
remaining_failures: AtomicUsize::new(5),
attempts: Mutex::new(Vec::new()),
};
tokio::spawn(async move {
let _ = TcpStream::connect(addr).await;
});
let mut backoff = Backoff::new();
let semaphore = Arc::new(tokio::sync::Semaphore::new(1));
accept_and_permit(&flaky, &mut backoff, &semaphore).await;
let recorded = flaky.attempts.lock().unwrap();
assert_eq!(recorded.len(), 6, "5 failures then 1 success");
let expected_gaps = [
ACCEPT_BACKOFF_INITIAL,
ACCEPT_BACKOFF_INITIAL * 2,
ACCEPT_BACKOFF_INITIAL * 4,
ACCEPT_BACKOFF_INITIAL * 8,
ACCEPT_BACKOFF_INITIAL * 16,
];
for (i, expected) in expected_gaps.iter().enumerate() {
let gap = recorded[i + 1] - recorded[i];
assert_eq!(
gap, *expected,
"gap between attempt {i} and {} should reflect the backoff delay, not a busy spin",
i + 1
);
}
}
#[tokio::test]
async fn error_handler_sanitizes_5xx_in_response_body() {
take_error_log(); let resp = error_response(StatusCode::INTERNAL_SERVER_ERROR, "raw db connection string leaked");
let (parts, body) = resp.into_parts();
assert_eq!(parts.status, StatusCode::INTERNAL_SERVER_ERROR);
let collected = body.collect().await.unwrap();
let bytes = collected.to_bytes();
let json: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
let msg = json.get("message").and_then(|v| v.as_str()).unwrap();
assert_eq!(msg, "internal server error", "5xx message should be sanitized");
}
#[test]
fn error_handler_captures_5xx_message_in_log() {
take_error_log(); let sensitive_msg = "raw db connection string leaked";
error_response(StatusCode::INTERNAL_SERVER_ERROR, sensitive_msg);
let log = take_error_log();
assert_eq!(log.len(), 1);
assert_eq!(log[0].0, 500);
assert_eq!(log[0].1, sensitive_msg);
}
#[tokio::test]
async fn error_handler_passes_through_4xx_in_response_body() {
take_error_log(); let msg = "bad request";
let resp = error_response(StatusCode::BAD_REQUEST, msg);
let (parts, body) = resp.into_parts();
assert_eq!(parts.status, StatusCode::BAD_REQUEST);
let collected = body.collect().await.unwrap();
let bytes = collected.to_bytes();
let json: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
let response_msg = json.get("message").and_then(|v| v.as_str()).unwrap();
assert_eq!(response_msg, msg, "4xx message should pass through");
}
#[test]
fn error_handler_logs_4xx_messages() {
take_error_log(); let msg = "bad request";
error_response(StatusCode::BAD_REQUEST, msg);
let log = take_error_log();
assert_eq!(log.len(), 1);
assert_eq!(log[0].0, 400);
assert_eq!(log[0].1, msg);
}
#[test]
fn a_failed_signal_install_reports_which_signal_and_why() {
let cause = std::io::Error::new(std::io::ErrorKind::PermissionDenied, "operation not permitted");
let err = signal_install_error("SIGTERM", cause);
assert_eq!(err.code, 500);
assert!(err.message.contains("SIGTERM"), "got: {}", err.message);
assert!(
err.message.contains("operation not permitted"),
"the underlying cause must survive, got: {}",
err.message
);
}