Skip to main content

borer_core/
trojan.rs

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
12/// Handles a single inbound Trojan connection.
13pub 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    /// Create a Trojan connection wrapper for the given authenticated user.
24    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    /// Proxy the Trojan stream to its requested upstream address.
33    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}