Skip to main content

eggress_core/
lib.rs

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/// A unique identifier for a protocol handler.
21#[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/// A unique identifier for a listener.
59pub type ListenerId = u64;
60
61/// A unique identifier for an upstream proxy.
62#[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/// The host of a target server, either an IP address or a domain name.
96#[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/// The address of a target server.
112#[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/// Client identity information.
173#[derive(Debug, Clone, PartialEq, Eq)]
174pub enum ClientIdentity {
175    Anonymous,
176    Username(String),
177    Opaque(String),
178}
179
180/// Context for a proxy session.
181#[derive(Debug, Clone)]
182pub struct SessionContext {
183    pub session_id: u64,
184    pub client_identity: ClientIdentity,
185    pub target_addr: TargetAddr,
186}
187
188/// Action to take for a routed connection.
189#[derive(Debug, Clone)]
190pub enum RouteAction {
191    Direct,
192    Upstream(UpstreamId),
193    Reject(RejectReason),
194}
195
196/// Reason for rejecting a connection.
197#[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
218/// A trait that combines AsyncRead and AsyncWrite for bidirectional streams.
219pub trait AsyncStream: AsyncRead + AsyncWrite + Send + Unpin {}
220impl<T: AsyncRead + AsyncWrite + Send + Unpin> AsyncStream for T {}
221
222/// A type alias for a boxed async stream.
223pub type BoxStream = Box<dyn AsyncStream>;
224
225/// Error types for connection operations.
226#[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/// Error types for protocol operations.
243#[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/// Error types for authentication operations.
258#[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/// Error types for relay operations.
271#[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}