mod common;
use std::str::FromStr;
use std::time::Duration;
use common::{
get_available_port, socks5_connect_ipv4, start_tunnel, start_tunnel_with_flags, TEST_TIMEOUT,
};
use rusnel::common::remote::RemoteRequest;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio::time::{sleep, timeout};
#[tokio::test]
async fn test_reverse_rejected_when_not_allowed() {
timeout(TEST_TIMEOUT, async {
let server_port = get_available_port();
let listen_port = get_available_port();
let target_port = get_available_port();
let _target = 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, false, vec![remote]).await;
sleep(Duration::from_millis(300)).await;
let connect_res = timeout(
Duration::from_secs(2),
TcpStream::connect(format!("127.0.0.1:{listen_port}")),
)
.await;
match connect_res {
Ok(Ok(_)) => panic!(
"expected reverse listener on port {listen_port} to be closed, but it accepted a connection"
),
Ok(Err(_)) => { }
Err(_) => panic!("connect attempt unexpectedly hung"),
}
})
.await
.expect("test_reverse_rejected_when_not_allowed timed out");
}
#[tokio::test]
async fn test_forward_socks_rejected_when_not_allowed() {
timeout(TEST_TIMEOUT, async {
let server_port = get_available_port();
let socks_port = get_available_port();
let remote = RemoteRequest::from_str(&format!("127.0.0.1:{socks_port}:socks")).unwrap();
let _env = start_tunnel_with_flags(server_port, false, false, vec![remote]).await;
sleep(Duration::from_millis(300)).await;
let connect_res = timeout(
Duration::from_secs(2),
TcpStream::connect(format!("127.0.0.1:{socks_port}")),
)
.await;
match connect_res {
Ok(Ok(_)) => panic!(
"expected local SOCKS listener on port {socks_port} to never bind, but it accepted a connection"
),
Ok(Err(_)) => { }
Err(_) => panic!("connect attempt unexpectedly hung"),
}
})
.await
.expect("test_forward_socks_rejected_when_not_allowed timed out");
}
#[tokio::test]
async fn test_reverse_socks_requires_both_flags() {
timeout(TEST_TIMEOUT, async {
let server_port = get_available_port();
let socks_listen_port = get_available_port();
let remote =
RemoteRequest::from_str(&format!("R:127.0.0.1:{socks_listen_port}:socks")).unwrap();
let _env = start_tunnel_with_flags(server_port, true, false, vec![remote]).await;
sleep(Duration::from_millis(300)).await;
let connect_res = timeout(
Duration::from_secs(2),
TcpStream::connect(format!("127.0.0.1:{socks_listen_port}")),
)
.await;
match connect_res {
Ok(Ok(_)) => panic!(
"expected reverse SOCKS listener on port {socks_listen_port} to be closed, but it accepted a connection"
),
Ok(Err(_)) => { }
Err(_) => panic!("connect attempt unexpectedly hung"),
}
})
.await
.expect("test_reverse_socks_requires_both_flags timed out");
}
#[tokio::test]
async fn test_tcp_forward_half_close() {
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 mut client_conn = TcpStream::connect(format!("127.0.0.1:{local_port}"))
.await
.unwrap();
let (mut target_stream, _) = target_listener.accept().await.unwrap();
client_conn.write_all(b"req").await.unwrap();
client_conn.shutdown().await.unwrap();
let mut buf = vec![0u8; 64];
let n = target_stream.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"req");
let n = target_stream.read(&mut buf).await.unwrap();
assert_eq!(n, 0, "expected EOF from target after client shutdown");
target_stream.write_all(b"resp").await.unwrap();
target_stream.shutdown().await.unwrap();
let mut got = Vec::new();
client_conn.read_to_end(&mut got).await.unwrap();
assert_eq!(got, b"resp");
})
.await
.expect("test_tcp_forward_half_close timed out");
}
#[tokio::test]
async fn test_tcp_forward_empty_transfer() {
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 mut client_conn = TcpStream::connect(format!("127.0.0.1:{local_port}"))
.await
.unwrap();
let (mut target_stream, _) = target_listener.accept().await.unwrap();
client_conn.shutdown().await.unwrap();
drop(client_conn);
let mut buf = vec![0u8; 64];
let n = timeout(Duration::from_secs(3), target_stream.read(&mut buf))
.await
.expect("target never observed EOF after empty transfer")
.unwrap();
assert_eq!(n, 0, "expected EOF, got {n} bytes");
})
.await
.expect("test_tcp_forward_empty_transfer timed out");
}
#[tokio::test]
async fn test_tcp_forward_small_immediate_write_loses_data() {
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 = b"hello-world";
let target_task = tokio::spawn(async move {
let (mut stream, _) = target_listener.accept().await.unwrap();
let mut got = Vec::new();
stream.read_to_end(&mut got).await.unwrap();
got
});
let mut conn = TcpStream::connect(format!("127.0.0.1:{local_port}"))
.await
.unwrap();
conn.write_all(payload).await.unwrap();
conn.shutdown().await.unwrap();
let got = target_task.await.unwrap();
assert_eq!(
got,
payload,
"payload was corrupted by the remote_start race (got {} bytes)",
got.len()
);
})
.await
.expect("test_tcp_forward_small_immediate_write_loses_data timed out");
}
#[tokio::test]
async fn test_socks5_unsupported_command() {
timeout(TEST_TIMEOUT, async {
let server_port = get_available_port();
let socks_port = get_available_port();
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 mut conn = TcpStream::connect(format!("127.0.0.1:{socks_port}"))
.await
.unwrap();
conn.write_all(&[0x05, 0x01, 0x00]).await.unwrap();
let mut greet = [0u8; 2];
conn.read_exact(&mut greet).await.unwrap();
assert_eq!(greet, [0x05, 0x00]);
let mut req = vec![0x05, 0x02, 0x00, 0x01, 127, 0, 0, 1];
req.extend_from_slice(&9999u16.to_be_bytes());
conn.write_all(&req).await.unwrap();
let mut reply = [0u8; 2];
conn.read_exact(&mut reply).await.unwrap();
assert_eq!(
reply,
[0x05, 0x07],
"expected SOCKS5 'command not supported' reply"
);
let target_port = get_available_port();
let target_listener = TcpListener::bind(format!("127.0.0.1:{target_port}"))
.await
.unwrap();
let mut ok_conn = socks5_connect_ipv4(
&format!("127.0.0.1:{socks_port}"),
[127, 0, 0, 1],
target_port,
)
.await;
ok_conn.write_all(b"after-bind").await.unwrap();
ok_conn.shutdown().await.unwrap();
let (mut srv, _) = target_listener.accept().await.unwrap();
let mut buf = vec![0u8; 64];
let n = srv.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"after-bind");
})
.await
.expect("test_socks5_unsupported_command timed out");
}
#[tokio::test]
async fn test_socks5_unsupported_address_type() {
timeout(TEST_TIMEOUT, async {
let server_port = get_available_port();
let socks_port = get_available_port();
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 mut conn = TcpStream::connect(format!("127.0.0.1:{socks_port}"))
.await
.unwrap();
conn.write_all(&[0x05, 0x01, 0x00]).await.unwrap();
let mut greet = [0u8; 2];
conn.read_exact(&mut greet).await.unwrap();
assert_eq!(greet, [0x05, 0x00]);
let mut req = vec![0x05, 0x01, 0x00, 0x05];
req.extend_from_slice(&[0u8; 4]);
req.extend_from_slice(&80u16.to_be_bytes());
conn.write_all(&req).await.unwrap();
let mut reply = [0u8; 2];
conn.read_exact(&mut reply).await.unwrap();
assert_eq!(
reply,
[0x05, 0x08],
"expected SOCKS5 'address type not supported' reply"
);
})
.await
.expect("test_socks5_unsupported_address_type timed out");
}
#[tokio::test]
async fn test_socks5_invalid_version_drops_connection() {
timeout(TEST_TIMEOUT, async {
let server_port = get_available_port();
let socks_port = get_available_port();
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 mut conn = TcpStream::connect(format!("127.0.0.1:{socks_port}"))
.await
.unwrap();
conn.write_all(&[0x04, 0x01, 0x00, 0x50, 127, 0, 0, 1, 0])
.await
.unwrap();
let mut buf = Vec::new();
let res = timeout(Duration::from_secs(3), conn.read_to_end(&mut buf)).await;
match res {
Ok(Ok(_)) => {
assert!(
buf.first().copied() != Some(0x05),
"got unexpected SOCKS5-looking reply: {buf:?}"
);
}
Ok(Err(_)) => { }
Err(_) => panic!("server did not close the bogus SOCKS4 connection in time"),
}
})
.await
.expect("test_socks5_invalid_version_drops_connection timed out");
}
#[tokio::test]
async fn test_tcp_forward_dead_target_closes_connection() {
timeout(TEST_TIMEOUT, async {
let server_port = get_available_port();
let local_port = get_available_port();
let dead_target_port = get_available_port();
let remote = RemoteRequest::from_str(&format!(
"127.0.0.1:{local_port}:127.0.0.1:{dead_target_port}"
))
.unwrap();
let _env = start_tunnel(server_port, false, vec![remote]).await;
let mut conn = TcpStream::connect(format!("127.0.0.1:{local_port}"))
.await
.unwrap();
let _ = conn.write_all(b"x").await;
conn.shutdown().await.ok();
let mut buf = vec![0u8; 16];
let n = timeout(Duration::from_secs(5), conn.read(&mut buf))
.await
.expect("client never observed tunnel teardown for dead target")
.unwrap_or(0);
assert_eq!(n, 0, "expected EOF from tunnel pointing at dead target");
})
.await
.expect("test_tcp_forward_dead_target_closes_connection timed out");
}
#[tokio::test]
async fn test_socks5_many_sequential_connections() {
timeout(TEST_TIMEOUT, async {
let server_port = get_available_port();
let socks_port = get_available_port();
let remote = RemoteRequest::from_str(&format!("127.0.0.1:{socks_port}:socks")).unwrap();
let _env = start_tunnel(server_port, false, vec![remote]).await;
for i in 0..6 {
let target_port = get_available_port();
let listener = TcpListener::bind(format!("127.0.0.1:{target_port}"))
.await
.unwrap();
let mut conn = socks5_connect_ipv4(
&format!("127.0.0.1:{socks_port}"),
[127, 0, 0, 1],
target_port,
)
.await;
let payload = format!("seq-socks-{i}");
conn.write_all(payload.as_bytes()).await.unwrap();
conn.shutdown().await.unwrap();
let (mut srv, _) = listener.accept().await.unwrap();
let mut got = Vec::new();
srv.read_to_end(&mut got).await.unwrap();
assert_eq!(got, payload.as_bytes());
}
})
.await
.expect("test_socks5_many_sequential_connections timed out");
}