Skip to main content

eggress_protocol_socks/socks5/
udp_codec.rs

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        // ATYP_IPV4 but missing address bytes
254        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}