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/// Failure to convert URI syntax into a dispatchable runtime protocol.
59#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
60pub enum ProtocolConversionError {
61    /// The syntax concept has no listener/runtime role (e.g. SSH upstream-only).
62    #[error("SSH is an upstream-only transport, not a listener protocol")]
63    UpstreamOnlyTransport,
64}
65
66impl ProtocolId {
67    /// Central typed conversion from URI syntax to runtime disposition.
68    ///
69    /// Exhaustive over [`eggress_uri::ProtocolSpec`]; unsupported and
70    /// non-dispatchable variants fail explicitly instead of silent fallback:
71    ///
72    /// - `HttpOnly` is an upstream request adapter; listeners serve `Http`.
73    /// - `Unix` is a transport concept; listeners serve it as `Raw` TCP
74    ///   semantics with a unix-socket bind (see `ListenerUdpConfig`/unix listener).
75    /// - `Ssh` is upstream-only and returns [`ProtocolConversionError::UpstreamOnlyTransport`].
76    /// - `Echo` / `Reverse` are runtime-only: they have no `ProtocolSpec`
77    ///   counterpart and never flow through this conversion.
78    /// - `Raw` covers both `raw` and `tunnel` URI aliases (canonicalized in
79    ///   `ProtocolSpec`), and `WebSocket` covers `ws`/`wss`.
80    /// - `H3`/`Quic` map directly; feature gating (optional `quic` build)
81    ///   stays in config compilation, not here.
82    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
104/// A unique identifier for a listener.
105pub type ListenerId = u64;
106
107/// A unique identifier for an upstream proxy.
108#[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/// The host of a target server, either an IP address or a domain name.
142#[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/// The address of a target server.
158#[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/// Client identity information.
219#[derive(Debug, Clone, PartialEq, Eq)]
220pub enum ClientIdentity {
221    Anonymous,
222    Username(String),
223    Opaque(String),
224}
225
226/// Context for a proxy session.
227#[derive(Debug, Clone)]
228pub struct SessionContext {
229    pub session_id: u64,
230    pub client_identity: ClientIdentity,
231    pub target_addr: TargetAddr,
232}
233
234/// Action to take for a routed connection.
235#[derive(Debug, Clone)]
236pub enum RouteAction {
237    Direct,
238    Upstream(UpstreamId),
239    Reject(RejectReason),
240}
241
242/// Reason for rejecting a connection.
243#[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
264/// A trait that combines AsyncRead and AsyncWrite for bidirectional streams.
265pub trait AsyncStream: AsyncRead + AsyncWrite + Send + Unpin {}
266impl<T: AsyncRead + AsyncWrite + Send + Unpin> AsyncStream for T {}
267
268/// A type alias for a boxed async stream.
269pub type BoxStream = Box<dyn AsyncStream>;
270
271/// Error types for connection operations.
272#[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/// Error types for protocol operations.
289#[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/// Error types for authentication operations.
304#[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/// Error types for relay operations.
317#[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        // Every current `ProtocolSpec` variant has an explicit disposition.
410        assert_eq!(eggress_uri::ProtocolSpec::all_variants().len(), 14);
411        let cases: &[(S, Option<ProtocolId>)] = &[
412            (S::Http, Some(ProtocolId::Http)),
413            // Upstream adapter collapses to Http for listener dispatch.
414            (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            // Transport concept: Unix listeners serve Raw semantics.
426            (S::Unix, Some(ProtocolId::Raw)),
427            // Upstream-only transport fails explicitly.
428            (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}