Skip to main content

eggress_protocol_socks/
lib.rs

1//! SOCKS4/5 proxy protocol implementation.
2//!
3//! This crate provides the SOCKS4 and SOCKS5 proxy protocol handlers.
4
5pub 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
18/// SOCKS5 protocol detector.
19///
20/// Checks if the first byte is 0x05 (SOCKS5 version).
21pub 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        // Send SOCKS4 request in fragments.
362        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        // Read reply.
379        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); // granted
384
385        // Send data and verify echo.
386        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}