1use crate::socks5::server::{SocksAddr, ATYP_DOMAIN, ATYP_IPV4, ATYP_IPV6};
2
3pub const MAX_UDP_DATAGRAM_SIZE: usize = 65535;
4
5pub struct Socks5UdpRequest<'a> {
6 pub target: SocksAddr,
7 pub payload: &'a [u8],
8}
9
10#[derive(Debug, thiserror::Error)]
11pub enum UdpCodecError {
12 #[error("packet too short")]
13 TooShort,
14 #[error("non-zero reserved field")]
15 BadReserved,
16 #[error("fragmentation not supported")]
17 FragmentationUnsupported,
18 #[error("unknown address type: {0}")]
19 UnknownAddressType(u8),
20 #[error("zero domain length")]
21 BadDomainLength,
22 #[error("malformed domain name")]
23 MalformedDomain,
24 #[error("missing port")]
25 MissingPort,
26 #[error("packet too large: {0} > {1}")]
27 PacketTooLarge(usize, usize),
28}
29
30pub fn decode_socks5_udp_datagram(packet: &[u8]) -> Result<Socks5UdpRequest<'_>, UdpCodecError> {
31 if packet.len() < 4 {
32 return Err(UdpCodecError::TooShort);
33 }
34
35 if packet[0] != 0x00 || packet[1] != 0x00 {
36 return Err(UdpCodecError::BadReserved);
37 }
38
39 if packet[2] != 0x00 {
40 return Err(UdpCodecError::FragmentationUnsupported);
41 }
42
43 let atyp = packet[3];
44 let mut offset = 4;
45
46 let target = match atyp {
47 ATYP_IPV4 => {
48 if packet.len() < offset + 4 + 2 {
49 return Err(UdpCodecError::TooShort);
50 }
51 let mut addr = [0u8; 4];
52 addr.copy_from_slice(&packet[offset..offset + 4]);
53 offset += 4;
54 let port = u16::from_be_bytes([packet[offset], packet[offset + 1]]);
55 offset += 2;
56 SocksAddr::IPv4(addr, port)
57 }
58 ATYP_IPV6 => {
59 if packet.len() < offset + 16 + 2 {
60 return Err(UdpCodecError::TooShort);
61 }
62 let mut addr = [0u8; 16];
63 addr.copy_from_slice(&packet[offset..offset + 16]);
64 offset += 16;
65 let port = u16::from_be_bytes([packet[offset], packet[offset + 1]]);
66 offset += 2;
67 SocksAddr::IPv6(addr, port)
68 }
69 ATYP_DOMAIN => {
70 if packet.len() < offset + 1 {
71 return Err(UdpCodecError::TooShort);
72 }
73 let domain_len = packet[offset] as usize;
74 offset += 1;
75 if domain_len == 0 {
76 return Err(UdpCodecError::BadDomainLength);
77 }
78 if packet.len() < offset + domain_len + 2 {
79 return Err(UdpCodecError::TooShort);
80 }
81 let domain = std::str::from_utf8(&packet[offset..offset + domain_len])
82 .map_err(|_| UdpCodecError::MalformedDomain)?;
83 offset += domain_len;
84 let port = u16::from_be_bytes([packet[offset], packet[offset + 1]]);
85 offset += 2;
86 SocksAddr::Domain(domain.to_string(), port)
87 }
88 _ => return Err(UdpCodecError::UnknownAddressType(atyp)),
89 };
90
91 let payload = &packet[offset..];
92 Ok(Socks5UdpRequest { target, payload })
93}
94
95pub fn encode_socks5_udp_datagram(
96 target: &SocksAddr,
97 payload: &[u8],
98 out: &mut Vec<u8>,
99) -> Result<(), UdpCodecError> {
100 out.clear();
101 out.extend_from_slice(&[0x00, 0x00]);
102 out.push(0x00);
103 let encoded = target
104 .encode_reply()
105 .map_err(|_| UdpCodecError::MalformedDomain)?;
106 out.extend_from_slice(&encoded);
107 out.extend_from_slice(payload);
108 Ok(())
109}
110
111pub fn decode_socks5_udp_request(packet: &[u8]) -> Result<Socks5UdpRequest<'_>, UdpCodecError> {
112 decode_socks5_udp_datagram(packet)
113}
114
115pub fn encode_socks5_udp_response(
116 target: &SocksAddr,
117 payload: &[u8],
118 out: &mut Vec<u8>,
119) -> Result<(), UdpCodecError> {
120 encode_socks5_udp_datagram(target, payload, out)
121}
122
123#[cfg(test)]
124mod tests {
125 use super::*;
126
127 fn ipv4_packet(target: [u8; 4], port: u16, payload: &[u8]) -> Vec<u8> {
128 let mut pkt = vec![0x00, 0x00, 0x00, ATYP_IPV4];
129 pkt.extend_from_slice(&target);
130 pkt.extend_from_slice(&port.to_be_bytes());
131 pkt.extend_from_slice(payload);
132 pkt
133 }
134
135 fn ipv6_packet(target: [u8; 16], port: u16, payload: &[u8]) -> Vec<u8> {
136 let mut pkt = vec![0x00, 0x00, 0x00, ATYP_IPV6];
137 pkt.extend_from_slice(&target);
138 pkt.extend_from_slice(&port.to_be_bytes());
139 pkt.extend_from_slice(payload);
140 pkt
141 }
142
143 fn domain_packet(domain: &str, port: u16, payload: &[u8]) -> Vec<u8> {
144 let mut pkt = vec![0x00, 0x00, 0x00, ATYP_DOMAIN, domain.len() as u8];
145 pkt.extend_from_slice(domain.as_bytes());
146 pkt.extend_from_slice(&port.to_be_bytes());
147 pkt.extend_from_slice(payload);
148 pkt
149 }
150
151 #[test]
152 fn decode_ipv4_target() {
153 let pkt = ipv4_packet([192, 168, 1, 1], 8080, b"hello");
154 let req = decode_socks5_udp_datagram(&pkt).unwrap();
155 assert_eq!(req.target, SocksAddr::IPv4([192, 168, 1, 1], 8080));
156 assert_eq!(req.payload, b"hello");
157 }
158
159 #[test]
160 fn decode_ipv6_target() {
161 let addr = [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1];
162 let pkt = ipv6_packet(addr, 443, b"data");
163 let req = decode_socks5_udp_datagram(&pkt).unwrap();
164 assert_eq!(req.target, SocksAddr::IPv6(addr, 443));
165 assert_eq!(req.payload, b"data");
166 }
167
168 #[test]
169 fn decode_domain_target() {
170 let pkt = domain_packet("example.com", 53, b"\x01\x02");
171 let req = decode_socks5_udp_datagram(&pkt).unwrap();
172 assert_eq!(req.target, SocksAddr::Domain("example.com".to_string(), 53));
173 assert_eq!(req.payload, b"\x01\x02");
174 }
175
176 #[test]
177 fn encode_decode_roundtrip_ipv4() {
178 let target = SocksAddr::IPv4([10, 0, 0, 1], 9999);
179 let payload = b"roundtrip";
180 let mut buf = Vec::new();
181 encode_socks5_udp_datagram(&target, payload, &mut buf).unwrap();
182
183 let req = decode_socks5_udp_datagram(&buf).unwrap();
184 assert_eq!(req.target, target);
185 assert_eq!(req.payload, payload);
186 }
187
188 #[test]
189 fn encode_decode_roundtrip_ipv6() {
190 let addr = [0x20, 0x01, 0x0d, 0xb8, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1];
191 let target = SocksAddr::IPv6(addr, 80);
192 let payload = b"test payload";
193 let mut buf = Vec::new();
194 encode_socks5_udp_datagram(&target, payload, &mut buf).unwrap();
195
196 let req = decode_socks5_udp_datagram(&buf).unwrap();
197 assert_eq!(req.target, target);
198 assert_eq!(req.payload, payload);
199 }
200
201 #[test]
202 fn encode_decode_roundtrip_domain() {
203 let target = SocksAddr::Domain("localhost".to_string(), 5353);
204 let payload = b"dns query";
205 let mut buf = Vec::new();
206 encode_socks5_udp_datagram(&target, payload, &mut buf).unwrap();
207
208 let req = decode_socks5_udp_datagram(&buf).unwrap();
209 assert_eq!(req.target, target);
210 assert_eq!(req.payload, payload);
211 }
212
213 #[test]
214 fn reject_bad_rsv() {
215 let mut pkt = ipv4_packet([1, 2, 3, 4], 80, b"");
216 pkt[0] = 0x01;
217 assert!(matches!(
218 decode_socks5_udp_datagram(&pkt),
219 Err(UdpCodecError::BadReserved)
220 ));
221 }
222
223 #[test]
224 fn reject_bad_rsv_second_byte() {
225 let mut pkt = ipv4_packet([1, 2, 3, 4], 80, b"");
226 pkt[1] = 0x01;
227 assert!(matches!(
228 decode_socks5_udp_datagram(&pkt),
229 Err(UdpCodecError::BadReserved)
230 ));
231 }
232
233 #[test]
234 fn reject_frag_nonzero() {
235 let mut pkt = ipv4_packet([1, 2, 3, 4], 80, b"");
236 pkt[2] = 0x01;
237 assert!(matches!(
238 decode_socks5_udp_datagram(&pkt),
239 Err(UdpCodecError::FragmentationUnsupported)
240 ));
241 }
242
243 #[test]
244 fn reject_short_packet_too_few_bytes() {
245 assert!(matches!(
246 decode_socks5_udp_datagram(&[0x00, 0x00]),
247 Err(UdpCodecError::TooShort)
248 ));
249 }
250
251 #[test]
252 fn reject_short_packet_ipv4() {
253 let pkt = vec![0x00, 0x00, 0x00, ATYP_IPV4, 1, 2, 3];
255 assert!(matches!(
256 decode_socks5_udp_datagram(&pkt),
257 Err(UdpCodecError::TooShort)
258 ));
259 }
260
261 #[test]
262 fn reject_short_packet_ipv6() {
263 let pkt = vec![0x00, 0x00, 0x00, ATYP_IPV6, 0, 0, 0, 0];
264 assert!(matches!(
265 decode_socks5_udp_datagram(&pkt),
266 Err(UdpCodecError::TooShort)
267 ));
268 }
269
270 #[test]
271 fn reject_short_packet_domain_no_len() {
272 let pkt = vec![0x00, 0x00, 0x00, ATYP_DOMAIN];
273 assert!(matches!(
274 decode_socks5_udp_datagram(&pkt),
275 Err(UdpCodecError::TooShort)
276 ));
277 }
278
279 #[test]
280 fn reject_short_packet_domain_truncated() {
281 let pkt = vec![0x00, 0x00, 0x00, ATYP_DOMAIN, 5, b'a', b'b'];
282 assert!(matches!(
283 decode_socks5_udp_datagram(&pkt),
284 Err(UdpCodecError::TooShort)
285 ));
286 }
287
288 #[test]
289 fn reject_unknown_atyp() {
290 let pkt = vec![0x00, 0x00, 0x00, 0x05, 0x01, 0x02, 0x03, 0x04, 0x00, 0x50];
291 assert!(matches!(
292 decode_socks5_udp_datagram(&pkt),
293 Err(UdpCodecError::UnknownAddressType(0x05))
294 ));
295 }
296
297 #[test]
298 fn reject_zero_domain_length() {
299 let pkt = vec![0x00, 0x00, 0x00, ATYP_DOMAIN, 0x00];
300 assert!(matches!(
301 decode_socks5_udp_datagram(&pkt),
302 Err(UdpCodecError::BadDomainLength)
303 ));
304 }
305
306 #[test]
307 fn preserve_payload_bytes_exactly() {
308 let payload: Vec<u8> = (0..=255).collect();
309 let pkt = ipv4_packet([10, 0, 0, 1], 1234, &payload);
310 let req = decode_socks5_udp_datagram(&pkt).unwrap();
311 assert_eq!(req.payload.len(), 256);
312 for (i, &b) in req.payload.iter().enumerate() {
313 assert_eq!(b, i as u8);
314 }
315 }
316
317 #[test]
318 fn zero_length_payload() {
319 let pkt = ipv4_packet([10, 0, 0, 1], 80, b"");
320 let req = decode_socks5_udp_datagram(&pkt).unwrap();
321 assert_eq!(req.payload.len(), 0);
322 }
323
324 #[test]
325 fn encode_response_format() {
326 let target = SocksAddr::IPv4([192, 168, 1, 1], 80);
327 let payload = b"test";
328 let mut buf = Vec::new();
329 encode_socks5_udp_datagram(&target, payload, &mut buf).unwrap();
330
331 assert_eq!(buf[0], 0x00);
332 assert_eq!(buf[1], 0x00);
333 assert_eq!(buf[2], 0x00);
334 assert_eq!(buf[3], ATYP_IPV4);
335 assert_eq!(&buf[4..8], &[192, 168, 1, 1]);
336 assert_eq!(&buf[8..10], &80u16.to_be_bytes());
337 assert_eq!(&buf[10..], b"test");
338 }
339}