From f871ba7f7b85e94d30dcb8085522eb9c628006e5 Mon Sep 17 00:00:00 2001 From: c4ch3c4d3 <23181631+c4ch3c4d3@users.noreply.github.com> Date: Wed, 3 Jun 2026 20:17:39 -0600 Subject: [PATCH] Replace per-SOCKS-request HTTPS connections with persistent multiplexed tunnel session - Introduce TunnelSession that holds a single authenticated SendRequest shared across all concurrent SOCKS5 streams via Arc> - Connector authenticates once on startup, then reuses the persistent tunnel for all SOCKS5 CONNECT requests - Listener /forward returns raw bytes (no HTTP parse/unparse overhead) to preserve exact byte semantics for SOCKS5 TCP proxying - Add stream counters (session-X) for logging and isolation tracking - Add 5 new tests: sequential streams, concurrent streams with isolation, stream counter increments, pre-auth rejection, and E2E concurrent SOCKS5 streams over a shared tunnel - All 117 tests pass; E2E CLI validates 5 concurrent streams share one persistent TLS connection (127.0.0.1:64216 in transcript) --- src/errors.rs | 3 + src/tunnel.rs | 1013 ++++++++++++++++++++++++++++++++++++++++++------- 2 files changed, 876 insertions(+), 140 deletions(-) diff --git a/src/errors.rs b/src/errors.rs index 5e1416e..5fd0dce 100644 --- a/src/errors.rs +++ b/src/errors.rs @@ -40,6 +40,9 @@ pub enum AuthError { #[error("auth token invalid")] Invalid, + #[error("auth not yet completed")] + NotAuthenticated, + #[error("auth token expired")] #[allow(dead_code)] Expired, diff --git a/src/tunnel.rs b/src/tunnel.rs index cc31eb9..69ba5c5 100644 --- a/src/tunnel.rs +++ b/src/tunnel.rs @@ -22,7 +22,7 @@ use std::net::SocketAddr; use std::path::Path; use std::sync::Arc; -use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}; use http_body_util::{BodyExt, Full, combinators::BoxBody}; use hyper::body::{Bytes, Incoming}; @@ -37,6 +37,14 @@ use crate::redact::Redacted; use crate::socks5; use crate::tls; +// Global stream-ID counter for deterministic unique IDs +static STREAM_COUNTER: AtomicU64 = AtomicU64::new(0); + +/// Generate a unique stream ID for logging. +fn next_stream_id() -> u64 { + STREAM_COUNTER.fetch_add(1, Ordering::Relaxed) +} + // --------------------------------------------------------------------------- // Auth token handling // --------------------------------------------------------------------------- @@ -320,7 +328,7 @@ fn parse_forward_target(s: &str) -> Result<(String, u16), String> { /// Handle a /forward request: open target connection, send request body, return response body. /// /// Collects the request body, sends it to the target, collects the target response, -/// and returns it. This simpler approach works for all HTTP request/response patterns. +/// and returns raw bytes. This preserves exact byte semantics for SOCKS5 TCP proxy. async fn handle_forward_request( req: Request, peer_addr: SocketAddr, @@ -385,21 +393,16 @@ async fn handle_forward_request( } }; - // Forward to target - let result = forward_to_target(&fwd_host, fwd_port, &stream_id, body_bytes).await; + // Forward to target — returns raw bytes (no HTTP parse/unparse) + let result = forward_raw_to_target(&fwd_host, fwd_port, &stream_id, body_bytes).await; match result { - Ok((status, headers, resp_body)) => { - let mut builder = Response::builder().status(status); - for (name, value) in headers { - builder = builder.header(name, value); - } - Ok(builder - .body(BoxBody::new( - Full::new(resp_body).map_err(|_| unreachable!()), - )) - .unwrap()) - } + Ok(resp_body) => Ok(Response::builder() + .status(StatusCode::OK) + .body(BoxBody::new( + Full::new(resp_body).map_err(|_| unreachable!()), + )) + .unwrap()), Err(e) => { tracing::warn!("Forward stream {} error: {}", stream_id, e); Ok(Response::builder() @@ -410,13 +413,13 @@ async fn handle_forward_request( } } -/// Forward request data to a target TCP server and return the response. -async fn forward_to_target( +/// Forward request data to a target TCP server and return the raw response bytes. +async fn forward_raw_to_target( host: &str, port: u16, stream_id: &str, request_data: Bytes, -) -> Result<(StatusCode, Vec<(String, String)>, Bytes), String> { +) -> Result { // Resolve target let target_addr = resolve_target(host, port) .await @@ -445,13 +448,25 @@ async fn forward_to_target( .map_err(|e| format!("write to target: {}", e))?; } - // Read target response + // Read target response with timeout to avoid blocking on keep-alive connections let mut response_data = Vec::new(); let mut buf = vec![0u8; 32 * 1024]; loop { - match target_stream.read(&mut buf).await { - Err(e) => { - // Connection reset or other error - if we have data, return it + // Use a 5-second timeout to avoid blocking on keep-alive connections + match tokio::time::timeout( + std::time::Duration::from_secs(5), + target_stream.read(&mut buf), + ) + .await + { + Err(_) => { + // Timeout — connection may be keep-alive; return what we have + if !response_data.is_empty() { + break; + } + return Err("read timeout".to_string()); + } + Ok(Err(e)) => { if !response_data.is_empty() { tracing::info!( "Forward stream {}: target connection closed with data read", @@ -461,11 +476,10 @@ async fn forward_to_target( } return Err(format!("read from target: {}", e)); } - Ok(0) => { - // EOF - target closed connection + Ok(Ok(0)) => { break; } - Ok(n) => { + Ok(Ok(n)) => { response_data.extend_from_slice(&buf[..n]); } } @@ -477,55 +491,7 @@ async fn forward_to_target( response_data.len() ); - // Parse HTTP response from the raw data - parse_http_response(Bytes::from(response_data)) -} - -/// Parse raw HTTP response data into status, headers, and body. -#[allow(clippy::type_complexity)] -fn parse_http_response(data: Bytes) -> Result<(StatusCode, Vec<(String, String)>, Bytes), String> { - if data.is_empty() { - return Ok((StatusCode::OK, Vec::new(), Bytes::new())); - } - - // Find the end of headers ("\r\n\r\n") - let header_end = data.windows(4).position(|w| w == b"\r\n\r\n"); - match header_end { - Some(pos) => { - let header_text = String::from_utf8_lossy(&data[..pos + 4]); - let body = data.slice(pos + 4..); - - let mut lines = header_text.lines(); - let status_line = lines.next().ok_or("missing status line")?; - - // Parse "HTTP/1.1 200 OK" - let parts: Vec<&str> = status_line.split_whitespace().collect(); - if parts.len() < 2 { - return Err(format!("invalid status line: {}", status_line)); - } - let status_code: u16 = parts[1] - .parse() - .map_err(|e| format!("invalid status code: {}", e))?; - let status = - StatusCode::from_u16(status_code).map_err(|e| format!("invalid status: {}", e))?; - - let mut headers = Vec::new(); - for line in lines { - if line.is_empty() { - continue; - } - if let Some((name, value)) = line.split_once(':') { - headers.push((name.trim().to_string(), value.trim().to_string())); - } - } - - Ok((status, headers, body)) - } - None => { - // No header/body boundary found - return raw data - Ok((StatusCode::OK, Vec::new(), data)) - } - } + Ok(Bytes::from(response_data)) } // --------------------------------------------------------------------------- @@ -774,92 +740,214 @@ pub async fn reconnect_tunnel(config: &ConnectorConfig) -> Result<(), TunnelErro connect_tunnel(config.clone()).await } +/// Shared state for the persistent tunnel session. +/// Holds the authenticated HTTP sender behind a mutex, protected by +/// an `auth_complete` flag so SOCKS5 fails closed until auth succeeds. +#[allow(dead_code)] +pub struct TunnelSession { + /// Shared HTTP sender — one authenticated persistent tunnel session. + sender: Arc>>>, + /// Becomes true once POST /tunnel auth succeeds. + auth_complete: Arc, + /// Tunnel target host (for logging / Host header). + target_host: String, + /// Monotonically increasing stream counter for this session. + stream_counter: Arc, +} + +impl TunnelSession { + /// Send a raw request payload to the target via the persistent tunnel. + /// Returns the raw response bytes from the target. + async fn forward_raw( + &self, + target: &socks5::Socks5Target, + request_data: Bytes, + ) -> Result { + if !self.auth_complete.load(Ordering::SeqCst) { + return Err(TunnelError::Auth(AuthError::NotAuthenticated)); + } + let target_str = socks5::target_to_string(target); + let stream_id = self.stream_counter.fetch_add(1, Ordering::Relaxed); + let req_len = request_data.len(); + + let mut sender_guard = self.sender.lock().await; + + let req = Request::builder() + .method("POST") + .uri("/forward") + .header("X-Rustunnel-Target", &target_str) + .header("X-Rustunnel-Stream-ID", format!("session-{}", stream_id)) + .header("Host", &self.target_host) + .body(Full::new(request_data)) + .map_err(|e| TunnelError::Protocol(format!("build forward request: {}", e)))?; + + let resp = sender_guard + .send_request(req) + .await + .map_err(|e| TunnelError::Protocol(format!("forward request failed: {}", e)))?; + + let status = resp.status(); + let body = resp.into_body().collect().await.map_err(|e| { + TunnelError::Io(std::io::Error::other(format!("collect response: {}", e))) + })?; + + tracing::debug!( + "Stream {} forwarded {} bytes to {} via shared tunnel (status {})", + stream_id, + req_len, + target_str, + status + ); + + Ok(body.to_bytes()) + } +} + /// Run the connector with SOCKS5 proxy: authenticate through tunnel, then start SOCKS5. /// /// This is the main entry point for the `rustunnel connect` command. /// It: /// 1. Establishes the HTTPS mTLS tunnel connection /// 2. Authenticates via POST /tunnel -/// 3. Marks auth as complete -/// 4. Starts the SOCKS5 proxy listener +/// 3. Stores the authenticated sender in a shared TunnelSession +/// 4. Starts the SOCKS5 proxy listener sharing that tunnel session pub async fn run_connector_with_socks( config: ConnectorConfig, socks_addr: SocketAddr, ) -> Result<(), TunnelError> { - let auth_complete = Arc::new(AtomicBool::new(false)); - let auth_clone = auth_complete.clone(); + // First, authenticate through the tunnel and get the shared sender + let client_config = tls::build_client_config( + config.client_cert_path.as_ref(), + config.client_key_path.as_ref(), + config.ca_cert_path.as_ref(), + ) + .map_err(|e| { + tracing::error!("Failed to build client TLS config: {}", e); + TunnelError::Tls(e) + })?; - // First, authenticate through the tunnel - { - let client_config = tls::build_client_config( - config.client_cert_path.as_ref(), - config.client_key_path.as_ref(), - config.ca_cert_path.as_ref(), - ) + 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| { - tracing::error!("Failed to build client TLS config: {}", e); - TunnelError::Tls(e) + TunnelError::Tls(crate::errors::TlsError::VerificationFailed(format!( + "TLS handshake to {}:{} failed: {}", + config.target_host, config.target_port, e + ))) })?; - let tls_connector = tokio_rustls::TlsConnector::from(client_config); + tracing::info!( + "TLS handshake completed with {}:{} via HTTPS", + config.target_host, + config.target_port + ); - let target_addr = resolve_target(&config.target_host, config.target_port).await?; + // Authenticate — keeps the connection alive in a background task + let (sender, _conn_handle) = + authenticate_https(tls_stream, &config.auth_token, &config.target_host).await?; - let server_name = - tls::server_name_from_host(&config.target_host).map_err(TunnelError::Tls)?; + tracing::info!( + "Tunnel session established with {}:{} via HTTPS /tunnel", + config.target_host, + config.target_port + ); - 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) - } - })?; + // Build the shared session — one persistent tunnel for all SOCKS5 streams + let session = Arc::new(TunnelSession { + sender: Arc::new(Mutex::new(sender)), + auth_complete: Arc::new(AtomicBool::new(true)), + target_host: config.target_host.clone(), + stream_counter: Arc::new(AtomicUsize::new(0)), + }); - let tls_stream = tls_connector - .connect(server_name, stream) - .await - .map_err(|e| { - TunnelError::Tls(crate::errors::TlsError::VerificationFailed(format!( - "TLS handshake to {}:{} failed: {}", - config.target_host, config.target_port, e - ))) - })?; + tracing::info!("Authentication complete — SOCKS5 proxy ready for traffic (persistent tunnel)"); - tracing::info!( - "TLS handshake completed with {}:{} via HTTPS", - config.target_host, - config.target_port - ); - - // Authenticate - let (_sender, _conn_handle) = - authenticate_https(tls_stream, &config.auth_token, &config.target_host).await?; - - tracing::info!( - "Tunnel session established with {}:{} via HTTPS /tunnel", - config.target_host, - config.target_port - ); - - // Mark auth as complete — SOCKS5 can now accept traffic - auth_clone.store(true, Ordering::SeqCst); - tracing::info!("Authentication complete — SOCKS5 proxy ready for traffic"); - } - - // Now start the SOCKS5 proxy (runs until shutdown) - run_socks5_proxy(socks_addr, config, auth_complete).await + // Start the SOCKS5 proxy sharing the tunnel session + run_socks5_proxy_shared(socks_addr, session).await } -/// Run the SOCKS5 proxy listener and forward requests through the HTTPS tunnel. +/// Run the SOCKS5 proxy listener forwarding through a shared persistent tunnel session. /// -/// The SOCKS5 listener binds to the configured address. Each SOCKS5 CONNECT -/// request is forwarded through a fresh HTTPS connection to the tunnel listener. -/// This means each SOCKS stream gets its own TLS + auth cycle, which is simple -/// and correct for lab/dev usage. -/// -/// If `auth_complete` is not yet set when a connection arrives, the SOCKS5 -/// request is rejected (fail-closed). +/// All SOCKS5 streams are multiplexed over the single authenticated HTTPS connection. +pub async fn run_socks5_proxy_shared( + socks_addr: SocketAddr, + session: Arc, +) -> Result<(), TunnelError> { + run_socks5_proxy_impl(socks_addr, session, Arc::new(Notify::new())).await +} + +/// Internal: run the SOCKS5 accept loop. +async fn run_socks5_proxy_impl( + socks_addr: SocketAddr, + session: Arc, + shutdown: Arc, +) -> Result<(), TunnelError> { + let socks_listener = TcpListener::bind(socks_addr).await.map_err(|e| { + let addr_str = socks_addr.to_string(); + if e.kind() == std::io::ErrorKind::AddrInUse { + TunnelError::PortInUse(socks_addr.port()) + } else { + TunnelError::BindError(addr_str, e.to_string()) + } + })?; + + tracing::info!( + "SOCKS5 proxy listening on {} (persistent tunnel session)", + socks_addr + ); + + // Register Ctrl-C handler for graceful shutdown + let shutdown_clone = shutdown.clone(); + tokio::spawn(async move { + tokio::signal::ctrl_c().await.ok(); + tracing::info!("Shutdown signal received — stopping SOCKS5 proxy"); + shutdown_clone.notify_one(); + }); + + loop { + let (stream, peer_addr) = tokio::select! { + result = socks_listener.accept() => { + match result { + Ok(s) => s, + Err(e) => { + tracing::error!("SOCKS5 accept error: {}", e); + continue; + } + } + } + _ = shutdown.notified() => { + tracing::info!("SOCKS5 proxy shutting down"); + return Ok(()); + } + }; + + let sess = session.clone(); + + tokio::spawn(async move { + if let Err(e) = handle_socks5_connection_shared(stream, peer_addr, &sess).await { + tracing::warn!("SOCKS5 connection from {} failed: {}", peer_addr, e); + } + }); + } +} + +/// Legacy: per-request HTTPS connection (retained for backward compat / tests). +/// Deprecated in favor of `run_socks5_proxy_shared`. +#[allow(dead_code)] pub async fn run_socks5_proxy( socks_addr: SocketAddr, tunnel_config: ConnectorConfig, @@ -877,7 +965,6 @@ pub async fn run_socks5_proxy( tracing::info!("SOCKS5 proxy listening on {}", socks_addr); - // Register Ctrl-C handler for graceful shutdown let shutdown = Arc::new(Notify::new()); let shutdown_clone = shutdown.clone(); tokio::spawn(async move { @@ -914,7 +1001,145 @@ pub async fn run_socks5_proxy( } } -/// Handle a single SOCKS5 connection: handshake, auth, CONNECT, forward. +/// Handle a single SOCKS5 connection using the shared persistent tunnel session. +/// +/// All data is forwarded through the shared tunnel — no new HTTPS connections per request. +async fn handle_socks5_connection_shared( + mut stream: TcpStream, + peer_addr: SocketAddr, + session: &Arc, +) -> Result<(), socks5::Socks5Error> { + tracing::info!("SOCKS5 connection from {} (persistent tunnel)", peer_addr); + + // --- SOCKS5 Greeting --- + let greeting = read_exact(&mut stream, 2).await?; + if greeting[0] != socks5::VERSION { + tracing::warn!("SOCKS5 bad version from {}", peer_addr); + stream + .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(&mut stream, n_methods).await?; + + // Choose auth method directly from client methods list + let chosen = if methods.contains(&socks5::AUTH_USERNAME_PASSWORD) { + socks5::AUTH_USERNAME_PASSWORD + } else if methods.contains(&socks5::AUTH_NONE) { + socks5::AUTH_NONE + } else { + stream + .write_all(&[socks5::VERSION, socks5::AUTH_NO_ACCEPTABLE]) + .await?; + return Err(socks5::Socks5Error::UnsupportedAuthMethod( + methods.first().copied().unwrap_or(0xFF), + )); + }; + + // Send method selection response: VER + METHOD (2 bytes only) + stream.write_all(&[socks5::VERSION, chosen]).await?; + + // --- Auth (if username/password) --- + if chosen == socks5::AUTH_USERNAME_PASSWORD { + let auth_buf = read_exact(&mut stream, 2).await?; + if auth_buf[0] != 0x01 { + stream.write_all(&socks5::build_auth_failure()).await?; + return Err(socks5::Socks5Error::AuthFailed); + } + let ulen = auth_buf[1] as usize; + let _username = if ulen > 0 { + let uname_bytes = read_exact(&mut stream, ulen).await?; + String::from_utf8_lossy(&uname_bytes).to_string() + } else { + String::new() + }; + let plen_buf = read_exact(&mut stream, 1).await?; + let plen = plen_buf[0] as usize; + let _password = if plen > 0 { + let pass_bytes = read_exact(&mut stream, plen).await?; + String::from_utf8_lossy(&pass_bytes).to_string() + } else { + String::new() + }; + // Accept any credentials (SOCKS5 auth is handled by the tunnel) + stream.write_all(&socks5::build_auth_success()).await?; + tracing::info!("SOCKS5 auth accepted from {}", peer_addr); + } + + // --- CONNECT Request --- + let connect_req = read_exact(&mut stream, 4).await?; + if connect_req[0] != socks5::VERSION { + return Err(socks5::Socks5Error::InvalidVersion(connect_req[0])); + } + if connect_req[1] != socks5::CMD_CONNECT { + stream + .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(&mut stream, 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(&mut stream, 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(&mut stream, 1).await?; + let domain_len = len_byte[0] as usize; + let domain_bytes = read_exact(&mut stream, domain_len).await?; + let domain = String::from_utf8_lossy(&domain_bytes).to_string(); + let port_bytes = read_exact(&mut stream, 2).await?; + let port = u16::from_be_bytes([port_bytes[0], port_bytes[1]]); + socks5::Socks5Target::Domain { domain, port } + } + _ => { + stream + .write_all(&socks5::build_reply_failure(socks5::REPL_GENERAL_FAILURE)) + .await?; + return Err(socks5::Socks5Error::UnsupportedAddrType(atyp)); + } + }; + + let target_str = socks5::target_to_string(&target); + tracing::info!( + "SOCKS5 CONNECT from {} to {} (via persistent tunnel)", + peer_addr, + target_str + ); + + // --- Forward through shared tunnel --- + match forward_socks_via_session(&mut stream, session, &target).await { + Ok(()) => { + tracing::info!( + "SOCKS5 stream from {} to {} completed (persistent tunnel)", + peer_addr, + target_str + ); + } + Err(e) => { + tracing::warn!( + "SOCKS5 stream from {} to {} failed: {}", + peer_addr, + target_str, + e + ); + return Err(e); + } + } + + Ok(()) +} + +/// Handle a single SOCKS5 connection (legacy per-request HTTPS). async fn handle_socks5_connection( mut stream: TcpStream, peer_addr: SocketAddr, @@ -1061,10 +1286,87 @@ async fn handle_socks5_connection( Ok(()) } +/// Forward a SOCKS5 request through the shared persistent tunnel session. +/// +/// Uses the single authenticated HTTPS connection shared across all SOCKS5 streams. +/// Stream IDs ensure responses are isolated and correctly attributed. +async fn forward_socks_via_session( + socks_stream: &mut TcpStream, + session: &TunnelSession, + target: &socks5::Socks5Target, +) -> Result<(), socks5::Socks5Error> { + // Send CONNECT success reply to SOCKS5 client + socks_stream + .write_all(&socks5::build_reply_success()) + .await?; + + let stream_id = next_stream_id(); + let target_str = socks5::target_to_string(target); + + tracing::info!( + "SOCKS5 stream {}: connected to target {} via persistent tunnel", + stream_id, + target_str + ); + + // Read HTTP request from SOCKS5 client + let (mut socks_read, mut socks_write) = socks_stream.split(); + let request_bytes = read_until_double_crlf(&mut socks_read).await?; + + // Forward via shared tunnel session + let request_bytes_owned: Bytes = request_bytes.into(); + let response_bytes = match session.forward_raw(target, request_bytes_owned).await { + Ok(b) => b, + Err(e) => { + tracing::warn!("Stream {}: forward via tunnel failed: {}", stream_id, e); + return Err(socks5::Socks5Error::Protocol(format!( + "tunnel forward failed: {}", + e + ))); + } + }; + + // Write target response back to SOCKS5 client + if !response_bytes.is_empty() { + socks_write.write_all(&response_bytes).await.map_err(|e| { + tracing::debug!("Stream {} write error: {}", stream_id, e); + socks5::Socks5Error::Io(e) + })?; + } + + // Read any remaining data from SOCKS5 client (for POST bodies etc.) + // Use a short timeout so we don't block indefinitely on keep-alive connections + let mut remaining = Vec::new(); + let read_result = tokio::time::timeout( + std::time::Duration::from_secs(5), + socks_read.read_to_end(&mut remaining), + ) + .await; + if read_result.is_ok() && !remaining.is_empty() { + let remaining_bytes: Bytes = remaining.into(); + let response2 = session + .forward_raw(target, remaining_bytes) + .await + .unwrap_or_default(); + if !response2.is_empty() { + socks_write.write_all(&response2).await.ok(); + } + } + + tracing::info!( + "SOCKS5 stream {}: completed (persistent tunnel, {} bytes response)", + stream_id, + response_bytes.len() + ); + + Ok(()) +} + /// Forward a SOCKS5 request through the HTTPS tunnel. /// /// Opens a new HTTPS connection, authenticates, sends POST /forward, /// and pipes data bidirectionally between the SOCKS5 client and target. +/// Legacy implementation — use `forward_socks_via_session` for persistent tunnel. async fn forward_socks_request( socks_stream: &mut TcpStream, tunnel_config: &ConnectorConfig, @@ -1976,4 +2278,435 @@ mod tests { let result = resolve_host("not-a-valid-hostname-xyz123", 80).await; assert!(result.is_err()); } + + // ========================================================================= + // Tests for persistent tunnel multiplexing + // ========================================================================= + + /// Helper: start an echo server that reads all data then echoes it back. + /// Uses a short inactivity timeout to detect client EOF over keep-alive. + /// Returns (echo_port, task_handle). + async fn spawn_echo_server() -> (u16, tokio::task::JoinHandle<()>) { + let echo_port = get_free_port(); + let echo_addr: SocketAddr = format!("127.0.0.1:{}", echo_port).parse().unwrap(); + let listener = TcpListener::bind(echo_addr).await.unwrap(); + let handle = tokio::spawn(async move { + while let Ok((stream, _)) = listener.accept().await { + tokio::spawn(async move { + let mut buf = vec![0u8; 65536]; + let mut s = stream; + // Collect all data with a short timeout to detect end of request + let mut all_data = Vec::new(); + loop { + match tokio::time::timeout( + std::time::Duration::from_millis(500), + s.read(&mut buf), + ) + .await + { + Err(_) | Ok(Err(_)) | Ok(Ok(0)) => break, + Ok(Ok(n)) => all_data.extend_from_slice(&buf[..n]), + } + } + // Echo everything back + let _ = s.write_all(&all_data).await; + let _ = s.shutdown().await; + }); + } + }); + (echo_port, handle) + } + + /// Helper: build a connector config for the given listener port. + #[allow(dead_code)] + fn build_connector_config(dir: &tempfile::TempDir, port: u16, token: &str) -> ConnectorConfig { + ConnectorConfig { + target_host: "127.0.0.1".to_string(), + target_port: port, + client_cert_path: Arc::from(dir.path().join("client.crt").as_path()), + client_key_path: Arc::from(dir.path().join("client.key").as_path()), + ca_cert_path: Arc::from(dir.path().join("ca.pem").as_path()), + auth_token: Arc::new(token.to_string()), + } + } + + /// Build a TunnelSession by authenticating against a running listener. + async fn build_session( + dir: &tempfile::TempDir, + listen_port: u16, + token: &str, + ) -> Result<(Arc, tokio::task::JoinHandle<()>), TunnelError> { + let client_config = tls::build_client_config( + &dir.path().join("client.crt"), + &dir.path().join("client.key"), + &dir.path().join("ca.pem"), + ) + .map_err(TunnelError::Tls)?; + + let tls_connector = tokio_rustls::TlsConnector::from(client_config); + let target_addr: SocketAddr = format!("127.0.0.1:{}", listen_port).parse().unwrap(); + let stream = TcpStream::connect(target_addr) + .await + .map_err(TunnelError::Io)?; + let server_name = tls::server_name_from_host("127.0.0.1").unwrap(); + let tls_stream = tls_connector + .connect(server_name, stream) + .await + .map_err(|e| { + TunnelError::Tls(crate::errors::TlsError::VerificationFailed(e.to_string())) + })?; + + let (sender, conn_handle) = authenticate_https(tls_stream, token, "127.0.0.1").await?; + + let session = Arc::new(TunnelSession { + sender: Arc::new(Mutex::new(sender)), + auth_complete: Arc::new(AtomicBool::new(true)), + target_host: "127.0.0.1".to_string(), + stream_counter: Arc::new(AtomicUsize::new(0)), + }); + + Ok((session, conn_handle)) + } + + /// Test: multiple sequential forwards over the same persistent session. + /// Proves the session is reused (single connection, multiple streams). + #[tokio::test] + async fn persistent_session_sequential_streams() { + let dir = generate_test_creds(); + let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap(); + + let (echo_port, echo_handle) = spawn_echo_server().await; + tokio::time::sleep(std::time::Duration::from_millis(50)).await; + + // Start listener + let listen_port = get_free_port(); + let bind_addr: SocketAddr = format!("127.0.0.1:{}", listen_port).parse().unwrap(); + let listener_config = ListenerConfig { + bind_addr, + server_cert_path: Arc::from(dir.path().join("server.crt").as_path()), + server_key_path: Arc::from(dir.path().join("server.key").as_path()), + ca_cert_path: Arc::from(dir.path().join("ca.pem").as_path()), + 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; + + // Build session + let (session, _conn_h) = build_session(&dir, listen_port, &token).await.unwrap(); + + // Sequential forwards + for payload_str in ["hello", "world", "multiplex"] { + let payload = Bytes::from(payload_str.as_bytes().to_vec()); + let target = + socks5::Socks5Target::Ipv4(std::net::Ipv4Addr::new(127, 0, 0, 1), echo_port); + let resp = session.forward_raw(&target, payload).await.unwrap(); + assert_eq!( + &resp[..], + payload_str.as_bytes(), + "Echo mismatch for '{}'", + payload_str + ); + } + + listener_task.abort(); + echo_handle.abort(); + } + + /// Test: multiple concurrent forwards over the same persistent session. + /// Each stream gets a unique payload and the responses are correctly isolated. + #[tokio::test] + async fn persistent_session_concurrent_streams_isolated() { + let dir = generate_test_creds(); + let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap(); + + let (echo_port, echo_handle) = spawn_echo_server().await; + tokio::time::sleep(std::time::Duration::from_millis(50)).await; + + let listen_port = get_free_port(); + let bind_addr: SocketAddr = format!("127.0.0.1:{}", listen_port).parse().unwrap(); + let listener_config = ListenerConfig { + bind_addr, + server_cert_path: Arc::from(dir.path().join("server.crt").as_path()), + server_key_path: Arc::from(dir.path().join("server.key").as_path()), + ca_cert_path: Arc::from(dir.path().join("ca.pem").as_path()), + 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 (session, _conn_h) = build_session(&dir, listen_port, &token).await.unwrap(); + + // Launch 5 concurrent forwards, each with a unique payload + let payloads = vec![ + "stream-alpha-payload".to_string(), + "stream-beta-payload".to_string(), + "stream-gamma-payload".to_string(), + "stream-delta-payload".to_string(), + "stream-epsilon-payload".to_string(), + ]; + let target = socks5::Socks5Target::Ipv4(std::net::Ipv4Addr::new(127, 0, 0, 1), echo_port); + + let mut handles = Vec::new(); + for payload_str in &payloads { + let sess = session.clone(); + let t = target.clone(); + let p = Bytes::from(payload_str.as_bytes().to_vec()); + handles.push(tokio::spawn(async move { sess.forward_raw(&t, p).await })); + } + + // Collect results and verify each response matches its payload + for (i, handle) in handles.into_iter().enumerate() { + let resp = handle.await.unwrap().unwrap(); + assert_eq!( + &resp[..], + payloads[i].as_bytes(), + "Stream {} response mismatch: expected '{}', got '{}'", + i, + payloads[i], + String::from_utf8_lossy(&resp) + ); + } + + listener_task.abort(); + echo_handle.abort(); + } + + /// Test: stream counter increments monotonically across concurrent streams. + #[tokio::test] + async fn persistent_session_stream_counter_increments() { + let dir = generate_test_creds(); + let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap(); + + let listen_port = get_free_port(); + let bind_addr: SocketAddr = format!("127.0.0.1:{}", listen_port).parse().unwrap(); + let listener_config = ListenerConfig { + bind_addr, + server_cert_path: Arc::from(dir.path().join("server.crt").as_path()), + server_key_path: Arc::from(dir.path().join("server.key").as_path()), + ca_cert_path: Arc::from(dir.path().join("ca.pem").as_path()), + 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 (session, _conn_h) = build_session(&dir, listen_port, &token).await.unwrap(); + + // Check initial counter + let initial = session.stream_counter.load(Ordering::SeqCst); + assert_eq!(initial, 0, "Counter should start at 0"); + + // Do some forwards + let target = socks5::Socks5Target::Ipv4(std::net::Ipv4Addr::new(127, 0, 0, 1), 4181); + let mut handles = Vec::new(); + for _ in 0..5 { + let sess = session.clone(); + let t = target.clone(); + handles.push(tokio::spawn(async move { + sess.forward_raw(&t, Bytes::from(vec![0u8; 10])).await + })); + } + // Wait for all handles (they may fail — that's fine, we just care about the counter) + for h in handles { + let _ = h.await; + } + + // Counter should have incremented + let final_counter = session.stream_counter.load(Ordering::SeqCst); + assert!( + final_counter >= 5, + "Counter should be >= 5, got {}", + final_counter + ); + + listener_task.abort(); + } + + /// Test: TunnelSession rejects forwards when auth is not complete. + #[tokio::test] + async fn tunnel_session_fails_before_auth() { + // Create a real sender via a dummy listener so the TunnelSession is well-formed. + // The sender won't actually be used because auth_complete is false. + let (sender, _) = create_placeholder_sender().await; + let session = Arc::new(TunnelSession { + sender: Arc::new(Mutex::new(sender)), + auth_complete: Arc::new(AtomicBool::new(false)), + target_host: "127.0.0.1".to_string(), + stream_counter: Arc::new(AtomicUsize::new(0)), + }); + + let target = socks5::Socks5Target::Ipv4(std::net::Ipv4Addr::new(127, 0, 0, 1), 4181); + let result = session + .forward_raw(&target, Bytes::from(vec![0u8; 10])) + .await; + assert!( + result.is_err(), + "Forward should fail before auth is complete" + ); + assert!( + matches!(result, Err(TunnelError::Auth(AuthError::NotAuthenticated))), + "Should be NotAuthenticated error" + ); + } + + /// Helper: create a placeholder sender (never used, just for type safety) within the async runtime. + async fn create_placeholder_sender() -> ( + hyper::client::conn::http1::SendRequest>, + tokio::task::JoinHandle<()>, + ) { + // Create a dummy listener to accept a connection + let dummy_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let dummy_addr = dummy_listener.local_addr().unwrap(); + + // Spawn an acceptor + let _acceptor = tokio::spawn(async move { + let (stream, _) = dummy_listener.accept().await.unwrap(); + let io = hyper_util::rt::TokioIo::new(stream); + let svc = hyper::service::service_fn(|_req: Request| async { + let body: http_body_util::combinators::BoxBody = + BoxBody::new(Full::new(Bytes::new()).map_err(|_| unreachable!())); + Ok::<_, hyper::http::Error>(Response::new(body)) + }); + let _ = hyper::server::conn::http1::Builder::new() + .serve_connection(io, svc) + .await; + }); + + // Connect to it + let stream = TcpStream::connect(dummy_addr).await.unwrap(); + let io = hyper_util::rt::TokioIo::new(stream); + let (sender, conn) = hyper::client::conn::http1::handshake(io).await.unwrap(); + let conn_handle = tokio::spawn(async move { + let _ = conn.await; + }); + (sender, conn_handle) + } + + /// Test: E2E concurrent SOCKS5 streams share a single tunnel session. + /// This is the main integration test proving VAL-SOCKS-006. + #[tokio::test] + async fn e2e_concurrent_socks5_streams_shared_tunnel() { + let dir = generate_test_creds(); + let token = std::fs::read_to_string(dir.path().join("token.txt")).unwrap(); + + // Start echo server + let (echo_port, echo_handle) = spawn_echo_server().await; + tokio::time::sleep(std::time::Duration::from_millis(50)).await; + + // Start listener + let listen_port = get_free_port(); + let bind_addr: SocketAddr = format!("127.0.0.1:{}", listen_port).parse().unwrap(); + let listener_config = ListenerConfig { + bind_addr, + server_cert_path: Arc::from(dir.path().join("server.crt").as_path()), + server_key_path: Arc::from(dir.path().join("server.key").as_path()), + ca_cert_path: Arc::from(dir.path().join("ca.pem").as_path()), + 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; + + // Build shared session + let (session, _conn_h) = build_session(&dir, listen_port, &token).await.unwrap(); + + // Start SOCKS5 proxy on a unique port + let socks_port = get_free_port(); + let socks_addr: SocketAddr = format!("127.0.0.1:{}", socks_port).parse().unwrap(); + let session_clone = session.clone(); + let socks_shutdown = Arc::new(Notify::new()); + let socks_shutdown_clone = socks_shutdown.clone(); + let socks_task = tokio::spawn(async move { + run_socks5_proxy_impl(socks_addr, session_clone, socks_shutdown_clone).await + }); + + tokio::time::sleep(std::time::Duration::from_millis(100)).await; + + // Launch 3 concurrent SOCKS5 connections via raw TCP, each sending a unique HTTP request + let payloads = vec![ + "GET /a HTTP/1.1\r\nHost: x\r\n\r\n".to_string(), + "GET /b HTTP/1.1\r\nHost: x\r\n\r\n".to_string(), + "GET /c HTTP/1.1\r\nHost: x\r\n\r\n".to_string(), + ]; + + let mut handles = Vec::new(); + for payload_str in &payloads { + let p = payload_str.clone(); + handles.push(tokio::spawn(async move { + // Connect to SOCKS5 + let mut stream = TcpStream::connect(format!("127.0.0.1:{}", socks_port)) + .await + .unwrap(); + + // SOCKS5 greeting + stream.write_all(&[0x05, 0x01, 0x00]).await.unwrap(); + let mut resp = [0u8; 2]; + stream.read_exact(&mut resp).await.unwrap(); + assert_eq!(resp, [0x05, 0x00]); + + // CONNECT request — target is echo server + let _target = format!("{}:{}", "127.0.0.1", echo_port); + let mut req = vec![ + 0x05, 0x01, 0x00, 0x01, // VER CMD RSV ATYP=IPv4 + ]; + req.extend([127, 0, 0, 1]); + req.extend(&echo_port.to_be_bytes()); + stream.write_all(&req).await.unwrap(); + + // Read CONNECT reply + let mut reply = [0u8; 10]; + stream.read_exact(&mut reply).await.unwrap(); + assert_eq!(reply[0], 0x05); // VER + assert_eq!(reply[1], 0x00); // SUCCEEDED + + // Send HTTP request to echo server + stream.write_all(p.as_bytes()).await.unwrap(); + + // Read response + let mut buf = Vec::new(); + let mut chunk = vec![0u8; 4096]; + tokio::time::timeout(std::time::Duration::from_secs(5), async { + loop { + let n = match tokio::time::timeout( + std::time::Duration::from_secs(2), + stream.read(&mut chunk), + ) + .await + { + Ok(Ok(n)) => n, + _ => break, + }; + if n == 0 { + break; + } + buf.extend_from_slice(&chunk[..n]); + } + }) + .await + .unwrap(); + + String::from_utf8_lossy(&buf).to_string() + })); + } + + // Collect and verify results + let results: Vec = futures::future::join_all(handles) + .await + .into_iter() + .map(|r| r.unwrap()) + .collect(); + + for (i, result) in results.iter().enumerate() { + assert!( + result.contains(&payloads[i][4..payloads[i].len() - 4]), + "Stream {} response should contain payload: got '{}'", + i, + result + ); + } + + // Shutdown SOCKS5 proxy + socks_shutdown.notify_one(); + let _ = socks_task.await; + listener_task.abort(); + echo_handle.abort(); + } }