1use std::net::SocketAddr;
2
3use anyhow::Context as _;
4use tokio::io::{AsyncRead, AsyncWrite, copy_bidirectional};
5
6use crate::{
7 dial::{Dial as _, DirectDial},
8 proto::{padding::Padding, trojan},
9 stream::stats::StatsStream,
10};
11
12pub struct TrojanConnection<T> {
14 inner: T,
15 user: String,
16 peer_addr: SocketAddr,
17}
18
19impl<T> TrojanConnection<T>
20where
21 T: AsyncRead + AsyncWrite + Unpin,
22{
23 pub fn new(ts: T, user: String, peer_addr: SocketAddr) -> Self {
25 Self {
26 inner: ts,
27 user,
28 peer_addr,
29 }
30 }
31
32 pub async fn handle(mut self) -> anyhow::Result<()> {
34 let stream = &mut self.inner;
35 let req = trojan::Request::read_from(stream)
36 .await
37 .context("trojan request read failed")?;
38 let addr = req.address.clone();
39
40 let padding = req.is_padding();
41 if padding {
42 Padding::default().write_to(stream).await?;
43 }
44
45 debug!("trojan start connect {addr}");
46
47 let mut stats_stream = StatsStream::new(
48 stream,
49 self.user,
50 req.hash,
51 self.peer_addr.to_string(),
52 req.address.to_string(),
53 padding,
54 );
55 let mut remote_ts = DirectDial::new(std::time::Duration::from_secs(3))
56 .dial(addr.clone())
57 .await
58 .context(format!("failed to dial {addr}"))?;
59
60 info!("[{padding}] Trojan Connect to {addr}. ",);
61 if let Ok((a, b)) = copy_bidirectional(&mut stats_stream, &mut remote_ts).await {
62 debug!(
63 "trojan copy end for {} traffic: {}<=>{} total: {}",
64 req.address,
65 a,
66 b,
67 a + b
68 );
69 }
70 Ok(())
71 }
72}
73
74#[cfg(test)]
75mod tests {
76 use std::{
77 net::{IpAddr, Ipv4Addr, SocketAddr},
78 time::{SystemTime, UNIX_EPOCH},
79 };
80
81 use bytes::BytesMut;
82 use socks5_proto::Address;
83 use tokio::{
84 io::{AsyncReadExt, AsyncWriteExt},
85 net::TcpListener,
86 };
87
88 use super::TrojanConnection;
89 use crate::proto::{
90 padding::Padding,
91 trojan::{Command, Request},
92 };
93
94 fn unique_hash() -> String {
95 format!(
96 "{:056x}",
97 SystemTime::now()
98 .duration_since(UNIX_EPOCH)
99 .unwrap()
100 .as_nanos()
101 )
102 }
103
104 async fn local_listener() -> (TcpListener, SocketAddr) {
105 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
106 let addr = listener.local_addr().unwrap();
107 (listener, addr)
108 }
109
110 #[tokio::test]
111 async fn handle_connect_proxies_payload_in_both_directions() {
112 let (listener, addr) = local_listener().await;
113 let (server_side, mut client_side) = tokio::io::duplex(1024);
114 let peer_addr = SocketAddr::from((Ipv4Addr::LOCALHOST, 40000));
115 let hash = unique_hash();
116
117 let handle_task = tokio::spawn(async move {
118 TrojanConnection::new(server_side, "alice".into(), peer_addr)
119 .handle()
120 .await
121 });
122 let target_task = tokio::spawn(async move {
123 let (mut target, _) = listener.accept().await.unwrap();
124 let mut inbound = [0u8; 4];
125 target.read_exact(&mut inbound).await.unwrap();
126 target.write_all(b"pong").await.unwrap();
127 target.shutdown().await.unwrap();
128 inbound
129 });
130
131 let request = Request::new(hash, Command::Connect, Address::SocketAddress(addr));
132 let mut raw = BytesMut::new();
133 request.write_to_buf(&mut raw);
134 raw.extend_from_slice(b"ping");
135
136 client_side.write_all(&raw).await.unwrap();
137 let mut response = [0u8; 4];
138 client_side.read_exact(&mut response).await.unwrap();
139 client_side.shutdown().await.unwrap();
140
141 assert_eq!(&response, b"pong");
142 assert_eq!(target_task.await.unwrap(), *b"ping");
143 handle_task.await.unwrap().unwrap();
144 }
145
146 #[tokio::test]
147 async fn handle_padding_request_replies_with_padding_then_proxies_payload() {
148 let (listener, addr) = local_listener().await;
149 let (server_side, mut client_side) = tokio::io::duplex(4096);
150 let peer_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 40001);
151 let hash = unique_hash();
152
153 let handle_task = tokio::spawn(async move {
154 TrojanConnection::new(server_side, "alice".into(), peer_addr)
155 .handle()
156 .await
157 });
158 let target_task = tokio::spawn(async move {
159 let (mut target, _) = listener.accept().await.unwrap();
160 let mut inbound = [0u8; 5];
161 target.read_exact(&mut inbound).await.unwrap();
162 target.write_all(b"reply").await.unwrap();
163 target.shutdown().await.unwrap();
164 inbound
165 });
166
167 let request = Request::new(hash, Command::Padding, Address::SocketAddress(addr));
168 let mut raw = BytesMut::new();
169 request.write_to_buf(&mut raw);
170 raw.extend_from_slice(b"hello");
171
172 client_side.write_all(&raw).await.unwrap();
173 let padding = Padding::read_from(&mut client_side).await.unwrap();
174 assert!(padding.serialized_len() >= 258);
175
176 let mut response = [0u8; 5];
177 client_side.read_exact(&mut response).await.unwrap();
178 client_side.shutdown().await.unwrap();
179
180 assert_eq!(&response, b"reply");
181 assert_eq!(target_task.await.unwrap(), *b"hello");
182 handle_task.await.unwrap().unwrap();
183 }
184}