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::timeout;
#[tokio::test]
async fn test_tcp_forward_target_initiated_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 app = TcpStream::connect(format!("127.0.0.1:{local_port}"))
.await
.unwrap();
let (mut target, _) = target_listener.accept().await.unwrap();
target.write_all(b"server-hello").await.unwrap();
target.shutdown().await.unwrap();
let mut got = Vec::new();
app.read_to_end(&mut got).await.unwrap();
assert_eq!(got, b"server-hello", "app must see target payload + EOF");
app.write_all(b"client-trailer").await.unwrap();
app.shutdown().await.unwrap();
let mut buf = Vec::new();
target.read_to_end(&mut buf).await.unwrap();
assert_eq!(buf, b"client-trailer", "target must drain app's trailer");
})
.await
.expect("test_tcp_forward_target_initiated_half_close timed out");
}
#[tokio::test]
async fn test_tcp_reverse_app_to_target_half_close() {
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 mut app = TcpStream::connect(format!("127.0.0.1:{listen_port}"))
.await
.unwrap();
let (mut target, _) = target_listener.accept().await.unwrap();
app.write_all(b"req").await.unwrap();
app.shutdown().await.unwrap();
let mut buf = vec![0u8; 64];
let n = target.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"req");
let n = target.read(&mut buf).await.unwrap();
assert_eq!(n, 0, "target must observe EOF after app shutdown");
target.write_all(b"resp").await.unwrap();
target.shutdown().await.unwrap();
let mut got = Vec::new();
app.read_to_end(&mut got).await.unwrap();
assert_eq!(got, b"resp");
})
.await
.expect("test_tcp_reverse_app_to_target_half_close timed out");
}
#[tokio::test]
async fn test_tcp_reverse_target_initiated_half_close() {
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 mut app = TcpStream::connect(format!("127.0.0.1:{listen_port}"))
.await
.unwrap();
let (mut target, _) = target_listener.accept().await.unwrap();
target.write_all(b"server-hello").await.unwrap();
target.shutdown().await.unwrap();
let mut got = Vec::new();
app.read_to_end(&mut got).await.unwrap();
assert_eq!(got, b"server-hello");
app.write_all(b"client-trailer").await.unwrap();
app.shutdown().await.unwrap();
let mut buf = Vec::new();
target.read_to_end(&mut buf).await.unwrap();
assert_eq!(buf, b"client-trailer");
})
.await
.expect("test_tcp_reverse_target_initiated_half_close timed out");
}
#[tokio::test]
async fn test_socks5_forward_half_close() {
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_with_flags(server_port, false, true, vec![remote]).await;
let mut app = socks5_connect_ipv4(
&format!("127.0.0.1:{socks_port}"),
[127, 0, 0, 1],
target_port,
)
.await;
let (mut target, _) = target_listener.accept().await.unwrap();
app.write_all(b"req").await.unwrap();
app.shutdown().await.unwrap();
let mut buf = vec![0u8; 64];
let n = target.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"req");
let n = timeout(Duration::from_secs(3), target.read(&mut buf))
.await
.expect("SOCKS path never propagated app's half-close as EOF")
.unwrap();
assert_eq!(n, 0, "expected EOF after app shutdown via SOCKS");
target.write_all(b"resp").await.unwrap();
target.shutdown().await.unwrap();
let mut got = Vec::new();
app.read_to_end(&mut got).await.unwrap();
assert_eq!(got, b"resp");
})
.await
.expect("test_socks5_forward_half_close timed out");
}