#![deny(unsafe_code)]
use std::io::{Read, Write};
use std::net::{TcpListener, TcpStream, UdpSocket};
use std::time::Duration;
const TEST_DATA: &[u8] = b"zenith-l4-relay-test-payload-1234567890";
fn test_tcp_to_tcp_relay() -> (u32, u32) {
println!("\n============================================================");
println!("测试 1: TCP → TCP 端口转发 (TcpRelay 统一 API)");
println!("============================================================");
let mut pass = 0u32;
let mut fail = 0u32;
let backend = TcpListener::bind("127.0.0.1:0").unwrap();
let backend_addr = backend.local_addr().unwrap();
println!(" 后端 TCP: {}", backend_addr);
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let listen_addr = listener.local_addr().unwrap();
println!(" 转发监听 TCP: {} → {}", listen_addr, backend_addr);
let backend_thread = std::thread::spawn(move || {
let (mut conn, _peer) = backend.accept().unwrap();
let mut buf = [0u8; 4096];
let n = conn.read(&mut buf).unwrap();
conn.write_all(&buf[..n]).unwrap();
});
let relay_thread = std::thread::spawn(move || {
let (client_conn, _) = listener.accept().unwrap();
let upstream = TcpStream::connect(backend_addr).unwrap();
let relay = zenith_forward::TcpRelay::default();
match relay.relay(&client_conn, &upstream) {
Ok((c2u, u2c)) => {
println!(" TcpRelay 完成: c2u={}B u2c={}B", c2u, u2c);
}
Err(e) => {
println!(" TcpRelay 错误: {}", e);
}
}
});
std::thread::sleep(Duration::from_millis(50));
let mut client = TcpStream::connect_timeout(
&listen_addr,
Duration::from_secs(2),
).unwrap();
client.set_read_timeout(Some(Duration::from_secs(5))).unwrap();
client.write_all(TEST_DATA).unwrap();
let mut buf = [0u8; 4096];
let n = client.read(&mut buf).unwrap();
if &buf[..n] == TEST_DATA {
println!(" [PASS] TCP→TCP 回显正确 ({} 字节)", n);
pass += 1;
} else {
println!(" [FAIL] TCP→TCP 回显不匹配");
fail += 1;
}
drop(client);
let _ = relay_thread.join();
let _ = backend_thread.join();
println!(" TCP→TCP: {} 通过, {} 失败", pass, fail);
(pass, fail)
}
fn test_udp_to_udp_relay() -> (u32, u32) {
println!("\n============================================================");
println!("测试 2: UDP → UDP 端口转发 (recv_from → send_to)");
println!("============================================================");
let mut pass = 0u32;
let mut fail = 0u32;
let backend = UdpSocket::bind("127.0.0.1:0").unwrap();
let backend_addr = backend.local_addr().unwrap();
println!(" 后端 UDP: {}", backend_addr);
let relay_sock = UdpSocket::bind("127.0.0.1:0").unwrap();
let relay_addr = relay_sock.local_addr().unwrap();
println!(" 转发监听 UDP: {} → {}", relay_addr, backend_addr);
let backend_thread = std::thread::spawn(move || {
let mut buf = [0u8; 4096];
let (n, src) = backend.recv_from(&mut buf).unwrap();
backend.send_to(&buf[..n], src).unwrap();
});
let relay_thread = std::thread::spawn(move || {
let mut buf = [0u8; 4096];
let (n, client_addr) = relay_sock.recv_from(&mut buf).unwrap();
relay_sock.send_to(&buf[..n], backend_addr).unwrap();
let (n2, _) = relay_sock.recv_from(&mut buf).unwrap();
relay_sock.send_to(&buf[..n2], client_addr).unwrap();
});
std::thread::sleep(Duration::from_millis(50));
let client = UdpSocket::bind("127.0.0.1:0").unwrap();
client.set_read_timeout(Some(Duration::from_secs(5))).unwrap();
client.send_to(TEST_DATA, relay_addr).unwrap();
let mut buf = [0u8; 4096];
let (n, _) = client.recv_from(&mut buf).unwrap();
if &buf[..n] == TEST_DATA {
println!(" [PASS] UDP→UDP 回显正确 ({} 字节)", n);
pass += 1;
} else {
println!(" [FAIL] UDP→UDP 回显不匹配");
fail += 1;
}
let _ = relay_thread.join();
let _ = backend_thread.join();
println!(" UDP→UDP: {} 通过, {} 失败", pass, fail);
(pass, fail)
}
fn test_tcp_to_udp_relay() -> (u32, u32) {
println!("\n============================================================");
println!("测试 3: TCP → UDP 协议转换转发");
println!("============================================================");
let mut pass = 0u32;
let mut fail = 0u32;
let backend = UdpSocket::bind("127.0.0.1:0").unwrap();
let backend_addr = backend.local_addr().unwrap();
println!(" 后端 UDP: {}", backend_addr);
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let listen_addr = listener.local_addr().unwrap();
println!(" 转发监听 TCP: {} → UDP {}", listen_addr, backend_addr);
let backend_thread = std::thread::spawn(move || {
let mut buf = [0u8; 4096];
let (n, src) = backend.recv_from(&mut buf).unwrap();
backend.send_to(&buf[..n], src).unwrap();
});
let relay_thread = std::thread::spawn(move || {
let (mut tcp_conn, _) = listener.accept().unwrap();
tcp_conn.set_read_timeout(Some(Duration::from_secs(5))).unwrap();
let udp_sock = UdpSocket::bind("127.0.0.1:0").unwrap();
let udp_local = udp_sock.local_addr().unwrap();
let mut buf = [0u8; 4096];
let n = tcp_conn.read(&mut buf).unwrap();
udp_sock.send_to(&buf[..n], backend_addr).unwrap();
let (n2, _) = udp_sock.recv_from(&mut buf).unwrap();
tcp_conn.write_all(&buf[..n2]).unwrap();
});
std::thread::sleep(Duration::from_millis(50));
let mut client = TcpStream::connect_timeout(&listen_addr, Duration::from_secs(2)).unwrap();
client.set_read_timeout(Some(Duration::from_secs(5))).unwrap();
client.write_all(TEST_DATA).unwrap();
let mut buf = [0u8; 4096];
let n = client.read(&mut buf).unwrap();
if &buf[..n] == TEST_DATA {
println!(" [PASS] TCP→UDP→TCP 回显正确 ({} 字节)", n);
pass += 1;
} else {
println!(" [FAIL] TCP→UDP→TCP 回显不匹配");
fail += 1;
}
drop(client);
let _ = relay_thread.join();
let _ = backend_thread.join();
println!(" TCP→UDP: {} 通过, {} 失败", pass, fail);
(pass, fail)
}
fn test_udp_to_tcp_relay() -> (u32, u32) {
println!("\n============================================================");
println!("测试 4: UDP → TCP 协议转换转发");
println!("============================================================");
let mut pass = 0u32;
let mut fail = 0u32;
let backend = TcpListener::bind("127.0.0.1:0").unwrap();
let backend_addr = backend.local_addr().unwrap();
println!(" 后端 TCP: {}", backend_addr);
let relay_udp = UdpSocket::bind("127.0.0.1:0").unwrap();
let relay_addr = relay_udp.local_addr().unwrap();
println!(" 转发监听 UDP: {} → TCP {}", relay_addr, backend_addr);
let backend_thread = std::thread::spawn(move || {
let (mut conn, _) = backend.accept().unwrap();
let mut buf = [0u8; 4096];
let n = conn.read(&mut buf).unwrap();
conn.write_all(&buf[..n]).unwrap();
});
let relay_thread = std::thread::spawn(move || {
let mut buf = [0u8; 4096];
let (n, client_addr) = relay_udp.recv_from(&mut buf).unwrap();
let mut tcp_conn = TcpStream::connect_timeout(&backend_addr, Duration::from_secs(2)).unwrap();
tcp_conn.set_read_timeout(Some(Duration::from_secs(5))).unwrap();
tcp_conn.write_all(&buf[..n]).unwrap();
let n2 = tcp_conn.read(&mut buf).unwrap();
relay_udp.send_to(&buf[..n2], client_addr).unwrap();
});
std::thread::sleep(Duration::from_millis(50));
let client = UdpSocket::bind("127.0.0.1:0").unwrap();
client.set_read_timeout(Some(Duration::from_secs(5))).unwrap();
client.send_to(TEST_DATA, relay_addr).unwrap();
let mut buf = [0u8; 4096];
let (n, _) = client.recv_from(&mut buf).unwrap();
if &buf[..n] == TEST_DATA {
println!(" [PASS] UDP→TCP→UDP 回显正确 ({} 字节)", n);
pass += 1;
} else {
println!(" [FAIL] UDP→TCP→UDP 回显不匹配");
fail += 1;
}
let _ = relay_thread.join();
let _ = backend_thread.join();
println!(" UDP→TCP: {} 通过, {} 失败", pass, fail);
(pass, fail)
}
fn test_admission_udp_filter() -> (u32, u32) {
use zenith_net::source_admission::*;
println!("\n============================================================");
println!("测试 5: SourceAdmissionEngine UDP 包过滤");
println!("============================================================");
let mut pass = 0u32;
let mut fail = 0u32;
let mut engine = SourceAdmissionEngine::deny_all();
engine.add_rule(AdmissionRule {
id: 1,
src_ip: IpAddr::V4([127, 0, 0, 1]),
prefix_len: 0,
src_port: 0,
dst_port: 8080,
proto: ProtoMatch::Udp,
action: AdmissionAction::Allow,
enabled: true,
}).unwrap();
engine.add_rule(AdmissionRule {
id: 2,
src_ip: IpAddr::V4([127, 0, 0, 1]),
prefix_len: 0,
src_port: 0,
dst_port: 443,
proto: ProtoMatch::Tcp,
action: AdmissionAction::Allow,
enabled: true,
}).unwrap();
let a = engine.evaluate(IpAddr::V4([127, 0, 0, 1]), 12345, 8080, 17);
if a == AdmissionAction::Allow {
println!(" [PASS] 127.0.0.1→UDP/8080 允许");
pass += 1;
} else { println!(" [FAIL] 127.0.0.1→UDP/8080 应允许"); fail += 1; }
let a = engine.evaluate(IpAddr::V4([127, 0, 0, 1]), 12345, 443, 6);
if a == AdmissionAction::Allow {
println!(" [PASS] 127.0.0.1→TCP/443 允许");
pass += 1;
} else { println!(" [FAIL] 127.0.0.1→TCP/443 应允许"); fail += 1; }
let a = engine.evaluate(IpAddr::V4([127, 0, 0, 1]), 12345, 443, 17);
if a == AdmissionAction::Deny {
println!(" [PASS] 127.0.0.1→UDP/443 被拒绝 (仅 TCP/443 允许)");
pass += 1;
} else { println!(" [FAIL] 127.0.0.1→UDP/443 应拒绝"); fail += 1; }
let a = engine.evaluate(IpAddr::V4([127, 0, 0, 1]), 12345, 8080, 6);
if a == AdmissionAction::Deny {
println!(" [PASS] 127.0.0.1→TCP/8080 被拒绝 (仅 UDP/8080 允许)");
pass += 1;
} else { println!(" [FAIL] 127.0.0.1→TCP/8080 应拒绝"); fail += 1; }
let a = engine.evaluate(IpAddr::V4([10, 0, 0, 1]), 12345, 8080, 17);
if a == AdmissionAction::Deny {
println!(" [PASS] 10.0.0.1→UDP/8080 被拒绝 (IP 不在白名单)");
pass += 1;
} else { println!(" [FAIL] 10.0.0.1→UDP/8080 应拒绝"); fail += 1; }
let a = engine.evaluate(IpAddr::V4([127, 0, 0, 1]), 0, 0, 1);
if a == AdmissionAction::Deny {
println!(" [PASS] 127.0.0.1→ICMP 被拒绝 (无 ICMP 规则)");
pass += 1;
} else { println!(" [FAIL] 127.0.0.1→ICMP 应拒绝"); fail += 1; }
println!(" UDP 过滤: {} 通过, {} 失败", pass, fail);
(pass, fail)
}
fn main() {
println!("Zenith L4 端口转发/协议转换全矩阵测试");
println!("测试数据: {} 字节", TEST_DATA.len());
let (p1, f1) = test_tcp_to_tcp_relay();
let (p2, f2) = test_udp_to_udp_relay();
let (p3, f3) = test_tcp_to_udp_relay();
let (p4, f4) = test_udp_to_tcp_relay();
let (p5, f5) = test_admission_udp_filter();
let total_pass = p1 + p2 + p3 + p4 + p5;
let total_fail = f1 + f2 + f3 + f4 + f5;
println!("\n============================================================");
println!(" 总计: {} 通过, {} 失败", total_pass, total_fail);
println!("============================================================");
if total_fail > 0 {
std::process::exit(1);
}
}