use std::io::Write;
use std::sync::atomic::Ordering;
use std::time::Duration;
use tempfile::NamedTempFile;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
fn write_config(content: &str) -> NamedTempFile {
let mut f = NamedTempFile::new().unwrap();
f.write_all(content.as_bytes()).unwrap();
f.flush().unwrap();
f
}
async fn start_slow_backend() -> (std::net::SocketAddr, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let jh = tokio::spawn(async move {
loop {
let (stream, _) = match listener.accept().await {
Ok(s) => s,
Err(_) => break,
};
tokio::spawn(async move {
let (mut rd, mut wr) = stream.into_split();
let _ = tokio::io::copy(&mut rd, &mut wr).await;
});
}
});
(addr, jh)
}
async fn socks5_handshake(
stream: &mut tokio::net::TcpStream,
target: std::net::SocketAddr,
) -> std::io::Result<()> {
stream.write_all(&[0x05, 0x01, 0x00]).await?;
let mut resp = [0u8; 2];
stream.read_exact(&mut resp).await?;
if resp[0] != 0x05 || resp[1] != 0x00 {
return Err(std::io::Error::other("SOCKS5 method negotiation failed"));
}
let octets = match target.ip() {
std::net::IpAddr::V4(v4) => v4.octets(),
_ => {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"only IPv4 supported",
))
}
};
let port = target.port().to_be_bytes();
stream
.write_all(&[
0x05, 0x01, 0x00, 0x01, octets[0], octets[1], octets[2], octets[3],
])
.await?;
stream.write_all(&port).await?;
let mut reply = [0u8; 10];
stream.read_exact(&mut reply).await?;
if reply[1] != 0x00 {
return Err(std::io::Error::other(format!(
"SOCKS5 connect failed: {:#04x}",
reply[1]
)));
}
Ok(())
}
#[tokio::test]
async fn readiness_transitions_to_false_on_shutdown() {
let config = r#"
version = 1
[[listeners]]
name = "http-in"
bind = "127.0.0.1:0"
protocols = ["http"]
"#;
let f = write_config(config);
let path = f.path().to_str().unwrap();
let mut sup = eggress_runtime::ServiceSupervisor::start(path).unwrap();
let state = sup.state().clone();
let token = sup.shutdown_token();
let jh = tokio::task::spawn_blocking(move || sup.run());
for _ in 0..50 {
if state.readiness.load(Ordering::Relaxed) {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
assert!(state.readiness.load(Ordering::Relaxed), "should be ready");
token.cancel();
jh.await.ok();
assert!(
!state.readiness.load(Ordering::Relaxed),
"readiness should be false after shutdown"
);
}
#[tokio::test]
async fn shutdown_drains_active_connections() {
let config = r#"
version = 1
[[listeners]]
name = "http-in"
bind = "127.0.0.1:0"
protocols = ["http"]
"#;
let f = write_config(config);
let path = f.path().to_str().unwrap();
let mut sup = eggress_runtime::ServiceSupervisor::start(path).unwrap();
let state = sup.state().clone();
let token = sup.shutdown_token();
let jh = tokio::task::spawn_blocking(move || sup.run());
for _ in 0..50 {
if state.readiness.load(Ordering::Relaxed) {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
assert!(state.readiness.load(Ordering::Relaxed));
let start = std::time::Instant::now();
token.cancel();
jh.await.ok();
let elapsed = start.elapsed();
assert!(
elapsed < Duration::from_secs(10),
"shutdown took too long: {:?}",
elapsed
);
}
#[tokio::test]
async fn shutdown_generation_remains_consistent() {
let config = r#"
version = 1
[[listeners]]
name = "http-in"
bind = "127.0.0.1:0"
protocols = ["http"]
"#;
let f = write_config(config);
let path = f.path().to_str().unwrap();
let mut sup = eggress_runtime::ServiceSupervisor::start(path).unwrap();
let state = sup.state().clone();
let gen_before = state.generation();
assert_eq!(gen_before, 0);
let token = sup.shutdown_token();
let jh = tokio::task::spawn_blocking(move || sup.run());
for _ in 0..50 {
if state.readiness.load(Ordering::Relaxed) {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
token.cancel();
jh.await.ok();
let gen_after = state.generation();
assert_eq!(
gen_before, gen_after,
"generation should not change during shutdown"
);
}
#[tokio::test]
async fn shutdown_active_connections_returns_to_zero() {
let config = r#"
version = 1
[[listeners]]
name = "http-in"
bind = "127.0.0.1:0"
protocols = ["http"]
"#;
let f = write_config(config);
let path = f.path().to_str().unwrap();
let mut sup = eggress_runtime::ServiceSupervisor::start(path).unwrap();
let state = sup.state().clone();
assert_eq!(state.active_connections.load(Ordering::Relaxed), 0);
let token = sup.shutdown_token();
let jh = tokio::task::spawn_blocking(move || sup.run());
for _ in 0..50 {
if state.readiness.load(Ordering::Relaxed) {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
token.cancel();
jh.await.ok();
assert_eq!(
state.active_connections.load(Ordering::Relaxed),
0,
"active connections should be zero after shutdown"
);
}
#[tokio::test]
async fn shutdown_stops_accepting_new_connections() {
let config = r#"
version = 1
[[listeners]]
name = "http-in"
bind = "127.0.0.1:0"
protocols = ["http"]
"#;
let f = write_config(config);
let path = f.path().to_str().unwrap();
let mut sup = eggress_runtime::ServiceSupervisor::start(path).unwrap();
let state = sup.state().clone();
let token = sup.shutdown_token();
let jh = tokio::task::spawn_blocking(move || sup.run());
for _ in 0..50 {
if state.readiness.load(Ordering::Relaxed) {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
assert!(state.readiness.load(Ordering::Relaxed));
token.cancel();
jh.await.ok();
assert_eq!(
state.active_connections.load(Ordering::Relaxed),
0,
"active connections should be zero after shutdown"
);
}
#[tokio::test]
async fn shutdown_completes_within_grace_period() {
let config = r#"
version = 1
[[listeners]]
name = "http-in"
bind = "127.0.0.1:0"
protocols = ["http"]
"#;
let f = write_config(config);
let path = f.path().to_str().unwrap();
let mut sup = eggress_runtime::ServiceSupervisor::start(path).unwrap();
let state = sup.state().clone();
let token = sup.shutdown_token();
let jh = tokio::task::spawn_blocking(move || sup.run());
for _ in 0..50 {
if state.readiness.load(Ordering::Relaxed) {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
assert!(state.readiness.load(Ordering::Relaxed));
let start = std::time::Instant::now();
token.cancel();
jh.await.ok();
let elapsed = start.elapsed();
assert!(
elapsed < Duration::from_secs(2),
"empty shutdown took too long: {:?}",
elapsed
);
assert!(!state.readiness.load(Ordering::Relaxed));
}
#[tokio::test]
async fn shutdown_force_cancels_after_deadline() {
let (backend_addr, _backend_jh) = start_slow_backend().await;
let config = r#"
version = 1
[process]
shutdown_grace = "2s"
[[listeners]]
name = "socks-in"
bind = "127.0.0.1:0"
protocols = ["socks5"]
[[rules]]
id = "route-all"
any = true
direct = true
"#;
let f = write_config(config);
let path = f.path().to_str().unwrap();
let mut sup = eggress_runtime::ServiceSupervisor::start(path).unwrap();
let state = sup.state().clone();
let token = sup.shutdown_token();
let jh = tokio::task::spawn_blocking(move || sup.run());
for _ in 0..100 {
if state.readiness.load(Ordering::Relaxed) {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
assert!(state.readiness.load(Ordering::Relaxed), "should be ready");
let listener_addr = {
let addrs = state.listener_addrs.lock().unwrap();
assert!(!addrs.is_empty(), "should have at least one listener");
addrs[0].unwrap()
};
let mut stream = tokio::net::TcpStream::connect(listener_addr)
.await
.expect("failed to connect to listener");
socks5_handshake(&mut stream, backend_addr)
.await
.expect("SOCS5 handshake failed");
tokio::time::sleep(Duration::from_millis(100)).await;
let active = state.active_connections.load(Ordering::Relaxed);
assert!(
active >= 1,
"should have at least 1 active connection, got {active}"
);
let start = std::time::Instant::now();
token.cancel();
jh.await.ok();
let elapsed = start.elapsed();
assert!(
elapsed < Duration::from_secs(6),
"shutdown took too long with active connection: {:?}",
elapsed
);
assert_eq!(
state.active_connections.load(Ordering::Relaxed),
0,
"active connections should be zero after forced shutdown"
);
let mut buf = [0u8; 1];
let result = tokio::time::timeout(Duration::from_millis(500), stream.read(&mut buf)).await;
assert!(
result.is_err() || matches!(result, Ok(Ok(0))),
"client stream should be dead after forced shutdown"
);
}
#[tokio::test]
async fn admin_responds_during_shutdown_drain() {
let (backend_addr, _backend_jh) = start_slow_backend().await;
let config = r#"
version = 1
[process]
shutdown_grace = "5s"
[[listeners]]
name = "socks-in"
bind = "127.0.0.1:0"
protocols = ["socks5"]
[[rules]]
id = "route-all"
any = true
direct = true
[admin]
bind = "127.0.0.1:0"
enabled = true
"#;
let f = write_config(config);
let path = f.path().to_str().unwrap();
let mut sup = eggress_runtime::ServiceSupervisor::start(path).unwrap();
let state = sup.state().clone();
let token = sup.shutdown_token();
let jh = tokio::task::spawn_blocking(move || sup.run());
for _ in 0..100 {
if state.readiness.load(Ordering::Relaxed) {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
assert!(state.readiness.load(Ordering::Relaxed));
let listener_addr = state.listener_addrs.lock().unwrap()[0].unwrap();
let admin_addr = state
.admin_local_addr
.lock()
.unwrap()
.expect("admin should have bound")
.to_string();
let mut client = tokio::net::TcpStream::connect(listener_addr)
.await
.expect("connect listener");
socks5_handshake(&mut client, backend_addr)
.await
.expect("socks5 handshake");
tokio::time::sleep(Duration::from_millis(100)).await;
let active = state.active_connections.load(Ordering::Relaxed);
assert!(
active >= 1,
"should have one active connection, got {active}"
);
token.cancel();
tokio::time::sleep(Duration::from_millis(100)).await;
let mut stream = tokio::net::TcpStream::connect(&admin_addr)
.await
.expect("admin should still be listening during drain");
let req = b"GET /-/ready HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n";
tokio::io::AsyncWriteExt::write_all(&mut stream, req)
.await
.unwrap();
tokio::io::AsyncWriteExt::flush(&mut stream).await.unwrap();
let mut buf = Vec::new();
let _ = tokio::time::timeout(
Duration::from_secs(2),
tokio::io::AsyncReadExt::read_to_end(&mut stream, &mut buf),
)
.await;
let response = String::from_utf8_lossy(&buf);
assert!(
response.starts_with("HTTP/1.1 503"),
"admin /-/ready should respond 503 during drain, got: {response}"
);
drop(client);
jh.await.ok();
}
#[tokio::test]
async fn admin_metrics_visible_during_drain() {
let (backend_addr, _backend_jh) = start_slow_backend().await;
let config = r#"
version = 1
[process]
shutdown_grace = "5s"
[[listeners]]
name = "socks-in"
bind = "127.0.0.1:0"
protocols = ["socks5"]
[[rules]]
id = "route-all"
any = true
direct = true
[admin]
bind = "127.0.0.1:0"
enabled = true
"#;
let f = write_config(config);
let path = f.path().to_str().unwrap();
let mut sup = eggress_runtime::ServiceSupervisor::start(path).unwrap();
let state = sup.state().clone();
let token = sup.shutdown_token();
let jh = tokio::task::spawn_blocking(move || sup.run());
for _ in 0..100 {
if state.readiness.load(Ordering::Relaxed) {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
assert!(state.readiness.load(Ordering::Relaxed));
let listener_addr = state.listener_addrs.lock().unwrap()[0].unwrap();
let admin_addr = state
.admin_local_addr
.lock()
.unwrap()
.expect("admin should have bound")
.to_string();
let mut client = tokio::net::TcpStream::connect(listener_addr)
.await
.expect("connect listener");
socks5_handshake(&mut client, backend_addr)
.await
.expect("socks5 handshake");
tokio::time::sleep(Duration::from_millis(100)).await;
let active = state.active_connections.load(Ordering::Relaxed);
assert!(
active >= 1,
"should have one active connection, got {active}"
);
token.cancel();
tokio::time::sleep(Duration::from_millis(100)).await;
let mut stream = tokio::net::TcpStream::connect(&admin_addr)
.await
.expect("admin should still be listening during drain");
let req = b"GET /metrics HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n";
tokio::io::AsyncWriteExt::write_all(&mut stream, req)
.await
.unwrap();
tokio::io::AsyncWriteExt::flush(&mut stream).await.unwrap();
let mut buf = Vec::new();
let _ = tokio::time::timeout(
Duration::from_secs(2),
tokio::io::AsyncReadExt::read_to_end(&mut stream, &mut buf),
)
.await;
let response = String::from_utf8_lossy(&buf);
assert!(
response.contains("eggress_connections_active"),
"/metrics should be available during drain, got: {response}"
);
drop(client);
jh.await.ok();
}
#[tokio::test]
async fn new_connections_refused_after_shutdown_begins() {
let config = r#"
version = 1
[[listeners]]
name = "http-in"
bind = "127.0.0.1:0"
protocols = ["http"]
"#;
let f = write_config(config);
let path = f.path().to_str().unwrap();
let mut sup = eggress_runtime::ServiceSupervisor::start(path).unwrap();
let state = sup.state().clone();
let token = sup.shutdown_token();
let jh = tokio::task::spawn_blocking(move || sup.run());
for _ in 0..50 {
if state.readiness.load(Ordering::Relaxed) {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
assert!(state.readiness.load(Ordering::Relaxed));
let listener_addr = state.listener_addrs.lock().unwrap()[0].unwrap();
token.cancel();
tokio::time::sleep(Duration::from_millis(100)).await;
let result = tokio::time::timeout(Duration::from_secs(2), async {
tokio::net::TcpStream::connect(listener_addr).await
})
.await;
match result {
Ok(Ok(_)) => {
assert!(
!state.readiness.load(Ordering::Relaxed),
"if connect succeeds during shutdown, readiness must be false"
);
}
Ok(Err(_)) => {
}
Err(_) => {
}
}
jh.await.ok();
}
#[tokio::test]
async fn malformed_handshake_does_not_corrupt_listener() {
let config = r#"
version = 1
[[listeners]]
name = "socks-in"
bind = "127.0.0.1:0"
protocols = ["socks5"]
[[rules]]
id = "route-all"
any = true
direct = true
"#;
let f = write_config(config);
let path = f.path().to_str().unwrap();
let mut sup = eggress_runtime::ServiceSupervisor::start(path).unwrap();
let state = sup.state().clone();
let token = sup.shutdown_token();
let jh = tokio::task::spawn_blocking(move || sup.run());
for _ in 0..50 {
if state.readiness.load(Ordering::Relaxed) {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
assert!(state.readiness.load(Ordering::Relaxed));
let listener_addr = state.listener_addrs.lock().unwrap()[0].unwrap();
{
let mut bad_stream = tokio::net::TcpStream::connect(listener_addr)
.await
.expect("connect for malformed handshake");
let _ = bad_stream.write_all(b"NOT_A_SOCKS5_PROTOCOL").await;
drop(bad_stream);
}
tokio::time::sleep(Duration::from_millis(200)).await;
{
let mut good_stream = tokio::net::TcpStream::connect(listener_addr)
.await
.expect("connect for valid handshake after malformed");
good_stream.write_all(&[0x05, 0x01, 0x00]).await.unwrap();
let mut resp = [0u8; 2];
let result =
tokio::time::timeout(Duration::from_secs(2), good_stream.read_exact(&mut resp)).await;
assert!(result.is_ok(), "valid handshake must get a response");
assert_eq!(resp, [0x05, 0x00], "server must accept valid SOCKS5");
}
assert!(state.readiness.load(Ordering::Relaxed));
token.cancel();
jh.await.ok();
}