Skip to main content

eggress_protocol_shadowsocks/
address.rs

1use eggress_core::{TargetAddr, TargetHost};
2
3use crate::error::ShadowsocksError;
4
5/// ATYP values for Shadowsocks address format.
6const ATYP_IPV4: u8 = 0x01;
7const ATYP_DOMAIN: u8 = 0x03;
8const ATYP_IPV6: u8 = 0x04;
9
10/// Encode a TargetAddr into Shadowsocks wire format.
11pub fn encode_address(target: &TargetAddr) -> Result<Vec<u8>, ShadowsocksError> {
12    let mut buf = Vec::with_capacity(1 + 4 + 2); // ATYP + max IP + port
13    match &target.host {
14        TargetHost::Ip(std::net::IpAddr::V4(ip)) => {
15            buf.push(ATYP_IPV4);
16            buf.extend_from_slice(&ip.octets());
17        }
18        TargetHost::Ip(std::net::IpAddr::V6(ip)) => {
19            buf.push(ATYP_IPV6);
20            buf.extend_from_slice(&ip.octets());
21        }
22        TargetHost::Domain(domain) => {
23            if domain.len() > 255 {
24                return Err(ShadowsocksError::InvalidAddress(format!(
25                    "domain too long: {} bytes (max 255)",
26                    domain.len()
27                )));
28            }
29            buf.push(ATYP_DOMAIN);
30            buf.push(domain.len() as u8);
31            buf.extend_from_slice(domain.as_bytes());
32        }
33    }
34    buf.extend_from_slice(&target.port.to_be_bytes());
35    Ok(buf)
36}
37
38/// Decode a Shadowsocks address from a byte slice.
39///
40/// Returns the decoded TargetAddr and the number of bytes consumed.
41pub fn decode_address(data: &[u8]) -> Result<(TargetAddr, usize), ShadowsocksError> {
42    if data.is_empty() {
43        return Err(ShadowsocksError::InvalidAddress("empty data".into()));
44    }
45
46    let atyp = data[0];
47
48    let (host, mut pos) = match atyp {
49        ATYP_IPV4 => {
50            if data.len() < 5 {
51                return Err(ShadowsocksError::InvalidAddress(
52                    "truncated IPv4 address".into(),
53                ));
54            }
55            let ip = std::net::Ipv4Addr::new(data[1], data[2], data[3], data[4]);
56            (TargetHost::Ip(std::net::IpAddr::V4(ip)), 5)
57        }
58        ATYP_IPV6 => {
59            if data.len() < 17 {
60                return Err(ShadowsocksError::InvalidAddress(
61                    "truncated IPv6 address".into(),
62                ));
63            }
64            let mut octets = [0u8; 16];
65            octets.copy_from_slice(&data[1..17]);
66            let ip = std::net::Ipv6Addr::from(octets);
67            (TargetHost::Ip(std::net::IpAddr::V6(ip)), 17)
68        }
69        ATYP_DOMAIN => {
70            if data.len() < 2 {
71                return Err(ShadowsocksError::InvalidAddress(
72                    "truncated domain length".into(),
73                ));
74            }
75            let len = data[1] as usize;
76            if data.len() < 2 + len {
77                return Err(ShadowsocksError::InvalidAddress("truncated domain".into()));
78            }
79            let domain = String::from_utf8(data[2..2 + len].to_vec())
80                .map_err(|e| ShadowsocksError::InvalidAddress(format!("invalid UTF-8: {}", e)))?;
81            (TargetHost::Domain(domain), 2 + len)
82        }
83        _ => {
84            return Err(ShadowsocksError::InvalidAddress(format!(
85                "unknown ATYP: {:#04x}",
86                atyp
87            )));
88        }
89    };
90
91    if data.len() < pos + 2 {
92        return Err(ShadowsocksError::InvalidAddress("truncated port".into()));
93    }
94
95    let port = u16::from_be_bytes([data[pos], data[pos + 1]]);
96    pos += 2;
97
98    Ok((TargetAddr { host, port }, pos))
99}
100
101#[cfg(test)]
102mod tests {
103    use super::*;
104
105    #[test]
106    fn test_encode_decode_ipv4() {
107        let target = TargetAddr {
108            host: TargetHost::Ip("192.168.1.1".parse().unwrap()),
109            port: 8080,
110        };
111        let encoded = encode_address(&target).unwrap();
112        assert_eq!(encoded[0], ATYP_IPV4);
113        let (decoded, consumed) = decode_address(&encoded).unwrap();
114        assert_eq!(decoded, target);
115        assert_eq!(consumed, encoded.len());
116    }
117
118    #[test]
119    fn test_encode_decode_ipv6() {
120        let target = TargetAddr {
121            host: TargetHost::Ip("::1".parse().unwrap()),
122            port: 443,
123        };
124        let encoded = encode_address(&target).unwrap();
125        assert_eq!(encoded[0], ATYP_IPV6);
126        let (decoded, consumed) = decode_address(&encoded).unwrap();
127        assert_eq!(decoded, target);
128        assert_eq!(consumed, encoded.len());
129    }
130
131    #[test]
132    fn test_encode_decode_domain() {
133        let target = TargetAddr {
134            host: TargetHost::Domain("example.com".to_string()),
135            port: 443,
136        };
137        let encoded = encode_address(&target).unwrap();
138        assert_eq!(encoded[0], ATYP_DOMAIN);
139        assert_eq!(encoded[1], 11); // "example.com".len()
140        let (decoded, consumed) = decode_address(&encoded).unwrap();
141        assert_eq!(decoded, target);
142        assert_eq!(consumed, encoded.len());
143    }
144
145    #[test]
146    fn test_decode_empty() {
147        assert!(decode_address(&[]).is_err());
148    }
149
150    #[test]
151    fn test_decode_truncated_ipv4() {
152        assert!(decode_address(&[ATYP_IPV4, 192, 168]).is_err());
153    }
154
155    #[test]
156    fn test_decode_truncated_port() {
157        let data = vec![ATYP_IPV4, 192, 168, 1, 1]; // missing port
158        assert!(decode_address(&data).is_err());
159    }
160
161    #[test]
162    fn test_decode_unknown_atyp() {
163        assert!(decode_address(&[0xFF]).is_err());
164    }
165
166    #[test]
167    fn test_encode_decode_roundtrip_various_ports() {
168        let ports = [0, 1, 80, 443, 8080, 65535];
169        for port in ports {
170            let target = TargetAddr {
171                host: TargetHost::Ip("10.0.0.1".parse().unwrap()),
172                port,
173            };
174            let encoded = encode_address(&target).unwrap();
175            let (decoded, _) = decode_address(&encoded).unwrap();
176            assert_eq!(decoded.port, port);
177        }
178    }
179}