1use std::fmt;
2use std::net::IpAddr;
3use std::sync::Arc;
4
5use tokio::io::{AsyncRead, AsyncWrite};
6
7pub use capability::{
8 classify_upstream_chain, CapabilityResult, TransportCapability, UpstreamCapabilities,
9};
10
11pub mod capability;
12pub mod chain;
13pub mod connector;
14pub mod detect;
15pub mod dispatch;
16pub mod listener;
17pub mod relay;
18pub mod replay;
19
20#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
22pub enum ProtocolId {
23 Http,
24 Socks4,
25 Socks5,
26 Shadowsocks,
27 ShadowsocksR,
28 Trojan,
29 Http2,
30 Http3,
31 Quic,
32 WebSocket,
33 Raw,
34 Echo,
35 Reverse,
36}
37
38impl fmt::Display for ProtocolId {
39 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
40 match self {
41 ProtocolId::Http => write!(f, "http"),
42 ProtocolId::Socks4 => write!(f, "socks4"),
43 ProtocolId::Socks5 => write!(f, "socks5"),
44 ProtocolId::Shadowsocks => write!(f, "shadowsocks"),
45 ProtocolId::ShadowsocksR => write!(f, "ssr"),
46 ProtocolId::Trojan => write!(f, "trojan"),
47 ProtocolId::Http2 => write!(f, "h2"),
48 ProtocolId::Http3 => write!(f, "h3"),
49 ProtocolId::Quic => write!(f, "quic"),
50 ProtocolId::WebSocket => write!(f, "websocket"),
51 ProtocolId::Raw => write!(f, "raw"),
52 ProtocolId::Echo => write!(f, "echo"),
53 ProtocolId::Reverse => write!(f, "reverse"),
54 }
55 }
56}
57
58pub type ListenerId = u64;
60
61#[derive(Debug, Clone, PartialEq, Eq, Hash)]
63pub struct UpstreamId(Arc<str>);
64
65impl serde::Serialize for UpstreamId {
66 fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
67 serializer.serialize_str(&self.0)
68 }
69}
70
71impl UpstreamId {
72 pub fn new(id: impl Into<Arc<str>>) -> Self {
73 Self(id.into())
74 }
75
76 pub fn as_str(&self) -> &str {
77 &self.0
78 }
79}
80
81impl fmt::Display for UpstreamId {
82 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
83 write!(f, "{}", self.0)
84 }
85}
86
87impl std::str::FromStr for UpstreamId {
88 type Err = std::convert::Infallible;
89
90 fn from_str(s: &str) -> Result<Self, Self::Err> {
91 Ok(Self::new(s))
92 }
93}
94
95#[derive(Debug, Clone, PartialEq, Eq, Hash)]
97pub enum TargetHost {
98 Ip(IpAddr),
99 Domain(String),
100}
101
102impl fmt::Display for TargetHost {
103 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
104 match self {
105 TargetHost::Ip(ip) => write!(f, "{}", ip),
106 TargetHost::Domain(domain) => write!(f, "{}", domain),
107 }
108 }
109}
110
111#[derive(Debug, Clone, PartialEq, Eq, Hash)]
113pub struct TargetAddr {
114 pub host: TargetHost,
115 pub port: u16,
116}
117
118impl fmt::Display for TargetAddr {
119 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
120 match &self.host {
121 TargetHost::Ip(IpAddr::V6(_)) => write!(f, "[{}]:{}", self.host, self.port),
122 _ => write!(f, "{}:{}", self.host, self.port),
123 }
124 }
125}
126
127impl std::str::FromStr for TargetAddr {
128 type Err = String;
129
130 fn from_str(s: &str) -> Result<Self, Self::Err> {
131 if let Some(rest) = s.strip_prefix('[') {
132 let close = rest
133 .find(']')
134 .ok_or_else(|| format!("invalid target format: missing closing ']' in '{s}'"))?;
135 let host_str = &rest[..close];
136 let after = &rest[close + 1..];
137 let port_str = after.strip_prefix(':').ok_or_else(|| {
138 format!("invalid target format: missing ':port' after ']' in '{s}'")
139 })?;
140 let port: u16 = port_str
141 .parse()
142 .map_err(|e| format!("invalid port '{port_str}': {e}"))?;
143 let ip: IpAddr = host_str
144 .parse()
145 .map_err(|e| format!("invalid IPv6 address '{host_str}': {e}"))?;
146 Ok(TargetAddr {
147 host: TargetHost::Ip(ip),
148 port,
149 })
150 } else if s.matches(':').count() > 1 {
151 Err(format!(
152 "invalid target format: unbracketed IPv6 literal in '{s}' (use [addr]:port)"
153 ))
154 } else if let Some(idx) = s.rfind(':') {
155 let host_part = &s[..idx];
156 let port_part = &s[idx + 1..];
157 let port: u16 = port_part
158 .parse()
159 .map_err(|e| format!("invalid port '{port_part}': {e}"))?;
160 let host = if let Ok(ip) = host_part.parse::<IpAddr>() {
161 TargetHost::Ip(ip)
162 } else {
163 TargetHost::Domain(host_part.to_string())
164 };
165 Ok(TargetAddr { host, port })
166 } else {
167 Err(format!("invalid target format: {s}"))
168 }
169 }
170}
171
172#[derive(Debug, Clone, PartialEq, Eq)]
174pub enum ClientIdentity {
175 Anonymous,
176 Username(String),
177 Opaque(String),
178}
179
180#[derive(Debug, Clone)]
182pub struct SessionContext {
183 pub session_id: u64,
184 pub client_identity: ClientIdentity,
185 pub target_addr: TargetAddr,
186}
187
188#[derive(Debug, Clone)]
190pub enum RouteAction {
191 Direct,
192 Upstream(UpstreamId),
193 Reject(RejectReason),
194}
195
196#[derive(Debug, Clone, PartialEq, Eq)]
198pub enum RejectReason {
199 UnsupportedProtocol,
200 AuthRequired,
201 AccessDenied,
202 Blocked,
203 InternalError,
204}
205
206impl fmt::Display for RejectReason {
207 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
208 match self {
209 RejectReason::UnsupportedProtocol => write!(f, "unsupported protocol"),
210 RejectReason::AuthRequired => write!(f, "authentication required"),
211 RejectReason::AccessDenied => write!(f, "access denied"),
212 RejectReason::Blocked => write!(f, "target address blocked"),
213 RejectReason::InternalError => write!(f, "internal error"),
214 }
215 }
216}
217
218pub trait AsyncStream: AsyncRead + AsyncWrite + Send + Unpin {}
220impl<T: AsyncRead + AsyncWrite + Send + Unpin> AsyncStream for T {}
221
222pub type BoxStream = Box<dyn AsyncStream>;
224
225#[derive(Debug, thiserror::Error)]
227pub enum ConnectError {
228 #[error("connection refused")]
229 ConnectionRefused,
230 #[error("connection timed out")]
231 Timeout,
232 #[error("DNS resolution failed: {0}")]
233 DnsResolution(String),
234 #[error("TLS handshake failed: {0}")]
235 TlsHandshake(String),
236 #[error("reserved or private target IP: {0}")]
237 ReservedTarget(std::net::IpAddr),
238 #[error("IO error: {0}")]
239 Io(#[from] std::io::Error),
240}
241
242#[derive(Debug, thiserror::Error)]
244pub enum ProtocolError {
245 #[error("malformed message")]
246 MalformedMessage,
247 #[error("unsupported version")]
248 UnsupportedVersion,
249 #[error("method not supported")]
250 MethodNotSupported,
251 #[error("address type not supported")]
252 AddressTypeNotSupported,
253 #[error("IO error: {0}")]
254 Io(#[from] std::io::Error),
255}
256
257#[derive(Debug, thiserror::Error)]
259pub enum AuthError {
260 #[error("invalid credentials")]
261 InvalidCredentials,
262 #[error("authentication method not supported")]
263 MethodNotSupported,
264 #[error("authentication required")]
265 Required,
266 #[error("IO error: {0}")]
267 Io(#[from] std::io::Error),
268}
269
270#[derive(Debug, thiserror::Error)]
272pub enum RelayError {
273 #[error("connection closed")]
274 ConnectionClosed,
275 #[error("IO error: {0}")]
276 Io(#[from] std::io::Error),
277}
278
279#[cfg(test)]
280mod tests {
281 use super::*;
282
283 #[test]
284 fn test_target_host_display() {
285 let ip_host = TargetHost::Ip("127.0.0.1".parse().unwrap());
286 assert_eq!(ip_host.to_string(), "127.0.0.1");
287
288 let domain_host = TargetHost::Domain("example.com".to_string());
289 assert_eq!(domain_host.to_string(), "example.com");
290 }
291
292 #[test]
293 fn test_target_addr_display() {
294 let addr = TargetAddr {
295 host: TargetHost::Domain("example.com".to_string()),
296 port: 8080,
297 };
298 assert_eq!(addr.to_string(), "example.com:8080");
299 }
300
301 #[test]
302 fn test_reject_reason_display() {
303 assert_eq!(
304 RejectReason::UnsupportedProtocol.to_string(),
305 "unsupported protocol"
306 );
307 assert_eq!(
308 RejectReason::AuthRequired.to_string(),
309 "authentication required"
310 );
311 }
312
313 #[test]
314 fn test_target_addr_from_str_bracketed_ipv6() {
315 let addr: TargetAddr = "[::1]:443".parse().unwrap();
316 assert_eq!(addr.host, TargetHost::Ip("::1".parse::<IpAddr>().unwrap()));
317 assert_eq!(addr.port, 443);
318 }
319
320 #[test]
321 fn test_target_addr_from_str_full_ipv6() {
322 let addr: TargetAddr = "[2001:db8::1]:80".parse().unwrap();
323 assert_eq!(
324 addr.host,
325 TargetHost::Ip("2001:db8::1".parse::<IpAddr>().unwrap())
326 );
327 assert_eq!(addr.port, 80);
328 }
329
330 #[test]
331 fn test_target_addr_from_str_rejects_unbracketed_ipv6() {
332 let err = "::1:443".parse::<TargetAddr>().unwrap_err();
333 assert!(err.contains("unbracketed IPv6"));
334 }
335
336 #[test]
337 fn test_target_addr_from_str_rejects_unclosed_bracket() {
338 let err = "[::1:443".parse::<TargetAddr>().unwrap_err();
339 assert!(err.contains("closing ']'"));
340 }
341
342 #[test]
343 fn test_target_addr_display_brackets_ipv6() {
344 let addr = TargetAddr {
345 host: TargetHost::Ip("::1".parse().unwrap()),
346 port: 443,
347 };
348 assert_eq!(addr.to_string(), "[::1]:443");
349 }
350
351 #[test]
352 fn test_target_addr_display_does_not_bracket_ipv4() {
353 let addr = TargetAddr {
354 host: TargetHost::Ip("127.0.0.1".parse().unwrap()),
355 port: 80,
356 };
357 assert_eq!(addr.to_string(), "127.0.0.1:80");
358 }
359}