use super::pool;
use super::request::{perform_request, perform_request_with_target};
use super::ssrf::{Scheme, ValidatedTarget};
use crate::seqstring::global_string;
use crate::value::{MapKey, Value};
use may::net::TcpListener;
use std::io::{BufRead, BufReader, Read, Write};
use std::net::{IpAddr, Ipv4Addr};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Clone, Copy)]
enum ServerMode {
FixedBody,
Echo,
Chunked,
CloseOnce,
}
fn spawn_test_server(
listener: TcpListener,
accept_count: Arc<AtomicUsize>,
req_count: Arc<AtomicUsize>,
mode: ServerMode,
) {
unsafe {
may::coroutine::spawn(move || {
for stream in listener.incoming() {
let mut stream = match stream {
Ok(s) => s,
Err(_) => break,
};
accept_count.fetch_add(1, Ordering::SeqCst);
let req_count = req_count.clone();
may::coroutine::spawn(move || {
let stream_clone = match stream.try_clone() {
Ok(s) => s,
Err(_) => return,
};
let mut reader = BufReader::new(stream_clone);
loop {
let mut saw_request = false;
let mut content_length: usize = 0;
loop {
let mut line = String::new();
match reader.read_line(&mut line) {
Ok(0) => return,
Ok(_) => {}
Err(_) => return,
}
if line == "\r\n" || line == "\n" {
if saw_request {
break;
}
continue;
}
saw_request = true;
if let Some(rest) = line.strip_prefix("Content-Length:") {
content_length = rest.trim().parse::<usize>().unwrap_or(0);
} else if let Some(rest) = line.strip_prefix("content-length:") {
content_length = rest.trim().parse::<usize>().unwrap_or(0);
}
}
req_count.fetch_add(1, Ordering::SeqCst);
let body_in: Vec<u8> =
if content_length > 0 && matches!(mode, ServerMode::Echo) {
let mut buf = vec![0u8; content_length];
if reader.read_exact(&mut buf).is_err() {
return;
}
buf
} else {
Vec::new()
};
let write_ok = match mode {
ServerMode::FixedBody => {
let body = b"world";
let header = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: text/plain\r\n\
Content-Length: {}\r\n\
Connection: keep-alive\r\n\
\r\n",
body.len()
);
stream.write_all(header.as_bytes()).is_ok()
&& stream.write_all(body).is_ok()
}
ServerMode::Echo => {
let header = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: application/octet-stream\r\n\
Content-Length: {}\r\n\
Connection: keep-alive\r\n\
\r\n",
body_in.len()
);
stream.write_all(header.as_bytes()).is_ok()
&& stream.write_all(&body_in).is_ok()
}
ServerMode::Chunked => {
let payload = b"HTTP/1.1 200 OK\r\n\
Content-Type: text/plain\r\n\
Transfer-Encoding: chunked\r\n\
Connection: keep-alive\r\n\
\r\n\
5\r\nHello\r\n\
6\r\n world\r\n\
0\r\n\r\n";
stream.write_all(payload).is_ok()
}
ServerMode::CloseOnce => {
let body = b"closed";
let header = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: text/plain\r\n\
Content-Length: {}\r\n\
Connection: close\r\n\
\r\n",
body.len()
);
let _ = stream.write_all(header.as_bytes());
let _ = stream.write_all(body);
return;
}
};
if !write_ok {
return;
}
}
});
}
});
}
}
fn loopback_target(port: u16, path: &str) -> ValidatedTarget {
ValidatedTarget {
scheme: Scheme::Http,
host: "127.0.0.1".to_string(),
port,
addrs: vec![IpAddr::V4(Ipv4Addr::LOCALHOST)],
path_and_query: path.to_string(),
}
}
fn unwrap_response(value: &Value) -> (i64, Vec<u8>, bool) {
let map = match value {
Value::Map(m) => m,
other => panic!("expected response Map, got {:?}", other),
};
let status = match map.get(&MapKey::String(global_string("status".to_string()))) {
Some(Value::Int(n)) => *n,
other => panic!("status missing/wrong type: {:?}", other),
};
let body = match map.get(&MapKey::String(global_string("body".to_string()))) {
Some(Value::String(s)) => s.as_bytes().to_vec(),
other => panic!("body missing/wrong type: {:?}", other),
};
let ok = match map.get(&MapKey::String(global_string("ok".to_string()))) {
Some(Value::Bool(b)) => *b,
other => panic!("ok missing/wrong type: {:?}", other),
};
(status, body, ok)
}
#[test]
#[serial_test::serial(http_global_state)]
fn http_get_end_to_end_against_local_server() {
unsafe { crate::scheduler::scheduler_init() };
pool::clear_for_test();
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let port = listener.local_addr().unwrap().port();
let accept_count = Arc::new(AtomicUsize::new(0));
let req_count = Arc::new(AtomicUsize::new(0));
spawn_test_server(
listener,
accept_count.clone(),
req_count.clone(),
ServerMode::FixedBody,
);
let target = loopback_target(port, "/hello");
let resp = perform_request_with_target("GET", target, None);
let (status, body, ok) = unwrap_response(&resp);
assert_eq!(status, 200);
assert_eq!(body, b"world");
assert!(ok);
assert_eq!(req_count.load(Ordering::SeqCst), 1);
assert_eq!(accept_count.load(Ordering::SeqCst), 1);
}
#[test]
#[serial_test::serial(http_global_state)]
fn http_pool_reuses_connection_for_second_request() {
unsafe { crate::scheduler::scheduler_init() };
pool::clear_for_test();
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let port = listener.local_addr().unwrap().port();
let accept_count = Arc::new(AtomicUsize::new(0));
let req_count = Arc::new(AtomicUsize::new(0));
spawn_test_server(
listener,
accept_count.clone(),
req_count.clone(),
ServerMode::FixedBody,
);
let target1 = loopback_target(port, "/first");
let r1 = perform_request_with_target("GET", target1, None);
assert_eq!(unwrap_response(&r1).0, 200);
let target2 = loopback_target(port, "/second");
let r2 = perform_request_with_target("GET", target2, None);
assert_eq!(unwrap_response(&r2).0, 200);
assert_eq!(req_count.load(Ordering::SeqCst), 2);
assert_eq!(
accept_count.load(Ordering::SeqCst),
1,
"second request should reuse pooled connection (saw {} accepts)",
accept_count.load(Ordering::SeqCst)
);
}
#[test]
#[serial_test::serial(http_global_state)]
fn http_post_with_body_round_trips() {
unsafe { crate::scheduler::scheduler_init() };
pool::clear_for_test();
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let port = listener.local_addr().unwrap().port();
let accept_count = Arc::new(AtomicUsize::new(0));
let req_count = Arc::new(AtomicUsize::new(0));
spawn_test_server(
listener,
accept_count.clone(),
req_count.clone(),
ServerMode::Echo,
);
let target = loopback_target(port, "/create");
let body = b"{\"name\":\"alice\"}";
let resp =
perform_request_with_target("POST", target, Some(("application/json", body.as_slice())));
let (status, echoed, ok) = unwrap_response(&resp);
assert_eq!(status, 200);
assert!(ok);
assert_eq!(echoed, body, "POST body must round-trip byte-for-byte");
}
#[test]
#[serial_test::serial(http_global_state)]
fn http_chunked_response_decodes_end_to_end() {
unsafe { crate::scheduler::scheduler_init() };
pool::clear_for_test();
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let port = listener.local_addr().unwrap().port();
let accept_count = Arc::new(AtomicUsize::new(0));
let req_count = Arc::new(AtomicUsize::new(0));
spawn_test_server(
listener,
accept_count.clone(),
req_count.clone(),
ServerMode::Chunked,
);
let target = loopback_target(port, "/chunked");
let resp = perform_request_with_target("GET", target, None);
let (status, body, ok) = unwrap_response(&resp);
assert_eq!(status, 200);
assert!(ok);
assert_eq!(body, b"Hello world");
assert_eq!(req_count.load(Ordering::SeqCst), 1);
}
#[test]
#[serial_test::serial(http_global_state)]
fn http_connection_close_response_evicts_from_pool() {
unsafe { crate::scheduler::scheduler_init() };
pool::clear_for_test();
let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
let port = listener.local_addr().unwrap().port();
let accept_count = Arc::new(AtomicUsize::new(0));
let req_count = Arc::new(AtomicUsize::new(0));
spawn_test_server(
listener,
accept_count.clone(),
req_count.clone(),
ServerMode::CloseOnce,
);
let key = pool::PoolKey {
scheme: Scheme::Http,
host: "127.0.0.1".to_string(),
port,
};
let target1 = loopback_target(port, "/first");
let r1 = perform_request_with_target("GET", target1, None);
assert_eq!(unwrap_response(&r1).0, 200);
assert_eq!(
pool::idle_count_for_test(&key),
0,
"Connection: close response must not survive pool::release. \
An entry here means release ignored keep_alive=false."
);
let target2 = loopback_target(port, "/second");
let r2 = perform_request_with_target("GET", target2, None);
assert_eq!(unwrap_response(&r2).0, 200);
assert_eq!(req_count.load(Ordering::SeqCst), 2);
assert_eq!(
accept_count.load(Ordering::SeqCst),
2,
"Connection: close response must NOT be pooled: each request \
should trigger a fresh accept on the server."
);
}
#[test]
#[serial_test::serial(dns_global_state)]
fn ssrf_dns_rebinding_closure_holds_at_most_one_resolve_per_request() {
unsafe { crate::scheduler::scheduler_init() };
pool::clear_for_test();
crate::dns::clear_scripted_responses();
crate::dns::reset_resolve_call_count();
crate::dns::push_scripted_response(vec!["224.0.0.1".to_string()]);
let _resp = perform_request("GET", "http://example.com/", None);
let count = crate::dns::resolve_call_count();
assert_eq!(
count, 1,
"perform_request must call dns::resolve at most once per \
request. count={count} > 1 means the connect or pool path \
is re-resolving the hostname (DNS-rebinding regression): \
the SSRF-validated address list should flow into connect, \
not the original hostname string."
);
}
fn build_tls_test_pair() -> (Arc<rustls::ServerConfig>, Arc<rustls::ClientConfig>) {
let rcgen::CertifiedKey {
cert,
signing_key: key_pair,
} = rcgen::generate_simple_self_signed(vec!["localhost".to_string()])
.expect("generate self-signed cert");
let cert_der: rustls::pki_types::CertificateDer<'static> = cert.der().clone();
let key_der: rustls::pki_types::PrivateKeyDer<'static> =
rustls::pki_types::PrivatePkcs8KeyDer::from(key_pair.serialize_der()).into();
let server_config = rustls::ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)
.expect("server config");
let mut roots = rustls::RootCertStore::empty();
roots.add(cert_der).expect("add test cert as trust root");
let client_config = rustls::ClientConfig::builder()
.with_root_certificates(roots)
.with_no_client_auth();
(Arc::new(server_config), Arc::new(client_config))
}
fn spawn_tls_test_server(
listener: std::net::TcpListener,
server_config: Arc<rustls::ServerConfig>,
) {
std::thread::spawn(move || {
for stream in listener.incoming() {
let stream = match stream {
Ok(s) => s,
Err(_) => return,
};
let server_config = server_config.clone();
std::thread::spawn(move || {
let conn = match rustls::ServerConnection::new(server_config) {
Ok(c) => c,
Err(_) => return,
};
let mut tls = rustls::StreamOwned::new(conn, stream);
let mut reader = BufReader::new(&mut tls);
let mut saw_request = false;
loop {
let mut line = String::new();
match reader.read_line(&mut line) {
Ok(0) | Err(_) => return,
_ => {}
}
if line == "\r\n" || line == "\n" {
if saw_request {
break;
}
continue;
}
saw_request = true;
}
drop(reader);
let body = b"hello-over-tls";
let header = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: text/plain\r\n\
Content-Length: {}\r\n\
Connection: close\r\n\
\r\n",
body.len()
);
let _ = tls.write_all(header.as_bytes());
let _ = tls.write_all(body);
let _ = tls.flush();
});
}
});
}
struct TestTlsConfigGuard;
impl Drop for TestTlsConfigGuard {
fn drop(&mut self) {
crate::tls::clear_test_tls_config();
}
}
#[test]
#[serial_test::serial(tls_global_state)]
#[serial_test::serial(http_global_state)]
fn https_round_trip_against_same_process_rustls_server() {
unsafe { crate::scheduler::scheduler_init() };
pool::clear_for_test();
let (server_config, client_config) = build_tls_test_pair();
crate::tls::install_test_tls_config(client_config);
let _tls_guard = TestTlsConfigGuard;
let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind");
let port = listener.local_addr().unwrap().port();
spawn_tls_test_server(listener, server_config);
let target = ValidatedTarget {
scheme: Scheme::Https,
host: "localhost".to_string(),
port,
addrs: vec![IpAddr::V4(Ipv4Addr::LOCALHOST)],
path_and_query: "/hello".to_string(),
};
let resp = perform_request_with_target("GET", target, None);
let (status, body, ok) = unwrap_response(&resp);
assert_eq!(status, 200, "HTTPS round-trip status");
assert!(ok, "HTTPS round-trip ok flag");
assert_eq!(body, b"hello-over-tls", "HTTPS round-trip body");
pool::clear_for_test();
}
struct HttpRequestTimeoutGuard;
impl Drop for HttpRequestTimeoutGuard {
fn drop(&mut self) {
super::request::set_test_http_request_timeout(None);
}
}
#[test]
#[serial_test::serial(http_global_state)]
fn http_request_timeout_fires_on_silent_server() {
unsafe { crate::scheduler::scheduler_init() };
pool::clear_for_test();
super::request::set_test_http_request_timeout(Some(std::time::Duration::from_millis(200)));
let _timeout_guard = HttpRequestTimeoutGuard;
let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind");
let port = listener.local_addr().unwrap().port();
std::thread::spawn(move || {
if let Ok((stream, _)) = listener.accept() {
std::thread::sleep(std::time::Duration::from_secs(60));
drop(stream);
}
});
let target = loopback_target(port, "/silent");
let start = std::time::Instant::now();
let resp = perform_request_with_target("GET", target, None);
let elapsed = start.elapsed();
let (status, _body, ok) = unwrap_response(&resp);
assert_eq!(status, 0, "silent-server request must surface as error");
assert!(!ok, "ok flag must be false");
assert!(
elapsed < std::time::Duration::from_secs(2),
"HTTP request must respect the per-IO timeout (elapsed={elapsed:?}). \
Above 2s suggests the deadline isn't being plumbed through to \
wire::read_response."
);
}
struct TlsHandshakeTimeoutGuard;
impl Drop for TlsHandshakeTimeoutGuard {
fn drop(&mut self) {
crate::tls::set_test_tls_handshake_timeout(None);
}
}
#[test]
#[serial_test::serial(tls_global_state)]
fn tls_handshake_timeout_fires_on_silent_peer() {
unsafe { crate::scheduler::scheduler_init() };
pool::clear_for_test();
let (_server_config, client_config) = build_tls_test_pair();
crate::tls::install_test_tls_config(client_config);
let _tls_guard = TestTlsConfigGuard;
crate::tls::set_test_tls_handshake_timeout(Some(std::time::Duration::from_millis(200)));
let _handshake_guard = TlsHandshakeTimeoutGuard;
let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind");
let port = listener.local_addr().unwrap().port();
std::thread::spawn(move || {
if let Ok((stream, _)) = listener.accept() {
std::thread::sleep(std::time::Duration::from_secs(60));
drop(stream);
}
});
let target = ValidatedTarget {
scheme: Scheme::Https,
host: "localhost".to_string(),
port,
addrs: vec![IpAddr::V4(Ipv4Addr::LOCALHOST)],
path_and_query: "/".to_string(),
};
let start = std::time::Instant::now();
let resp = perform_request_with_target("GET", target, None);
let elapsed = start.elapsed();
let (status, _body, ok) = unwrap_response(&resp);
assert_eq!(status, 0, "silent-peer TLS upgrade must surface as error");
assert!(!ok);
assert!(
elapsed < std::time::Duration::from_secs(2),
"TLS handshake must respect the per-IO timeout (elapsed={elapsed:?}). \
Above 2s suggests build_tls isn't setting read/write timeouts on \
the TcpStream before complete_io."
);
pool::clear_for_test();
}
struct TcpConnectTimeoutGuard;
impl Drop for TcpConnectTimeoutGuard {
fn drop(&mut self) {
crate::tcp::set_test_tcp_connect_timeout(None);
}
}
#[test]
#[serial_test::serial(http_global_state)]
fn tcp_connect_timeout_fires_on_silent_route() {
unsafe { crate::scheduler::scheduler_init() };
pool::clear_for_test();
crate::tcp::set_test_tcp_connect_timeout(Some(std::time::Duration::from_millis(200)));
let _connect_guard = TcpConnectTimeoutGuard;
let target = ValidatedTarget {
scheme: Scheme::Http,
host: "192.0.2.1".to_string(),
port: 80,
addrs: vec![IpAddr::V4(Ipv4Addr::new(192, 0, 2, 1))],
path_and_query: "/".to_string(),
};
let start = std::time::Instant::now();
let resp = perform_request_with_target("GET", target, None);
let elapsed = start.elapsed();
let (status, _body, ok) = unwrap_response(&resp);
assert_eq!(status, 0, "silent-route connect must surface as error");
assert!(!ok, "ok flag must be false");
assert!(
elapsed < std::time::Duration::from_secs(2),
"TCP connect must respect the configured timeout (elapsed={elapsed:?}). \
Above 2s suggests connect_to_addrs is calling plain connect, \
not connect_timeout."
);
}