Skip to main content

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

1use super::super::{ParsingError, SerializeError};
2
3use bytes::{Buf, BufMut};
4use std::net::SocketAddrV4;
5
6/// +-----+-----+----+----+----+----+----+----+-------------+------+------------+------+
7/// |  VN |  CD | DSTPORT |        DSTIP      |    USERID   | NULL |   DOMAIN   | NULL |
8/// +-----+-----+----+----+----+----+----+----+-------------+------+------------+------+
9/// |  1  |  1  |    2    |         4         |   Variable  |  1   |  Variable  |   1  |
10/// +-----+-----+----+----+----+----+----+----+-------------+------+------------+------+
11///                                                                ^^^^^^^^^^^^^^^^^^^^^
12///                                                   optional: only do if IP is 0.0.0.X
13#[derive(Debug, PartialEq)]
14pub struct Request<'a>(pub &'a Address);
15
16/// +-----+-----+----+----+----+----+----+----+
17/// |  VN |  CD | DSTPORT |       DSTIP       |
18/// +-----+-----+----+----+----+----+----+----+
19/// |  1  |  1  |    2    |         4         |
20/// +-----+-----+----+----+----+----+----+----+
21///             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
22///              ignore: only for SOCKSv4 BIND
23#[derive(Debug, PartialEq)]
24pub struct Response(pub Status);
25
26#[derive(Debug, PartialEq)]
27pub enum Address {
28    Socket(SocketAddrV4),
29    Domain(String, u16),
30}
31
32#[derive(Debug, PartialEq)]
33pub enum Status {
34    Success = 90,
35    Failed = 91,
36    IdentFailure = 92,
37    IdentMismatch = 93,
38}
39
40impl Request<'_> {
41    pub fn write_to_buf<B: BufMut>(&self, mut buf: B) -> Result<usize, SerializeError> {
42        match self.0 {
43            Address::Socket(socket) => {
44                if buf.remaining_mut() < 9 {
45                    return Err(SerializeError::WouldOverflow);
46                }
47
48                buf.put_u8(0x04); // Version
49                buf.put_u8(0x01); // CONNECT
50
51                buf.put_u16(socket.port()); // Port
52                buf.put_slice(&socket.ip().octets()); // IP
53
54                buf.put_u8(0x00); // NULL terminating an empty USERID
55
56                Ok(9)
57            }
58
59            Address::Domain(domain, port) => {
60                if buf.remaining_mut() < 9 + domain.len() + 1 {
61                    return Err(SerializeError::WouldOverflow);
62                }
63
64                buf.put_u8(0x04); // Version
65                buf.put_u8(0x01); // CONNECT
66
67                buf.put_u16(*port); // Port
68                buf.put_slice(&[0x00, 0x00, 0x00, 0xFF]); // Invalid IP
69
70                buf.put_u8(0x00); // NULL terminating an empty USERID
71
72                buf.put_slice(domain.as_bytes()); // Domain
73                buf.put_u8(0x00); // NULL
74
75                Ok(9 + domain.len() + 1)
76            }
77        }
78    }
79}
80
81impl TryFrom<&mut &[u8]> for Response {
82    type Error = ParsingError;
83
84    fn try_from(buf: &mut &[u8]) -> Result<Self, Self::Error> {
85        if buf.remaining() < 8 {
86            return Err(ParsingError::Incomplete);
87        }
88
89        if buf.get_u8() != 0x00 {
90            return Err(ParsingError::Other);
91        }
92
93        let status = buf.get_u8().try_into()?;
94        let _addr = {
95            let port = buf.get_u16();
96            let mut ip = [0; 4];
97            buf.copy_to_slice(&mut ip);
98
99            SocketAddrV4::new(ip.into(), port)
100        };
101
102        Ok(Self(status))
103    }
104}
105
106impl TryFrom<u8> for Status {
107    type Error = ParsingError;
108
109    fn try_from(byte: u8) -> Result<Self, Self::Error> {
110        Ok(match byte {
111            90 => Self::Success,
112            91 => Self::Failed,
113            92 => Self::IdentFailure,
114            93 => Self::IdentMismatch,
115            _ => return Err(ParsingError::Other),
116        })
117    }
118}
119
120impl std::fmt::Display for Status {
121    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
122        f.write_str(match self {
123            Self::Success => "success",
124            Self::Failed => "server failed to execute command",
125            Self::IdentFailure => "server ident service failed",
126            Self::IdentMismatch => "server ident service did not recognise client identifier",
127        })
128    }
129}
130
131#[cfg(test)]
132mod test {
133    use super::*;
134    use bytes::BytesMut;
135    use std::net::Ipv4Addr;
136
137    #[test]
138    fn request_serialization_with_socket() {
139        let expected = [
140            0x04, // protocol version
141            0x01, // command: connect
142            0x1F, 0x90, // destination port: 8080
143            127, 0, 0, 1,    // destination address: 127.0.0.1
144            0x00, // null terminating an empty userid
145        ];
146
147        let addr = Address::Socket(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 8080));
148        let mut buf = BytesMut::with_capacity(expected.len());
149        let n = Request(&addr).write_to_buf(&mut buf).unwrap();
150        assert_eq!(n, buf.len());
151        assert_eq!(&buf[..], &expected[..]);
152    }
153
154    #[test]
155    fn request_serialization_with_domain() {
156        let expected = [
157            0x04, // protocol version
158            0x01, // command: connect
159            0x1F, 0x90, // destination port: 8080
160            0x00, 0x00, 0x00, 0xFF, // invalid IP: signals that a domain follows (SOCKS4a)
161            0x00, // null terminating an empty userid
162            b'e', b'x', b'a', b'm', b'p', b'l', b'e', b'.', b'c', b'o', b'm', // domain
163            0x00, // null terminator
164        ];
165
166        let addr = Address::Domain("example.com".into(), 8080);
167        let mut buf = BytesMut::with_capacity(expected.len());
168        let n = Request(&addr).write_to_buf(&mut buf).unwrap();
169        assert_eq!(n, buf.len());
170        assert_eq!(&buf[..], &expected[..]);
171    }
172
173    #[test]
174    fn response_deserialization() {
175        let raw = [
176            0x00, // reply version
177            90,   // status: request granted
178            0x1F, 0x90, // port: 8080 (ignored, only used for BIND)
179            127, 0, 0, 1, // address: 127.0.0.1 (ignored, only used for BIND)
180        ];
181        let mut view = &raw[..];
182
183        let res = Response::try_from(&mut view).unwrap();
184        assert_eq!(res, Response(Status::Success));
185        assert!(view.is_empty());
186    }
187
188    #[test]
189    fn response_incomplete() {
190        let raw = [
191            0x00, // reply version
192            90,   // status: request granted
193            0x00, // truncated mid-port
194        ];
195
196        let err = Response::try_from(&mut &raw[..]).unwrap_err();
197        assert!(matches!(err, ParsingError::Incomplete));
198    }
199
200    #[test]
201    fn request_serialization_would_overflow() {
202        let addr = Address::Socket(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 8080));
203        let mut short = [0u8; 5];
204        let err = Request(&addr).write_to_buf(&mut short[..]).unwrap_err();
205        assert!(matches!(err, SerializeError::WouldOverflow));
206    }
207}