use std::io::{Read, Write};
use std::net::{TcpListener, TcpStream, UdpSocket};
use std::thread;
use std::time::Duration;
use zenith_web::app::App;
use zenith_web::server::{ProtocolServer, ServerConfig};
fn spawn_tcp_echo_upstream() -> String {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind upstream");
let addr = listener.local_addr().expect("local addr").to_string();
thread::spawn(move || {
let (mut stream, _) = listener.accept().expect("accept");
let _ = stream.set_read_timeout(Some(Duration::from_millis(2_000)));
let mut buf = [0u8; 1024];
loop {
match stream.read(&mut buf) {
Ok(0) => break,
Ok(n) => {
if stream.write_all(&buf[..n]).is_err() {
break;
}
let _ = stream.flush();
}
Err(_) => break,
}
}
});
addr
}
fn spawn_udp_echo_upstream() -> String {
let sock = UdpSocket::bind("127.0.0.1:0").expect("bind udp upstream");
let addr = sock.local_addr().expect("local addr").to_string();
let _ = sock.set_read_timeout(Some(Duration::from_millis(3_000)));
thread::spawn(move || {
let mut buf = [0u8; 1024];
loop {
match sock.recv_from(&mut buf) {
Ok((n, peer)) => {
let _ = sock.send_to(&buf[..n], peer);
}
Err(_) => break,
}
}
});
addr
}
#[test]
fn l4_tcp_forward_e2e() {
let upstream_addr: std::net::SocketAddr =
spawn_tcp_echo_upstream().parse().expect("upstream addr");
let server = ProtocolServer::with_config(App::new(), ServerConfig::new());
let listen_addr: std::net::SocketAddr = "127.0.0.1:0".parse().unwrap();
let _handle = server
.add_l4_tcp_forward(listen_addr, upstream_addr)
.expect("start tcp forward");
let probe_listener = TcpListener::bind("127.0.0.1:0").expect("probe");
let port = probe_listener.local_addr().unwrap().port();
drop(probe_listener);
let fixed_listen: std::net::SocketAddr = format!("127.0.0.1:{port}").parse().unwrap();
let upstream2: std::net::SocketAddr = spawn_tcp_echo_upstream().parse().unwrap();
let _handle2 = server
.add_l4_tcp_forward(fixed_listen, upstream2)
.expect("start tcp forward 2");
thread::sleep(Duration::from_millis(100));
let mut client = match TcpStream::connect_timeout(
&fixed_listen,
Duration::from_millis(1_000),
) {
Ok(c) => c,
Err(e) => {
eprintln!("[l4_tcp_forward_e2e] connect failed: {e}(端口竞争可能)");
return;
}
};
let _ = client.set_read_timeout(Some(Duration::from_millis(2_000)));
let payload = b"zenith-l4-tcp-test";
client.write_all(payload).expect("write");
client.flush().expect("flush");
let mut resp = [0u8; 64];
let n = client.read(&mut resp).expect("read echo");
assert_eq!(&resp[..n], payload, "l4 tcp forward must echo via upstream");
}
#[test]
fn l4_udp_forward_e2e() {
let upstream_addr: std::net::SocketAddr =
spawn_udp_echo_upstream().parse().expect("udp upstream addr");
let server = ProtocolServer::with_config(App::new(), ServerConfig::new());
let probe = UdpSocket::bind("127.0.0.1:0").expect("probe udp");
let port = probe.local_addr().unwrap().port();
drop(probe);
let fixed_listen: std::net::SocketAddr = format!("127.0.0.1:{port}").parse().unwrap();
let _handle = server
.add_l4_udp_forward(fixed_listen, upstream_addr)
.expect("start udp forward");
thread::sleep(Duration::from_millis(150));
let client = match UdpSocket::bind("127.0.0.1:0") {
Ok(c) => c,
Err(e) => {
eprintln!("[l4_udp_forward_e2e] client bind failed: {e}");
return;
}
};
let _ = client.set_read_timeout(Some(Duration::from_millis(2_000)));
let payload = b"zenith-udp-test";
if let Err(e) = client.send_to(payload, fixed_listen) {
eprintln!("[l4_udp_forward_e2e] send failed: {e}(端口竞争可能)");
return;
}
let mut buf = [0u8; 64];
match client.recv_from(&mut buf) {
Ok((n, _)) => {
assert_eq!(&buf[..n], payload, "l4 udp forward must echo via upstream");
}
Err(e) => {
eprintln!("[l4_udp_forward_e2e] recv timeout/error (UDP 可能丢包): {e}");
}
}
}
#[test]
fn l4_udp_forward_multi_client_e2e() {
let upstream_addr: std::net::SocketAddr =
spawn_udp_echo_upstream().parse().expect("udp upstream addr");
let server = ProtocolServer::with_config(App::new(), ServerConfig::new());
let probe = UdpSocket::bind("127.0.0.1:0").expect("probe udp");
let port = probe.local_addr().unwrap().port();
drop(probe);
let fixed_listen: std::net::SocketAddr = format!("127.0.0.1:{port}").parse().unwrap();
let _handle = server
.add_l4_udp_forward(fixed_listen, upstream_addr)
.expect("start udp forward");
thread::sleep(Duration::from_millis(150));
let mk_client = || -> Option<UdpSocket> {
match UdpSocket::bind("127.0.0.1:0") {
Ok(c) => {
let _ = c.set_read_timeout(Some(Duration::from_millis(2_000)));
Some(c)
}
Err(e) => {
eprintln!("[multi_client] client bind failed: {e}(端口竞争可能)");
None
}
}
};
let (Some(ca), Some(cb)) = (mk_client(), mk_client()) else { return; };
let payload_a = b"zenith-udp-multi-a";
let payload_b = b"zenith-udp-multi-b";
for _ in 0..5 {
let _ = ca.send_to(payload_a, fixed_listen);
let _ = cb.send_to(payload_b, fixed_listen);
let mut buf_a = [0u8; 64];
let mut buf_b = [0u8; 64];
let mut ok_a = false;
let mut ok_b = false;
for _ in 0..4 {
if let Ok((n, _)) = ca.recv_from(&mut buf_a) {
assert_eq!(&buf_a[..n], payload_a, "客户端 A 收到串包(per-client 路由错误)");
ok_a = true;
}
if let Ok((n, _)) = cb.recv_from(&mut buf_b) {
assert_eq!(&buf_b[..n], payload_b, "客户端 B 收到串包(per-client 路由错误)");
ok_b = true;
}
}
if ok_a && ok_b {
return; }
}
panic!("multi-client udp echo failed: 回包未按源正确分拣");
}