Co-authored-by: factory-droid[bot] <138933559+factory-droid[bot]@users.noreply.github.com>
3527 lines
128 KiB
Rust
3527 lines
128 KiB
Rust
/// HTTPS mTLS tunnel with application-level authentication and SOCKS5 forwarding.
|
|
///
|
|
/// After the mTLS + auth handshake completes, both sides switch from HTTP/1.1
|
|
/// to a binary framing protocol that carries multiple bidirectional streams
|
|
/// over a single persistent authenticated tunnel.
|
|
///
|
|
/// Protocol flow:
|
|
/// 1. Client connects over HTTPS with mTLS
|
|
/// 2. Client sends HTTP POST /tunnel with auth token
|
|
/// 3. Server validates and responds 200 OK
|
|
/// 4. Both sides switch to binary framing on the same connection
|
|
/// 5. Client sends CONNECT frames; server opens target TCP connections
|
|
/// 6. DATA frames flow bidirectionally, keyed by stream ID
|
|
/// 7. CLOSE/ERROR frames propagate deterministically
|
|
///
|
|
/// Security gates (all must pass before framing):
|
|
/// 1. HTTPS connection (TLS acceptor/connector)
|
|
/// 2. Client certificate validation (server side, mTLS)
|
|
/// 3. Server certificate validation (client side)
|
|
/// 4. Application auth token validation
|
|
use std::collections::HashMap;
|
|
use std::net::SocketAddr;
|
|
use std::path::Path;
|
|
use std::sync::Arc;
|
|
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
|
|
|
|
use bytes::Bytes;
|
|
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
|
use tokio::net::{TcpListener, TcpStream};
|
|
use tokio::sync::{mpsc, oneshot, watch};
|
|
|
|
use crate::errors::{AuthError, HostError, TunnelError};
|
|
use crate::framing::{
|
|
CONNECT_FAILED, CONNECT_OK, FRAME_CLOSE, FRAME_CONNECT, FRAME_CONNECT_REPLY, FRAME_DATA,
|
|
FRAME_ERROR, Frame, FrameReader, FrameWriter, decode_target, encode_target,
|
|
};
|
|
use crate::redact::Redacted;
|
|
use crate::socks5;
|
|
use crate::tls;
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// Auth token handling
|
|
// ---------------------------------------------------------------------------
|
|
|
|
/// Load an auth token from a file or use the provided value.
|
|
pub fn load_auth_token(
|
|
token_value: Option<&str>,
|
|
token_path: Option<&Path>,
|
|
) -> Result<String, AuthError> {
|
|
if let Some(path) = token_path {
|
|
let content = std::fs::read_to_string(path).map_err(|_| AuthError::Missing)?;
|
|
let trimmed = content.trim().to_string();
|
|
if trimmed.is_empty() {
|
|
return Err(AuthError::Missing);
|
|
}
|
|
Ok(trimmed)
|
|
} else if let Some(t) = token_value {
|
|
if t.is_empty() {
|
|
return Err(AuthError::Missing);
|
|
}
|
|
Ok(t.to_string())
|
|
} else {
|
|
Err(AuthError::Missing)
|
|
}
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// Constant-time comparison
|
|
// ---------------------------------------------------------------------------
|
|
|
|
fn constant_time_compare(a: &str, b: &str) -> bool {
|
|
let a_bytes = a.as_bytes();
|
|
let b_bytes = b.as_bytes();
|
|
if a_bytes.len() != b_bytes.len() {
|
|
return false;
|
|
}
|
|
let mut result: u8 = 0;
|
|
for (x, y) in a_bytes.iter().zip(b_bytes.iter()) {
|
|
result |= x ^ y;
|
|
}
|
|
result == 0
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// Hostname resolution
|
|
// ---------------------------------------------------------------------------
|
|
|
|
/// Resolve a hostname to a list of SocketAddrs.
|
|
pub async fn resolve_host(host: &str, port: u16) -> Result<Vec<SocketAddr>, HostError> {
|
|
if host.is_empty() {
|
|
return Err(HostError::InvalidHostname(host.to_string()));
|
|
}
|
|
let lookup_addr = format!("{}:{}", host, port);
|
|
let addrs: Vec<SocketAddr> = tokio::net::lookup_host(&lookup_addr)
|
|
.await
|
|
.map_err(|e| HostError::DnsResolution(host.to_string(), e.to_string()))?
|
|
.collect();
|
|
if addrs.is_empty() {
|
|
return Err(HostError::NoAddresses(host.to_string()));
|
|
}
|
|
Ok(addrs)
|
|
}
|
|
|
|
/// Resolve a host:port to a single SocketAddr.
|
|
pub async fn resolve_target(host: &str, port: u16) -> Result<SocketAddr, HostError> {
|
|
let addrs = resolve_host(host, port).await?;
|
|
Ok(addrs[0])
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// Configs
|
|
// ---------------------------------------------------------------------------
|
|
|
|
/// Configuration for the tunnel listener.
|
|
#[derive(Debug, Clone)]
|
|
pub struct ListenerConfig {
|
|
pub bind_addr: SocketAddr,
|
|
pub socks_addr: Option<SocketAddr>,
|
|
pub tls: ServerTlsMaterial,
|
|
pub auth_token: Arc<String>,
|
|
pub insecure_skip_tls_verify: bool,
|
|
}
|
|
|
|
/// Configuration for the tunnel connector.
|
|
#[derive(Debug, Clone)]
|
|
pub struct ConnectorConfig {
|
|
pub target_host: String,
|
|
pub target_port: u16,
|
|
pub tls: ClientTlsMaterial,
|
|
pub auth_token: Arc<String>,
|
|
pub insecure_skip_tls_verify: bool,
|
|
}
|
|
|
|
#[derive(Debug, Clone)]
|
|
pub enum ServerTlsMaterial {
|
|
Paths {
|
|
cert: Arc<std::path::PathBuf>,
|
|
key: Arc<std::path::PathBuf>,
|
|
ca_cert: Arc<std::path::PathBuf>,
|
|
},
|
|
Pem {
|
|
cert: Arc<String>,
|
|
key: Arc<String>,
|
|
ca_cert: Arc<String>,
|
|
},
|
|
}
|
|
|
|
#[derive(Debug, Clone)]
|
|
pub enum ClientTlsMaterial {
|
|
Paths {
|
|
cert: Arc<std::path::PathBuf>,
|
|
key: Arc<std::path::PathBuf>,
|
|
ca_cert: Arc<std::path::PathBuf>,
|
|
},
|
|
Pem {
|
|
cert: Arc<String>,
|
|
key: Arc<String>,
|
|
ca_cert: Arc<String>,
|
|
},
|
|
}
|
|
|
|
impl ServerTlsMaterial {
|
|
pub fn from_paths(
|
|
cert: impl Into<std::path::PathBuf>,
|
|
key: impl Into<std::path::PathBuf>,
|
|
ca_cert: impl Into<std::path::PathBuf>,
|
|
) -> Self {
|
|
Self::Paths {
|
|
cert: Arc::new(cert.into()),
|
|
key: Arc::new(key.into()),
|
|
ca_cert: Arc::new(ca_cert.into()),
|
|
}
|
|
}
|
|
|
|
fn build_config(
|
|
&self,
|
|
insecure: bool,
|
|
) -> Result<Arc<rustls::ServerConfig>, crate::errors::TlsError> {
|
|
if insecure {
|
|
match self {
|
|
Self::Paths { cert, key, .. } => {
|
|
tls::build_server_config_insecure(cert.as_ref(), key.as_ref())
|
|
}
|
|
Self::Pem { cert, key, .. } => {
|
|
tls::build_server_config_insecure_from_pem(cert, key)
|
|
}
|
|
}
|
|
} else {
|
|
match self {
|
|
Self::Paths { cert, key, ca_cert } => {
|
|
tls::build_server_config(cert.as_ref(), key.as_ref(), ca_cert.as_ref())
|
|
}
|
|
Self::Pem { cert, key, ca_cert } => {
|
|
tls::build_server_config_from_pem(cert, key, ca_cert)
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
fn server_certs(
|
|
&self,
|
|
) -> Result<Vec<rustls::pki_types::CertificateDer<'static>>, crate::errors::TlsError> {
|
|
match self {
|
|
Self::Paths { cert, .. } => tls::load_certs(cert.as_ref()),
|
|
Self::Pem { cert, .. } => {
|
|
tls::load_certs_from_pem(cert.as_bytes(), "server certificate")
|
|
}
|
|
}
|
|
}
|
|
|
|
fn ca_certs(
|
|
&self,
|
|
) -> Result<Vec<rustls::pki_types::CertificateDer<'static>>, crate::errors::TlsError> {
|
|
match self {
|
|
Self::Paths { ca_cert, .. } => tls::load_certs(ca_cert.as_ref()),
|
|
Self::Pem { ca_cert, .. } => {
|
|
tls::load_certs_from_pem(ca_cert.as_bytes(), "CA certificate")
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
impl ClientTlsMaterial {
|
|
pub fn from_paths(
|
|
cert: impl Into<std::path::PathBuf>,
|
|
key: impl Into<std::path::PathBuf>,
|
|
ca_cert: impl Into<std::path::PathBuf>,
|
|
) -> Self {
|
|
Self::Paths {
|
|
cert: Arc::new(cert.into()),
|
|
key: Arc::new(key.into()),
|
|
ca_cert: Arc::new(ca_cert.into()),
|
|
}
|
|
}
|
|
|
|
fn build_config(
|
|
&self,
|
|
insecure: bool,
|
|
) -> Result<Arc<rustls::ClientConfig>, crate::errors::TlsError> {
|
|
if insecure {
|
|
match self {
|
|
Self::Paths { cert, key, .. } => {
|
|
tls::build_client_config_insecure(cert.as_ref(), key.as_ref())
|
|
}
|
|
Self::Pem { cert, key, .. } => {
|
|
tls::build_client_config_insecure_from_pem(cert, key)
|
|
}
|
|
}
|
|
} else {
|
|
match self {
|
|
Self::Paths { cert, key, ca_cert } => {
|
|
tls::build_client_config(cert.as_ref(), key.as_ref(), ca_cert.as_ref())
|
|
}
|
|
Self::Pem { cert, key, ca_cert } => {
|
|
tls::build_client_config_from_pem(cert, key, ca_cert)
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
fn client_certs(
|
|
&self,
|
|
) -> Result<Vec<rustls::pki_types::CertificateDer<'static>>, crate::errors::TlsError> {
|
|
match self {
|
|
Self::Paths { cert, .. } => tls::load_certs(cert.as_ref()),
|
|
Self::Pem { cert, .. } => {
|
|
tls::load_certs_from_pem(cert.as_bytes(), "client certificate")
|
|
}
|
|
}
|
|
}
|
|
|
|
fn ca_certs(
|
|
&self,
|
|
) -> Result<Vec<rustls::pki_types::CertificateDer<'static>>, crate::errors::TlsError> {
|
|
match self {
|
|
Self::Paths { ca_cert, .. } => tls::load_certs(ca_cert.as_ref()),
|
|
Self::Pem { ca_cert, .. } => {
|
|
tls::load_certs_from_pem(ca_cert.as_bytes(), "CA certificate")
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// HTTP helpers — manual parsing for the auth handshake
|
|
// ---------------------------------------------------------------------------
|
|
|
|
/// Parse an HTTP/1.1 request from the stream.
|
|
/// Returns (method, path, headers_map).
|
|
async fn parse_http_request<R>(
|
|
reader: &mut R,
|
|
) -> Result<(String, String, HashMap<String, String>), TunnelError>
|
|
where
|
|
R: AsyncReadExt + Unpin,
|
|
{
|
|
let mut buf = Vec::with_capacity(1024);
|
|
let mut crlf_count = 0;
|
|
|
|
loop {
|
|
let mut byte = [0u8; 1];
|
|
reader.read_exact(&mut byte).await?;
|
|
buf.push(byte[0]);
|
|
if byte[0] == b'\r' || byte[0] == b'\n' {
|
|
crlf_count += 1;
|
|
if crlf_count >= 4 && buf.ends_with(b"\r\n\r\n") {
|
|
break;
|
|
}
|
|
} else {
|
|
crlf_count = 0;
|
|
}
|
|
// Safety valve: prevent reading forever for malformed requests
|
|
if buf.len() > 8192 {
|
|
return Err(TunnelError::Protocol("request too large".into()));
|
|
}
|
|
}
|
|
|
|
let text = String::from_utf8_lossy(&buf);
|
|
let lines: Vec<&str> = text.lines().collect();
|
|
if lines.is_empty() {
|
|
return Err(TunnelError::Protocol("empty request".into()));
|
|
}
|
|
|
|
let parts: Vec<&str> = lines[0].split_whitespace().collect();
|
|
if parts.len() < 2 {
|
|
return Err(TunnelError::Protocol("invalid request line".into()));
|
|
}
|
|
let method = parts[0].to_string();
|
|
let path = parts[1].to_string();
|
|
|
|
let mut headers = HashMap::new();
|
|
for line in lines.iter().skip(1) {
|
|
if let Some((key, value)) = line.split_once(':') {
|
|
headers.insert(key.trim().to_lowercase(), value.trim().to_string());
|
|
}
|
|
}
|
|
|
|
Ok((method, path, headers))
|
|
}
|
|
|
|
/// Send a minimal HTTP/1.1 response (no body, Content-Length: 0).
|
|
async fn send_http_response<W>(
|
|
writer: &mut W,
|
|
status: u16,
|
|
status_text: &str,
|
|
) -> Result<(), TunnelError>
|
|
where
|
|
W: AsyncWriteExt + Unpin,
|
|
{
|
|
let response = format!(
|
|
"HTTP/1.1 {} {}\r\nContent-Length: 0\r\n\r\n",
|
|
status, status_text
|
|
);
|
|
writer.write_all(response.as_bytes()).await?;
|
|
writer.flush().await?;
|
|
Ok(())
|
|
}
|
|
|
|
/// Read an HTTP/1.1 response status line. Returns the status code.
|
|
async fn read_http_response_status<R>(reader: &mut R) -> Result<u16, TunnelError>
|
|
where
|
|
R: AsyncReadExt + Unpin,
|
|
{
|
|
let mut buf = Vec::with_capacity(256);
|
|
loop {
|
|
let mut byte = [0u8; 1];
|
|
reader.read_exact(&mut byte).await?;
|
|
buf.push(byte[0]);
|
|
if buf.ends_with(b"\r\n\r\n") {
|
|
break;
|
|
}
|
|
if buf.len() > 4096 {
|
|
return Err(TunnelError::Protocol("response headers too large".into()));
|
|
}
|
|
}
|
|
|
|
let text = String::from_utf8_lossy(&buf);
|
|
let first_line = text
|
|
.lines()
|
|
.next()
|
|
.ok_or_else(|| TunnelError::Protocol("empty response".into()))?;
|
|
let parts: Vec<&str> = first_line.split_whitespace().collect();
|
|
if parts.len() < 2 {
|
|
return Err(TunnelError::Protocol("invalid response line".into()));
|
|
}
|
|
parts[1]
|
|
.parse::<u16>()
|
|
.map_err(|_| TunnelError::Protocol("invalid status code".into()))
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// Listener
|
|
// ---------------------------------------------------------------------------
|
|
|
|
/// Start the HTTPS tunnel listener.
|
|
pub async fn run_listener(config: ListenerConfig) -> Result<(), TunnelError> {
|
|
let addr = config.bind_addr;
|
|
let tls_material = config.tls.clone();
|
|
let auth_token = config.auth_token.clone();
|
|
|
|
let tls_config = tls_material
|
|
.build_config(config.insecure_skip_tls_verify)
|
|
.map_err(|e| {
|
|
tracing::error!("Failed to build TLS config: {}", e);
|
|
e
|
|
})?;
|
|
|
|
let server_certs = tls_material.server_certs().map_err(TunnelError::Tls)?;
|
|
let ca_certs = tls_material.ca_certs().map_err(TunnelError::Tls)?;
|
|
let bind_host = addr.ip().to_string();
|
|
if !addr.ip().is_unspecified() {
|
|
tls::validate_cert_identity(&server_certs[0], &bind_host).map_err(|e| {
|
|
TunnelError::Tls(crate::errors::TlsError::IdentityMismatch(format!(
|
|
"server cert does not match bind address {}: {}",
|
|
bind_host, e
|
|
)))
|
|
})?;
|
|
}
|
|
|
|
let tls_acceptor = tokio_rustls::TlsAcceptor::from(tls_config);
|
|
let listener = TcpListener::bind(addr).await.map_err(|e| {
|
|
if e.kind() == std::io::ErrorKind::AddrInUse {
|
|
TunnelError::PortInUse(addr.port())
|
|
} else {
|
|
TunnelError::BindError(addr.to_string(), e.to_string())
|
|
}
|
|
})?;
|
|
|
|
tracing::info!(
|
|
"HTTPS tunnel listener bound to {} — endpoint: /tunnel (HTTPS default transport)",
|
|
addr
|
|
);
|
|
tracing::info!(
|
|
"Effective config: listener={}, auth_token={}",
|
|
addr,
|
|
Redacted::new(&*auth_token)
|
|
);
|
|
tracing::info!(
|
|
"TLS material fingerprints: ca={}, server={}",
|
|
tls::cert_fingerprint_sha256(&ca_certs[0]),
|
|
tls::cert_fingerprint_sha256(&server_certs[0])
|
|
);
|
|
|
|
let (mux_sender, mux_receiver) = watch::channel::<Option<Arc<StreamMux>>>(None);
|
|
let socks_enabled = config.socks_addr.is_some();
|
|
let socks_handle = if let Some(socks_addr) = config.socks_addr {
|
|
tracing::info!(
|
|
"Server-side SOCKS5 proxy configured on {} for connector-side network access",
|
|
socks_addr
|
|
);
|
|
Some(tokio::spawn(run_socks5_proxy(socks_addr, mux_receiver)))
|
|
} else {
|
|
None
|
|
};
|
|
|
|
// Listen for shutdown signal
|
|
let shutdown = crate::signal::shutdown_signal();
|
|
tokio::pin!(shutdown);
|
|
|
|
loop {
|
|
tokio::select! {
|
|
_ = &mut shutdown => {
|
|
tracing::info!("Shutdown signal received — stopping HTTPS tunnel listener");
|
|
break;
|
|
}
|
|
result = listener.accept() => {
|
|
match result {
|
|
Ok((stream, peer_addr)) => {
|
|
tracing::info!("New connection from {}", peer_addr);
|
|
let acceptor = tls_acceptor.clone();
|
|
let token = auth_token.clone();
|
|
let session_mux_sender = if socks_enabled {
|
|
Some(mux_sender.clone())
|
|
} else {
|
|
None
|
|
};
|
|
|
|
let insecure = config.insecure_skip_tls_verify;
|
|
tokio::spawn(async move {
|
|
if let Err(e) = handle_listener_connection(stream, peer_addr, acceptor, token, session_mux_sender, insecure).await {
|
|
tracing::warn!("Connection from {} failed: {}", peer_addr, e);
|
|
}
|
|
});
|
|
}
|
|
Err(e) => {
|
|
tracing::error!("Accept error: {}", e);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
if let Some(handle) = socks_handle {
|
|
handle.abort();
|
|
}
|
|
|
|
tracing::info!(
|
|
"HTTPS tunnel listener shutting down — port {} released",
|
|
addr.port()
|
|
);
|
|
Ok(())
|
|
}
|
|
|
|
/// Handle a single incoming connection: TLS → mTLS → Auth → Framing.
|
|
async fn handle_listener_connection(
|
|
stream: TcpStream,
|
|
peer_addr: SocketAddr,
|
|
acceptor: tokio_rustls::TlsAcceptor,
|
|
auth_token: Arc<String>,
|
|
mux_sender: Option<watch::Sender<Option<Arc<StreamMux>>>>,
|
|
insecure_skip_tls_verify: bool,
|
|
) -> Result<(), TunnelError> {
|
|
// TLS accept
|
|
let tls_stream = acceptor.accept(stream).await.map_err(|e| {
|
|
tracing::warn!("TLS accept failed for {}: {}", peer_addr, e);
|
|
TunnelError::Tls(crate::errors::TlsError::VerificationFailed(e.to_string()))
|
|
})?;
|
|
|
|
if insecure_skip_tls_verify {
|
|
tracing::info!(
|
|
"TLS connection accepted from {} (insecure mode — skipping client certificate verification)",
|
|
peer_addr
|
|
);
|
|
} else {
|
|
// Verify mTLS — client must present a valid certificate
|
|
let (_, server_conn) = tls_stream.get_ref();
|
|
let peer_certs = server_conn.peer_certificates().ok_or_else(|| {
|
|
tracing::warn!("No client certificate from {}", peer_addr);
|
|
TunnelError::Tls(crate::errors::TlsError::Missing("client certificate"))
|
|
})?;
|
|
if peer_certs.is_empty() {
|
|
return Err(TunnelError::Tls(crate::errors::TlsError::Missing(
|
|
"client certificate",
|
|
)));
|
|
}
|
|
tracing::info!("mTLS established with {}", peer_addr);
|
|
}
|
|
|
|
// Pin the stream so we can use it with io::split and generic async functions
|
|
let mut tls_stream: std::pin::Pin<Box<tokio_rustls::server::TlsStream<TcpStream>>> =
|
|
Box::pin(tls_stream);
|
|
|
|
// Phase 1: HTTP auth handshake
|
|
let (method, path, headers) = parse_http_request(&mut *tls_stream).await?;
|
|
|
|
if method == "POST" && path == "/tunnel" {
|
|
let received_token = headers
|
|
.get("x-rustunnel-token")
|
|
.map(|s| s.as_str())
|
|
.unwrap_or("");
|
|
|
|
if constant_time_compare(received_token, &auth_token) {
|
|
tracing::info!("Tunnel established with {} via HTTPS /tunnel", peer_addr);
|
|
send_http_response(&mut *tls_stream, 200, "OK").await?;
|
|
} else {
|
|
tracing::warn!("Auth rejected for {} — invalid token on /tunnel", peer_addr);
|
|
send_http_response(&mut *tls_stream, 401, "Unauthorized").await?;
|
|
return Ok(()); // close connection on auth failure
|
|
}
|
|
} else if method == "GET" && path == "/tunnel" {
|
|
// Health check — respond and close
|
|
send_http_response(&mut *tls_stream, 200, "OK").await?;
|
|
return Ok(());
|
|
} else {
|
|
tracing::debug!(
|
|
"Rejected {} {} from {} — only POST /tunnel accepted",
|
|
method,
|
|
path,
|
|
peer_addr
|
|
);
|
|
send_http_response(&mut *tls_stream, 404, "Not Found").await?;
|
|
return Ok(());
|
|
}
|
|
|
|
// Phase 2: Switch to binary framing on the same connection
|
|
tracing::info!(
|
|
"Upgraded connection from {} to binary framing protocol",
|
|
peer_addr
|
|
);
|
|
handle_framing_server(tls_stream, peer_addr, mux_sender).await
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// Framing server (listener side)
|
|
// ---------------------------------------------------------------------------
|
|
|
|
/// Entry for an active stream on the listener side.
|
|
struct ServerStreamEntry {
|
|
/// Write half of the target TCP connection
|
|
target_write: tokio::net::tcp::OwnedWriteHalf,
|
|
}
|
|
|
|
/// Handle the binary framing protocol on the server side.
|
|
async fn handle_framing_server<S>(
|
|
stream: S,
|
|
peer_addr: SocketAddr,
|
|
mux_sender: Option<watch::Sender<Option<Arc<StreamMux>>>>,
|
|
) -> Result<(), TunnelError>
|
|
where
|
|
S: AsyncReadExt + AsyncWriteExt + Unpin + Send + 'static,
|
|
{
|
|
let mut stream = stream;
|
|
let mut reader = FrameReader::new();
|
|
let (outbound_tx, mut outbound_rx) = mpsc::channel::<Frame>(256);
|
|
let mux = mux_sender
|
|
.as_ref()
|
|
.map(|_| Arc::new(StreamMux::new(outbound_tx.clone())));
|
|
|
|
// Active streams: stream_id → target write half
|
|
let mut streams: HashMap<u64, ServerStreamEntry> = HashMap::new();
|
|
|
|
tracing::info!("Framing session started with {}", peer_addr);
|
|
if let (Some(sender), Some(mux)) = (&mux_sender, &mux) {
|
|
let _ = sender.send(Some(mux.clone()));
|
|
}
|
|
|
|
loop {
|
|
tokio::select! {
|
|
// Write outbound frames (CONNECT_REPLY, DATA from target→tunnel, CLOSE, ERROR)
|
|
Some(frame) = outbound_rx.recv() => {
|
|
if let Err(e) = FrameWriter::write_frame(&mut stream, &frame).await {
|
|
tracing::debug!("Write error for stream {}: {}", frame.stream_id, e);
|
|
break;
|
|
}
|
|
}
|
|
|
|
// Read inbound frames from the tunnel (CONNECT, DATA from client, CLOSE)
|
|
result = reader.read_frame(&mut stream) => {
|
|
match result {
|
|
Ok(Some(frame)) => {
|
|
let result = if matches!(
|
|
frame.frame_type,
|
|
FRAME_CONNECT_REPLY | FRAME_DATA | FRAME_CLOSE | FRAME_ERROR
|
|
) {
|
|
if let Some(mux) = &mux {
|
|
mux.dispatch_frame(&frame).await;
|
|
Ok(())
|
|
} else {
|
|
handle_inbound_frame(
|
|
&frame,
|
|
&mut stream,
|
|
&outbound_tx,
|
|
&mut streams,
|
|
peer_addr,
|
|
).await
|
|
}
|
|
} else {
|
|
handle_inbound_frame(
|
|
&frame,
|
|
&mut stream,
|
|
&outbound_tx,
|
|
&mut streams,
|
|
peer_addr,
|
|
).await
|
|
};
|
|
if let Err(e) = result {
|
|
tracing::warn!(
|
|
"Stream {} inbound error: {}",
|
|
frame.stream_id,
|
|
e
|
|
);
|
|
// Send error frame and continue
|
|
let _ = outbound_tx
|
|
.send(Frame::error(frame.stream_id, e.to_string()))
|
|
.await;
|
|
}
|
|
}
|
|
Ok(None) => {
|
|
tracing::info!("Framing session with {} closed (EOF)", peer_addr);
|
|
break;
|
|
}
|
|
Err(e) => {
|
|
tracing::warn!("Framing read error from {}: {}", peer_addr, e);
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
// Clean up all streams
|
|
for (sid, mut entry) in streams.drain() {
|
|
let _ = entry.target_write.shutdown().await;
|
|
tracing::debug!("Closed stream {} on shutdown", sid);
|
|
}
|
|
if let Some(sender) = mux_sender {
|
|
let _ = sender.send(None);
|
|
}
|
|
|
|
tracing::info!("Framing session with {} ended", peer_addr);
|
|
Ok(())
|
|
}
|
|
|
|
/// Handle an inbound frame on the server side.
|
|
async fn handle_inbound_frame<S>(
|
|
frame: &Frame,
|
|
_stream: &mut S,
|
|
outbound_tx: &mpsc::Sender<Frame>,
|
|
streams: &mut HashMap<u64, ServerStreamEntry>,
|
|
peer_addr: SocketAddr,
|
|
) -> Result<(), TunnelError>
|
|
where
|
|
S: AsyncWriteExt + Unpin,
|
|
{
|
|
match frame.frame_type {
|
|
FRAME_CONNECT => {
|
|
if streams.contains_key(&frame.stream_id) {
|
|
let _ = outbound_tx
|
|
.send(Frame::error(frame.stream_id, "duplicate stream ID"))
|
|
.await;
|
|
return Ok(());
|
|
}
|
|
|
|
let target = decode_target(&frame.payload)
|
|
.map_err(|e| TunnelError::Protocol(format!("decode target: {}", e)))?;
|
|
|
|
let target_str = socks5::target_to_string(&target);
|
|
tracing::info!(
|
|
"Stream {} CONNECT from {} to {} (listener-side target)",
|
|
frame.stream_id,
|
|
peer_addr,
|
|
target_str
|
|
);
|
|
|
|
// Resolve and connect to target from the listener side
|
|
let target_addr = match socks5::resolve_target(&target).await {
|
|
Ok(addr) => addr,
|
|
Err(e) => {
|
|
tracing::warn!(
|
|
"Stream {} resolve failed for {}: {}",
|
|
frame.stream_id,
|
|
target_str,
|
|
e
|
|
);
|
|
let _ = outbound_tx
|
|
.send(Frame::connect_reply(frame.stream_id, CONNECT_FAILED))
|
|
.await;
|
|
let _ = outbound_tx
|
|
.send(Frame::error(frame.stream_id, format!("resolve: {}", e)))
|
|
.await;
|
|
return Ok(());
|
|
}
|
|
};
|
|
|
|
match TcpStream::connect(target_addr).await {
|
|
Ok(target_stream) => {
|
|
tracing::info!(
|
|
"Stream {} connected to target {} (listener side)",
|
|
frame.stream_id,
|
|
target_addr
|
|
);
|
|
|
|
let (mut target_read, target_write) = target_stream.into_split();
|
|
let sid = frame.stream_id;
|
|
let outbound = outbound_tx.clone();
|
|
|
|
// Spawn target → tunnel direction
|
|
tokio::spawn(async move {
|
|
let mut buf = [0u8; 32 * 1024];
|
|
loop {
|
|
match target_read.read(&mut buf).await {
|
|
Ok(0) => break, // Target closed (EOF)
|
|
Ok(n) => {
|
|
if outbound.send(Frame::data(sid, &buf[..n])).await.is_err() {
|
|
break;
|
|
}
|
|
}
|
|
Err(e) => {
|
|
tracing::debug!("Stream {} target read error: {}", sid, e);
|
|
let _ = outbound.send(Frame::error(sid, e.to_string())).await;
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
tracing::info!("Stream {} target → tunnel EOF", sid);
|
|
let _ = outbound.send(Frame::close(sid)).await;
|
|
});
|
|
|
|
streams.insert(frame.stream_id, ServerStreamEntry { target_write });
|
|
let _ = outbound_tx
|
|
.send(Frame::connect_reply(frame.stream_id, CONNECT_OK))
|
|
.await;
|
|
}
|
|
Err(e) => {
|
|
let err_msg = if e.kind() == std::io::ErrorKind::ConnectionRefused {
|
|
format!("connection refused to {}", target_addr)
|
|
} else {
|
|
format!("connect to {}: {}", target_addr, e)
|
|
};
|
|
tracing::warn!(
|
|
"Stream {} connect to {} failed: {}",
|
|
frame.stream_id,
|
|
target_addr,
|
|
err_msg
|
|
);
|
|
let _ = outbound_tx
|
|
.send(Frame::connect_reply(frame.stream_id, CONNECT_FAILED))
|
|
.await;
|
|
let _ = outbound_tx
|
|
.send(Frame::error(frame.stream_id, err_msg))
|
|
.await;
|
|
}
|
|
}
|
|
}
|
|
|
|
FRAME_DATA => {
|
|
// Forward data to the target
|
|
if let Some(entry) = streams.get_mut(&frame.stream_id)
|
|
&& let Err(e) = entry.target_write.write_all(&frame.payload).await
|
|
{
|
|
tracing::debug!("Stream {} target write error: {}", frame.stream_id, e);
|
|
streams.remove(&frame.stream_id);
|
|
}
|
|
}
|
|
|
|
FRAME_CLOSE => {
|
|
if let Some(mut entry) = streams.remove(&frame.stream_id) {
|
|
let _ = entry.target_write.shutdown().await;
|
|
tracing::info!(
|
|
"Stream {} closed by client (sender-side EOF)",
|
|
frame.stream_id
|
|
);
|
|
}
|
|
}
|
|
|
|
FRAME_ERROR => {
|
|
let msg = String::from_utf8_lossy(&frame.payload);
|
|
tracing::warn!("Stream {} error from client: {}", frame.stream_id, msg);
|
|
if let Some(mut entry) = streams.remove(&frame.stream_id) {
|
|
let _ = entry.target_write.shutdown().await;
|
|
}
|
|
}
|
|
|
|
_ => {
|
|
tracing::debug!(
|
|
"Unknown frame type {} on stream {} from {}",
|
|
frame.frame_type,
|
|
frame.stream_id,
|
|
peer_addr
|
|
);
|
|
}
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// Connector: StreamMux for multiplexed streams
|
|
// ---------------------------------------------------------------------------
|
|
|
|
/// Entry for an active stream on the connector side.
|
|
struct ConnectorStreamEntry {
|
|
/// Channel for tunnel → SOCKS5 data (frame → SOCKS5 client)
|
|
data_tx: mpsc::Sender<Bytes>,
|
|
/// One-shot channel for CONNECT_REPLY (Option so we can take it for send)
|
|
reply_tx: Option<oneshot::Sender<u8>>,
|
|
}
|
|
|
|
/// Stream multiplexer for the connector side.
|
|
/// Manages multiple concurrent SOCKS5 streams over one authenticated tunnel.
|
|
pub struct StreamMux {
|
|
/// Channel for all outbound frames to the tunnel
|
|
outbound_tx: mpsc::Sender<Frame>,
|
|
/// Active streams registry
|
|
streams: Arc<tokio::sync::Mutex<HashMap<u64, ConnectorStreamEntry>>>,
|
|
/// Monotonically increasing stream ID counter (odd IDs for client-initiated)
|
|
next_id: Arc<AtomicU64>,
|
|
/// Auth state
|
|
auth_complete: Arc<AtomicBool>,
|
|
}
|
|
|
|
impl StreamMux {
|
|
/// Create a new StreamMux.
|
|
pub fn new(outbound_tx: mpsc::Sender<Frame>) -> Self {
|
|
Self {
|
|
outbound_tx,
|
|
streams: Arc::new(tokio::sync::Mutex::new(HashMap::new())),
|
|
next_id: Arc::new(AtomicU64::new(1)),
|
|
auth_complete: Arc::new(AtomicBool::new(true)),
|
|
}
|
|
}
|
|
|
|
/// Open a new stream to the target.
|
|
/// Returns (stream_id, data_rx) where data_rx receives tunnel → SOCKS5 data.
|
|
pub async fn open_stream(
|
|
&self,
|
|
target: &socks5::Socks5Target,
|
|
) -> Result<(u64, mpsc::Receiver<Bytes>), TunnelError> {
|
|
if !self.auth_complete.load(Ordering::SeqCst) {
|
|
return Err(TunnelError::Auth(AuthError::NotAuthenticated));
|
|
}
|
|
|
|
let stream_id = self.next_id.fetch_add(2, Ordering::Relaxed); // odd IDs
|
|
let (reply_tx, reply_rx) = oneshot::channel();
|
|
let (data_tx, data_rx) = mpsc::channel::<Bytes>(64);
|
|
|
|
self.streams.lock().await.insert(
|
|
stream_id,
|
|
ConnectorStreamEntry {
|
|
data_tx,
|
|
reply_tx: Some(reply_tx),
|
|
},
|
|
);
|
|
|
|
let target_bytes = encode_target(target);
|
|
self.outbound_tx
|
|
.send(Frame::connect(stream_id, target_bytes))
|
|
.await
|
|
.map_err(|_| TunnelError::Protocol("outbound channel closed".into()))?;
|
|
|
|
let status = reply_rx
|
|
.await
|
|
.map_err(|_| TunnelError::Protocol("connect reply channel closed".into()))?;
|
|
|
|
if status != CONNECT_OK {
|
|
self.streams.lock().await.remove(&stream_id);
|
|
return Err(TunnelError::Protocol(format!(
|
|
"CONNECT failed with status {}",
|
|
status
|
|
)));
|
|
}
|
|
|
|
tracing::info!(
|
|
"Stream {} opened to {} via tunnel",
|
|
stream_id,
|
|
socks5::target_to_string(target)
|
|
);
|
|
Ok((stream_id, data_rx))
|
|
}
|
|
|
|
/// Dispatch an inbound frame to the correct stream.
|
|
pub async fn dispatch_frame(&self, frame: &Frame) {
|
|
match frame.frame_type {
|
|
FRAME_CONNECT_REPLY => {
|
|
let mut streams = self.streams.lock().await;
|
|
if let Some(entry) = streams.get_mut(&frame.stream_id) {
|
|
let status = frame.payload.first().copied().unwrap_or(CONNECT_FAILED);
|
|
if let Some(tx) = entry.reply_tx.take() {
|
|
let _ = tx.send(status);
|
|
}
|
|
}
|
|
}
|
|
FRAME_DATA => {
|
|
let mut streams = self.streams.lock().await;
|
|
if let Some(entry) = streams.get(&frame.stream_id)
|
|
&& entry.data_tx.send(frame.payload.clone()).await.is_err()
|
|
{
|
|
// SOCKS5 handler dropped — clean up
|
|
streams.remove(&frame.stream_id);
|
|
let _ = self.outbound_tx.send(Frame::close(frame.stream_id)).await;
|
|
}
|
|
}
|
|
FRAME_CLOSE => {
|
|
let mut streams = self.streams.lock().await;
|
|
if streams.remove(&frame.stream_id).is_some() {
|
|
tracing::info!("Stream {} closed by server", frame.stream_id);
|
|
}
|
|
}
|
|
FRAME_ERROR => {
|
|
let msg = String::from_utf8_lossy(&frame.payload);
|
|
tracing::warn!("Stream {} error from server: {}", frame.stream_id, msg);
|
|
let mut streams = self.streams.lock().await;
|
|
streams.remove(&frame.stream_id);
|
|
}
|
|
_ => {
|
|
tracing::debug!(
|
|
"Unknown frame type {} on stream {}",
|
|
frame.frame_type,
|
|
frame.stream_id
|
|
);
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Close a stream from the connector side.
|
|
pub async fn close_stream(&self, stream_id: u64) {
|
|
self.streams.lock().await.remove(&stream_id);
|
|
let _ = self.outbound_tx.send(Frame::close(stream_id)).await;
|
|
tracing::info!("Stream {} closed from connector side", stream_id);
|
|
}
|
|
|
|
/// Check if auth is complete.
|
|
pub fn is_authenticated(&self) -> bool {
|
|
self.auth_complete.load(Ordering::SeqCst)
|
|
}
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// Connector with SOCKS5 (reconnect-aware, graceful shutdown)
|
|
// ---------------------------------------------------------------------------
|
|
|
|
/// Run the connector with SOCKS5 proxy: connect, auth, framing, then SOCKS5.
|
|
/// The SOCKS5 listener stays bound permanently while a reconnect loop
|
|
/// maintains the tunnel connection. On tunnel failure, reconnect retries
|
|
/// auth and swaps the mux; new SOCKS5 requests work with the fresh tunnel.
|
|
/// Existing streams get EOF/closed deterministically.
|
|
pub async fn run_connector_with_socks(
|
|
config: ConnectorConfig,
|
|
socks_addr: SocketAddr,
|
|
) -> Result<(), TunnelError> {
|
|
let client_config = config
|
|
.tls
|
|
.build_config(config.insecure_skip_tls_verify)
|
|
.map_err(|e| {
|
|
tracing::error!("Failed to build client TLS config: {}", e);
|
|
TunnelError::Tls(e)
|
|
})?;
|
|
let client_certs = config.tls.client_certs().map_err(TunnelError::Tls)?;
|
|
let ca_certs = config.tls.ca_certs().map_err(TunnelError::Tls)?;
|
|
|
|
let target_addr = resolve_target(&config.target_host, config.target_port).await?;
|
|
|
|
// Log effective config (secrets redacted)
|
|
tracing::info!(
|
|
"Effective config: target={}:{}, socks5={}, auth_token={}",
|
|
config.target_host,
|
|
config.target_port,
|
|
socks_addr,
|
|
Redacted::new(&*config.auth_token)
|
|
);
|
|
tracing::info!(
|
|
"TLS material fingerprints: ca={}, client={}",
|
|
tls::cert_fingerprint_sha256(&ca_certs[0]),
|
|
tls::cert_fingerprint_sha256(&client_certs[0])
|
|
);
|
|
|
|
// Create a watch channel for the active StreamMux.
|
|
// Initial value is None — first tunnel session will set it.
|
|
let (mux_sender, mux_receiver) = watch::channel::<Option<Arc<StreamMux>>>(None);
|
|
|
|
// Spawn the reconnect loop (runs in background)
|
|
let reconnect_config = config.clone();
|
|
let reconnect_sender = mux_sender.clone();
|
|
let reconnect_handle = tokio::spawn(async move {
|
|
run_reconnect_loop(
|
|
reconnect_config,
|
|
target_addr,
|
|
client_config,
|
|
reconnect_sender,
|
|
)
|
|
.await;
|
|
});
|
|
|
|
// Start the SOCKS5 proxy (stays bound permanently)
|
|
let socks_result = run_socks5_proxy(socks_addr, mux_receiver).await;
|
|
|
|
// Shutdown: abort reconnect loop
|
|
tracing::info!("Connector shutting down — aborting reconnect loop");
|
|
reconnect_handle.abort();
|
|
|
|
match socks_result {
|
|
Ok(()) => {
|
|
tracing::info!(
|
|
"Connector shut down gracefully — SOCKS5 port {} released",
|
|
socks_addr.port()
|
|
);
|
|
Ok(())
|
|
}
|
|
Err(e) => Err(e),
|
|
}
|
|
}
|
|
|
|
/// Reconnect loop: continuously attempts to establish and maintain the tunnel.
|
|
/// On success, publishes the StreamMux via the watch channel.
|
|
/// On failure, logs the error and retries with exponential backoff.
|
|
async fn run_reconnect_loop(
|
|
config: ConnectorConfig,
|
|
target_addr: SocketAddr,
|
|
client_config: Arc<rustls::ClientConfig>,
|
|
mux_sender: watch::Sender<Option<Arc<StreamMux>>>,
|
|
) {
|
|
let server_name = match tls::server_name_from_host(&config.target_host) {
|
|
Ok(sn) => sn,
|
|
Err(e) => {
|
|
tracing::error!("Invalid target hostname: {}", e);
|
|
return;
|
|
}
|
|
};
|
|
|
|
let mut reconnect_delay = std::time::Duration::from_secs(1);
|
|
let max_delay = std::time::Duration::from_secs(30);
|
|
|
|
loop {
|
|
tracing::info!(
|
|
"Tunnel connecting to {}:{} (HTTPS default transport)",
|
|
config.target_host,
|
|
config.target_port
|
|
);
|
|
|
|
// Connect with fresh TLS (re-runs all security gates)
|
|
let stream = match TcpStream::connect(target_addr).await {
|
|
Ok(s) => s,
|
|
Err(e) => {
|
|
tracing::warn!(
|
|
"Tunnel connect failed to {}: {} — retrying in {:?}",
|
|
target_addr,
|
|
e,
|
|
reconnect_delay
|
|
);
|
|
tokio::time::sleep(reconnect_delay).await;
|
|
reconnect_delay = (reconnect_delay * 2).min(max_delay);
|
|
continue;
|
|
}
|
|
};
|
|
|
|
let tls_connector = tokio_rustls::TlsConnector::from(client_config.clone());
|
|
let tls_stream = match tls_connector.connect(server_name.clone(), stream).await {
|
|
Ok(s) => s,
|
|
Err(e) => {
|
|
if let Some(hint) = tls_handshake_hint(&e.to_string()) {
|
|
tracing::warn!("{}", hint);
|
|
}
|
|
tracing::warn!(
|
|
"TLS handshake to {}:{} failed: {} — retrying in {:?}",
|
|
config.target_host,
|
|
config.target_port,
|
|
e,
|
|
reconnect_delay
|
|
);
|
|
tokio::time::sleep(reconnect_delay).await;
|
|
reconnect_delay = (reconnect_delay * 2).min(max_delay);
|
|
continue;
|
|
}
|
|
};
|
|
|
|
tracing::info!(
|
|
"TLS handshake completed with {}:{} via HTTPS",
|
|
config.target_host,
|
|
config.target_port
|
|
);
|
|
|
|
// Authenticate (re-runs auth check each reconnect)
|
|
let auth_result =
|
|
match authenticate_tunnel(&config.auth_token, &config.target_host, tls_stream).await {
|
|
Ok(r) => r,
|
|
Err(TunnelError::Auth(e)) => {
|
|
// Auth failures are terminal — don't retry endlessly
|
|
tracing::error!("Auth failed (terminal): {}", e);
|
|
mux_sender.send(None).ok();
|
|
return;
|
|
}
|
|
Err(e) => {
|
|
tracing::warn!(
|
|
"Connection error: {} — retrying in {:?}",
|
|
e,
|
|
reconnect_delay
|
|
);
|
|
tokio::time::sleep(reconnect_delay).await;
|
|
reconnect_delay = (reconnect_delay * 2).min(max_delay);
|
|
continue;
|
|
}
|
|
};
|
|
|
|
// Reset reconnect delay on success
|
|
reconnect_delay = std::time::Duration::from_secs(1);
|
|
|
|
let mux = auth_result.mux.clone();
|
|
|
|
tracing::info!(
|
|
"Tunnel session established with {}:{} via HTTPS /tunnel — state: connected",
|
|
config.target_host,
|
|
config.target_port
|
|
);
|
|
|
|
// Publish the mux via watch channel
|
|
mux_sender.send(Some(mux.clone())).ok();
|
|
|
|
// Wait for the framing task to finish (tunnel disconnect)
|
|
let _ = auth_result.framing_task.await;
|
|
|
|
tracing::info!(
|
|
"Tunnel disconnected from {}:{} — state: disconnected, will reconnect",
|
|
config.target_host,
|
|
config.target_port
|
|
);
|
|
|
|
// Notify that tunnel is down (SOCKS5 handler will see mux changes)
|
|
// Don't clear the mux — let existing streams drain gracefully
|
|
}
|
|
}
|
|
|
|
fn tls_handshake_hint(error: &str) -> Option<&'static str> {
|
|
if error.contains("BadSignature") {
|
|
return Some(
|
|
"TLS hint: the listener certificate is not signed by the CA trusted by this connector; use the same connection key on both sides or matching --ca-cert/server cert material",
|
|
);
|
|
}
|
|
|
|
if error.contains("NotValidForName")
|
|
|| error.contains("CertNotValidForName")
|
|
|| error.contains("NotValidForNameContext")
|
|
|| error.contains("certificate does not match")
|
|
{
|
|
return Some(
|
|
"TLS hint: the listener certificate does not match the target host; generate material with --advertise/--target set to the public listener address",
|
|
);
|
|
}
|
|
|
|
None
|
|
}
|
|
|
|
/// Result of the authenticate_tunnel function.
|
|
struct AuthResult {
|
|
mux: Arc<StreamMux>,
|
|
framing_task: tokio::task::JoinHandle<()>,
|
|
}
|
|
|
|
/// Authenticate over HTTP then switch to binary framing.
|
|
async fn authenticate_tunnel(
|
|
auth_token: &str,
|
|
target_host: &str,
|
|
tls_stream: tokio_rustls::client::TlsStream<TcpStream>,
|
|
) -> Result<AuthResult, TunnelError> {
|
|
let mut stream: std::pin::Pin<Box<tokio_rustls::client::TlsStream<TcpStream>>> =
|
|
Box::pin(tls_stream);
|
|
|
|
// Send HTTP auth request
|
|
let request = format!(
|
|
"POST /tunnel HTTP/1.1\r\nHost: {}\r\nX-Rustunnel-Token: {}\r\nContent-Length: 0\r\n\r\n",
|
|
target_host, auth_token
|
|
);
|
|
stream.write_all(request.as_bytes()).await?;
|
|
|
|
// Read HTTP response
|
|
let status = read_http_response_status(&mut *stream).await?;
|
|
|
|
if status != 200 {
|
|
tracing::warn!(
|
|
"Auth rejected — HTTP status {} from {}",
|
|
status,
|
|
target_host
|
|
);
|
|
return Err(TunnelError::Auth(AuthError::Invalid));
|
|
}
|
|
|
|
tracing::info!("Authenticated with {} via HTTPS /tunnel", target_host);
|
|
|
|
// Switch to binary framing
|
|
tracing::info!(
|
|
"Upgraded connection to {} to binary framing protocol",
|
|
target_host
|
|
);
|
|
|
|
// Set up the StreamMux and framing loop
|
|
let (outbound_tx, mut outbound_rx) = mpsc::channel::<Frame>(256);
|
|
let mux = Arc::new(StreamMux::new(outbound_tx));
|
|
let mux_dispatch = mux.clone();
|
|
let peer_name = target_host.to_string();
|
|
|
|
// Spawn the framing loop
|
|
let framing_task = tokio::spawn(async move {
|
|
run_framing_client(stream, mux_dispatch, &mut outbound_rx, peer_name).await;
|
|
});
|
|
|
|
Ok(AuthResult { mux, framing_task })
|
|
}
|
|
|
|
/// Run the framing loop on the client side.
|
|
async fn run_framing_client<S>(
|
|
mut stream: S,
|
|
mux: Arc<StreamMux>,
|
|
outbound_rx: &mut mpsc::Receiver<Frame>,
|
|
peer_name: String,
|
|
) where
|
|
S: AsyncReadExt + AsyncWriteExt + Unpin,
|
|
{
|
|
let mut reader = FrameReader::new();
|
|
let (target_tx, mut target_rx) = mpsc::channel::<Frame>(256);
|
|
let mut target_streams: HashMap<u64, ServerStreamEntry> = HashMap::new();
|
|
|
|
loop {
|
|
tokio::select! {
|
|
// Write outbound frames
|
|
Some(frame) = outbound_rx.recv() => {
|
|
if let Err(e) = FrameWriter::write_frame(&mut stream, &frame).await {
|
|
tracing::debug!("Client write error: {}", e);
|
|
break;
|
|
}
|
|
}
|
|
|
|
Some(frame) = target_rx.recv() => {
|
|
if let Err(e) = FrameWriter::write_frame(&mut stream, &frame).await {
|
|
tracing::debug!("Client target write error for stream {}: {}", frame.stream_id, e);
|
|
break;
|
|
}
|
|
}
|
|
|
|
// Read inbound frames and dispatch
|
|
result = reader.read_frame(&mut stream) => {
|
|
match result {
|
|
Ok(Some(frame)) => {
|
|
if frame.frame_type == FRAME_CONNECT
|
|
|| target_streams.contains_key(&frame.stream_id)
|
|
{
|
|
if let Err(e) = handle_inbound_frame(
|
|
&frame,
|
|
&mut stream,
|
|
&target_tx,
|
|
&mut target_streams,
|
|
SocketAddr::from(([0, 0, 0, 0], 0)),
|
|
).await {
|
|
tracing::warn!(
|
|
"Connector-side stream {} inbound error from {}: {}",
|
|
frame.stream_id,
|
|
peer_name,
|
|
e
|
|
);
|
|
let _ = target_tx
|
|
.send(Frame::error(frame.stream_id, e.to_string()))
|
|
.await;
|
|
}
|
|
} else {
|
|
mux.dispatch_frame(&frame).await;
|
|
}
|
|
}
|
|
Ok(None) => {
|
|
tracing::info!("Tunnel connection closed (EOF)");
|
|
break;
|
|
}
|
|
Err(e) => {
|
|
tracing::warn!("Client framing read error: {}", e);
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
for (sid, mut entry) in target_streams.drain() {
|
|
let _ = entry.target_write.shutdown().await;
|
|
tracing::debug!("Closed connector-side target stream {} on shutdown", sid);
|
|
}
|
|
|
|
tracing::info!("Client framing loop ended");
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// SOCKS5 handler with StreamMux (reconnect-aware)
|
|
// ---------------------------------------------------------------------------
|
|
|
|
/// Run the SOCKS5 proxy that stays bound permanently.
|
|
/// Receives mux updates via watch channel to support reconnect.
|
|
async fn run_socks5_proxy(
|
|
socks_addr: SocketAddr,
|
|
mux_receiver: watch::Receiver<Option<Arc<StreamMux>>>,
|
|
) -> Result<(), TunnelError> {
|
|
let socks_listener = TcpListener::bind(socks_addr).await.map_err(|e| {
|
|
if e.kind() == std::io::ErrorKind::AddrInUse {
|
|
TunnelError::PortInUse(socks_addr.port())
|
|
} else {
|
|
TunnelError::BindError(socks_addr.to_string(), e.to_string())
|
|
}
|
|
})?;
|
|
|
|
tracing::info!(
|
|
"SOCKS5 proxy listening on {} (persistent tunnel, binary framing)",
|
|
socks_addr
|
|
);
|
|
|
|
let shutdown = crate::signal::shutdown_signal();
|
|
tokio::pin!(shutdown);
|
|
|
|
loop {
|
|
tokio::select! {
|
|
biased; // prioritize shutdown
|
|
|
|
_ = &mut shutdown => {
|
|
tracing::info!("Shutdown signal received — stopping SOCKS5 proxy");
|
|
break;
|
|
}
|
|
|
|
result = socks_listener.accept() => {
|
|
match result {
|
|
Ok((stream, peer_addr)) => {
|
|
let mux_opt = mux_receiver.borrow().clone();
|
|
match mux_opt {
|
|
Some(ref m) => {
|
|
let m = m.clone();
|
|
tokio::spawn(async move {
|
|
if let Err(e) = handle_socks5_connection(stream, peer_addr, m).await {
|
|
tracing::warn!("SOCKS5 connection from {} failed: {}", peer_addr, e);
|
|
}
|
|
});
|
|
}
|
|
None => {
|
|
tracing::warn!(
|
|
"SOCKS5 connection from {} rejected: tunnel not connected yet",
|
|
peer_addr
|
|
);
|
|
}
|
|
}
|
|
}
|
|
Err(e) => {
|
|
tracing::error!("SOCKS5 accept error: {}", e);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
tracing::info!(
|
|
"SOCKS5 proxy shutting down — port {} released",
|
|
socks_addr.port()
|
|
);
|
|
Ok(())
|
|
}
|
|
|
|
/// Handle a single SOCKS5 connection using the shared StreamMux.
|
|
async fn handle_socks5_connection(
|
|
stream: TcpStream,
|
|
peer_addr: SocketAddr,
|
|
mux: Arc<StreamMux>,
|
|
) -> Result<(), socks5::Socks5Error> {
|
|
tracing::info!("SOCKS5 connection from {}", peer_addr);
|
|
|
|
if !mux.is_authenticated() {
|
|
tracing::warn!(
|
|
"SOCKS5 request from {} rejected: tunnel not authenticated yet",
|
|
peer_addr
|
|
);
|
|
return Err(socks5::Socks5Error::NotAuthenticated);
|
|
}
|
|
|
|
// --- SOCKS5 Handshake ---
|
|
let (mut socks_read, mut socks_write) = stream.into_split();
|
|
let target = socks5_handshake(&mut socks_read, &mut socks_write).await?;
|
|
|
|
let target_str = socks5::target_to_string(&target);
|
|
tracing::info!(
|
|
"SOCKS5 CONNECT from {} to {} (via multiplexed tunnel)",
|
|
peer_addr,
|
|
target_str
|
|
);
|
|
|
|
// --- Open stream via tunnel ---
|
|
let (stream_id, mut data_rx) = match mux.open_stream(&target).await {
|
|
Ok(r) => r,
|
|
Err(e) => {
|
|
tracing::warn!("Stream open failed for {}: {}", target_str, e);
|
|
socks_write
|
|
.write_all(&socks5::build_reply_failure(socks5::REPL_GENERAL_FAILURE))
|
|
.await?;
|
|
return Err(socks5::Socks5Error::Protocol(format!(
|
|
"tunnel stream open failed: {}",
|
|
e
|
|
)));
|
|
}
|
|
};
|
|
|
|
tracing::info!(
|
|
"SOCKS5 stream {}: target {} connected via tunnel",
|
|
stream_id,
|
|
target_str
|
|
);
|
|
|
|
// Send SOCKS5 success reply
|
|
socks_write
|
|
.write_all(&socks5::build_reply_success())
|
|
.await?;
|
|
|
|
// --- Bidirectional copy: SOCKS5 ↔ Tunnel (raw bytes, no buffering) ---
|
|
|
|
// SOCKS5 → Tunnel: read from SOCKS5 client, send DATA frames
|
|
let outbound = mux.outbound_tx.clone();
|
|
let sid = stream_id;
|
|
let socks_to_tunnel = tokio::spawn(async move {
|
|
let mut buf = [0u8; 32 * 1024];
|
|
loop {
|
|
match socks_read.read(&mut buf).await {
|
|
Ok(0) => break, // SOCKS5 client closed
|
|
Ok(n) => {
|
|
if outbound.send(Frame::data(sid, &buf[..n])).await.is_err() {
|
|
break;
|
|
}
|
|
}
|
|
Err(e) => {
|
|
tracing::debug!("Stream {} SOCKS5 read error: {}", sid, e);
|
|
let _ = outbound.send(Frame::error(sid, e.to_string())).await;
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
tracing::info!("Stream {} SOCKS5 → tunnel EOF", sid);
|
|
let _ = outbound.send(Frame::close(sid)).await;
|
|
});
|
|
|
|
// Tunnel → SOCKS5: read DATA frames, write to SOCKS5 client
|
|
while let Some(data) = data_rx.recv().await {
|
|
if let Err(e) = socks_write.write_all(&data).await {
|
|
tracing::debug!("Stream {} SOCKS5 write error: {}", stream_id, e);
|
|
break;
|
|
}
|
|
}
|
|
|
|
// Clean up
|
|
socks_to_tunnel.abort();
|
|
mux.close_stream(stream_id).await;
|
|
|
|
tracing::info!(
|
|
"Stream {} closed: SOCKS5 connection from {} completed",
|
|
stream_id,
|
|
peer_addr
|
|
);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
/// Parse the SOCKS5 handshake (greeting + auth + CONNECT request).
|
|
/// Returns the target address.
|
|
async fn socks5_handshake(
|
|
reader: &mut tokio::net::tcp::OwnedReadHalf,
|
|
writer: &mut tokio::net::tcp::OwnedWriteHalf,
|
|
) -> Result<socks5::Socks5Target, socks5::Socks5Error> {
|
|
// --- Greeting ---
|
|
let greeting = read_exact_2(reader).await?;
|
|
if greeting[0] != socks5::VERSION {
|
|
writer
|
|
.write_all(&[socks5::VERSION, socks5::AUTH_NO_ACCEPTABLE])
|
|
.await?;
|
|
return Err(socks5::Socks5Error::InvalidVersion(greeting[0]));
|
|
}
|
|
let n_methods = greeting[1] as usize;
|
|
if n_methods == 0 || n_methods > 255 {
|
|
return Err(socks5::Socks5Error::InvalidAuthMethodCount(n_methods));
|
|
}
|
|
let methods = read_exact(reader, n_methods).await?;
|
|
|
|
let chosen = if methods.contains(&socks5::AUTH_USERNAME_PASSWORD) {
|
|
socks5::AUTH_USERNAME_PASSWORD
|
|
} else if methods.contains(&socks5::AUTH_NONE) {
|
|
socks5::AUTH_NONE
|
|
} else {
|
|
writer
|
|
.write_all(&[socks5::VERSION, socks5::AUTH_NO_ACCEPTABLE])
|
|
.await?;
|
|
return Err(socks5::Socks5Error::UnsupportedAuthMethod(
|
|
methods.first().copied().unwrap_or(0xFF),
|
|
));
|
|
};
|
|
|
|
writer.write_all(&[socks5::VERSION, chosen]).await?;
|
|
|
|
// --- Auth (if username/password) ---
|
|
if chosen == socks5::AUTH_USERNAME_PASSWORD {
|
|
let auth_buf = read_exact_2(reader).await?;
|
|
if auth_buf[0] != 0x01 {
|
|
writer.write_all(&socks5::build_auth_failure()).await?;
|
|
return Err(socks5::Socks5Error::AuthFailed);
|
|
}
|
|
let ulen = auth_buf[1] as usize;
|
|
if ulen > 0 {
|
|
let _username = read_exact(reader, ulen).await?;
|
|
}
|
|
let plen_buf = read_exact(reader, 1).await?;
|
|
let plen = plen_buf[0] as usize;
|
|
if plen > 0 {
|
|
let _password = read_exact(reader, plen).await?;
|
|
}
|
|
writer.write_all(&socks5::build_auth_success()).await?;
|
|
}
|
|
|
|
// --- CONNECT Request ---
|
|
let connect_req = read_exact(reader, 4).await?;
|
|
if connect_req[0] != socks5::VERSION {
|
|
return Err(socks5::Socks5Error::InvalidVersion(connect_req[0]));
|
|
}
|
|
if connect_req[1] != socks5::CMD_CONNECT {
|
|
writer
|
|
.write_all(&socks5::build_reply_failure(socks5::REPL_CONN_NOT_ALLOWED))
|
|
.await?;
|
|
return Err(socks5::Socks5Error::UnsupportedCommand(connect_req[1]));
|
|
}
|
|
|
|
let atyp = connect_req[3];
|
|
let target = match atyp {
|
|
socks5::ATYP_IPV4 => {
|
|
let addr_bytes = read_exact(reader, 4).await?;
|
|
let ipv4 =
|
|
std::net::Ipv4Addr::new(addr_bytes[0], addr_bytes[1], addr_bytes[2], addr_bytes[3]);
|
|
let port_bytes = read_exact(reader, 2).await?;
|
|
let port = u16::from_be_bytes([port_bytes[0], port_bytes[1]]);
|
|
socks5::Socks5Target::Ipv4(ipv4, port)
|
|
}
|
|
socks5::ATYP_DOMAIN => {
|
|
let len_byte = read_exact(reader, 1).await?;
|
|
let domain_len = len_byte[0] as usize;
|
|
let domain_bytes = read_exact(reader, domain_len).await?;
|
|
let domain = String::from_utf8_lossy(&domain_bytes).to_string();
|
|
let port_bytes = read_exact(reader, 2).await?;
|
|
let port = u16::from_be_bytes([port_bytes[0], port_bytes[1]]);
|
|
socks5::Socks5Target::Domain { domain, port }
|
|
}
|
|
_ => {
|
|
writer
|
|
.write_all(&socks5::build_reply_failure(socks5::REPL_GENERAL_FAILURE))
|
|
.await?;
|
|
return Err(socks5::Socks5Error::UnsupportedAddrType(atyp));
|
|
}
|
|
};
|
|
|
|
Ok(target)
|
|
}
|
|
|
|
/// Read exactly `len` bytes.
|
|
async fn read_exact(
|
|
reader: &mut tokio::net::tcp::OwnedReadHalf,
|
|
len: usize,
|
|
) -> Result<Vec<u8>, socks5::Socks5Error> {
|
|
let mut buf = vec![0u8; len];
|
|
reader.read_exact(&mut buf).await?;
|
|
Ok(buf)
|
|
}
|
|
|
|
/// Read exactly 2 bytes.
|
|
async fn read_exact_2(
|
|
reader: &mut tokio::net::tcp::OwnedReadHalf,
|
|
) -> Result<[u8; 2], socks5::Socks5Error> {
|
|
let buf = read_exact(reader, 2).await?;
|
|
Ok([buf[0], buf[1]])
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// connect_tunnel: establishes a tunnel session (kept for test compatibility)
|
|
// ---------------------------------------------------------------------------
|
|
|
|
/// Connect to the HTTPS tunnel and establish a session.
|
|
/// Used by tests to verify tunnel auth/connect paths.
|
|
#[cfg(test)]
|
|
pub async fn connect_tunnel(config: ConnectorConfig) -> Result<(), TunnelError> {
|
|
let client_config = config
|
|
.tls
|
|
.build_config(config.insecure_skip_tls_verify)
|
|
.map_err(|e| {
|
|
tracing::error!("Failed to build client TLS config: {}", e);
|
|
TunnelError::Tls(e)
|
|
})?;
|
|
|
|
let tls_connector = tokio_rustls::TlsConnector::from(client_config);
|
|
let target_addr = resolve_target(&config.target_host, config.target_port).await?;
|
|
let server_name = tls::server_name_from_host(&config.target_host).map_err(TunnelError::Tls)?;
|
|
|
|
let stream = TcpStream::connect(target_addr).await.map_err(|e| {
|
|
if e.kind() == std::io::ErrorKind::ConnectionRefused {
|
|
TunnelError::ConnectionRefused(target_addr.to_string())
|
|
} else {
|
|
TunnelError::Io(e)
|
|
}
|
|
})?;
|
|
|
|
let tls_stream = tls_connector
|
|
.connect(server_name.clone(), stream)
|
|
.await
|
|
.map_err(|e| {
|
|
TunnelError::Tls(crate::errors::TlsError::VerificationFailed(format!(
|
|
"TLS handshake to {}:{} failed: {}",
|
|
config.target_host, config.target_port, e
|
|
)))
|
|
})?;
|
|
|
|
// Authenticate
|
|
let mut stream: std::pin::Pin<Box<tokio_rustls::client::TlsStream<TcpStream>>> =
|
|
Box::pin(tls_stream);
|
|
|
|
let request = format!(
|
|
"POST /tunnel HTTP/1.1\r\nHost: {}\r\nX-Rustunnel-Token: {}\r\nContent-Length: 0\r\n\r\n",
|
|
config.target_host, config.auth_token
|
|
);
|
|
stream.write_all(request.as_bytes()).await?;
|
|
let status = read_http_response_status(&mut *stream).await?;
|
|
|
|
if status != 200 {
|
|
return Err(TunnelError::Auth(AuthError::Invalid));
|
|
}
|
|
|
|
tracing::info!(
|
|
"Tunnel session established with {}:{} via HTTPS /tunnel",
|
|
config.target_host,
|
|
config.target_port
|
|
);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// Security gate state machine
|
|
// ---------------------------------------------------------------------------
|
|
|
|
/// Represents the progression of security gates during tunnel establishment.
|
|
/// Used by tests to verify that forwarding starts only after all gates pass.
|
|
#[cfg(test)]
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
|
pub enum SecurityGateStatus {
|
|
TlsPending,
|
|
ClientCertPending,
|
|
AuthPending,
|
|
Authenticated,
|
|
#[allow(dead_code)]
|
|
Failed,
|
|
}
|
|
|
|
#[cfg(test)]
|
|
#[derive(Debug, Clone)]
|
|
pub struct SecurityGateState {
|
|
status: Arc<tokio::sync::Mutex<SecurityGateStatus>>,
|
|
}
|
|
|
|
#[cfg(test)]
|
|
impl SecurityGateState {
|
|
pub fn new() -> Self {
|
|
Self {
|
|
status: Arc::new(tokio::sync::Mutex::new(SecurityGateStatus::TlsPending)),
|
|
}
|
|
}
|
|
|
|
pub async fn advance(&self, to: SecurityGateStatus) {
|
|
let mut status = self.status.lock().await;
|
|
*status = to;
|
|
}
|
|
|
|
pub async fn is_authenticated(&self) -> bool {
|
|
let status = self.status.lock().await;
|
|
*status == SecurityGateStatus::Authenticated
|
|
}
|
|
|
|
pub async fn get_status(&self) -> SecurityGateStatus {
|
|
*self.status.lock().await
|
|
}
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// Tests
|
|
// ---------------------------------------------------------------------------
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
fn get_free_port() -> u16 {
|
|
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
|
|
listener.local_addr().unwrap().port()
|
|
}
|
|
|
|
fn generate_test_creds() -> tempfile::TempDir {
|
|
let _ = rustls::crypto::ring::default_provider().install_default();
|
|
let dir = tempfile::tempdir().unwrap();
|
|
crate::generate::generate(dir.path(), "test-ca", "test-server", "test-client").unwrap();
|
|
dir
|
|
}
|
|
|
|
/// Helper: quick-connect that authenticates then immediately returns.
|
|
async fn quick_connect(config: &ConnectorConfig) -> Result<(), TunnelError> {
|
|
connect_tunnel(config.clone()).await
|
|
}
|
|
|
|
#[test]
|
|
fn tls_handshake_hint_explains_bad_signature() {
|
|
let hint = tls_handshake_hint("invalid peer certificate: BadSignature").unwrap();
|
|
assert!(hint.contains("not signed by the CA"));
|
|
}
|
|
|
|
#[test]
|
|
fn tls_handshake_hint_explains_name_mismatch() {
|
|
let hint = tls_handshake_hint("invalid peer certificate: NotValidForName").unwrap();
|
|
assert!(hint.contains("does not match the target host"));
|
|
}
|
|
|
|
// =========================================================================
|
|
// Auth / mTLS tests
|
|
// =========================================================================
|
|
|
|
#[tokio::test]
|
|
async fn valid_mtls_and_auth_establishes_tunnel() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let port = get_free_port();
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", port).parse().unwrap();
|
|
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
|
|
let result = quick_connect(&connector_config).await;
|
|
assert!(
|
|
result.is_ok(),
|
|
"Valid mTLS+auth should establish tunnel: {:?}",
|
|
result
|
|
);
|
|
|
|
listener_task.abort();
|
|
}
|
|
|
|
/// Test: insecure listener accepts a connector without valid client certificate.
|
|
#[tokio::test]
|
|
async fn insecure_listener_accepts_no_client_cert() {
|
|
let dir1 = generate_test_creds();
|
|
let dir2 = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir1.path().join("token.txt")).unwrap();
|
|
let port = get_free_port();
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", port).parse().unwrap();
|
|
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: true,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir1.path().join("server.crt"),
|
|
dir1.path().join("server.key"),
|
|
dir1.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
// Connector uses client cert from a DIFFERENT CA but verifies server with listener CA
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir2.path().join("client.crt"),
|
|
dir2.path().join("client.key"),
|
|
dir1.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
|
|
let result = quick_connect(&connector_config).await;
|
|
assert!(
|
|
result.is_ok(),
|
|
"Insecure listener should accept connector with untrusted client cert: {:?}",
|
|
result
|
|
);
|
|
|
|
listener_task.abort();
|
|
}
|
|
|
|
/// Test: insecure connector accepts a server certificate signed by an unknown CA.
|
|
#[tokio::test]
|
|
async fn insecure_connector_accepts_unknown_server_cert() {
|
|
let dir1 = generate_test_creds();
|
|
let dir2 = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir1.path().join("token.txt")).unwrap();
|
|
let port = get_free_port();
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", port).parse().unwrap();
|
|
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir1.path().join("server.crt"),
|
|
dir1.path().join("server.key"),
|
|
dir1.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
// Connector trusts a DIFFERENT CA (would normally fail server cert verification)
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: true,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir1.path().join("client.crt"),
|
|
dir1.path().join("client.key"),
|
|
dir2.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
|
|
let result = quick_connect(&connector_config).await;
|
|
assert!(
|
|
result.is_ok(),
|
|
"Insecure connector should accept server with unknown CA: {:?}",
|
|
result
|
|
);
|
|
|
|
listener_task.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn invalid_auth_token_fails() {
|
|
let dir = generate_test_creds();
|
|
let port = get_free_port();
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", port).parse().unwrap();
|
|
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new("correct-token".to_string()),
|
|
};
|
|
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new("wrong-token".to_string()),
|
|
};
|
|
|
|
let result = quick_connect(&connector_config).await;
|
|
assert!(
|
|
result.is_err(),
|
|
"Invalid auth token should fail: {:?}",
|
|
result
|
|
);
|
|
|
|
listener_task.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn missing_client_cert_fails() {
|
|
let dir1 = generate_test_creds();
|
|
let dir2 = generate_test_creds();
|
|
let port = get_free_port();
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", port).parse().unwrap();
|
|
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir1.path().join("server.crt"),
|
|
dir1.path().join("server.key"),
|
|
dir1.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new("token".to_string()),
|
|
};
|
|
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir2.path().join("client.crt"),
|
|
dir2.path().join("client.key"),
|
|
dir2.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new("token".to_string()),
|
|
};
|
|
|
|
let result = quick_connect(&connector_config).await;
|
|
assert!(
|
|
result.is_err(),
|
|
"Different CA client cert should fail: {:?}",
|
|
result
|
|
);
|
|
|
|
listener_task.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn port_conflict_fails_without_fallback() {
|
|
let dir = generate_test_creds();
|
|
let port = get_free_port();
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", port).parse().unwrap();
|
|
|
|
let _listener = TcpListener::bind(bind_addr).await.unwrap();
|
|
|
|
let config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new("token".to_string()),
|
|
};
|
|
|
|
let result = run_listener(config).await;
|
|
assert!(matches!(result, Err(TunnelError::PortInUse(_))));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn missing_tls_config_fails_closed() {
|
|
let port = get_free_port();
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", port).parse().unwrap();
|
|
|
|
let config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
"/nonexistent/cert.pem",
|
|
"/nonexistent/key.pem",
|
|
"/nonexistent/ca.pem",
|
|
),
|
|
auth_token: Arc::new("token".to_string()),
|
|
};
|
|
|
|
let result = run_listener(config).await;
|
|
assert!(result.is_err(), "Missing TLS config should fail closed");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn invalid_server_cert_rejected_by_client() {
|
|
let dir1 = generate_test_creds();
|
|
let dir2 = generate_test_creds();
|
|
let port = get_free_port();
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", port).parse().unwrap();
|
|
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir1.path().join("server.crt"),
|
|
dir1.path().join("server.key"),
|
|
dir1.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new("token".to_string()),
|
|
};
|
|
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir2.path().join("client.crt"),
|
|
dir2.path().join("client.key"),
|
|
dir2.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new("token".to_string()),
|
|
};
|
|
|
|
let result = quick_connect(&connector_config).await;
|
|
assert!(
|
|
result.is_err(),
|
|
"Client should reject server cert not signed by expected CA: {:?}",
|
|
result
|
|
);
|
|
|
|
listener_task.abort();
|
|
}
|
|
|
|
// =========================================================================
|
|
// Utility tests
|
|
// =========================================================================
|
|
|
|
#[test]
|
|
fn load_auth_token_from_value() {
|
|
let token = load_auth_token(Some("my-token"), None).unwrap();
|
|
assert_eq!(token, "my-token");
|
|
}
|
|
|
|
#[test]
|
|
fn load_auth_token_from_path() {
|
|
let dir = tempfile::tempdir().unwrap();
|
|
let token_path = dir.path().join("token.txt");
|
|
std::fs::write(&token_path, " file-token \n").unwrap();
|
|
|
|
let token = load_auth_token(None, Some(&token_path)).unwrap();
|
|
assert_eq!(token, "file-token");
|
|
}
|
|
|
|
#[test]
|
|
fn load_auth_token_empty_value_fails() {
|
|
let result = load_auth_token(Some(""), None);
|
|
assert!(result.is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn constant_time_compare_equal() {
|
|
assert!(constant_time_compare("abc", "abc"));
|
|
assert!(constant_time_compare("", ""));
|
|
}
|
|
|
|
#[test]
|
|
fn constant_time_compare_not_equal() {
|
|
assert!(!constant_time_compare("abc", "abd"));
|
|
assert!(!constant_time_compare("abc", "abcd"));
|
|
assert!(!constant_time_compare("abc", ""));
|
|
}
|
|
|
|
// =========================================================================
|
|
// Resolution tests
|
|
// =========================================================================
|
|
|
|
#[tokio::test]
|
|
async fn resolve_hostname_success() {
|
|
let addrs = resolve_host("localhost", 80).await.unwrap();
|
|
assert!(!addrs.is_empty(), "localhost should resolve");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn resolve_empty_hostname_fails() {
|
|
let result = resolve_host("", 80).await;
|
|
assert!(result.is_err());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn resolve_invalid_hostname_fails() {
|
|
let result = resolve_host("this.is.not.a.real.host.example", 80).await;
|
|
assert!(result.is_err());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn resolve_target_success() {
|
|
let addr = resolve_target("127.0.0.1", 80).await.unwrap();
|
|
assert_eq!(addr.ip().to_string(), "127.0.0.1");
|
|
assert_eq!(addr.port(), 80);
|
|
}
|
|
|
|
// =========================================================================
|
|
// Security gate tests
|
|
// =========================================================================
|
|
|
|
#[tokio::test]
|
|
async fn security_gate_state_transitions() {
|
|
let state = SecurityGateState::new();
|
|
assert!(!state.is_authenticated().await);
|
|
assert_eq!(state.get_status().await, SecurityGateStatus::TlsPending);
|
|
|
|
state.advance(SecurityGateStatus::Authenticated).await;
|
|
assert!(state.is_authenticated().await);
|
|
assert_eq!(state.get_status().await, SecurityGateStatus::Authenticated);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn forwarding_starts_only_after_all_gates() {
|
|
let state = SecurityGateState::new();
|
|
assert!(!state.is_authenticated().await);
|
|
state.advance(SecurityGateStatus::ClientCertPending).await;
|
|
assert!(!state.is_authenticated().await);
|
|
state.advance(SecurityGateStatus::AuthPending).await;
|
|
assert!(!state.is_authenticated().await);
|
|
state.advance(SecurityGateStatus::Authenticated).await;
|
|
assert!(state.is_authenticated().await);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn auth_required_in_addition_to_mtls() {
|
|
let dir = generate_test_creds();
|
|
let port = get_free_port();
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", port).parse().unwrap();
|
|
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new("correct-token".to_string()),
|
|
};
|
|
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new("wrong-token".to_string()),
|
|
};
|
|
|
|
let result = quick_connect(&connector_config).await;
|
|
assert!(
|
|
result.is_err(),
|
|
"Valid mTLS + wrong auth should fail: {:?}",
|
|
result
|
|
);
|
|
|
|
listener_task.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn reconnect_revalidates_security_gates() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let port = get_free_port();
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", port).parse().unwrap();
|
|
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
|
|
let result = quick_connect(&connector_config).await;
|
|
assert!(result.is_ok());
|
|
|
|
let bad_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new("bad-token".to_string()),
|
|
};
|
|
let result = quick_connect(&bad_config).await;
|
|
assert!(
|
|
result.is_err(),
|
|
"Reconnect with bad auth should fail: {:?}",
|
|
result
|
|
);
|
|
|
|
let result = quick_connect(&connector_config).await;
|
|
assert!(
|
|
result.is_ok(),
|
|
"Reconnect with valid auth should succeed: {:?}",
|
|
result
|
|
);
|
|
|
|
listener_task.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn https_is_default_carrier() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let port = get_free_port();
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", port).parse().unwrap();
|
|
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
// Plain TCP connects but the server expects TLS
|
|
let stream_result = TcpStream::connect(bind_addr).await;
|
|
assert!(stream_result.is_ok(), "TCP should connect to TLS listener");
|
|
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
|
|
let result = quick_connect(&connector_config).await;
|
|
assert!(result.is_ok(), "HTTPS mTLS should work");
|
|
|
|
listener_task.abort();
|
|
}
|
|
|
|
// =========================================================================
|
|
// HTTP helper tests
|
|
// =========================================================================
|
|
|
|
#[tokio::test]
|
|
async fn parse_http_request_basic() {
|
|
let data = b"POST /tunnel HTTP/1.1\r\nHost: localhost\r\nX-Rustunnel-Token: abc123\r\n\r\n";
|
|
let mut cursor = std::io::Cursor::new(data.as_slice());
|
|
let (method, path, headers) = parse_http_request(&mut cursor).await.unwrap();
|
|
assert_eq!(method, "POST");
|
|
assert_eq!(path, "/tunnel");
|
|
assert_eq!(headers.get("x-rustunnel-token").unwrap(), "abc123");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn parse_http_request_empty_fails() {
|
|
let mut cursor = std::io::Cursor::new(&[] as &[u8]);
|
|
let result = parse_http_request(&mut cursor).await;
|
|
assert!(result.is_err());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn read_http_response_status_ok() {
|
|
let data = b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
|
|
let mut cursor = std::io::Cursor::new(data.as_slice());
|
|
let status = read_http_response_status(&mut cursor).await.unwrap();
|
|
assert_eq!(status, 200);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn read_http_response_status_unauthorized() {
|
|
let data = b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 0\r\n\r\n";
|
|
let mut cursor = std::io::Cursor::new(data.as_slice());
|
|
let status = read_http_response_status(&mut cursor).await.unwrap();
|
|
assert_eq!(status, 401);
|
|
}
|
|
|
|
// =========================================================================
|
|
// Framing E2E tests
|
|
// =========================================================================
|
|
|
|
/// Test: full E2E flow with binary framing — listener, connector, target.
|
|
#[tokio::test]
|
|
async fn e2e_framing_basic_forward() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let listener_port = get_free_port();
|
|
let socks_port = get_free_port();
|
|
let target_port = get_free_port();
|
|
|
|
// Start a simple echo target
|
|
let target_listener = TcpListener::bind(format!("127.0.0.1:{}", target_port))
|
|
.await
|
|
.unwrap();
|
|
let target_task = tokio::spawn(async move {
|
|
if let Ok((stream, _)) = target_listener.accept().await {
|
|
let (mut rd, mut wr) = stream.into_split();
|
|
let mut buf = [0u8; 4096];
|
|
if let Ok(n) = rd.read(&mut buf).await
|
|
&& n > 0
|
|
{
|
|
let _ = wr.write_all(&buf[..n]).await;
|
|
}
|
|
}
|
|
});
|
|
|
|
// Start listener
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", listener_port).parse().unwrap();
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
// Start connector with SOCKS5
|
|
let socks_addr: SocketAddr = format!("127.0.0.1:{}", socks_port).parse().unwrap();
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: listener_port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let connector_task =
|
|
tokio::spawn(
|
|
async move { run_connector_with_socks(connector_config, socks_addr).await },
|
|
);
|
|
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
|
|
|
|
// Connect to SOCKS5 and send a request
|
|
let mut socks_stream = TcpStream::connect(socks_addr).await.unwrap();
|
|
{
|
|
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
|
|
|
|
// SOCKS5 greeting
|
|
socks_stream.write_all(&[0x05, 0x01, 0x00]).await.unwrap();
|
|
let mut resp = [0u8; 2];
|
|
socks_stream.read_exact(&mut resp).await.unwrap();
|
|
assert_eq!(resp, [0x05, 0x00]);
|
|
|
|
// CONNECT request (127.0.0.1:target_port)
|
|
let mut req = vec![0x05, 0x01, 0x00, 0x01, 127, 0, 0, 1];
|
|
req.extend_from_slice(&target_port.to_be_bytes());
|
|
socks_stream.write_all(&req).await.unwrap();
|
|
|
|
// Read reply
|
|
let mut reply = [0u8; 10];
|
|
socks_stream.read_exact(&mut reply).await.unwrap();
|
|
assert_eq!(reply[0], 0x05);
|
|
assert_eq!(reply[1], 0x00); // succeeded
|
|
|
|
// Send HTTP request
|
|
let request = format!("GET / HTTP/1.1\r\nHost: 127.0.0.1:{}\r\n\r\n", target_port);
|
|
socks_stream.write_all(request.as_bytes()).await.unwrap();
|
|
|
|
// Read response (echo)
|
|
let mut response = vec![0u8; 4096];
|
|
let n = socks_stream.read(&mut response).await.unwrap();
|
|
let response_str = String::from_utf8_lossy(&response[..n]);
|
|
assert!(
|
|
response_str.contains("GET / HTTP/1.1"),
|
|
"Echo response should contain request: {}",
|
|
response_str
|
|
);
|
|
}
|
|
|
|
// Cleanup
|
|
connector_task.abort();
|
|
listener_task.abort();
|
|
target_task.abort();
|
|
}
|
|
|
|
/// Test: server-side SOCKS uses the connector to open targets.
|
|
#[tokio::test]
|
|
async fn e2e_server_side_socks_reaches_connector_side_target() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let listener_port = get_free_port();
|
|
let socks_port = get_free_port();
|
|
let target_port = get_free_port();
|
|
|
|
let target_listener = TcpListener::bind(format!("127.0.0.1:{}", target_port))
|
|
.await
|
|
.unwrap();
|
|
let target_task = tokio::spawn(async move {
|
|
if let Ok((stream, _)) = target_listener.accept().await {
|
|
let (mut rd, mut wr) = stream.into_split();
|
|
let mut buf = [0u8; 4096];
|
|
if let Ok(n) = rd.read(&mut buf).await
|
|
&& n > 0
|
|
{
|
|
let _ = wr.write_all(&buf[..n]).await;
|
|
}
|
|
}
|
|
});
|
|
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", listener_port).parse().unwrap();
|
|
let socks_addr: SocketAddr = format!("127.0.0.1:{}", socks_port).parse().unwrap();
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: Some(socks_addr),
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: listener_port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let client_socks_addr: SocketAddr =
|
|
format!("127.0.0.1:{}", get_free_port()).parse().unwrap();
|
|
let connector_task = tokio::spawn(async move {
|
|
run_connector_with_socks(connector_config, client_socks_addr).await
|
|
});
|
|
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
|
|
|
|
let mut stream = TcpStream::connect(socks_addr).await.unwrap();
|
|
socks5_greeting(&mut stream).await.unwrap();
|
|
socks5_connect(&mut stream, "127.0.0.1", target_port)
|
|
.await
|
|
.unwrap();
|
|
let echoed = socks5_echo(&mut stream, "reverse-socks").await;
|
|
assert_eq!(echoed, "reverse-socks");
|
|
|
|
connector_task.abort();
|
|
listener_task.abort();
|
|
target_task.abort();
|
|
}
|
|
|
|
/// Test: multiple concurrent streams share one tunnel session.
|
|
#[tokio::test]
|
|
async fn e2e_concurrent_streams_isolation() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let listener_port = get_free_port();
|
|
let socks_port = get_free_port();
|
|
|
|
// Start two echo targets on different ports
|
|
let target_port1 = get_free_port();
|
|
let target_port2 = get_free_port();
|
|
|
|
let target1_listener = TcpListener::bind(format!("127.0.0.1:{}", target_port1))
|
|
.await
|
|
.unwrap();
|
|
let target2_listener = TcpListener::bind(format!("127.0.0.1:{}", target_port2))
|
|
.await
|
|
.unwrap();
|
|
|
|
// Echo server helper
|
|
fn echo_task(listener: TcpListener) -> tokio::task::JoinHandle<()> {
|
|
tokio::spawn(async move {
|
|
if let Ok((stream, _)) = listener.accept().await {
|
|
let (mut rd, mut wr) = stream.into_split();
|
|
let mut buf = [0u8; 4096];
|
|
if let Ok(n) = rd.read(&mut buf).await
|
|
&& n > 0
|
|
{
|
|
let _ = wr.write_all(&buf[..n]).await;
|
|
}
|
|
}
|
|
})
|
|
}
|
|
|
|
let target1_task = echo_task(target1_listener);
|
|
let target2_task = echo_task(target2_listener);
|
|
|
|
// Start listener
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", listener_port).parse().unwrap();
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
// Start connector
|
|
let socks_addr: SocketAddr = format!("127.0.0.1:{}", socks_port).parse().unwrap();
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: listener_port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let connector_task =
|
|
tokio::spawn(
|
|
async move { run_connector_with_socks(connector_config, socks_addr).await },
|
|
);
|
|
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
|
|
|
|
// Connect two concurrent SOCKS5 streams
|
|
let msg1 = "STREAM1-UNIQUE-DATA";
|
|
let msg2 = "STREAM2-UNIQUE-DATA";
|
|
|
|
let (result1, result2) = tokio::join!(
|
|
async {
|
|
let mut stream = TcpStream::connect(socks_addr).await.unwrap();
|
|
socks5_greeting(&mut stream).await.unwrap();
|
|
socks5_connect(&mut stream, "127.0.0.1", target_port1)
|
|
.await
|
|
.unwrap();
|
|
socks5_echo(&mut stream, msg1).await
|
|
},
|
|
async {
|
|
let mut stream = TcpStream::connect(socks_addr).await.unwrap();
|
|
socks5_greeting(&mut stream).await.unwrap();
|
|
socks5_connect(&mut stream, "127.0.0.1", target_port2)
|
|
.await
|
|
.unwrap();
|
|
socks5_echo(&mut stream, msg2).await
|
|
}
|
|
);
|
|
|
|
// Verify isolation: each stream got its own echo back
|
|
assert_eq!(result1, msg1, "Stream 1 data should be isolated");
|
|
assert_eq!(result2, msg2, "Stream 2 data should be isolated");
|
|
|
|
connector_task.abort();
|
|
listener_task.abort();
|
|
target1_task.abort();
|
|
target2_task.abort();
|
|
}
|
|
|
|
/// Test: target connection failure is reported without crashing.
|
|
#[tokio::test]
|
|
async fn e2e_target_failure_handled() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let listener_port = get_free_port();
|
|
let socks_port = get_free_port();
|
|
let bad_target_port = get_free_port(); // no target listening
|
|
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", listener_port).parse().unwrap();
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
let socks_addr: SocketAddr = format!("127.0.0.1:{}", socks_port).parse().unwrap();
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: listener_port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let connector_task =
|
|
tokio::spawn(
|
|
async move { run_connector_with_socks(connector_config, socks_addr).await },
|
|
);
|
|
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
|
|
|
|
// Try to connect to a target that doesn't exist
|
|
let mut socks_stream = TcpStream::connect(socks_addr).await.unwrap();
|
|
|
|
socks5_greeting(&mut socks_stream).await.unwrap();
|
|
// The connect should fail with a failure reply
|
|
socks5_connect(&mut socks_stream, "127.0.0.1", bad_target_port)
|
|
.await
|
|
.unwrap_err();
|
|
|
|
// Verify connector is still alive
|
|
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
|
|
assert!(
|
|
!connector_task.is_finished(),
|
|
"Connector should still be running after target failure"
|
|
);
|
|
|
|
connector_task.abort();
|
|
listener_task.abort();
|
|
}
|
|
|
|
/// Test: listener-side target origin — verify the listener opens the target connection.
|
|
#[tokio::test]
|
|
async fn e2e_listener_side_target_origin() {
|
|
// This test verifies that the target TCP connection is opened from the
|
|
// listener side, not the connector side. We do this by verifying that
|
|
// the full tunnel path works correctly.
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let listener_port = get_free_port();
|
|
let socks_port = get_free_port();
|
|
let target_port = get_free_port();
|
|
|
|
// Start target that records the source IP
|
|
let target_listener = TcpListener::bind(format!("127.0.0.1:{}", target_port))
|
|
.await
|
|
.unwrap();
|
|
let target_task = tokio::spawn(async move {
|
|
if let Ok((stream, _)) = target_listener.accept().await {
|
|
// The source should be localhost (listener opened it)
|
|
let peer = stream.peer_addr().unwrap();
|
|
assert_eq!(peer.ip().to_string(), "127.0.0.1");
|
|
// Echo back
|
|
let (mut rd, mut wr) = stream.into_split();
|
|
let mut buf = [0u8; 4096];
|
|
if let Ok(n) = rd.read(&mut buf).await
|
|
&& n > 0
|
|
{
|
|
let _ = wr.write_all(&buf[..n]).await;
|
|
}
|
|
}
|
|
});
|
|
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", listener_port).parse().unwrap();
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
let socks_addr: SocketAddr = format!("127.0.0.1:{}", socks_port).parse().unwrap();
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: listener_port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let connector_task =
|
|
tokio::spawn(
|
|
async move { run_connector_with_socks(connector_config, socks_addr).await },
|
|
);
|
|
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
|
|
|
|
// Send data through the tunnel
|
|
let result = tokio::time::timeout(std::time::Duration::from_secs(5), async {
|
|
let mut stream = TcpStream::connect(socks_addr).await.unwrap();
|
|
socks5_greeting(&mut stream).await.unwrap();
|
|
socks5_connect(&mut stream, "127.0.0.1", target_port)
|
|
.await
|
|
.unwrap();
|
|
socks5_echo(&mut stream, "test-origin").await
|
|
})
|
|
.await;
|
|
|
|
assert!(result.is_ok(), "Echo should succeed");
|
|
assert_eq!(result.unwrap(), "test-origin");
|
|
|
|
connector_task.abort();
|
|
listener_task.abort();
|
|
target_task.abort();
|
|
}
|
|
|
|
// =========================================================================
|
|
// Slow/interleaved stream tests
|
|
// =========================================================================
|
|
|
|
/// Test: slow bidirectional streaming — data is forwarded incrementally.
|
|
#[tokio::test]
|
|
async fn e2e_slow_stream_incremental() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let listener_port = get_free_port();
|
|
let socks_port = get_free_port();
|
|
let target_port = get_free_port();
|
|
|
|
// Target: slow echo server that responds byte by byte
|
|
let target_listener = TcpListener::bind(format!("127.0.0.1:{}", target_port))
|
|
.await
|
|
.unwrap();
|
|
let target_task = tokio::spawn(async move {
|
|
if let Ok((stream, _)) = target_listener.accept().await {
|
|
let (mut rd, mut wr) = stream.into_split();
|
|
let mut buf = [0u8; 1];
|
|
loop {
|
|
match rd.read(&mut buf).await {
|
|
Ok(0) => break,
|
|
Ok(n) => {
|
|
if n == 0 {
|
|
break;
|
|
}
|
|
if wr.write_all(&buf[..n]).await.is_err() {
|
|
break;
|
|
}
|
|
// Small delay to simulate slow target
|
|
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
|
|
}
|
|
Err(_) => break,
|
|
}
|
|
}
|
|
}
|
|
});
|
|
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", listener_port).parse().unwrap();
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
let socks_addr: SocketAddr = format!("127.0.0.1:{}", socks_port).parse().unwrap();
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: listener_port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let connector_task =
|
|
tokio::spawn(
|
|
async move { run_connector_with_socks(connector_config, socks_addr).await },
|
|
);
|
|
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
|
|
|
|
// Send data through the slow echo server
|
|
let result = tokio::time::timeout(std::time::Duration::from_secs(10), async {
|
|
let mut stream = TcpStream::connect(socks_addr).await.unwrap();
|
|
socks5_greeting(&mut stream).await.unwrap();
|
|
socks5_connect(&mut stream, "127.0.0.1", target_port)
|
|
.await
|
|
.unwrap();
|
|
|
|
// Send data incrementally
|
|
use tokio::io::AsyncWriteExt as _;
|
|
for chunk in b"SLOW-STREAM-TEST-DATA".chunks(2) {
|
|
stream.write_all(chunk).await.unwrap();
|
|
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
|
|
}
|
|
|
|
// Read response incrementally
|
|
use tokio::io::AsyncReadExt as _;
|
|
let mut full_response = Vec::new();
|
|
let mut buf = [0u8; 2];
|
|
loop {
|
|
let n = tokio::time::timeout(
|
|
std::time::Duration::from_millis(500),
|
|
stream.read(&mut buf),
|
|
)
|
|
.await;
|
|
match n {
|
|
Ok(Ok(0)) => break,
|
|
Ok(Ok(n)) => full_response.extend_from_slice(&buf[..n]),
|
|
Ok(Err(_)) | Err(_) => break,
|
|
}
|
|
}
|
|
String::from_utf8_lossy(&full_response).to_string()
|
|
})
|
|
.await;
|
|
|
|
assert!(
|
|
result.is_ok(),
|
|
"Slow stream echo should complete within timeout"
|
|
);
|
|
let response = result.unwrap();
|
|
assert_eq!(
|
|
response, "SLOW-STREAM-TEST-DATA",
|
|
"Incremental echo should preserve data: got '{}'",
|
|
response
|
|
);
|
|
|
|
connector_task.abort();
|
|
listener_task.abort();
|
|
target_task.abort();
|
|
}
|
|
|
|
// =========================================================================
|
|
// Stream close/error frame propagation tests
|
|
// =========================================================================
|
|
|
|
/// Test: close frames propagate deterministically.
|
|
#[tokio::test]
|
|
async fn e2e_close_frame_propagation() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let listener_port = get_free_port();
|
|
let socks_port = get_free_port();
|
|
let target_port = get_free_port();
|
|
|
|
// Target: responds immediately then closes
|
|
let target_listener = TcpListener::bind(format!("127.0.0.1:{}", target_port))
|
|
.await
|
|
.unwrap();
|
|
let target_task = tokio::spawn(async move {
|
|
if let Ok((stream, _)) = target_listener.accept().await {
|
|
let (mut rd, mut wr) = stream.into_split();
|
|
let mut buf = [0u8; 4096];
|
|
if let Ok(n) = rd.read(&mut buf).await
|
|
&& n > 0
|
|
{
|
|
let _ = wr.write_all(&buf[..n]).await;
|
|
}
|
|
// Close immediately after echo
|
|
let _ = wr.shutdown().await;
|
|
}
|
|
});
|
|
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", listener_port).parse().unwrap();
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
let socks_addr: SocketAddr = format!("127.0.0.1:{}", socks_port).parse().unwrap();
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: listener_port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let connector_task =
|
|
tokio::spawn(
|
|
async move { run_connector_with_socks(connector_config, socks_addr).await },
|
|
);
|
|
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
|
|
|
|
// Send data, get echo, then verify close propagation
|
|
let result = tokio::time::timeout(std::time::Duration::from_secs(5), async {
|
|
let mut stream = TcpStream::connect(socks_addr).await.unwrap();
|
|
socks5_greeting(&mut stream).await.unwrap();
|
|
socks5_connect(&mut stream, "127.0.0.1", target_port)
|
|
.await
|
|
.unwrap();
|
|
|
|
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
|
|
stream.write_all(b"CLOSE-TEST").await.unwrap();
|
|
|
|
// Read echo
|
|
let mut buf = [0u8; 4096];
|
|
let n = stream.read(&mut buf).await.unwrap();
|
|
let echo = String::from_utf8_lossy(&buf[..n]).to_string();
|
|
|
|
// After target closes, the next read should return 0 (EOF)
|
|
// This verifies CLOSE frame propagation
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
let n2 = stream.read(&mut buf).await.unwrap();
|
|
(echo, n2)
|
|
})
|
|
.await;
|
|
|
|
assert!(result.is_ok(), "Close test should complete");
|
|
let (echo, remaining) = result.unwrap();
|
|
assert_eq!(echo, "CLOSE-TEST");
|
|
assert_eq!(remaining, 0, "After target close, read should return 0");
|
|
|
|
connector_task.abort();
|
|
listener_task.abort();
|
|
target_task.abort();
|
|
}
|
|
|
|
// =========================================================================
|
|
// SOCKS5 helper functions for tests
|
|
// =========================================================================
|
|
|
|
async fn socks5_greeting(stream: &mut TcpStream) -> Result<(), std::io::Error> {
|
|
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
|
|
stream.write_all(&[0x05, 0x01, 0x00]).await?;
|
|
let mut resp = [0u8; 2];
|
|
stream.read_exact(&mut resp).await?;
|
|
assert_eq!(resp[0], 0x05);
|
|
assert_eq!(resp[1], 0x00);
|
|
Ok(())
|
|
}
|
|
|
|
async fn socks5_connect(
|
|
stream: &mut TcpStream,
|
|
host: &str,
|
|
port: u16,
|
|
) -> Result<(), std::io::Error> {
|
|
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
|
|
let octets: Vec<u8> = host
|
|
.split('.')
|
|
.map(|s| s.parse::<u8>().unwrap_or(0))
|
|
.collect();
|
|
let mut req = vec![0x05, 0x01, 0x00, 0x01];
|
|
req.extend_from_slice(&octets);
|
|
req.extend_from_slice(&port.to_be_bytes());
|
|
stream.write_all(&req).await?;
|
|
let mut reply = [0u8; 10];
|
|
stream.read_exact(&mut reply).await?;
|
|
if reply[1] != 0x00 {
|
|
return Err(std::io::Error::new(
|
|
std::io::ErrorKind::ConnectionRefused,
|
|
format!("SOCKS5 connect failed with code {}", reply[1]),
|
|
));
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
async fn socks5_echo(stream: &mut TcpStream, msg: &str) -> String {
|
|
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
|
|
stream.write_all(msg.as_bytes()).await.unwrap();
|
|
let mut buf = vec![0u8; msg.len() + 100];
|
|
let n = tokio::time::timeout(std::time::Duration::from_secs(5), stream.read(&mut buf))
|
|
.await
|
|
.unwrap()
|
|
.unwrap();
|
|
String::from_utf8_lossy(&buf[..n]).to_string()
|
|
}
|
|
|
|
// =========================================================================
|
|
// Operations tests: reconnect, shutdown, logs
|
|
// =========================================================================
|
|
|
|
/// VAL-OPS-003: Reconnect restores new traffic after tunnel interruption.
|
|
/// Simulates a tunnel drop and verifies that new SOCKS5 requests work
|
|
/// after the reconnect loop re-establishes the session.
|
|
#[tokio::test]
|
|
async fn e2e_reconnect_restores_traffic() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let listener_port = get_free_port();
|
|
let socks_port = get_free_port();
|
|
let target_port = get_free_port();
|
|
|
|
// Start echo target
|
|
let target_listener = TcpListener::bind(format!("127.0.0.1:{}", target_port))
|
|
.await
|
|
.unwrap();
|
|
let target_task = tokio::spawn(async move {
|
|
// Accept multiple connections (for before/after reconnect)
|
|
loop {
|
|
if let Ok((stream, _)) = target_listener.accept().await {
|
|
let (mut rd, mut wr) = stream.into_split();
|
|
let mut buf = [0u8; 4096];
|
|
if let Ok(n) = rd.read(&mut buf).await
|
|
&& n > 0
|
|
{
|
|
let _ = wr.write_all(&buf[..n]).await;
|
|
}
|
|
}
|
|
}
|
|
});
|
|
|
|
// Start listener
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", listener_port).parse().unwrap();
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
// Start connector with SOCKS5 (has reconnect loop built-in)
|
|
let socks_addr: SocketAddr = format!("127.0.0.1:{}", socks_port).parse().unwrap();
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: listener_port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let connector_task =
|
|
tokio::spawn(
|
|
async move { run_connector_with_socks(connector_config, socks_addr).await },
|
|
);
|
|
tokio::time::sleep(std::time::Duration::from_millis(400)).await;
|
|
|
|
// Step 1: Verify initial connection works
|
|
let result1 = tokio::time::timeout(std::time::Duration::from_secs(5), async {
|
|
let mut stream = TcpStream::connect(socks_addr).await.unwrap();
|
|
socks5_greeting(&mut stream).await.unwrap();
|
|
socks5_connect(&mut stream, "127.0.0.1", target_port)
|
|
.await
|
|
.unwrap();
|
|
socks5_echo(&mut stream, "BEFORE-RECONNECT").await
|
|
})
|
|
.await;
|
|
assert!(result1.is_ok(), "Initial connection should work");
|
|
assert_eq!(result1.unwrap(), "BEFORE-RECONNECT");
|
|
|
|
// Step 2: Force listener to stop (simulates tunnel drop)
|
|
listener_task.abort();
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
// Step 3: Restart listener
|
|
let bind_addr2: SocketAddr = format!("127.0.0.1:{}", listener_port).parse().unwrap();
|
|
let listener_config2 = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr: bind_addr2,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let listener_task2 = tokio::spawn(async move { run_listener(listener_config2).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
// Step 4: Wait for reconnect loop to re-establish the tunnel
|
|
// The reconnect loop has exponential backoff starting at 1s, so wait a bit
|
|
tokio::time::sleep(std::time::Duration::from_secs(3)).await;
|
|
|
|
// Step 5: Verify new SOCKS5 request works after reconnect
|
|
let result2 = tokio::time::timeout(std::time::Duration::from_secs(8), async {
|
|
let mut stream = TcpStream::connect(socks_addr).await.unwrap();
|
|
socks5_greeting(&mut stream).await.unwrap();
|
|
socks5_connect(&mut stream, "127.0.0.1", target_port)
|
|
.await
|
|
.unwrap();
|
|
socks5_echo(&mut stream, "AFTER-RECONNECT").await
|
|
})
|
|
.await;
|
|
|
|
assert!(
|
|
result2.is_ok(),
|
|
"After reconnect, new traffic should work: {:?}",
|
|
result2
|
|
);
|
|
assert_eq!(result2.unwrap(), "AFTER-RECONNECT");
|
|
|
|
// Cleanup
|
|
connector_task.abort();
|
|
listener_task2.abort();
|
|
target_task.abort();
|
|
}
|
|
|
|
/// VAL-OPS-004: Shutdown is graceful — listener exits cleanly.
|
|
#[tokio::test]
|
|
async fn e2e_graceful_shutdown_listener() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let port = get_free_port();
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", port).parse().unwrap();
|
|
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
|
|
// Run listener with a timeout — if it hangs on shutdown, the test fails
|
|
let handle = tokio::spawn(async move { run_listener(listener_config).await });
|
|
|
|
// Let it start
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
// Abort should cause graceful shutdown within reasonable time
|
|
handle.abort();
|
|
let result = tokio::time::timeout(std::time::Duration::from_secs(3), handle).await;
|
|
|
|
assert!(
|
|
result.is_ok(),
|
|
"Listener should shut down gracefully within 3 seconds"
|
|
);
|
|
|
|
// Verify port is released — a new listener should be able to bind
|
|
let bind_result = TcpListener::bind(bind_addr).await;
|
|
assert!(
|
|
bind_result.is_ok(),
|
|
"Port should be released after graceful shutdown"
|
|
);
|
|
}
|
|
|
|
/// VAL-OPS-004: Shutdown is graceful — connector exits cleanly.
|
|
#[tokio::test]
|
|
async fn e2e_graceful_shutdown_connector() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let listener_port = get_free_port();
|
|
let socks_port = get_free_port();
|
|
|
|
// Start listener
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", listener_port).parse().unwrap();
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
|
|
// Start connector
|
|
let socks_addr: SocketAddr = format!("127.0.0.1:{}", socks_port).parse().unwrap();
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: listener_port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let connector_task =
|
|
tokio::spawn(
|
|
async move { run_connector_with_socks(connector_config, socks_addr).await },
|
|
);
|
|
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
|
|
|
|
// Abort connector — should shut down gracefully
|
|
connector_task.abort();
|
|
let result = tokio::time::timeout(std::time::Duration::from_secs(3), connector_task).await;
|
|
|
|
assert!(
|
|
result.is_ok(),
|
|
"Connector should shut down gracefully within 3 seconds"
|
|
);
|
|
|
|
// Verify SOCKS5 port is released
|
|
let bind_result = TcpListener::bind(socks_addr).await;
|
|
assert!(
|
|
bind_result.is_ok(),
|
|
"SOCKS5 port should be released after graceful shutdown"
|
|
);
|
|
|
|
listener_task.abort();
|
|
}
|
|
|
|
/// VAL-OPS-005: Shutdown during active streams is graceful.
|
|
/// Verify that stopping the connector during active streams produces
|
|
/// clean EOF errors and doesn't mix stream data.
|
|
#[tokio::test]
|
|
async fn e2e_shutdown_during_active_streams() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let listener_port = get_free_port();
|
|
let socks_port = get_free_port();
|
|
let target_port = get_free_port();
|
|
|
|
// Start slow echo target
|
|
let target_listener = TcpListener::bind(format!("127.0.0.1:{}", target_port))
|
|
.await
|
|
.unwrap();
|
|
let target_task = tokio::spawn(async move {
|
|
loop {
|
|
if let Ok((stream, _)) = target_listener.accept().await {
|
|
let (mut rd, mut wr) = stream.into_split();
|
|
let mut buf = [0u8; 4096];
|
|
if let Ok(n) = rd.read(&mut buf).await
|
|
&& n > 0
|
|
{
|
|
let _ = wr.write_all(&buf[..n]).await;
|
|
}
|
|
}
|
|
}
|
|
});
|
|
|
|
// Start listener
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", listener_port).parse().unwrap();
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
// Start connector
|
|
let socks_addr: SocketAddr = format!("127.0.0.1:{}", socks_port).parse().unwrap();
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: listener_port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let connector_task =
|
|
tokio::spawn(
|
|
async move { run_connector_with_socks(connector_config, socks_addr).await },
|
|
);
|
|
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
|
|
|
|
// Open two concurrent streams and send unique data
|
|
let (r1, r2) = tokio::join!(
|
|
async {
|
|
let mut stream = TcpStream::connect(socks_addr).await.unwrap();
|
|
socks5_greeting(&mut stream).await.unwrap();
|
|
socks5_connect(&mut stream, "127.0.0.1", target_port)
|
|
.await
|
|
.unwrap();
|
|
socks5_echo(&mut stream, "STREAM-A-DATA").await
|
|
},
|
|
async {
|
|
let mut stream = TcpStream::connect(socks_addr).await.unwrap();
|
|
socks5_greeting(&mut stream).await.unwrap();
|
|
socks5_connect(&mut stream, "127.0.0.1", target_port)
|
|
.await
|
|
.unwrap();
|
|
socks5_echo(&mut stream, "STREAM-B-DATA").await
|
|
}
|
|
);
|
|
|
|
// Verify stream data was not mixed
|
|
assert_eq!(r1, "STREAM-A-DATA", "Stream A data should be preserved");
|
|
assert_eq!(r2, "STREAM-B-DATA", "Stream B data should be preserved");
|
|
|
|
// Now shutdown the connector
|
|
connector_task.abort();
|
|
let shutdown_result =
|
|
tokio::time::timeout(std::time::Duration::from_secs(3), connector_task).await;
|
|
assert!(
|
|
shutdown_result.is_ok(),
|
|
"Connector should shut down gracefully even after active streams"
|
|
);
|
|
|
|
// Verify ports are released
|
|
let socks_rebind = TcpListener::bind(socks_addr).await;
|
|
assert!(socks_rebind.is_ok(), "SOCKS5 port should be released");
|
|
|
|
// Cleanup
|
|
listener_task.abort();
|
|
target_task.abort();
|
|
}
|
|
|
|
/// VAL-OPS-001: Logs report effective configuration and state transitions.
|
|
/// This test verifies that the connector logs the effective config at startup
|
|
/// including target, SOCKS5 address, and redacted auth token.
|
|
#[tokio::test]
|
|
async fn ops_logs_effective_config() {
|
|
let dir = generate_test_creds();
|
|
let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap();
|
|
let listener_port = get_free_port();
|
|
let socks_port = get_free_port();
|
|
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", listener_port).parse().unwrap();
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
let socks_addr: SocketAddr = format!("127.0.0.1:{}", socks_port).parse().unwrap();
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: listener_port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(token.clone()),
|
|
};
|
|
let connector_task =
|
|
tokio::spawn(
|
|
async move { run_connector_with_socks(connector_config, socks_addr).await },
|
|
);
|
|
tokio::time::sleep(std::time::Duration::from_millis(400)).await;
|
|
|
|
// The logs should have been emitted (tracing::info!) by the connector
|
|
// We verify the connector is running, which means the logs were produced
|
|
assert!(
|
|
!connector_task.is_finished(),
|
|
"Connector should be running and have logged effective config"
|
|
);
|
|
|
|
// Verify the logs don't contain the raw token
|
|
let token_display = format!("{}", crate::redact::Redacted::new(&token));
|
|
assert!(
|
|
token_display.contains("REDACTED"),
|
|
"Token should be redacted in logs"
|
|
);
|
|
assert!(
|
|
!token_display.contains(&token),
|
|
"Raw token should not appear in redacted output"
|
|
);
|
|
|
|
connector_task.abort();
|
|
listener_task.abort();
|
|
}
|
|
|
|
/// VAL-OPS-002: Auth and certificate failures are actionable and redacted.
|
|
#[tokio::test]
|
|
async fn ops_auth_failure_is_actionable_and_redacted() {
|
|
let dir = generate_test_creds();
|
|
let port = get_free_port();
|
|
let bind_addr: SocketAddr = format!("127.0.0.1:{}", port).parse().unwrap();
|
|
|
|
let listener_config = ListenerConfig {
|
|
insecure_skip_tls_verify: false,
|
|
bind_addr,
|
|
socks_addr: None,
|
|
tls: ServerTlsMaterial::from_paths(
|
|
dir.path().join("server.crt"),
|
|
dir.path().join("server.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new("secret-token-abc".to_string()),
|
|
};
|
|
let listener_task = tokio::spawn(async move { run_listener(listener_config).await });
|
|
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
|
|
|
// Connect with wrong token
|
|
let bad_token = "wrong-secret-token";
|
|
let connector_config = ConnectorConfig {
|
|
insecure_skip_tls_verify: false,
|
|
target_host: "127.0.0.1".to_string(),
|
|
target_port: port,
|
|
tls: ClientTlsMaterial::from_paths(
|
|
dir.path().join("client.crt"),
|
|
dir.path().join("client.key"),
|
|
dir.path().join("ca.pem"),
|
|
),
|
|
auth_token: Arc::new(bad_token.to_string()),
|
|
};
|
|
let result = connect_tunnel(connector_config).await;
|
|
assert!(result.is_err(), "Bad auth token should fail to connect");
|
|
|
|
// Verify error message doesn't leak the actual token values
|
|
let err_msg = result.unwrap_err().to_string();
|
|
assert!(
|
|
!err_msg.contains("secret-token-abc"),
|
|
"Error should not contain the actual server token: {}",
|
|
err_msg
|
|
);
|
|
assert!(
|
|
!err_msg.contains("wrong-secret-token"),
|
|
"Error should not contain the client's bad token: {}",
|
|
err_msg
|
|
);
|
|
|
|
listener_task.abort();
|
|
}
|
|
}
|