Skip to main content

hyper_util/client/legacy/connect/proxy/socks/v5/
messages.rs

1use super::super::{ParsingError, SerializeError};
2
3use bytes::{Buf, BufMut, BytesMut};
4use std::net::SocketAddr;
5
6///  +----+----------+----------+
7/// |VER | NMETHODS | METHODS  |
8/// +----+----------+----------+
9/// | 1  |    1     | 1 to 255 |
10/// +----+----------+----------+
11#[derive(Debug, PartialEq)]
12pub struct NegotiationReq<'a>(pub &'a AuthMethod);
13
14/// +----+--------+
15/// |VER | METHOD |
16/// +----+--------+
17/// | 1  |   1    |
18/// +----+--------+
19#[derive(Debug, PartialEq)]
20pub struct NegotiationRes(pub AuthMethod);
21
22/// +----+------+----------+------+----------+
23/// |VER | ULEN |  UNAME   | PLEN |  PASSWD  |
24/// +----+------+----------+------+----------+
25/// | 1  |  1   | 1 to 255 |  1   | 1 to 255 |
26/// +----+------+----------+------+----------+
27#[derive(Debug, PartialEq)]
28pub struct AuthenticationReq<'a>(pub &'a str, pub &'a str);
29
30/// +----+--------+
31/// |VER | STATUS |
32/// +----+--------+
33/// | 1  |   1    |
34/// +----+--------+
35#[derive(Debug, PartialEq)]
36pub struct AuthenticationRes(pub bool);
37
38/// +----+-----+-------+------+----------+----------+
39/// |VER | CMD |  RSV  | ATYP | DST.ADDR | DST.PORT |
40/// +----+-----+-------+------+----------+----------+
41/// | 1  |  1  | X'00' |  1   | Variable |    2     |
42/// +----+-----+-------+------+----------+----------+
43#[derive(Debug, PartialEq)]
44pub struct ProxyReq<'a>(pub &'a Address);
45
46/// +----+-----+-------+------+----------+----------+
47/// |VER | REP |  RSV  | ATYP | BND.ADDR | BND.PORT |
48/// +----+-----+-------+------+----------+----------+
49/// | 1  |  1  | X'00' |  1   | Variable |    2     |
50/// +----+-----+-------+------+----------+----------+
51#[derive(Debug, PartialEq)]
52pub struct ProxyRes(pub Status);
53
54#[repr(u8)]
55#[derive(Debug, Copy, Clone, PartialEq)]
56pub enum AuthMethod {
57    NoAuth = 0x00,
58    UserPass = 0x02,
59    NoneAcceptable = 0xFF,
60}
61
62#[derive(Debug, PartialEq)]
63pub enum Address {
64    Socket(SocketAddr),
65    Domain(String, u16),
66}
67
68#[derive(Debug, Copy, Clone, PartialEq)]
69pub enum Status {
70    Success,
71    GeneralServerFailure,
72    ConnectionNotAllowed,
73    NetworkUnreachable,
74    HostUnreachable,
75    ConnectionRefused,
76    TtlExpired,
77    CommandNotSupported,
78    AddressTypeNotSupported,
79}
80
81impl NegotiationReq<'_> {
82    pub fn write_to_buf(&self, buf: &mut BytesMut) -> Result<usize, SerializeError> {
83        if buf.capacity() - buf.len() < 3 {
84            return Err(SerializeError::WouldOverflow);
85        }
86
87        buf.put_u8(0x05); // Version
88        buf.put_u8(0x01); // Number of authentication methods
89        buf.put_u8(*self.0 as u8); // Authentication method
90
91        Ok(3)
92    }
93}
94
95impl TryFrom<&mut &[u8]> for NegotiationRes {
96    type Error = ParsingError;
97
98    fn try_from(buf: &mut &[u8]) -> Result<Self, ParsingError> {
99        if buf.remaining() < 2 {
100            return Err(ParsingError::Incomplete);
101        }
102
103        if buf.get_u8() != 0x05 {
104            return Err(ParsingError::Other);
105        }
106
107        let method = buf.get_u8().try_into()?;
108        Ok(Self(method))
109    }
110}
111
112impl AuthenticationReq<'_> {
113    pub fn write_to_buf(&self, buf: &mut BytesMut) -> Result<usize, SerializeError> {
114        if buf.capacity() - buf.len() < 3 + self.0.len() + self.1.len() {
115            return Err(SerializeError::WouldOverflow);
116        }
117
118        buf.put_u8(0x01); // Version
119
120        buf.put_u8(self.0.len() as u8); // Username length (guaranteed to be 255 or less)
121        buf.put_slice(self.0.as_bytes()); // Username
122
123        buf.put_u8(self.1.len() as u8); // Password length (guaranteed to be 255 or less)
124        buf.put_slice(self.1.as_bytes()); // Password
125
126        Ok(3 + self.0.len() + self.1.len())
127    }
128}
129
130impl TryFrom<&mut &[u8]> for AuthenticationRes {
131    type Error = ParsingError;
132
133    fn try_from(buf: &mut &[u8]) -> Result<Self, ParsingError> {
134        if buf.remaining() < 2 {
135            return Err(ParsingError::Incomplete);
136        }
137
138        if buf.get_u8() != 0x01 {
139            return Err(ParsingError::Other);
140        }
141
142        if buf.get_u8() == 0 {
143            Ok(Self(true))
144        } else {
145            Ok(Self(false))
146        }
147    }
148}
149
150impl ProxyReq<'_> {
151    pub fn write_to_buf(&self, buf: &mut BytesMut) -> Result<usize, SerializeError> {
152        let addr_len = match self.0 {
153            Address::Socket(SocketAddr::V4(_)) => 1 + 4 + 2,
154            Address::Socket(SocketAddr::V6(_)) => 1 + 16 + 2,
155            Address::Domain(domain, _) => 1 + 1 + domain.len() + 2,
156        };
157
158        if buf.capacity() - buf.len() < 3 + addr_len {
159            return Err(SerializeError::WouldOverflow);
160        }
161
162        buf.put_u8(0x05); // Version
163        buf.put_u8(0x01); // TCP tunneling command
164        buf.put_u8(0x00); // Reserved
165        let _ = self.0.write_to_buf(buf); // Address
166
167        Ok(3 + addr_len)
168    }
169}
170
171impl TryFrom<&mut &[u8]> for ProxyRes {
172    type Error = ParsingError;
173
174    fn try_from(buf: &mut &[u8]) -> Result<Self, ParsingError> {
175        if buf.remaining() < 3 {
176            return Err(ParsingError::Incomplete);
177        }
178
179        // VER
180        if buf.get_u8() != 0x05 {
181            return Err(ParsingError::Other);
182        }
183
184        // REP
185        let status = buf.get_u8().try_into()?;
186
187        // RSV
188        if buf.get_u8() != 0x00 {
189            return Err(ParsingError::Other);
190        }
191
192        // ATYP + ADDR
193        Address::try_from(buf)?;
194
195        Ok(Self(status))
196    }
197}
198
199impl Address {
200    pub fn write_to_buf(&self, buf: &mut BytesMut) -> Result<usize, SerializeError> {
201        match self {
202            Self::Socket(SocketAddr::V4(v4)) => {
203                if buf.capacity() - buf.len() < 1 + 4 + 2 {
204                    return Err(SerializeError::WouldOverflow);
205                }
206
207                buf.put_u8(0x01);
208                buf.put_slice(&v4.ip().octets());
209                buf.put_u16(v4.port()); // Network Order/BigEndian for port
210
211                Ok(7)
212            }
213
214            Self::Socket(SocketAddr::V6(v6)) => {
215                if buf.capacity() - buf.len() < 1 + 16 + 2 {
216                    return Err(SerializeError::WouldOverflow);
217                }
218
219                buf.put_u8(0x04);
220                buf.put_slice(&v6.ip().octets());
221                buf.put_u16(v6.port()); // Network Order/BigEndian for port
222
223                Ok(19)
224            }
225
226            Self::Domain(domain, port) => {
227                if buf.capacity() - buf.len() < 1 + 1 + domain.len() + 2 {
228                    return Err(SerializeError::WouldOverflow);
229                }
230
231                buf.put_u8(0x03);
232                buf.put_u8(domain.len() as u8); // Guarenteed to be less than 255
233                buf.put_slice(domain.as_bytes());
234                buf.put_u16(*port);
235
236                Ok(4 + domain.len())
237            }
238        }
239    }
240}
241
242impl TryFrom<&mut &[u8]> for Address {
243    type Error = ParsingError;
244
245    fn try_from(buf: &mut &[u8]) -> Result<Self, Self::Error> {
246        if buf.remaining() < 2 {
247            return Err(ParsingError::Incomplete);
248        }
249
250        Ok(match buf.get_u8() {
251            // IPv4
252            0x01 => {
253                let mut ip = [0; 4];
254
255                if buf.remaining() < 6 {
256                    return Err(ParsingError::Incomplete);
257                }
258
259                buf.copy_to_slice(&mut ip);
260                let port = buf.get_u16();
261
262                Self::Socket(SocketAddr::new(ip.into(), port))
263            }
264            // Domain
265            0x03 => {
266                let len = buf.get_u8() as usize;
267
268                if len == 0 {
269                    return Err(ParsingError::Other);
270                } else if buf.remaining() < len + 2 {
271                    return Err(ParsingError::Incomplete);
272                }
273
274                let domain = std::str::from_utf8(&buf.chunk()[..len])
275                    .map_err(|_| ParsingError::Other)?
276                    .to_string();
277                buf.advance(len);
278
279                let port = buf.get_u16();
280
281                Self::Domain(domain, port)
282            }
283            // IPv6
284            0x04 => {
285                let mut ip = [0; 16];
286
287                if buf.remaining() < 18 {
288                    return Err(ParsingError::Incomplete);
289                }
290                buf.copy_to_slice(&mut ip);
291                let port = buf.get_u16();
292
293                Self::Socket(SocketAddr::new(ip.into(), port))
294            }
295
296            _ => return Err(ParsingError::Other),
297        })
298    }
299}
300
301impl TryFrom<u8> for Status {
302    type Error = ParsingError;
303
304    fn try_from(byte: u8) -> Result<Self, Self::Error> {
305        Ok(match byte {
306            0x00 => Self::Success,
307
308            0x01 => Self::GeneralServerFailure,
309            0x02 => Self::ConnectionNotAllowed,
310            0x03 => Self::NetworkUnreachable,
311            0x04 => Self::HostUnreachable,
312            0x05 => Self::ConnectionRefused,
313            0x06 => Self::TtlExpired,
314            0x07 => Self::CommandNotSupported,
315            0x08 => Self::AddressTypeNotSupported,
316            _ => return Err(ParsingError::Other),
317        })
318    }
319}
320
321impl TryFrom<u8> for AuthMethod {
322    type Error = ParsingError;
323
324    fn try_from(byte: u8) -> Result<Self, Self::Error> {
325        Ok(match byte {
326            0x00 => Self::NoAuth,
327            0x02 => Self::UserPass,
328            0xFF => Self::NoneAcceptable,
329
330            _ => return Err(ParsingError::Other),
331        })
332    }
333}
334
335impl std::fmt::Display for Status {
336    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
337        f.write_str(match self {
338            Self::Success => "success",
339            Self::GeneralServerFailure => "general server failure",
340            Self::ConnectionNotAllowed => "connection not allowed",
341            Self::NetworkUnreachable => "network unreachable",
342            Self::HostUnreachable => "host unreachable",
343            Self::ConnectionRefused => "connection refused",
344            Self::TtlExpired => "ttl expired",
345            Self::CommandNotSupported => "command not supported",
346            Self::AddressTypeNotSupported => "address type not supported",
347        })
348    }
349}
350
351#[cfg(test)]
352mod test {
353    use super::*;
354    use std::net::{Ipv4Addr, Ipv6Addr};
355
356    #[test]
357    fn negotiation_req_serialization() {
358        let expected = [
359            0x05, // protocol version
360            0x01, // number of authentication methods: 1
361            0x00, // method: no authentication
362        ];
363
364        let mut buf = BytesMut::with_capacity(expected.len());
365        let n = NegotiationReq(&AuthMethod::NoAuth)
366            .write_to_buf(&mut buf)
367            .unwrap();
368        assert_eq!(n, buf.len());
369        assert_eq!(&buf[..], &expected[..]);
370    }
371
372    #[test]
373    fn negotiation_res_deserialization() {
374        let raw = [
375            0x05, // protocol version
376            0x02, // selected method: username/password
377        ];
378
379        let mut view = &raw[..];
380
381        let res = NegotiationRes::try_from(&mut view).unwrap();
382        assert_eq!(res, NegotiationRes(AuthMethod::UserPass));
383        assert!(view.is_empty());
384    }
385
386    #[test]
387    fn authentication_req_serialization() {
388        let expected = [
389            0x01, // authentication version
390            0x04, // username length: 4
391            b'u', b's', b'e', b'r', // username
392            0x04, // password length: 4
393            b'p', b'a', b's', b's', // password
394        ];
395
396        let mut buf = BytesMut::with_capacity(expected.len());
397        let n = AuthenticationReq("user", "pass")
398            .write_to_buf(&mut buf)
399            .unwrap();
400        assert_eq!(n, buf.len());
401        assert_eq!(&buf[..], &expected[..]);
402    }
403
404    #[test]
405    fn authentication_res_deserialization() {
406        let raw = [
407            0x01, // authentication version
408            0x00, // status: success
409        ];
410        assert_eq!(
411            AuthenticationRes::try_from(&mut &raw[..]).unwrap(),
412            AuthenticationRes(true)
413        );
414
415        let raw = [
416            0x01, // authentication version
417            0x01, // status: failure
418        ];
419        assert_eq!(
420            AuthenticationRes::try_from(&mut &raw[..]).unwrap(),
421            AuthenticationRes(false)
422        );
423    }
424
425    #[test]
426    fn proxy_req_serialization() {
427        let expected = [
428            0x05, // protocol version
429            0x01, // command: connect
430            0x00, // reserved
431            0x01, // address type: IPv4
432            127, 0, 0, 1, // destination address: 127.0.0.1
433            0x1F, 0x90, // destination port: 8080
434        ];
435
436        let addr = Address::Socket(SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 8080));
437
438        let mut buf = BytesMut::with_capacity(expected.len());
439        let n = ProxyReq(&addr).write_to_buf(&mut buf).unwrap();
440        assert_eq!(n, buf.len());
441        assert_eq!(&buf[..], &expected[..]);
442    }
443
444    #[test]
445    fn proxy_res_deserialization() {
446        let raw = [
447            0x05, // protocol version
448            0x00, // reply: success
449            0x00, // reserved
450            0x01, // address type: IPv4
451            127, 0, 0, 1, // bound address: 127.0.0.1
452            0x1F, 0x90, // bound port: 8080
453        ];
454        let mut view = &raw[..];
455
456        let res = ProxyRes::try_from(&mut view).unwrap();
457        assert_eq!(res, ProxyRes(Status::Success));
458        assert!(view.is_empty());
459    }
460
461    #[test]
462    fn serialization_would_overflow() {
463        let mut buf = BytesMut::with_capacity(2);
464
465        let err = NegotiationReq(&AuthMethod::NoAuth)
466            .write_to_buf(&mut buf)
467            .unwrap_err();
468        assert!(matches!(err, SerializeError::WouldOverflow));
469    }
470
471    fn assert_address_roundtrips(addr: Address) {
472        let mut buf = BytesMut::with_capacity(64);
473
474        let n = addr.write_to_buf(&mut buf).unwrap();
475        assert_eq!(n, buf.len());
476
477        let mut view = &buf[..];
478        assert_eq!(Address::try_from(&mut view).unwrap(), addr);
479        assert!(view.is_empty(), "address bytes should be fully consumed");
480    }
481
482    #[test]
483    fn address_roundtrip_ipv4() {
484        assert_address_roundtrips(Address::Socket(SocketAddr::new(
485            Ipv4Addr::LOCALHOST.into(),
486            8080,
487        )));
488    }
489
490    #[test]
491    fn address_roundtrip_ipv6() {
492        assert_address_roundtrips(Address::Socket(SocketAddr::new(
493            Ipv6Addr::LOCALHOST.into(),
494            8080,
495        )));
496    }
497
498    #[test]
499    fn address_roundtrip_domain() {
500        assert_address_roundtrips(Address::Domain("example.com".into(), 8080));
501    }
502}