use std::net::SocketAddr;
use anyhow::Context as _;
use tokio::io::{AsyncRead, AsyncWrite, copy_bidirectional};
use crate::{
dial::{Dial as _, DirectDial},
proto::{padding::Padding, trojan},
stream::stats::StatsStream,
};
pub struct TrojanConnection<T> {
inner: T,
user: String,
peer_addr: SocketAddr,
}
impl<T> TrojanConnection<T>
where
T: AsyncRead + AsyncWrite + Unpin,
{
pub fn new(ts: T, user: String, peer_addr: SocketAddr) -> Self {
Self {
inner: ts,
user,
peer_addr,
}
}
pub async fn handle(mut self) -> anyhow::Result<()> {
let stream = &mut self.inner;
let req = trojan::Request::read_from(stream)
.await
.context("trojan request read failed")?;
let addr = req.address.clone();
let padding = req.is_padding();
if padding {
Padding::default().write_to(stream).await?;
}
debug!("trojan start connect {addr}");
let mut stats_stream = StatsStream::new(
stream,
self.user,
req.hash,
self.peer_addr.to_string(),
req.address.to_string(),
padding,
);
let mut remote_ts = DirectDial::new(std::time::Duration::from_secs(3))
.dial(addr.clone())
.await
.context(format!("failed to dial {addr}"))?;
info!("[{padding}] Trojan Connect to {addr}. ",);
if let Ok((a, b)) = copy_bidirectional(&mut stats_stream, &mut remote_ts).await {
debug!(
"trojan copy end for {} traffic: {}<=>{} total: {}",
req.address,
a,
b,
a + b
);
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use std::{
net::{IpAddr, Ipv4Addr, SocketAddr},
time::{SystemTime, UNIX_EPOCH},
};
use bytes::BytesMut;
use socks5_proto::Address;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpListener,
};
use super::TrojanConnection;
use crate::proto::{
padding::Padding,
trojan::{Command, Request},
};
fn unique_hash() -> String {
format!(
"{:056x}",
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
)
}
async fn local_listener() -> (TcpListener, SocketAddr) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
(listener, addr)
}
#[tokio::test]
async fn handle_connect_proxies_payload_in_both_directions() {
let (listener, addr) = local_listener().await;
let (server_side, mut client_side) = tokio::io::duplex(1024);
let peer_addr = SocketAddr::from((Ipv4Addr::LOCALHOST, 40000));
let hash = unique_hash();
let handle_task = tokio::spawn(async move {
TrojanConnection::new(server_side, "alice".into(), peer_addr)
.handle()
.await
});
let target_task = tokio::spawn(async move {
let (mut target, _) = listener.accept().await.unwrap();
let mut inbound = [0u8; 4];
target.read_exact(&mut inbound).await.unwrap();
target.write_all(b"pong").await.unwrap();
target.shutdown().await.unwrap();
inbound
});
let request = Request::new(hash, Command::Connect, Address::SocketAddress(addr));
let mut raw = BytesMut::new();
request.write_to_buf(&mut raw);
raw.extend_from_slice(b"ping");
client_side.write_all(&raw).await.unwrap();
let mut response = [0u8; 4];
client_side.read_exact(&mut response).await.unwrap();
client_side.shutdown().await.unwrap();
assert_eq!(&response, b"pong");
assert_eq!(target_task.await.unwrap(), *b"ping");
handle_task.await.unwrap().unwrap();
}
#[tokio::test]
async fn handle_padding_request_replies_with_padding_then_proxies_payload() {
let (listener, addr) = local_listener().await;
let (server_side, mut client_side) = tokio::io::duplex(4096);
let peer_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 40001);
let hash = unique_hash();
let handle_task = tokio::spawn(async move {
TrojanConnection::new(server_side, "alice".into(), peer_addr)
.handle()
.await
});
let target_task = tokio::spawn(async move {
let (mut target, _) = listener.accept().await.unwrap();
let mut inbound = [0u8; 5];
target.read_exact(&mut inbound).await.unwrap();
target.write_all(b"reply").await.unwrap();
target.shutdown().await.unwrap();
inbound
});
let request = Request::new(hash, Command::Padding, Address::SocketAddress(addr));
let mut raw = BytesMut::new();
request.write_to_buf(&mut raw);
raw.extend_from_slice(b"hello");
client_side.write_all(&raw).await.unwrap();
let padding = Padding::read_from(&mut client_side).await.unwrap();
assert!(padding.serialized_len() >= 258);
let mut response = [0u8; 5];
client_side.read_exact(&mut response).await.unwrap();
client_side.shutdown().await.unwrap();
assert_eq!(&response, b"reply");
assert_eq!(target_task.await.unwrap(), *b"hello");
handle_task.await.unwrap().unwrap();
}
}