diff --git a/CHANGELOG.md b/CHANGELOG.md index be2eb8e..fb9b3ea 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,8 @@ ### Added - Add `--channel-open-timeout` CLI flag. +- Add `--tcp-keepalive-time` CLI flag. +- Add `--tcp-keepalive-interval` CLI flag. ### Fixed diff --git a/Cargo.lock b/Cargo.lock index 47bf57f..46b1b8e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4454,7 +4454,7 @@ dependencies = [ [[package]] name = "sandhole" -version = "0.10.0" +version = "0.10.1" dependencies = [ "ahash", "async-speed-limit", @@ -4506,6 +4506,7 @@ dependencies = [ "sandhole_socket", "serde", "serde_json", + "socket2 0.6.3", "sysinfo", "test-log", "thiserror 2.0.18", @@ -4526,7 +4527,7 @@ dependencies = [ [[package]] name = "sandhole_socket" -version = "0.1.0" +version = "0.1.1" dependencies = [ "color-eyre", "socket2 0.6.3", @@ -4535,13 +4536,14 @@ dependencies = [ [[package]] name = "sandhole_udp_over_tcp" -version = "0.1.0" +version = "0.1.1" dependencies = [ "ahash", "clap", "color-eyre", "dashmap", "sandhole_socket", + "socket2 0.6.3", "tokio", ] diff --git a/Cargo.toml b/Cargo.toml index f3607aa..090a3f0 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "sandhole" -version = "0.10.0" +version = "0.10.1" edition = "2024" rust-version = "1.88.0" description = "Expose HTTP/SSH/TCP services through SSH port forwarding." @@ -21,7 +21,7 @@ ahash = "0.8.12" clap = { version = "4.6.1", features = ["derive", "string"] } color-eyre = "0.6.5" dashmap = "6.1.0" -sandhole_socket = { version = "0.1.0", path = "sandhole_socket" } +sandhole_socket = { version = "0.1.1", path = "sandhole_socket" } socket2 = "0.6.3" tokio = { version = "1.52.3", features = ["full"] } @@ -100,6 +100,7 @@ serde = { version = "1.0.228", default-features = false, features = [ "serde_derive", ] } serde_json = "1.0.149" +socket2 = { workspace = true } sysinfo = "0.38.4" # 0.39.x requires MSRV bump to 1.95.0 thiserror = "2.0.18" tokio = { workspace = true } diff --git a/book/src/cli.md b/book/src/cli.md index 89a8bab..92e0444 100644 --- a/book/src/cli.md +++ b/book/src/cli.md @@ -426,6 +426,23 @@ Expose HTTP/SSH/TCP services through SSH port forwarding. By default, these connections are not terminated by Sandhole. + --tcp-keepalive-time <DURATION> + If set, enables TCP keepalive on accepted client connections, + sending the first probe after the set idle time. + + Set this option to avoid "connection reset by peer" on + socket reuse. + + By default, keepalive is disabled. + + --tcp-keepalive-interval <DURATION> + Interval between TCP keepalive probes once `--tcp-keepalive-time` + has elapsed. + + Only applies when `--tcp-keepalive-time` is set. + + [default: 10s] + --udp-timeout <DURATION> How long until SSH channels from UDP sockets are automatically garbage-collected diff --git a/nix/default.nix b/nix/default.nix index 60258a6..d7981fb 100644 --- a/nix/default.nix +++ b/nix/default.nix @@ -42,7 +42,6 @@ let inherit cargoArtifacts; doCheck = false; postInstall = lib.optionalString (pkgs.stdenv.buildPlatform.canExecute pkgs.stdenv.hostPlatform) '' - $out/bin/sandhole --completions bash installShellCompletion --cmd sandhole \ --bash <($out/bin/sandhole --completions bash) \ --fish <($out/bin/sandhole --completions fish) \ diff --git a/nix/modules/sandhole.nix b/nix/modules/sandhole.nix index 203c6d4..a79fa2c 100644 --- a/nix/modules/sandhole.nix +++ b/nix/modules/sandhole.nix @@ -210,7 +210,10 @@ in systemd.services.sandhole = { description = "Sandhole - Expose HTTP/SSH/TCP services through SSH port forwarding"; wantedBy = [ "multi-user.target" ]; - after = [ "network-online.target" ]; + after = [ + "network-online.target" + "nss-lookup.target" + ]; wants = [ "network-online.target" ]; serviceConfig = { User = cfg.user; diff --git a/sandhole_socket/Cargo.toml b/sandhole_socket/Cargo.toml index 2d16874..2d2d270 100644 --- a/sandhole_socket/Cargo.toml +++ b/sandhole_socket/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "sandhole_socket" -version = "0.1.0" +version = "0.1.1" edition = "2024" rust-version = "1.88.0" description = "Sandhole utilities for handling TCP and UDP sockets." diff --git a/sandhole_socket/src/tcp_listener.rs b/sandhole_socket/src/tcp_listener.rs index 996cc2e..463f323 100644 --- a/sandhole_socket/src/tcp_listener.rs +++ b/sandhole_socket/src/tcp_listener.rs @@ -1,6 +1,6 @@ use std::{io, net::ToSocketAddrs}; -use socket2::{Domain, Socket, Type}; +use socket2::{Domain, Socket, TcpKeepalive, Type}; use tokio::net::TcpListener; /// Create an async TCP listener with Nagle's algorithm disabled @@ -50,3 +50,56 @@ pub fn get_tcp_listener(addr: A) -> io::Result { TcpListener::from_std(socket.into()) } + +/// Create an async TCP listener with Nagle's algorithm disabled, +/// TCP keepalive, and any necessary configurations for dualstack. +pub fn get_tcp_listener_with_keepalive( + addr: A, + keepalive: TcpKeepalive, +) -> io::Result { + let addr = addr.to_socket_addrs()?.next().ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidInput, + "could not resolve to any address", + ) + })?; + let is_ipv6 = addr.is_ipv6(); + + let socket = Socket::new( + if is_ipv6 { Domain::IPV6 } else { Domain::IPV4 }, + Type::STREAM, + None, + )?; + + socket.set_nonblocking(true)?; + socket.set_tcp_nodelay(true)?; + if is_ipv6 { + socket.set_only_v6(false)?; + } + + // On platforms with Berkeley-derived sockets, this allows to quickly + // rebind a socket, without needing to wait for the OS to clean up the + // previous one. + // + // On Windows, this allows rebinding sockets which are actively in use, + // which allows “socket hijacking”, so we explicitly don't set it here. + // https://docs.microsoft.com/en-us/windows/win32/winsock/using-so-reuseaddr-and-so-exclusiveaddruse + #[cfg(not(windows))] + socket.set_reuse_address(true)?; + + socket.set_tcp_keepalive(&keepalive)?; + + socket.bind(&addr.into())?; + socket.listen({ + #[cfg(not(windows))] + { + -1 + } + #[cfg(windows)] + { + 128 + } + })?; + + TcpListener::from_std(socket.into()) +} diff --git a/src/config.rs b/src/config.rs index a0815a7..77d24c4 100644 --- a/src/config.rs +++ b/src/config.rs @@ -540,6 +540,26 @@ pub struct ApplicationConfig { #[arg(long, value_parser = validate_duration, value_name = "DURATION")] pub tcp_connection_timeout: Option, + /// If set, enables TCP keepalive on accepted client connections, + /// sending the first probe after the set idle time. + /// + /// Set this option to avoid "connection reset by peer" on socket reuse. + /// + /// By default, keepalive is disabled. + #[arg(long, value_parser = validate_duration, value_name = "DURATION")] + pub tcp_keepalive_time: Option, + + /// Interval between TCP keepalive probes once `--tcp-keepalive-time` has elapsed. + /// + /// Only applies when `--tcp-keepalive-time` is set. + #[arg( + long, + default_value = "10s", + value_parser = validate_duration, + value_name = "DURATION" + )] + pub tcp_keepalive_interval: Duration, + /// How long until SSH channels from UDP sockets are automatically garbage-collected. #[arg( long, @@ -670,6 +690,8 @@ mod application_config_tests { authentication_request_timeout: Duration::from_secs(5), http_request_timeout: None, tcp_connection_timeout: None, + tcp_keepalive_time: None, + tcp_keepalive_interval: Duration::from_secs(10), udp_timeout: Duration::from_secs(60), } ) @@ -741,6 +763,8 @@ mod application_config_tests { "--authentication-request-timeout=6s", "--http-request-timeout=15s", "--tcp-connection-timeout=30s", + "--tcp-keepalive-time=15s", + "--tcp-keepalive-interval=5s", "--udp-timeout=30s", ]); assert_eq!( @@ -815,6 +839,8 @@ mod application_config_tests { authentication_request_timeout: Duration::from_secs(6), http_request_timeout: Some(Duration::from_secs(15)), tcp_connection_timeout: Some(Duration::from_secs(30)), + tcp_keepalive_time: Some(Duration::from_secs(15)), + tcp_keepalive_interval: Duration::from_secs(5), udp_timeout: Duration::from_secs(30), } ) diff --git a/src/entrypoint.rs b/src/entrypoint.rs index 00c4fea..c4cf7f4 100644 --- a/src/entrypoint.rs +++ b/src/entrypoint.rs @@ -25,7 +25,8 @@ use rustls::ServerConfig; #[cfg(feature = "acme")] use rustls_acme::acme::ACME_TLS_ALPN_NAME; use rustls_pki_types::ServerName; -use sandhole_socket::tcp_listener::get_tcp_listener; +use sandhole_socket::tcp_listener::{get_tcp_listener, get_tcp_listener_with_keepalive}; +use socket2::TcpKeepalive; use sysinfo::{CpuRefreshKind, MemoryRefreshKind, Networks, RefreshKind, System}; #[cfg_attr(not(feature = "acme"), allow(unused_imports))] use tokio::io::AsyncWriteExt; @@ -264,6 +265,11 @@ pub async fn entrypoint(config: ApplicationConfig) -> color_eyre::Result<()> { Some(max_quota) => Arc::new(Box::new(Arc::new(QuotaMap::new(max_quota.into())))), None => Arc::new(Box::new(DummyQuotaHandler)), }; + let tcp_keepalive = config.tcp_keepalive_time.map(|tcp_keepalive_time| { + TcpKeepalive::new() + .with_time(tcp_keepalive_time) + .with_interval(config.tcp_keepalive_interval) + }); let http_connections = Arc::new( ConnectionMap::builder() .strategy(config.load_balancing) @@ -327,6 +333,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> color_eyre::Result<()> { .ip_filter(Arc::clone(&ip_filter)) .buffer_size(buffer_size) .disable_tcp_logs(config.disable_tcp_logs) + .maybe_tcp_keepalive(tcp_keepalive.clone()) .maybe_tcp_connection_timeout(tcp_connection_timeout) .build(), ); @@ -818,8 +825,17 @@ pub async fn entrypoint(config: ApplicationConfig) -> color_eyre::Result<()> { let mut join_handle_http = if config.disable_http { DroppableHandle(tokio::spawn(future::pending())) } else { - let http_listener = get_tcp_listener((config.listen_address, config.http_port.into())) - .with_context(|| "Error listening to HTTP port")?; + let http_listener = { + if let Some(keepalive) = tcp_keepalive.clone() { + get_tcp_listener_with_keepalive( + (config.listen_address, config.http_port.into()), + keepalive, + ) + } else { + get_tcp_listener((config.listen_address, config.http_port.into())) + } + } + .with_context(|| "Error listening to HTTP port")?; #[cfg(not(coverage_nightly))] tracing::info!( "Listening for HTTP connections on port {}.", @@ -907,8 +923,15 @@ pub async fn entrypoint(config: ApplicationConfig) -> color_eyre::Result<()> { let mut join_handle_https = if config.disable_http || config.disable_https { DroppableHandle(tokio::spawn(future::pending())) } else { - let https_listener = get_tcp_listener((config.listen_address, config.https_port.into())) - .with_context(|| "Error listening to HTTPS port")?; + let https_listener = if let Some(keepalive) = tcp_keepalive.clone() { + get_tcp_listener_with_keepalive( + (config.listen_address, config.https_port.into()), + keepalive, + ) + } else { + get_tcp_listener((config.listen_address, config.https_port.into())) + } + .with_context(|| "Error listening to HTTPS port")?; #[cfg(not(coverage_nightly))] tracing::info!( "Listening for HTTPS connections on port {}.", @@ -993,8 +1016,12 @@ pub async fn entrypoint(config: ApplicationConfig) -> color_eyre::Result<()> { }; // Start Sandhole on SSH port - let ssh_listener = get_tcp_listener((config.listen_address, config.ssh_port.into())) - .with_context(|| "Error listening to SSH port")?; + let ssh_listener = if let Some(keepalive) = tcp_keepalive { + get_tcp_listener_with_keepalive((config.listen_address, config.ssh_port.into()), keepalive) + } else { + get_tcp_listener((config.listen_address, config.ssh_port.into())) + } + .with_context(|| "Error listening to SSH port")?; #[cfg(not(coverage_nightly))] tracing::info!("Listening for SSH connections on port {}.", config.ssh_port); #[cfg(not(coverage_nightly))] diff --git a/src/tcp.rs b/src/tcp.rs index 5cac2a2..cc73c70 100644 --- a/src/tcp.rs +++ b/src/tcp.rs @@ -14,7 +14,8 @@ use bon::Builder; use color_eyre::eyre::Context; use dashmap::DashMap; use metrics::counter; -use sandhole_socket::tcp_listener::get_tcp_listener; +use sandhole_socket::tcp_listener::{get_tcp_listener, get_tcp_listener_with_keepalive}; +use socket2::TcpKeepalive; use tokio::{io::copy_bidirectional_with_sizes, time::timeout}; // Service that handles creating TCP sockets for reverse forwarding connections. @@ -31,6 +32,8 @@ pub(crate) struct TcpHandler { ip_filter: Arc, // Buffer size for bidirectional copying. buffer_size: usize, + // Parameters (delay + interval) to keep TCP client connections alive. + tcp_keepalive: Option, // Optional duration to time out TCP connections. tcp_connection_timeout: Option, // Whether to send TCP logs to the SSH handles behind the forwarded connections. @@ -50,7 +53,11 @@ impl TcpPortHandler for Arc { return Ok(port); } // Check if we're able to bind to the given address and port. - let listener = get_tcp_listener((self.listen_address, port))?; + let listener = if let Some(keepalive) = self.tcp_keepalive.clone() { + get_tcp_listener_with_keepalive((self.listen_address, port), keepalive) + } else { + get_tcp_listener((self.listen_address, port)) + }?; let port = listener .local_addr() .with_context(|| "Missing local address when binding port")? diff --git a/udp_over_tcp/Cargo.toml b/udp_over_tcp/Cargo.toml index 096bbf8..5b5a559 100644 --- a/udp_over_tcp/Cargo.toml +++ b/udp_over_tcp/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "sandhole_udp_over_tcp" -version = "0.1.0" +version = "0.1.1" edition = "2024" rust-version = "1.88.0" description = "Proxy UDP traffic for Sandhole via SSH." @@ -18,4 +18,5 @@ clap = { workspace = true } color-eyre = { workspace = true } dashmap = { workspace = true } sandhole_socket = { workspace = true } +socket2 = { workspace = true } tokio = { workspace = true } diff --git a/udp_over_tcp/src/main.rs b/udp_over_tcp/src/main.rs index 7adf31b..a61e906 100644 --- a/udp_over_tcp/src/main.rs +++ b/udp_over_tcp/src/main.rs @@ -4,6 +4,7 @@ use std::{ num::NonZeroU16, pin::{Pin, pin}, sync::Arc, + time::Duration, }; use ahash::RandomState; @@ -11,10 +12,11 @@ use clap::Parser; use color_eyre::eyre::eyre; use dashmap::DashMap; use sandhole_socket::{ - tcp_listener::get_tcp_listener, + tcp_listener::get_tcp_listener_with_keepalive, udp_listener::get_udp_socket, udp_over_tcp::{datagram_buffer, deserialize_datagram, serialize_datagram}, }; +use socket2::TcpKeepalive; use tokio::{io::AsyncWriteExt, net::TcpSocket, select, sync::mpsc}; #[doc(hidden)] @@ -46,11 +48,16 @@ type FutType = Pin> + Send + Sync // Starts a TCP server that redirects requests to the UDP socket. async fn remote_forwarding(args: Cli) -> color_eyre::Result<()> { - let listener = get_tcp_listener(( - args.tcp_address - .unwrap_or(IpAddr::V6(Ipv6Addr::UNSPECIFIED)), - u16::from(args.tcp_port), - ))?; + let listener = get_tcp_listener_with_keepalive( + ( + args.tcp_address + .unwrap_or(IpAddr::V6(Ipv6Addr::UNSPECIFIED)), + u16::from(args.tcp_port), + ), + TcpKeepalive::new() + .with_time(Duration::from_secs(15)) + .with_interval(Duration::from_secs(15)), + )?; let udp_address = ( args.udp_address.unwrap_or(IpAddr::V4(Ipv4Addr::LOCALHOST)), u16::from(args.udp_port),