eggress-core 1.0.3

Core types, traits, and infrastructure for eggress proxy
Documentation
use std::fmt;
use std::net::IpAddr;
use std::sync::Arc;

use tokio::io::{AsyncRead, AsyncWrite};

pub use capability::{
    classify_upstream_chain, CapabilityResult, TransportCapability, UpstreamCapabilities,
};

pub mod capability;
pub mod chain;
pub mod connector;
pub mod detect;
pub mod dispatch;
pub mod listener;
pub mod relay;
pub mod replay;

/// A unique identifier for a protocol handler.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum ProtocolId {
    Http,
    Socks4,
    Socks5,
    Shadowsocks,
    ShadowsocksR,
    Trojan,
    Http2,
    Http3,
    Quic,
    WebSocket,
    Raw,
    Echo,
    Reverse,
}

impl fmt::Display for ProtocolId {
    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
        match self {
            ProtocolId::Http => write!(f, "http"),
            ProtocolId::Socks4 => write!(f, "socks4"),
            ProtocolId::Socks5 => write!(f, "socks5"),
            ProtocolId::Shadowsocks => write!(f, "shadowsocks"),
            ProtocolId::ShadowsocksR => write!(f, "ssr"),
            ProtocolId::Trojan => write!(f, "trojan"),
            ProtocolId::Http2 => write!(f, "h2"),
            ProtocolId::Http3 => write!(f, "h3"),
            ProtocolId::Quic => write!(f, "quic"),
            ProtocolId::WebSocket => write!(f, "websocket"),
            ProtocolId::Raw => write!(f, "raw"),
            ProtocolId::Echo => write!(f, "echo"),
            ProtocolId::Reverse => write!(f, "reverse"),
        }
    }
}

/// A unique identifier for a listener.
pub type ListenerId = u64;

/// A unique identifier for an upstream proxy.
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct UpstreamId(Arc<str>);

impl serde::Serialize for UpstreamId {
    fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
        serializer.serialize_str(&self.0)
    }
}

impl UpstreamId {
    pub fn new(id: impl Into<Arc<str>>) -> Self {
        Self(id.into())
    }

    pub fn as_str(&self) -> &str {
        &self.0
    }
}

impl fmt::Display for UpstreamId {
    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
        write!(f, "{}", self.0)
    }
}

impl std::str::FromStr for UpstreamId {
    type Err = std::convert::Infallible;

    fn from_str(s: &str) -> Result<Self, Self::Err> {
        Ok(Self::new(s))
    }
}

/// The host of a target server, either an IP address or a domain name.
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum TargetHost {
    Ip(IpAddr),
    Domain(String),
}

impl fmt::Display for TargetHost {
    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
        match self {
            TargetHost::Ip(ip) => write!(f, "{}", ip),
            TargetHost::Domain(domain) => write!(f, "{}", domain),
        }
    }
}

/// The address of a target server.
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct TargetAddr {
    pub host: TargetHost,
    pub port: u16,
}

impl fmt::Display for TargetAddr {
    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
        match &self.host {
            TargetHost::Ip(IpAddr::V6(_)) => write!(f, "[{}]:{}", self.host, self.port),
            _ => write!(f, "{}:{}", self.host, self.port),
        }
    }
}

impl std::str::FromStr for TargetAddr {
    type Err = String;

    fn from_str(s: &str) -> Result<Self, Self::Err> {
        if let Some(rest) = s.strip_prefix('[') {
            let close = rest
                .find(']')
                .ok_or_else(|| format!("invalid target format: missing closing ']' in '{s}'"))?;
            let host_str = &rest[..close];
            let after = &rest[close + 1..];
            let port_str = after.strip_prefix(':').ok_or_else(|| {
                format!("invalid target format: missing ':port' after ']' in '{s}'")
            })?;
            let port: u16 = port_str
                .parse()
                .map_err(|e| format!("invalid port '{port_str}': {e}"))?;
            let ip: IpAddr = host_str
                .parse()
                .map_err(|e| format!("invalid IPv6 address '{host_str}': {e}"))?;
            Ok(TargetAddr {
                host: TargetHost::Ip(ip),
                port,
            })
        } else if s.matches(':').count() > 1 {
            Err(format!(
                "invalid target format: unbracketed IPv6 literal in '{s}' (use [addr]:port)"
            ))
        } else if let Some(idx) = s.rfind(':') {
            let host_part = &s[..idx];
            let port_part = &s[idx + 1..];
            let port: u16 = port_part
                .parse()
                .map_err(|e| format!("invalid port '{port_part}': {e}"))?;
            let host = if let Ok(ip) = host_part.parse::<IpAddr>() {
                TargetHost::Ip(ip)
            } else {
                TargetHost::Domain(host_part.to_string())
            };
            Ok(TargetAddr { host, port })
        } else {
            Err(format!("invalid target format: {s}"))
        }
    }
}

/// Client identity information.
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ClientIdentity {
    Anonymous,
    Username(String),
    Opaque(String),
}

/// Context for a proxy session.
#[derive(Debug, Clone)]
pub struct SessionContext {
    pub session_id: u64,
    pub client_identity: ClientIdentity,
    pub target_addr: TargetAddr,
}

