mod common;
use std::str::FromStr;
use common::{
get_available_port, get_available_udp_port, socks5_connect_ipv4, start_tunnel, TEST_TIMEOUT,
};
use rusnel::common::remote::RemoteRequest;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream, UdpSocket};
use tokio::time::timeout;
const LARGE_PAYLOAD_LEN: usize = 256 * 1024;
const UDP_PACKET_LEN: usize = 1000;
const UDP_PACKET_COUNT: usize = 200;
fn xorshift_fill(buf: &mut [u8], seed: u64) {
let mut state = seed.max(1);
for chunk in buf.chunks_mut(8) {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
let bytes = state.to_le_bytes();
let n = chunk.len();
chunk.copy_from_slice(&bytes[..n]);
}
}
fn make_payload(len: usize, seed: u64) -> Vec<u8> {
let mut buf = vec![0u8; len];
xorshift_fill(&mut buf, seed);
buf
}
#[tokio::test]
async fn test_tcp_forward_large_payload() {
timeout(TEST_TIMEOUT, async {
let server_port = get_available_port();
let local_port = get_available_port();
let remote_port = get_available_port();
let target_listener = TcpListener::bind(format!("127.0.0.1:{remote_port}"))
.await
.unwrap();
let remote =
RemoteRequest::from_str(&format!("127.0.0.1:{local_port}:127.0.0.1:{remote_port}"))
.unwrap();
let _env = start_tunnel(server_port, false, vec![remote]).await;
let payload = make_payload(LARGE_PAYLOAD_LEN, 0xC0FFEE);
let target_task = {
let expected = payload.clone();
tokio::spawn(async move {
let (mut stream, _) = target_listener.accept().await.unwrap();
let mut received = Vec::with_capacity(expected.len());
stream.read_to_end(&mut received).await.unwrap();
assert_eq!(received.len(), expected.len(), "received length mismatch");
assert!(received == expected, "payload bytes mismatch");
})
};
let mut client_conn = TcpStream::connect(format!("127.0.0.1:{local_port}"))
.await
.unwrap();
client_conn.write_all(&payload).await.unwrap();
client_conn.shutdown().await.unwrap();
target_task.await.unwrap();
})
.await
.expect("test_tcp_forward_large_payload timed out");
}
#[tokio::test]
async fn test_tcp_reverse_large_payload() {
timeout(TEST_TIMEOUT, async {
let server_port = get_available_port();
let listen_port = get_available_port();
let target_port = get_available_port();
let target_listener = TcpListener::bind(format!("127.0.0.1:{target_port}"))
.await
.unwrap();
let remote = RemoteRequest::from_str(&format!(
"R:127.0.0.1:{listen_port}:127.0.0.1:{target_port}"
))
.unwrap();
let _env = start_tunnel(server_port, true, vec![remote]).await;
let payload = make_payload(LARGE_PAYLOAD_LEN, 0xBADCAFE);
let target_task = {
let expected = payload.clone();
tokio::spawn(async move {
let (mut stream, _) = target_listener.accept().await.unwrap();
let mut received = Vec::with_capacity(expected.len());
stream.read_to_end(&mut received).await.unwrap();
assert_eq!(received.len(), expected.len(), "received length mismatch");
assert!(received == expected, "payload bytes mismatch");
})
};
let mut client_conn = TcpStream::connect(format!("127.0.0.1:{listen_port}"))
.await
.unwrap();
client_conn.write_all(&payload).await.unwrap();
client_conn.shutdown().await.unwrap();
target_task.await.unwrap();
})
.await
.expect("test_tcp_reverse_large_payload timed out");
}
#[tokio::test]
async fn test_tcp_forward_large_bidirectional_echo() {
timeout(TEST_TIMEOUT, async {
let server_port = get_available_port();
let local_port = get_available_port();
let remote_port = get_available_port();
let target_listener = TcpListener::bind(format!("127.0.0.1:{remote_port}"))
.await
.unwrap();
let remote =
RemoteRequest::from_str(&format!("127.0.0.1:{local_port}:127.0.0.1:{remote_port}"))
.unwrap();
let _env = start_tunnel(server_port, false, vec![remote]).await;
let payload = make_payload(LARGE_PAYLOAD_LEN, 0xDECAFBAD);
let echo_task = tokio::spawn(async move {
let (stream, _) = target_listener.accept().await.unwrap();
let (mut r, mut w) = stream.into_split();
let mut buf = vec![0u8; 64 * 1024];
loop {
let n = r.read(&mut buf).await.unwrap();
if n == 0 {
break;
}
w.write_all(&buf[..n]).await.unwrap();
}
w.shutdown().await.unwrap();
});
let mut client_conn = TcpStream::connect(format!("127.0.0.1:{local_port}"))
.await
.unwrap();
let (mut client_r, mut client_w) = client_conn.split();
let payload_clone = payload.clone();
let writer = async move {
client_w.write_all(&payload_clone).await.unwrap();
client_w.shutdown().await.unwrap();
};
let expected = payload.clone();
let reader = async move {
let mut received = Vec::with_capacity(expected.len());
client_r.read_to_end(&mut received).await.unwrap();
assert_eq!(received.len(), expected.len(), "echoed length mismatch");
assert!(received == expected, "echoed payload bytes mismatch");
};
tokio::join!(writer, reader);
echo_task.await.unwrap();
})
.await
.expect("test_tcp_forward_large_bidirectional_echo timed out");
}
#[tokio::test]
async fn test_socks5_large_payload() {
timeout(TEST_TIMEOUT, async {
let server_port = get_available_port();
let socks_port = get_available_port();
let target_port = get_available_port();
let target_listener = TcpListener::bind(format!("127.0.0.1:{target_port}"))
.await
.unwrap();
let remote = RemoteRequest::from_str(&format!("127.0.0.1:{socks_port}:socks")).unwrap();
let _env = start_tunnel(server_port, false, vec![remote]).await;
let payload = make_payload(LARGE_PAYLOAD_LEN, 0x5_0CC_5_0CC);
let target_task = {
let expected = payload.clone();
tokio::spawn(async move {
let (mut stream, _) = target_listener.accept().await.unwrap();
let mut received = Vec::with_capacity(expected.len());
stream.read_to_end(&mut received).await.unwrap();
assert_eq!(received.len(), expected.len(), "received length mismatch");
assert!(received == expected, "payload bytes mismatch");
})
};
let mut socks_conn = socks5_connect_ipv4(
&format!("127.0.0.1:{socks_port}"),
[127, 0, 0, 1],
target_port,
)
.await;
socks_conn.write_all(&payload).await.unwrap();
socks_conn.shutdown().await.unwrap();
target_task.await.unwrap();
})
.await
.expect("test_socks5_large_payload timed out");
}
#[tokio::test]
async fn test_udp_forward_many_packets() {
timeout(TEST_TIMEOUT, async {
let server_port = get_available_port();
let local_port = get_available_udp_port();
let remote_port = get_available_udp_port();
let target_socket = UdpSocket::bind(format!("127.0.0.1:{remote_port}"))
.await
.unwrap();
let remote = RemoteRequest::from_str(&format!(
"127.0.0.1:{local_port}:127.0.0.1:{remote_port}/udp"
))
.unwrap();
let _env = start_tunnel(server_port, false, vec![remote]).await;
let recv_task = tokio::spawn(async move {
let mut seen = vec![false; UDP_PACKET_COUNT];
let mut buf = vec![0u8; 4096];
let mut received = 0usize;
while received < UDP_PACKET_COUNT {
let n = target_socket.recv(&mut buf).await.unwrap();
assert_eq!(n, UDP_PACKET_LEN, "unexpected datagram length");
let mut seq_bytes = [0u8; 8];
seq_bytes.copy_from_slice(&buf[..8]);
let seq = u64::from_le_bytes(seq_bytes) as usize;
assert!(seq < UDP_PACKET_COUNT, "sequence number out of range");
let mut expected = vec![0u8; UDP_PACKET_LEN];
expected[..8].copy_from_slice(&(seq as u64).to_le_bytes());
xorshift_fill(&mut expected[8..], (seq as u64).wrapping_add(1));
assert_eq!(&buf[..n], &expected[..], "datagram payload mismatch");
if !seen[seq] {
seen[seq] = true;
received += 1;
}
}
});
let sender = UdpSocket::bind("127.0.0.1:0").await.unwrap();
for seq in 0..UDP_PACKET_COUNT {
let mut pkt = vec![0u8; UDP_PACKET_LEN];
pkt[..8].copy_from_slice(&(seq as u64).to_le_bytes());
xorshift_fill(&mut pkt[8..], (seq as u64).wrapping_add(1));
sender
.send_to(&pkt, format!("127.0.0.1:{local_port}"))
.await
.unwrap();
}
recv_task.await.unwrap();
})
.await
.expect("test_udp_forward_many_packets timed out");
}