diff --git a/CHANGELOG.md b/CHANGELOG.md index 0da4dd0..a1906a7 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,11 @@ # Changelog +## Unreleased + +### Added + +- Add `--buffer-size` CLI flag. + ## 0.5.1 (2025-03-31) ### Added diff --git a/Cargo.lock b/Cargo.lock index a9e35ce..b8e8ea2 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -570,6 +570,12 @@ version = "1.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d71b6127be86fdcfddb610f7182ac57211d4b18a3e9c82eb2d17662f2227ad6a" +[[package]] +name = "bytesize" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a3c8f83209414aacf0eeae3cf730b18d6981697fba62f200fcfb92b9f082acba" + [[package]] name = "cassowary" version = "0.3.0" @@ -1817,12 +1823,6 @@ version = "1.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" -[[package]] -name = "human_bytes" -version = "0.4.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "91f255a4535024abf7640cb288260811fc14794f62b063652ed349f9a6c2348e" - [[package]] name = "humantime" version = "2.2.0" @@ -3785,6 +3785,7 @@ dependencies = [ "axum", "block-id", "bytes", + "bytesize", "chrono", "clap", "crossterm", @@ -3795,7 +3796,6 @@ dependencies = [ "hickory-resolver 0.25.1", "http", "http-body-util", - "human_bytes", "humantime", "hyper", "hyper-util", diff --git a/Cargo.toml b/Cargo.toml index e140bb5..d2f29b6 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -19,6 +19,7 @@ aws-lc-rs = "1.12.6" axum = { version = "0.8.1", default-features = false } block-id = "0.2.1" bytes = "1.10.1" +bytesize = "2.0.1" chrono = "0.4.39" clap = { version = "4.5.32", features = ["derive", "string"] } crossterm = { version = "0.28.1", default-features = false } @@ -31,7 +32,6 @@ env_logger = { version = "0.11.7", default-features = false, features = [ hickory-resolver = "0.25.1" http = "1.3.1" http-body-util = "0.1.3" -human_bytes = { version = "0.4.3", default-features = false } humantime = "2.2.0" hyper = { version = "1.6.0", features = ["full"] } hyper-util = { version = "0.1.10", features = ["full"] } diff --git a/book/src/configuration.md b/book/src/configuration.md index 212b3d7..6228161 100644 --- a/book/src/configuration.md +++ b/book/src/configuration.md @@ -26,7 +26,7 @@ sandhole --domain server.com --ssh-port 22 Without extra configuration, Sandhole will not let users bind to requested subdomains and ports, and will always allocate a random one instead. -If you wish to change the default behavior, and allow users to provide their own subdomains/ports to bind to, add the options `--allow-provided-subdomains` and `--allow-requested-ports`, respectively. +If you wish to change the default behavior, and allow users to provide their own subdomains/ports to bind to, add the options `--allow-requested-subdomains` and `--allow-requested-ports`, respectively. Otherwise, if you wish the subdomains to still be random, but persist between requests/disconnections, check out the `--random-subdomain-seed` option in the [command-line interface](./cli.md). diff --git a/src/admin.rs b/src/admin.rs index 39bd7d9..c684767 100644 --- a/src/admin.rs +++ b/src/admin.rs @@ -6,7 +6,7 @@ use std::{ time::Duration, }; -use human_bytes::human_bytes; +use bytesize::ByteSize; use itertools::Itertools; use ratatui::{ Terminal, TerminalOptions, Viewport, @@ -533,18 +533,18 @@ impl AdminState { " Memory ".bold().reversed(), format!( " {} / {}", - human_bytes(used_memory as f64), - human_bytes(total_memory as f64) + ByteSize::b(used_memory).display().iec_short(), + ByteSize::b(total_memory).display().iec_short(), ) .into(), ]); let network_tx = Line::from(vec![ " TX ".bold().reversed(), - format!(" {}/s", human_bytes(network_tx as f64)).into(), + format!(" {}/s", ByteSize::b(network_tx).display().iec_short()).into(), ]); let network_rx = Line::from(vec![ " RX ".bold().reversed(), - format!(" {}/s", human_bytes(network_rx as f64)).into(), + format!(" {}/s", ByteSize::b(network_rx).display().iec_short()).into(), ]); Widget::render(block, area, buf); Widget::render(cpu_usage, cpu_area, buf); diff --git a/src/config.rs b/src/config.rs index d7ffdbc..6bb9ace 100644 --- a/src/config.rs +++ b/src/config.rs @@ -284,6 +284,17 @@ pub struct ApplicationConfig { #[arg(long, value_delimiter = ',', value_name = "CIDR")] pub ip_blocklist: Option>, + /// Size to use for bidirectional buffers. + /// + /// A higher value will lead to higher memory consumption. + #[arg( + long, + default_value = "8KB", + value_parser = validate_byte_size, + value_name = "SIZE" + )] + pub buffer_size: usize, + /// Grace period for dangling/unauthenticated SSH connections before they are forcefully disconnected. /// /// A low value may cause valid proxy/tunnel connections to be erroneously removed. @@ -346,6 +357,14 @@ fn validate_duration(value: &str) -> Result { .into()) } +fn validate_byte_size(value: &str) -> Result { + Ok(bytesize::ByteSize::from_str(value) + .map_err(|_| "invalid byte size")? + .as_u64() + .try_into() + .map_err(|_| "cannot convert to usize")?) +} + #[cfg(test)] #[cfg_attr(coverage_nightly, coverage(off))] mod application_config_tests { @@ -397,6 +416,7 @@ mod application_config_tests { requested_domain_filter_profanities: false, ip_allowlist: None, ip_blocklist: None, + buffer_size: 8_000, idle_connection_timeout: Duration::from_secs(2), unproxied_connection_timeout: None, authentication_request_timeout: Duration::from_secs(5), @@ -445,6 +465,7 @@ mod application_config_tests { "--requested-domain-filter-profanities", "--ip-allowlist=10.0.0.0/8", "--ip-blocklist=10.1.0.0/16,10.2.0.0/16", + "--buffer-size=4KB", "--idle-connection-timeout=3s", "--unproxied-connection-timeout=4s", "--authentication-request-timeout=6s", @@ -492,6 +513,7 @@ mod application_config_tests { IpNet::from_str("10.1.0.0/16").unwrap(), IpNet::from_str("10.2.0.0/16").unwrap() ]), + buffer_size: 4_000, idle_connection_timeout: Duration::from_secs(3), unproxied_connection_timeout: Some(Duration::from_secs(4)), authentication_request_timeout: Duration::from_secs(6), @@ -529,4 +551,16 @@ mod application_config_tests { .is_err() ); } + + #[test] + fn fails_to_parse_if_invalid_byte_size() { + assert!( + ApplicationConfig::try_parse_from([ + "sandhole", + "--domain=foobar.tld", + "--buffer_size=42" + ]) + .is_err() + ); + } } diff --git a/src/http.rs b/src/http.rs index 79f1ca8..c6e525c 100644 --- a/src/http.rs +++ b/src/http.rs @@ -26,7 +26,7 @@ use hyper_util::rt::{TokioExecutor, TokioIo}; use log::{debug, warn}; use russh::keys::ssh_key::Fingerprint; use tokio::{ - io::{AsyncRead, AsyncWrite, copy_bidirectional}, + io::{AsyncRead, AsyncWrite, copy_bidirectional_with_sizes}, time::timeout, }; @@ -124,6 +124,8 @@ where pub(crate) protocol: Protocol, // Configuration on which type of channel to retrieve from the handler. pub(crate) proxy_type: ProxyType, + // Buffer size for bidirectional copying. + pub(crate) buffer_size: usize, // Optional duration until an outgoing request is canceled. pub(crate) http_request_timeout: Option, // Optional duration until an established Websocket connection is canceled. @@ -410,6 +412,7 @@ where // Retrieve the upgraded connection from the response let upgraded_response = hyper::upgrade::on(&mut response).await?; let websocket_timeout = proxy_data.websocket_timeout; + let buffer_size = proxy_data.buffer_size; // Start a task to copy data between the two Upgraded parts tokio::spawn(async move { let mut upgraded_request = @@ -419,9 +422,11 @@ where // If there is a Websocket timeout, copy until the deadline is reached. Some(duration) => { let _ = timeout(duration, async { - copy_bidirectional( + copy_bidirectional_with_sizes( &mut upgraded_response, &mut upgraded_request, + buffer_size, + buffer_size, ) .await }) @@ -429,9 +434,11 @@ where } // If there isn't a Websocket timeout, copy data between both sides unconditionally. None => { - let _ = copy_bidirectional( + let _ = copy_bidirectional_with_sizes( &mut upgraded_response, &mut upgraded_request, + buffer_size, + buffer_size, ) .await; } @@ -548,6 +555,7 @@ mod proxy_handler_tests { }), protocol: Protocol::Http { port: 80 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: None, websocket_timeout: None, disable_http_logs: false, @@ -591,6 +599,7 @@ mod proxy_handler_tests { }), protocol: Protocol::Http { port: 80 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: None, websocket_timeout: None, disable_http_logs: false, @@ -634,6 +643,7 @@ mod proxy_handler_tests { }), protocol: Protocol::Http { port: 80 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: None, websocket_timeout: None, disable_http_logs: false, @@ -699,6 +709,7 @@ mod proxy_handler_tests { }), protocol: Protocol::TlsRedirect { from: 80, to: 443 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: None, websocket_timeout: None, disable_http_logs: false, @@ -764,6 +775,7 @@ mod proxy_handler_tests { }), protocol: Protocol::TlsRedirect { from: 80, to: 8443 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: None, websocket_timeout: None, disable_http_logs: false, @@ -830,6 +842,7 @@ mod proxy_handler_tests { }), protocol: Protocol::Http { port: 80 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: None, websocket_timeout: None, disable_http_logs: false, @@ -895,6 +908,7 @@ mod proxy_handler_tests { }), protocol: Protocol::Http { port: 80 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: None, websocket_timeout: None, disable_http_logs: false, @@ -986,6 +1000,7 @@ mod proxy_handler_tests { }), protocol: Protocol::Https { port: 443 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: Some(Duration::from_millis(500)), websocket_timeout: None, disable_http_logs: false, @@ -1072,6 +1087,7 @@ mod proxy_handler_tests { }), protocol: Protocol::Https { port: 443 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: Some(Duration::from_millis(500)), websocket_timeout: None, disable_http_logs: false, @@ -1179,6 +1195,7 @@ mod proxy_handler_tests { }), protocol: Protocol::Https { port: 443 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: Some(Duration::from_millis(500)), websocket_timeout: None, disable_http_logs: false, @@ -1275,6 +1292,7 @@ mod proxy_handler_tests { }), protocol: Protocol::Https { port: 443 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: None, websocket_timeout: None, disable_http_logs: false, @@ -1371,6 +1389,7 @@ mod proxy_handler_tests { }), protocol: Protocol::Http { port: 80 }, proxy_type: ProxyType::Aliasing, + buffer_size: 8_000, http_request_timeout: None, websocket_timeout: None, disable_http_logs: false, @@ -1467,6 +1486,7 @@ mod proxy_handler_tests { }), protocol: Protocol::Https { port: 443 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: None, websocket_timeout: None, disable_http_logs: false, @@ -1552,6 +1572,7 @@ mod proxy_handler_tests { }), protocol: Protocol::Https { port: 443 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: None, websocket_timeout: None, disable_http_logs: false, @@ -1657,6 +1678,7 @@ mod proxy_handler_tests { }), protocol: Protocol::Https { port: 443 }, proxy_type: ProxyType::Tunneling, + buffer_size: 8_000, http_request_timeout: None, websocket_timeout: None, disable_http_logs: false, diff --git a/src/lib.rs b/src/lib.rs index a87bc4c..8bb6067 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -45,13 +45,13 @@ use rustls::ServerConfig; use rustls_acme::acme::ACME_TLS_ALPN_NAME; use rustrict::CensorStr; use sysinfo::{CpuRefreshKind, MemoryRefreshKind, Networks, RefreshKind, System}; -use tcp::TcpHandler; +use tcp::{TcpHandler, TcpHandlerConfig}; use tcp_alias::TcpAlias; use telemetry::Telemetry; use tls::peek_sni_and_alpn; use tokio::{ fs, - io::{AsyncWriteExt, copy_bidirectional}, + io::{AsyncWriteExt, copy_bidirectional_with_sizes}, net::{TcpListener, TcpStream}, pin, time::{sleep, timeout}, @@ -176,6 +176,8 @@ pub(crate) struct SandholeServer { pub(crate) disable_tcp: bool, // If true, aliasing is disabled, including SSH and all local forwarding connections. pub(crate) disable_aliasing: bool, + // Buffer size for bidirectional copying. + pub(crate) buffer_size: usize, // How long until a login API request is timed out. pub(crate) authentication_request_timeout: Duration, // How long until an unauthed connection is closed. @@ -326,14 +328,15 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { Arc::clone("a_handler), Some(AliasReactor(Arc::clone(&telemetry))), )); - let tcp_handler: Arc = Arc::new(TcpHandler::new( - config.listen_address, - Arc::clone(&tcp_connections), - Arc::clone(&telemetry), - Arc::clone(&ip_filter), + let tcp_handler: Arc = Arc::new(TcpHandler::new(TcpHandlerConfig { + listen_address: config.listen_address, + conn_manager: Arc::clone(&tcp_connections), + telemetry: Arc::clone(&telemetry), + ip_filter: Arc::clone(&ip_filter), + buffer_size: config.buffer_size, tcp_connection_timeout, - config.disable_tcp_logs, - )); + disable_tcp_logs: config.disable_tcp_logs, + })); // Add TCP handler service as a listener for TCP port updates. tcp_connections.update_reactor(Some(TcpReactor { handler: Arc::clone(&tcp_handler), @@ -536,6 +539,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { }, // Always use aliasing channels instead of tunneling channels. proxy_type: ProxyType::Aliasing, + buffer_size: config.buffer_size, http_request_timeout, websocket_timeout: tcp_connection_timeout, disable_http_logs: config.disable_http_logs, @@ -571,6 +575,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { disable_sni: config.disable_sni, disable_tcp: config.disable_tcp, disable_aliasing: config.disable_aliasing, + buffer_size: config.buffer_size, authentication_request_timeout: config.authentication_request_timeout, idle_connection_timeout: config.idle_connection_timeout, unproxied_connection_timeout: config @@ -608,6 +613,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { }, // Always use tunneling channels. proxy_type: ProxyType::Tunneling, + buffer_size: config.buffer_size, http_request_timeout, websocket_timeout: tcp_connection_timeout, disable_http_logs: config.disable_http_logs, @@ -686,6 +692,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { }, // Always use tunneling channels. proxy_type: ProxyType::Tunneling, + buffer_size: config.buffer_size, http_request_timeout, websocket_timeout: tcp_connection_timeout, disable_http_logs: config.disable_http_logs, @@ -846,12 +853,24 @@ fn handle_https_connection( match sandhole.tcp_connection_timeout { Some(duration) => { let _ = timeout(duration, async { - let _ = copy_bidirectional(&mut stream, &mut channel).await; + let _ = copy_bidirectional_with_sizes( + &mut stream, + &mut channel, + sandhole.buffer_size, + sandhole.buffer_size, + ) + .await; }) .await; } None => { - let _ = copy_bidirectional(&mut stream, &mut channel).await; + let _ = copy_bidirectional_with_sizes( + &mut stream, + &mut channel, + sandhole.buffer_size, + sandhole.buffer_size, + ) + .await; } } return; diff --git a/src/ssh.rs b/src/ssh.rs index 418b56b..3952df0 100644 --- a/src/ssh.rs +++ b/src/ssh.rs @@ -42,7 +42,7 @@ use russh::{ server::{Auth, Handler, Msg, Session}, }; use tokio::{ - io::copy_bidirectional, + io::copy_bidirectional_with_sizes, sync::{ Mutex, RwLock, mpsc::{self, UnboundedReceiver, UnboundedSender}, @@ -1854,17 +1854,30 @@ impl Handler for ServerHandler { self.server.unproxied_connection_timeout; let cancellation_token = self.cancellation_token.clone(); let tcp_connection_timeout = self.server.tcp_connection_timeout; + let buffer_size = self.server.buffer_size; tokio::spawn(async move { let mut stream = channel.into_stream(); match tcp_connection_timeout { Some(duration) => { let _ = timeout(duration, async { - copy_bidirectional(&mut stream, &mut io).await + copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await }) .await; } None => { - let _ = copy_bidirectional(&mut stream, &mut io).await; + let _ = copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await; } } if proxy_count.fetch_sub(1, Ordering::AcqRel) == 1 { @@ -1879,17 +1892,30 @@ impl Handler for ServerHandler { // Serve SSH normally for authed user _ => { let tcp_connection_timeout = self.server.tcp_connection_timeout; + let buffer_size = self.server.buffer_size; tokio::spawn(async move { let mut stream = channel.into_stream(); match tcp_connection_timeout { Some(duration) => { let _ = timeout(duration, async { - copy_bidirectional(&mut stream, &mut io).await + copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await }) .await; } None => { - let _ = copy_bidirectional(&mut stream, &mut io).await; + let _ = copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await; } } }); @@ -1922,16 +1948,29 @@ impl Handler for ServerHandler { .telemetry .add_sni_connection(host_to_connect.into()); let tcp_connection_timeout = self.server.tcp_connection_timeout; + let buffer_size = self.server.buffer_size; tokio::spawn(async move { match tcp_connection_timeout { Some(duration) => { let _ = timeout(duration, async { - let _ = copy_bidirectional(&mut stream, &mut channel).await; + let _ = copy_bidirectional_with_sizes( + &mut stream, + &mut channel, + buffer_size, + buffer_size, + ) + .await; }) .await; } None => { - let _ = copy_bidirectional(&mut stream, &mut channel).await; + let _ = copy_bidirectional_with_sizes( + &mut stream, + &mut channel, + buffer_size, + buffer_size, + ) + .await; } } }); @@ -2046,17 +2085,30 @@ impl Handler for ServerHandler { self.server.unproxied_connection_timeout; let cancellation_token = self.cancellation_token.clone(); let tcp_connection_timeout = self.server.tcp_connection_timeout; + let buffer_size = self.server.buffer_size; tokio::spawn(async move { let mut stream = channel.into_stream(); match tcp_connection_timeout { Some(duration) => { let _ = timeout(duration, async { - copy_bidirectional(&mut stream, &mut io).await + copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await }) .await; } None => { - let _ = copy_bidirectional(&mut stream, &mut io).await; + let _ = copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await; } } if proxy_count.fetch_sub(1, Ordering::AcqRel) == 1 { @@ -2071,17 +2123,30 @@ impl Handler for ServerHandler { // Serve TCP normally for authed user _ => { let tcp_connection_timeout = self.server.tcp_connection_timeout; + let buffer_size = self.server.buffer_size; tokio::spawn(async move { let mut stream = channel.into_stream(); match tcp_connection_timeout { Some(duration) => { let _ = timeout(duration, async { - copy_bidirectional(&mut stream, &mut io).await + copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await }) .await; } None => { - let _ = copy_bidirectional(&mut stream, &mut io).await; + let _ = copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await; } } }); @@ -2136,17 +2201,30 @@ impl Handler for ServerHandler { self.server.unproxied_connection_timeout; let cancellation_token = self.cancellation_token.clone(); let tcp_connection_timeout = self.server.tcp_connection_timeout; + let buffer_size = self.server.buffer_size; tokio::spawn(async move { let mut stream = channel.into_stream(); match tcp_connection_timeout { Some(duration) => { let _ = timeout(duration, async { - copy_bidirectional(&mut stream, &mut io).await + copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await }) .await; } None => { - let _ = copy_bidirectional(&mut stream, &mut io).await; + let _ = copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await; } } if proxy_count.fetch_sub(1, Ordering::AcqRel) == 1 { @@ -2161,17 +2239,30 @@ impl Handler for ServerHandler { // Serve TCP normally for authed user _ => { let tcp_connection_timeout = self.server.tcp_connection_timeout; + let buffer_size = self.server.buffer_size; tokio::spawn(async move { let mut stream = channel.into_stream(); match tcp_connection_timeout { Some(duration) => { let _ = timeout(duration, async { - copy_bidirectional(&mut stream, &mut io).await + copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await }) .await; } None => { - let _ = copy_bidirectional(&mut stream, &mut io).await; + let _ = copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await; } } }); diff --git a/src/tcp.rs b/src/tcp.rs index 4b03b2a..959ecd7 100644 --- a/src/tcp.rs +++ b/src/tcp.rs @@ -8,7 +8,25 @@ use crate::{ use anyhow::Context; use dashmap::DashMap; use log::{error, info, warn}; -use tokio::{io::copy_bidirectional, net::TcpListener, time::timeout}; +use tokio::{io::copy_bidirectional_with_sizes, net::TcpListener, time::timeout}; + +// Service that handles creating TCP sockets for reverse forwarding connections. +pub(crate) struct TcpHandlerConfig { + // Address to listen to when creating sockets. + pub(crate) listen_address: IpAddr, + // Connection map to assign a tunneling service for each incoming connection. + pub(crate) conn_manager: Arc, TcpReactor>>, + // Telemetry server to keep track of the total connections. + pub(crate) telemetry: Arc, + // Service that identifies whether to allow or block a given IP address. + pub(crate) ip_filter: Arc, + // Buffer size for bidirectional copying. + pub(crate) buffer_size: usize, + // Optional duration to time out TCP connections. + pub(crate) tcp_connection_timeout: Option, + // Whether to send TCP logs to the SSH handles behind the forwarded connections. + pub(crate) disable_tcp_logs: bool, +} // Service that handles creating TCP sockets for reverse forwarding connections. pub(crate) struct TcpHandler { @@ -18,9 +36,12 @@ pub(crate) struct TcpHandler { sockets: DashMap>, // Connection map to assign a tunneling service for each incoming connection. conn_manager: Arc, TcpReactor>>, + // Telemetry server to keep track of the total connections. telemetry: Arc, // Service that identifies whether to allow or block a given IP address. ip_filter: Arc, + // Buffer size for bidirectional copying. + buffer_size: usize, // Optional duration to time out TCP connections. tcp_connection_timeout: Option, // Whether to send TCP logs to the SSH handles behind the forwarded connections. @@ -29,12 +50,15 @@ pub(crate) struct TcpHandler { impl TcpHandler { pub(crate) fn new( - listen_address: IpAddr, - conn_manager: Arc, TcpReactor>>, - telemetry: Arc, - ip_filter: Arc, - tcp_connection_timeout: Option, - disable_tcp_logs: bool, + TcpHandlerConfig { + listen_address, + conn_manager, + telemetry, + ip_filter, + buffer_size, + tcp_connection_timeout, + disable_tcp_logs, + }: TcpHandlerConfig, ) -> Self { TcpHandler { listen_address, @@ -42,6 +66,7 @@ impl TcpHandler { conn_manager, telemetry, ip_filter, + buffer_size, tcp_connection_timeout, disable_tcp_logs, } @@ -103,12 +128,24 @@ impl PortHandler for Arc { match clone.tcp_connection_timeout { Some(duration) => { let _ = timeout(duration, async { - copy_bidirectional(&mut stream, &mut channel).await + copy_bidirectional_with_sizes( + &mut stream, + &mut channel, + clone.buffer_size, + clone.buffer_size, + ) + .await }) .await; } None => { - let _ = copy_bidirectional(&mut stream, &mut channel).await; + let _ = copy_bidirectional_with_sizes( + &mut stream, + &mut channel, + clone.buffer_size, + clone.buffer_size, + ) + .await; } } }