/// Action to take for a routed connection.
#[derive(Debug, Clone)]
pub enum RouteAction {
    Direct,
    Upstream(UpstreamId),
    Reject(RejectReason),
}

/// Reason for rejecting a connection.
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RejectReason {
    UnsupportedProtocol,
    AuthRequired,
    AccessDenied,
    Blocked,
    InternalError,
}

impl fmt::Display for RejectReason {
    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
        match self {
            RejectReason::UnsupportedProtocol => write!(f, "unsupported protocol"),
            RejectReason::AuthRequired => write!(f, "authentication required"),
            RejectReason::AccessDenied => write!(f, "access denied"),
            RejectReason::Blocked => write!(f, "target address blocked"),
            RejectReason::InternalError => write!(f, "internal error"),
        }
    }
}

/// A trait that combines AsyncRead and AsyncWrite for bidirectional streams.
pub trait AsyncStream: AsyncRead + AsyncWrite + Send + Unpin {}
impl<T: AsyncRead + AsyncWrite + Send + Unpin> AsyncStream for T {}

/// A type alias for a boxed async stream.
pub type BoxStream = Box<dyn AsyncStream>;

/// Error types for connection operations.
#[derive(Debug, thiserror::Error)]
pub enum ConnectError {
    #[error("connection refused")]
    ConnectionRefused,
    #[error("connection timed out")]
    Timeout,
    #[error("DNS resolution failed: {0}")]
    DnsResolution(String),
    #[error("TLS handshake failed: {0}")]
    TlsHandshake(String),
    #[error("reserved or private target IP: {0}")]
    ReservedTarget(std::net::IpAddr),
    #[error("IO error: {0}")]
    Io(#[from] std::io::Error),
}

/// Error types for protocol operations.
#[derive(Debug, thiserror::Error)]
pub enum ProtocolError {
    #[error("malformed message")]
    MalformedMessage,
    #[error("unsupported version")]
    UnsupportedVersion,
    #[error("method not supported")]
    MethodNotSupported,
    #[error("address type not supported")]
    AddressTypeNotSupported,
    #[error("IO error: {0}")]
    Io(#[from] std::io::Error),
}

/// Error types for authentication operations.
#[derive(Debug, thiserror::Error)]
pub enum AuthError {
    #[error("invalid credentials")]
    InvalidCredentials,
    #[error("authentication method not supported")]
    MethodNotSupported,
    #[error("authentication required")]
    Required,
    #[error("IO error: {0}")]
    Io(#[from] std::io::Error),
}

/// Error types for relay operations.
#[derive(Debug, thiserror::Error)]
pub enum RelayError {
    #[error("connection closed")]
    ConnectionClosed,
    #[error("IO error: {0}")]
    Io(#[from] std::io::Error),
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn test_target_host_display() {
        let ip_host = TargetHost::Ip("127.0.0.1".parse().unwrap());
        assert_eq!(ip_host.to_string(), "127.0.0.1");

        let domain_host = TargetHost::Domain("example.com".to_string());
        assert_eq!(domain_host.to_string(), "example.com");
    }

    #[test]
    fn test_target_addr_display() {
        let addr = TargetAddr {
            host: TargetHost::Domain("example.com".to_string()),
            port: 8080,
        };
        assert_eq!(addr.to_string(), "example.com:8080");
    }

    #[test]
    fn test_reject_reason_display() {
        assert_eq!(
            RejectReason::UnsupportedProtocol.to_string(),
            "unsupported protocol"
        );
        assert_eq!(
            RejectReason::AuthRequired.to_string(),
            "authentication required"
        );
    }

    #[test]
    fn test_target_addr_from_str_bracketed_ipv6() {
        let addr: TargetAddr = "[::1]:443".parse().unwrap();
        assert_eq!(addr.host, TargetHost::Ip("::1".parse::<IpAddr>().unwrap()));
        assert_eq!(addr.port, 443);
    }

    #[test]
    fn test_target_addr_from_str_full_ipv6() {
        let addr: TargetAddr = "[2001:db8::1]:80".parse().unwrap();
        assert_eq!(
            addr.host,
            TargetHost::Ip("2001:db8::1".parse::<IpAddr>().unwrap())
        );
        assert_eq!(addr.port, 80);
    }

    #[test]
    fn test_target_addr_from_str_rejects_unbracketed_ipv6() {
        let err = "::1:443".parse::<TargetAddr>().unwrap_err();
        assert!(err.contains("unbracketed IPv6"));
    }

    #[test]
    fn test_target_addr_from_str_rejects_unclosed_bracket() {
        let err = "[::1:443".parse::<TargetAddr>().unwrap_err();
        assert!(err.contains("closing ']'"));
    }

    #[test]
    fn test_target_addr_display_brackets_ipv6() {
        let addr = TargetAddr {
            host: TargetHost::Ip("::1".parse().unwrap()),
            port: 443,
        };
        assert_eq!(addr.to_string(), "[::1]:443");
    }

    #[test]
    fn test_target_addr_display_does_not_bracket_ipv4() {
        let addr = TargetAddr {
            host: TargetHost::Ip("127.0.0.1".parse().unwrap()),
            port: 80,
        };
        assert_eq!(addr.to_string(), "127.0.0.1:80");
    }
}