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;
#[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"),
}
}
}
pub type ListenerId = u64;
#[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))
}
}
#[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),
}
}
}
#[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}"))
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ClientIdentity {
Anonymous,
Username(String),
Opaque(String),
}
#[derive(Debug, Clone)]
pub struct SessionContext {
pub session_id: u64,
pub client_identity: ClientIdentity,
pub target_addr: TargetAddr,
}
#[derive(Debug, Clone)]
pub enum RouteAction {
Direct,
Upstream(UpstreamId),
Reject(RejectReason),
}
#[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"),
}
}
}
pub trait AsyncStream: AsyncRead + AsyncWrite + Send + Unpin {}
impl<T: AsyncRead + AsyncWrite + Send + Unpin> AsyncStream for T {}
pub type BoxStream = Box<dyn AsyncStream>;
#[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),
}
#[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),
}
#[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),
}
#[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");
}
}