1pub mod detector;
6pub mod error;
7pub mod socks4;
8pub mod socks5;
9
10pub use detector::{Socks4Detector, SOCKS4_PROTOCOL_ID};
11pub use error::Socks5Error;
12pub use socks4::server::{read_socks4_request, write_socks4_reply};
13pub use socks4::{socks4_connect, Socks4Error, Socks4Request, Socks4Status};
14
15use eggress_core::detect::{DetectResult, ProtocolDetector};
16use eggress_core::ProtocolId;
17
18pub struct Socks5Detector;
22
23impl ProtocolDetector for Socks5Detector {
24 fn id(&self) -> ProtocolId {
25 ProtocolId::Socks5
26 }
27
28 fn detect(&self, prefix: &[u8]) -> DetectResult {
29 if prefix.is_empty() {
30 DetectResult::NeedMore { minimum: 1 }
31 } else if prefix[0] == 0x05 {
32 DetectResult::Match { confidence: 100 }
33 } else {
34 DetectResult::NoMatch
35 }
36 }
37}
38
39#[cfg(test)]
40mod tests {
41 use super::*;
42 use eggress_core::detect::{DetectResult, ProtocolDetector};
43 use eggress_core::{BoxStream, TargetAddr, TargetHost};
44
45 #[tokio::test]
46 async fn test_detector_identifies_socks4() {
47 let detector = Socks4Detector;
48 assert_eq!(detector.id(), ProtocolId::Socks4);
49 assert_eq!(
50 detector.detect(b"\x04"),
51 DetectResult::Match { confidence: 100 }
52 );
53 }
54
55 #[test]
56 fn test_socks5_detector_match() {
57 let detector = Socks5Detector;
58 assert_eq!(detector.id(), ProtocolId::Socks5);
59 assert_eq!(
60 detector.detect(&[0x05]),
61 DetectResult::Match { confidence: 100 }
62 );
63 }
64
65 #[test]
66 fn test_socks5_detector_no_match() {
67 let detector = Socks5Detector;
68 assert_eq!(detector.detect(&[0x04]), DetectResult::NoMatch);
69 assert_eq!(detector.detect(&[0x00]), DetectResult::NoMatch);
70 }
71
72 #[test]
73 fn test_socks5_detector_need_more() {
74 let detector = Socks5Detector;
75 assert_eq!(detector.detect(&[]), DetectResult::NeedMore { minimum: 1 });
76 }
77
78 #[test]
79 fn test_socks5_detector_with_more_data() {
80 let detector = Socks5Detector;
81 assert_eq!(
82 detector.detect(&[0x05, 0x01, 0x00]),
83 DetectResult::Match { confidence: 100 }
84 );
85 }
86
87 #[tokio::test]
88 async fn test_full_socks4_roundtrip() {
89 let (addr, jh) = eggress_testkit::start_echo_server().await;
90 let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
91 let server_addr = server_listener.local_addr().unwrap();
92
93 let server_jh = tokio::spawn(async move {
94 let (mut client_stream, _) = server_listener.accept().await.unwrap();
95 let request = read_socks4_request(&mut client_stream).await.unwrap();
96 assert_eq!(request.command, 0x01);
97 assert_eq!(request.addr, addr);
98
99 let target_stream = tokio::net::TcpStream::connect(request.addr).await.unwrap();
100 let _ = write_socks4_reply(
101 &mut client_stream,
102 Socks4Status::Granted,
103 "127.0.0.1:0".parse().unwrap(),
104 )
105 .await;
106
107 let (mut cr, mut cw) = tokio::io::split(client_stream);
108 let (mut tr, mut tw) = tokio::io::split(target_stream);
109 tokio::spawn(async move {
110 let _ = tokio::io::copy(&mut cr, &mut tw).await;
111 });
112 tokio::spawn(async move {
113 let _ = tokio::io::copy(&mut tr, &mut cw).await;
114 });
115 });
116
117 let stream = tokio::net::TcpStream::connect(server_addr).await.unwrap();
118 let boxed: BoxStream = Box::new(stream);
119 let target = TargetAddr {
120 host: TargetHost::Ip(addr.ip()),
121 port: addr.port(),
122 };
123 let mut conn = socks4_connect(boxed, &target, None).await.unwrap();
124
125 use tokio::io::{AsyncReadExt, AsyncWriteExt};
126 conn.write_all(b"hello socks4").await.unwrap();
127 conn.shutdown().await.unwrap();
128
129 let mut buf = [0u8; 12];
130 conn.read_exact(&mut buf).await.unwrap();
131 assert_eq!(&buf, b"hello socks4");
132
133 server_jh.abort();
134 jh.abort();
135 }
136
137 #[tokio::test]
138 async fn test_socks4_with_user_id() {
139 let (addr, jh) = eggress_testkit::start_echo_server().await;
140 let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
141 let server_addr = server_listener.local_addr().unwrap();
142
143 let server_jh = tokio::spawn(async move {
144 let (mut client_stream, _) = server_listener.accept().await.unwrap();
145 let request = read_socks4_request(&mut client_stream).await.unwrap();
146 assert_eq!(request.user_id, "testuser");
147 assert_eq!(request.addr, addr);
148
149 let target_stream = tokio::net::TcpStream::connect(request.addr).await.unwrap();
150 let _ = write_socks4_reply(
151 &mut client_stream,
152 Socks4Status::Granted,
153 "127.0.0.1:0".parse().unwrap(),
154 )
155 .await;
156
157 let (mut cr, mut cw) = tokio::io::split(client_stream);
158 let (mut tr, mut tw) = tokio::io::split(target_stream);
159 tokio::spawn(async move {
160 let _ = tokio::io::copy(&mut cr, &mut tw).await;
161 });
162 tokio::spawn(async move {
163 let _ = tokio::io::copy(&mut tr, &mut cw).await;
164 });
165 });
166
167 let stream = tokio::net::TcpStream::connect(server_addr).await.unwrap();
168 let boxed: BoxStream = Box::new(stream);
169 let target = TargetAddr {
170 host: TargetHost::Ip(addr.ip()),
171 port: addr.port(),
172 };
173 let mut conn = socks4_connect(boxed, &target, Some("testuser"))
174 .await
175 .unwrap();
176
177 use tokio::io::{AsyncReadExt, AsyncWriteExt};
178 conn.write_all(b"hello user").await.unwrap();
179 conn.shutdown().await.unwrap();
180
181 let mut buf = [0u8; 10];
182 conn.read_exact(&mut buf).await.unwrap();
183 assert_eq!(&buf, b"hello user");
184
185 server_jh.abort();
186 jh.abort();
187 }
188
189 #[tokio::test]
190 async fn test_socks4a_domain_target() {
191 let (addr, jh) = eggress_testkit::start_echo_server().await;
192 let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
193 let server_addr = server_listener.local_addr().unwrap();
194
195 let server_jh = tokio::spawn(async move {
196 let (mut client_stream, _) = server_listener.accept().await.unwrap();
197 let request = read_socks4_request(&mut client_stream).await.unwrap();
198 assert_eq!(request.domain.as_deref(), Some("example.com"));
199
200 let target_stream = tokio::net::TcpStream::connect(addr).await.unwrap();
201 let _ = write_socks4_reply(
202 &mut client_stream,
203 Socks4Status::Granted,
204 "127.0.0.1:0".parse().unwrap(),
205 )
206 .await;
207
208 let (mut cr, mut cw) = tokio::io::split(client_stream);
209 let (mut tr, mut tw) = tokio::io::split(target_stream);
210 tokio::spawn(async move {
211 let _ = tokio::io::copy(&mut cr, &mut tw).await;
212 });
213 tokio::spawn(async move {
214 let _ = tokio::io::copy(&mut tr, &mut cw).await;
215 });
216 });
217
218 let stream = tokio::net::TcpStream::connect(server_addr).await.unwrap();
219 let boxed: BoxStream = Box::new(stream);
220 let target = TargetAddr {
221 host: TargetHost::Domain("example.com".to_string()),
222 port: 80,
223 };
224 let mut conn = socks4_connect(boxed, &target, None).await.unwrap();
225
226 use tokio::io::{AsyncReadExt, AsyncWriteExt};
227 conn.write_all(b"hello domain").await.unwrap();
228 conn.shutdown().await.unwrap();
229
230 let mut buf = [0u8; 12];
231 conn.read_exact(&mut buf).await.unwrap();
232 assert_eq!(&buf, b"hello domain");
233
234 server_jh.abort();
235 jh.abort();
236 }
237
238 #[tokio::test]
239 async fn test_invalid_version_rejection() {
240 let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
241 let server_addr = server_listener.local_addr().unwrap();
242
243 let server_jh = tokio::spawn(async move {
244 let (mut client_stream, _) = server_listener.accept().await.unwrap();
245 let result = read_socks4_request(&mut client_stream).await;
246 assert!(result.is_err());
247 match result.unwrap_err() {
248 Socks4Error::InvalidVersion(v) => assert_eq!(v, 0x05),
249 other => panic!("expected InvalidVersion, got {:?}", other),
250 }
251 });
252
253 let mut stream = tokio::net::TcpStream::connect(server_addr).await.unwrap();
254 use tokio::io::AsyncWriteExt;
255 stream
256 .write_all(&[0x05, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00])
257 .await
258 .unwrap();
259
260 server_jh.await.unwrap();
261 }
262
263 #[tokio::test]
264 async fn test_command_rejection_bind() {
265 let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
266 let server_addr = server_listener.local_addr().unwrap();
267
268 let server_jh = tokio::spawn(async move {
269 let (mut client_stream, _) = server_listener.accept().await.unwrap();
270 let result = read_socks4_request(&mut client_stream).await;
271 assert!(result.is_err());
272 match result.unwrap_err() {
273 Socks4Error::UnsupportedCommand(cmd) => assert_eq!(cmd, 0x02),
274 other => panic!("expected UnsupportedCommand, got {:?}", other),
275 }
276 });
277
278 let mut stream = tokio::net::TcpStream::connect(server_addr).await.unwrap();
279 use tokio::io::AsyncWriteExt;
280 stream
281 .write_all(&[0x04, 0x02, 0x00, 0x50, 127, 0, 0, 1, b'x', 0x00])
282 .await
283 .unwrap();
284
285 server_jh.await.unwrap();
286 }
287
288 #[tokio::test]
289 async fn test_user_id_too_long() {
290 let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
291 let server_addr = server_listener.local_addr().unwrap();
292
293 let server_jh = tokio::spawn(async move {
294 let (mut client_stream, _) = server_listener.accept().await.unwrap();
295 let result = read_socks4_request(&mut client_stream).await;
296 assert!(result.is_err());
297 assert!(matches!(result.unwrap_err(), Socks4Error::UserIdTooLong));
298 });
299
300 let mut stream = tokio::net::TcpStream::connect(server_addr).await.unwrap();
301 use tokio::io::AsyncWriteExt;
302 let mut payload = vec![0x04, 0x01, 0x00, 0x50, 127, 0, 0, 1];
303 payload.extend(std::iter::repeat_n(b'A', 256));
304 payload.push(0x00);
305 stream.write_all(&payload).await.unwrap();
306
307 server_jh.await.unwrap();
308 }
309
310 #[tokio::test]
311 async fn test_socks4_rejects_non_matching_protocol() {
312 let detector = Socks4Detector;
313 assert_eq!(detector.detect(b"\x05"), DetectResult::NoMatch);
314 assert_eq!(detector.detect(b"\x16"), DetectResult::NoMatch);
315 assert_eq!(detector.detect(b"G"), DetectResult::NoMatch);
316 }
317
318 #[tokio::test]
319 async fn test_user_id_too_long_client() {
320 let long_id = "A".repeat(256);
321 let result = client_too_long_uid_guard(&long_id).await;
322 assert!(matches!(result, Err(Socks4Error::UserIdTooLong)));
323 }
324
325 async fn client_too_long_uid_guard(uid: &str) -> Result<(), Socks4Error> {
326 if uid.len() > 255 {
327 return Err(Socks4Error::UserIdTooLong);
328 }
329 Ok(())
330 }
331
332 #[tokio::test]
333 async fn test_socks4_fragmented_read() {
334 let (addr, jh) = eggress_testkit::start_echo_server().await;
335 let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
336 let server_addr = server_listener.local_addr().unwrap();
337
338 let server_jh = tokio::spawn(async move {
339 let (mut client_stream, _) = server_listener.accept().await.unwrap();
340 let request = read_socks4_request(&mut client_stream).await.unwrap();
341 assert_eq!(request.command, 0x01);
342
343 let target_stream = tokio::net::TcpStream::connect(request.addr).await.unwrap();
344 let _ = write_socks4_reply(
345 &mut client_stream,
346 Socks4Status::Granted,
347 "127.0.0.1:0".parse().unwrap(),
348 )
349 .await;
350
351 let (mut cr, mut cw) = tokio::io::split(client_stream);
352 let (mut tr, mut tw) = tokio::io::split(target_stream);
353 tokio::spawn(async move {
354 let _ = tokio::io::copy(&mut cr, &mut tw).await;
355 });
356 tokio::spawn(async move {
357 let _ = tokio::io::copy(&mut tr, &mut cw).await;
358 });
359 });
360
361 let mut stream = tokio::net::TcpStream::connect(server_addr).await.unwrap();
363 use tokio::io::AsyncWriteExt;
364 stream.write_all(&[0x04]).await.unwrap();
365 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
366 stream.write_all(&[0x01]).await.unwrap();
367 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
368 stream.write_all(&addr.port().to_be_bytes()).await.unwrap();
369 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
370 let ip = match addr.ip() {
371 std::net::IpAddr::V4(v4) => v4.octets(),
372 _ => panic!("SOCKS4 test requires IPv4 address, got: {addr}"),
373 };
374 stream.write_all(&ip).await.unwrap();
375 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
376 stream.write_all(&[0x00]).await.unwrap();
377
378 let mut reply = [0u8; 8];
380 use tokio::io::AsyncReadExt;
381 stream.read_exact(&mut reply).await.unwrap();
382 assert_eq!(reply[0], 0x00);
383 assert_eq!(reply[1], 90); stream.write_all(b"frag").await.unwrap();
387 stream.shutdown().await.unwrap();
388 let mut buf = [0u8; 4];
389 stream.read_exact(&mut buf).await.unwrap();
390 assert_eq!(&buf, b"frag");
391
392 server_jh.abort();
393 jh.abort();
394 }
395}