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
58#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
60pub enum ProtocolConversionError {
61 #[error("SSH is an upstream-only transport, not a listener protocol")]
63 UpstreamOnlyTransport,
64}
65
66impl ProtocolId {
67 pub fn from_protocol_spec(
83 spec: eggress_uri::ProtocolSpec,
84 ) -> Result<Self, ProtocolConversionError> {
85 use eggress_uri::ProtocolSpec as S;
86 match spec {
87 S::Http | S::HttpOnly => Ok(ProtocolId::Http),
88 S::Socks4 => Ok(ProtocolId::Socks4),
89 S::Socks5 => Ok(ProtocolId::Socks5),
90 S::Shadowsocks => Ok(ProtocolId::Shadowsocks),
91 S::ShadowsocksR => Ok(ProtocolId::ShadowsocksR),
92 S::Trojan => Ok(ProtocolId::Trojan),
93 S::Http2 => Ok(ProtocolId::Http2),
94 S::Http3 => Ok(ProtocolId::Http3),
95 S::Quic => Ok(ProtocolId::Quic),
96 S::WebSocket => Ok(ProtocolId::WebSocket),
97 S::Raw => Ok(ProtocolId::Raw),
98 S::Unix => Ok(ProtocolId::Raw),
99 S::Ssh => Err(ProtocolConversionError::UpstreamOnlyTransport),
100 }
101 }
102}
103
104pub type ListenerId = u64;
106
107#[derive(Debug, Clone, PartialEq, Eq, Hash)]
109pub struct UpstreamId(Arc<str>);
110
111impl serde::Serialize for UpstreamId {
112 fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
113 serializer.serialize_str(&self.0)
114 }
115}
116
117impl UpstreamId {
118 pub fn new(id: impl Into<Arc<str>>) -> Self {
119 Self(id.into())
120 }
121
122 pub fn as_str(&self) -> &str {
123 &self.0
124 }
125}
126
127impl fmt::Display for UpstreamId {
128 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
129 write!(f, "{}", self.0)
130 }
131}
132
133impl std::str::FromStr for UpstreamId {
134 type Err = std::convert::Infallible;
135
136 fn from_str(s: &str) -> Result<Self, Self::Err> {
137 Ok(Self::new(s))
138 }
139}
140
141#[derive(Debug, Clone, PartialEq, Eq, Hash)]
143pub enum TargetHost {
144 Ip(IpAddr),
145 Domain(String),
146}
147
148impl fmt::Display for TargetHost {
149 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
150 match self {
151 TargetHost::Ip(ip) => write!(f, "{}", ip),
152 TargetHost::Domain(domain) => write!(f, "{}", domain),
153 }
154 }
155}
156
157#[derive(Debug, Clone, PartialEq, Eq, Hash)]
159pub struct TargetAddr {
160 pub host: TargetHost,
161 pub port: u16,
162}
163
164impl fmt::Display for TargetAddr {
165 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
166 match &self.host {
167 TargetHost::Ip(IpAddr::V6(_)) => write!(f, "[{}]:{}", self.host, self.port),
168 _ => write!(f, "{}:{}", self.host, self.port),
169 }
170 }
171}
172
173impl std::str::FromStr for TargetAddr {
174 type Err = String;
175
176 fn from_str(s: &str) -> Result<Self, Self::Err> {
177 if let Some(rest) = s.strip_prefix('[') {
178 let close = rest
179 .find(']')
180 .ok_or_else(|| format!("invalid target format: missing closing ']' in '{s}'"))?;
181 let host_str = &rest[..close];
182 let after = &rest[close + 1..];
183 let port_str = after.strip_prefix(':').ok_or_else(|| {
184 format!("invalid target format: missing ':port' after ']' in '{s}'")
185 })?;
186 let port: u16 = port_str
187 .parse()
188 .map_err(|e| format!("invalid port '{port_str}': {e}"))?;
189 let ip: IpAddr = host_str
190 .parse()
191 .map_err(|e| format!("invalid IPv6 address '{host_str}': {e}"))?;
192 Ok(TargetAddr {
193 host: TargetHost::Ip(ip),
194 port,
195 })
196 } else if s.matches(':').count() > 1 {
197 Err(format!(
198 "invalid target format: unbracketed IPv6 literal in '{s}' (use [addr]:port)"
199 ))
200 } else if let Some(idx) = s.rfind(':') {
201 let host_part = &s[..idx];
202 let port_part = &s[idx + 1..];
203 let port: u16 = port_part
204 .parse()
205 .map_err(|e| format!("invalid port '{port_part}': {e}"))?;
206 let host = if let Ok(ip) = host_part.parse::<IpAddr>() {
207 TargetHost::Ip(ip)
208 } else {
209 TargetHost::Domain(host_part.to_string())
210 };
211 Ok(TargetAddr { host, port })
212 } else {
213 Err(format!("invalid target format: {s}"))
214 }
215 }
216}
217
218#[derive(Debug, Clone, PartialEq, Eq)]
220pub enum ClientIdentity {
221 Anonymous,
222 Username(String),
223 Opaque(String),
224}
225
226#[derive(Debug, Clone)]
228pub struct SessionContext {
229 pub session_id: u64,
230 pub client_identity: ClientIdentity,
231 pub target_addr: TargetAddr,
232}
233
234#[derive(Debug, Clone)]
236pub enum RouteAction {
237 Direct,
238 Upstream(UpstreamId),
239 Reject(RejectReason),
240}
241
242#[derive(Debug, Clone, PartialEq, Eq)]
244pub enum RejectReason {
245 UnsupportedProtocol,
246 AuthRequired,
247 AccessDenied,
248 Blocked,
249 InternalError,
250}
251
252impl fmt::Display for RejectReason {
253 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
254 match self {
255 RejectReason::UnsupportedProtocol => write!(f, "unsupported protocol"),
256 RejectReason::AuthRequired => write!(f, "authentication required"),
257 RejectReason::AccessDenied => write!(f, "access denied"),
258 RejectReason::Blocked => write!(f, "target address blocked"),
259 RejectReason::InternalError => write!(f, "internal error"),
260 }
261 }
262}
263
264pub trait AsyncStream: AsyncRead + AsyncWrite + Send + Unpin {}
266impl<T: AsyncRead + AsyncWrite + Send + Unpin> AsyncStream for T {}
267
268pub type BoxStream = Box<dyn AsyncStream>;
270
271#[derive(Debug, thiserror::Error)]
273pub enum ConnectError {
274 #[error("connection refused")]
275 ConnectionRefused,
276 #[error("connection timed out")]
277 Timeout,
278 #[error("DNS resolution failed: {0}")]
279 DnsResolution(String),
280 #[error("TLS handshake failed: {0}")]
281 TlsHandshake(String),
282 #[error("reserved or private target IP: {0}")]
283 ReservedTarget(std::net::IpAddr),
284 #[error("IO error: {0}")]
285 Io(#[from] std::io::Error),
286}
287
288#[derive(Debug, thiserror::Error)]
290pub enum ProtocolError {
291 #[error("malformed message")]
292 MalformedMessage,
293 #[error("unsupported version")]
294 UnsupportedVersion,
295 #[error("method not supported")]
296 MethodNotSupported,
297 #[error("address type not supported")]
298 AddressTypeNotSupported,
299 #[error("IO error: {0}")]
300 Io(#[from] std::io::Error),
301}
302
303#[derive(Debug, thiserror::Error)]
305pub enum AuthError {
306 #[error("invalid credentials")]
307 InvalidCredentials,
308 #[error("authentication method not supported")]
309 MethodNotSupported,
310 #[error("authentication required")]
311 Required,
312 #[error("IO error: {0}")]
313 Io(#[from] std::io::Error),
314}
315
316#[derive(Debug, thiserror::Error)]
318pub enum RelayError {
319 #[error("connection closed")]
320 ConnectionClosed,
321 #[error("IO error: {0}")]
322 Io(#[from] std::io::Error),
323}
324
325#[cfg(test)]
326mod tests {
327 use super::*;
328
329 #[test]
330 fn test_target_host_display() {
331 let ip_host = TargetHost::Ip("127.0.0.1".parse().unwrap());
332 assert_eq!(ip_host.to_string(), "127.0.0.1");
333
334 let domain_host = TargetHost::Domain("example.com".to_string());
335 assert_eq!(domain_host.to_string(), "example.com");
336 }
337
338 #[test]
339 fn test_target_addr_display() {
340 let addr = TargetAddr {
341 host: TargetHost::Domain("example.com".to_string()),
342 port: 8080,
343 };
344 assert_eq!(addr.to_string(), "example.com:8080");
345 }
346
347 #[test]
348 fn test_reject_reason_display() {
349 assert_eq!(
350 RejectReason::UnsupportedProtocol.to_string(),
351 "unsupported protocol"
352 );
353 assert_eq!(
354 RejectReason::AuthRequired.to_string(),
355 "authentication required"
356 );
357 }
358
359 #[test]
360 fn test_target_addr_from_str_bracketed_ipv6() {
361 let addr: TargetAddr = "[::1]:443".parse().unwrap();
362 assert_eq!(addr.host, TargetHost::Ip("::1".parse::<IpAddr>().unwrap()));
363 assert_eq!(addr.port, 443);
364 }
365
366 #[test]
367 fn test_target_addr_from_str_full_ipv6() {
368 let addr: TargetAddr = "[2001:db8::1]:80".parse().unwrap();
369 assert_eq!(
370 addr.host,
371 TargetHost::Ip("2001:db8::1".parse::<IpAddr>().unwrap())
372 );
373 assert_eq!(addr.port, 80);
374 }
375
376 #[test]
377 fn test_target_addr_from_str_rejects_unbracketed_ipv6() {
378 let err = "::1:443".parse::<TargetAddr>().unwrap_err();
379 assert!(err.contains("unbracketed IPv6"));
380 }
381
382 #[test]
383 fn test_target_addr_from_str_rejects_unclosed_bracket() {
384 let err = "[::1:443".parse::<TargetAddr>().unwrap_err();
385 assert!(err.contains("closing ']'"));
386 }
387
388 #[test]
389 fn test_target_addr_display_brackets_ipv6() {
390 let addr = TargetAddr {
391 host: TargetHost::Ip("::1".parse().unwrap()),
392 port: 443,
393 };
394 assert_eq!(addr.to_string(), "[::1]:443");
395 }
396
397 #[test]
398 fn test_target_addr_display_does_not_bracket_ipv4() {
399 let addr = TargetAddr {
400 host: TargetHost::Ip("127.0.0.1".parse().unwrap()),
401 port: 80,
402 };
403 assert_eq!(addr.to_string(), "127.0.0.1:80");
404 }
405
406 #[test]
407 fn test_protocol_spec_runtime_disposition_is_exhaustive() {
408 use eggress_uri::ProtocolSpec as S;
409 assert_eq!(eggress_uri::ProtocolSpec::all_variants().len(), 14);
411 let cases: &[(S, Option<ProtocolId>)] = &[
412 (S::Http, Some(ProtocolId::Http)),
413 (S::HttpOnly, Some(ProtocolId::Http)),
415 (S::Socks4, Some(ProtocolId::Socks4)),
416 (S::Socks5, Some(ProtocolId::Socks5)),
417 (S::Shadowsocks, Some(ProtocolId::Shadowsocks)),
418 (S::ShadowsocksR, Some(ProtocolId::ShadowsocksR)),
419 (S::Trojan, Some(ProtocolId::Trojan)),
420 (S::Http2, Some(ProtocolId::Http2)),
421 (S::Http3, Some(ProtocolId::Http3)),
422 (S::Quic, Some(ProtocolId::Quic)),
423 (S::WebSocket, Some(ProtocolId::WebSocket)),
424 (S::Raw, Some(ProtocolId::Raw)),
425 (S::Unix, Some(ProtocolId::Raw)),
427 (S::Ssh, None),
429 ];
430 for (spec, expected) in cases {
431 assert_eq!(
432 ProtocolId::from_protocol_spec(*spec).ok(),
433 *expected,
434 "disposition for {spec:?}"
435 );
436 }
437 assert_eq!(
438 ProtocolId::from_protocol_spec(S::Ssh),
439 Err(ProtocolConversionError::UpstreamOnlyTransport)
440 );
441 }
442}