diff --git a/Cargo.toml b/Cargo.toml index cd133b7..e251a36 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -93,7 +93,7 @@ thiserror = "2.0.18" tokio = { version = "1.52.3", features = ["full"] } tokio-rustls = "0.26.4" tokio-stream = "0.1.18" -tokio-util = "0.7.18" +tokio-util = { version = "0.7.18", features = ["net"] } tracing = "0.1.44" tracing_duper = { version = "0.1.1", optional = true } tracing-error = "0.2.1" diff --git a/src/admin/interface.rs b/src/admin/interface.rs index 485f74d..c28f1f8 100644 --- a/src/admin/interface.rs +++ b/src/admin/interface.rs @@ -30,8 +30,8 @@ use crate::{ droppable_handle::DroppableHandle, fingerprints::{AuthenticationType, KeyData}, quota::TokenHolderUser, + sock_addr_alias::SockAddrAlias, ssh::ServerHandlerSender, - tcp_alias::TcpAlias, }; struct BufferedSender { @@ -57,6 +57,7 @@ enum Tab { Sni, Ssh, Tcp, + Udp, Alias, } @@ -68,6 +69,7 @@ impl Tab { Tab::Sni => Color::Cyan, Tab::Ssh => Color::Yellow, Tab::Tcp => Color::Green, + Tab::Udp => Color::Magenta, Tab::Alias => Color::Red, } } @@ -100,6 +102,7 @@ impl TabData { Tab::Sni => Line::from(" SNI ".black().bg(tab.color())), Tab::Ssh => Line::from(" SSH ".black().bg(tab.color())), Tab::Tcp => Line::from(" TCP ".black().bg(tab.color())), + Tab::Udp => Line::from(" UDP ".black().bg(tab.color())), Tab::Alias => Line::from(" Alias ".black().bg(tab.color())), }) .collect::>(), @@ -613,6 +616,50 @@ impl AdminState { &mut self.vertical_scroll, ); } + Tab::Udp => { + // Get data for UDP + let data = self.server.udp_data.lock().expect("not poisoned").clone(); + self.vertical_scroll = self.vertical_scroll.content_length(data.len()); + // Create rows for each socket or alias + let rows: Vec> = data + .iter() + .map(|(port, (connections, conns_per_min, current_conns))| { + let len = connections.len() as u16; + let (peers, users): (Vec<_>, Vec<_>) = connections.iter().unzip(); + Row::new(vec![ + port.to_string(), + conns_per_min.to_string(), + current_conns.to_string(), + users.iter().join("\n"), + peers.iter().map(to_socket_addr_string).join("\n"), + ]) + .height(len) + }) + .collect(); + let constraints = [ + Constraint::Length(5), + Constraint::Length(7), + Constraint::Length(6), + Constraint::Min(7), + Constraint::Length(47), + ]; + let header = Row::new(["Port", "Con/min", "#Conns", "User(s)", "Peer(s)"]) + .add_modifier(Modifier::UNDERLINED); + let title = + Block::new().title(Line::from("UDP services".fg(color).bold()).centered()); + let table = Table::new(rows, constraints) + .header(header) + .column_spacing(1) + .block(title) + .row_highlight_style(Style::new().fg(color).reversed()); + StatefulWidget::render(table, area, buf, &mut self.table_state); + StatefulWidget::render( + Scrollbar::default().orientation(ScrollbarOrientation::VerticalRight), + area.inner(Margin::new(0, 2)), + buf, + &mut self.vertical_scroll, + ); + } Tab::Alias => { // Get data for aliases let data = self.server.alias_data.lock().expect("not poisoned").clone(); @@ -621,7 +668,10 @@ impl AdminState { let rows: Vec> = data .iter() .map( - |(TcpAlias(alias, port), (connections, conns_per_min, current_conns))| { + |( + SockAddrAlias(alias, port), + (connections, conns_per_min, current_conns), + )| { let len = connections.len() as u16; let (peers, users): (Vec<_>, Vec<_>) = connections.iter().unzip(); Row::new(vec![ @@ -774,6 +824,9 @@ impl AdminInterface { if !server.disable_tcp { tabs.push(Tab::Tcp); } + if !server.disable_udp { + tabs.push(Tab::Udp); + } if !server.disable_aliasing { tabs.push(Tab::Alias); } @@ -1065,6 +1118,15 @@ impl AdminInterface { .values() .nth(row) .map(|value| value.0.values().cloned().collect()), + Tab::Udp => interface + .state + .server + .udp_data + .lock() + .expect("not poisoned") + .values() + .nth(row) + .map(|value| value.0.values().cloned().collect()), Tab::Alias => interface .state .server diff --git a/src/config.rs b/src/config.rs index e7e4f50..09d4a25 100644 --- a/src/config.rs +++ b/src/config.rs @@ -205,6 +205,10 @@ pub struct ApplicationConfig { #[arg(long, default_value_t = false)] pub disable_tcp_logs: bool, + /// Disable sending UDP logs to clients. + #[arg(long, default_value_t = false)] + pub disable_udp_logs: bool, + /// Contact e-mail to use with Let's Encrypt. If set, enables ACME for HTTPS certificates. /// /// By providing your e-mail, you agree to the Let's Encrypt Subscriber Agreement. @@ -306,6 +310,10 @@ pub struct ApplicationConfig { #[arg(long, default_value_t = false)] pub disable_tcp: bool, + /// Disable all UDP port tunneling. By default, this is enabled globally. + #[arg(long, default_value_t = false)] + pub disable_udp: bool, + /// Disable all aliasing (i.e. local forwarding). By default, this is enabled globally. #[arg(long, default_value_t = false)] pub disable_aliasing: bool, @@ -603,6 +611,7 @@ mod application_config_tests { force_https: false, disable_http_logs: false, disable_tcp_logs: false, + disable_udp_logs: false, acme_contact_email: None, acme_use_staging: false, password_authentication_url: None, @@ -617,6 +626,7 @@ mod application_config_tests { disable_https: false, disable_sni: false, disable_tcp: false, + disable_udp: false, disable_aliasing: false, disable_prometheus: false, quota_per_user: None, @@ -670,6 +680,7 @@ mod application_config_tests { "--force-https", "--disable-http-logs", "--disable-tcp-logs", + "--disable-udp-logs", "--acme-contact-email=admin@server.com", "--acme-use-staging", "--password-authentication-url=https://auth.server.com/validate", @@ -684,6 +695,7 @@ mod application_config_tests { "--disable-https", "--disable-sni", "--disable-tcp", + "--disable-udp", "--disable-aliasing", "--disable-prometheus", "--quota-per-user=10", @@ -737,6 +749,7 @@ mod application_config_tests { force_https: true, disable_http_logs: true, disable_tcp_logs: true, + disable_udp_logs: true, acme_contact_email: Some("admin@server.com".into()), acme_use_staging: true, password_authentication_url: Some("https://auth.server.com/validate".into()), @@ -751,6 +764,7 @@ mod application_config_tests { disable_https: true, disable_sni: true, disable_tcp: true, + disable_udp: true, disable_aliasing: true, disable_prometheus: true, quota_per_user: Some(10.try_into().unwrap()), diff --git a/src/connections.rs b/src/connections.rs index 5da91d5..ca0fa93 100644 --- a/src/connections.rs +++ b/src/connections.rs @@ -18,8 +18,8 @@ use crate::{ error::ServerError, quota::{QuotaHandler, QuotaToken, TokenHolder, TokenHolderUser}, reactor::{AliasReactor, ConnectionMapReactor, DummyConnectionMapReactor, HttpReactor}, + sock_addr_alias::{BorrowedSockAddrAlias, SockAddrAlias, SockAddrAliasKey}, ssh::connection_handler::SshTunnelHandler, - tcp_alias::{BorrowedTcpAlias, TcpAlias, TcpAliasKey}, }; // Data stored for a connection map entry. @@ -274,14 +274,16 @@ where #[derive(Builder)] pub(crate) struct HttpAliasingConnection { http: Arc, HttpReactor>>, - alias: Arc, AliasReactor>>, + alias: Arc, AliasReactor>>, } impl ConnectionGetByHttpHost> for Arc { fn get_by_http_host(&self, host: &str, ip: IpAddr) -> Option> { self.http.get(host, ip).or_else(|| { - self.alias - .get(&BorrowedTcpAlias(host, &80) as &dyn TcpAliasKey, ip) + self.alias.get( + &BorrowedSockAddrAlias(host, &80) as &dyn SockAddrAliasKey, + ip, + ) }) } } diff --git a/src/entrypoint.rs b/src/entrypoint.rs index 6e4b12a..d28cbc7 100644 --- a/src/entrypoint.rs +++ b/src/entrypoint.rs @@ -60,7 +60,7 @@ use crate::{ http::{DomainRedirect, Protocol, ProxyData, ProxyType, proxy_handler}, ip::{IpFilter, IpFilterConfig}, quota::{DummyQuotaHandler, QuotaHandler, QuotaMap}, - reactor::{AliasReactor, HttpReactor, SniReactor, SshReactor, TcpReactor}, + reactor::{AliasReactor, HttpReactor, SniReactor, SshReactor, TcpReactor, UdpReactor}, ssh::Server, tcp::TcpHandler, tcp_listener::get_tcp_listener, @@ -70,11 +70,12 @@ use crate::{ TELEMETRY_GAUGE_USED_MEMORY, TELEMETRY_KEY_HOSTNAME, Telemetry, }, tls::{TlsPeekData, peek_sni_and_alpn}, + udp::UdpHandler, }; #[cfg_attr(not(feature = "prometheus"), allow(unused_imports))] use crate::{ admin::{ADMIN_ALIAS_PORT, connection_handler::AdminAliasHandler}, - tcp_alias::TcpAlias, + sock_addr_alias::SockAddrAlias, }; #[doc(hidden)] @@ -297,6 +298,13 @@ pub async fn entrypoint(config: ApplicationConfig) -> color_eyre::Result<()> { .quota_handler(Arc::clone("a_handler)) .build(), ); + let udp_connections = Arc::new( + ConnectionMap::builder() + .strategy(config.load_balancing) + .algorithm(config.load_balancing_algorithm) + .quota_handler(Arc::clone("a_handler)) + .build(), + ); let admin_alias_connections = Arc::new( ConnectionMap::builder() .strategy(crate::LoadBalancingStrategy::Deny) @@ -327,6 +335,19 @@ pub async fn entrypoint(config: ApplicationConfig) -> color_eyre::Result<()> { handler: Arc::clone(&tcp_handler), telemetry: Arc::clone(&telemetry), })); + let udp_handler: Arc = Arc::new( + UdpHandler::builder() + .listen_address(config.listen_address) + .conn_manager(Arc::clone(&udp_connections)) + .ip_filter(Arc::clone(&ip_filter)) + .disable_udp_logs(config.disable_udp_logs) + .build(), + ); + // Add udp handler service as a listener for udp port updates. + udp_connections.update_reactor(Some(UdpReactor { + handler: Arc::clone(&udp_handler), + telemetry: Arc::clone(&telemetry), + })); // Add addressing service with optional profanity filtering let addressing = Arc::new({ let builder = AddressDelegator::builder() @@ -501,6 +522,37 @@ pub async fn entrypoint(config: ApplicationConfig) -> color_eyre::Result<()> { } }); } + let udp_data = Arc::new(Mutex::default()); + if !config.disable_udp { + let data_clone = Arc::clone(&udp_data); + let connections_clone = Arc::clone(&udp_connections); + let telemetry_clone = Arc::clone(&telemetry); + // Periodically update UDP data, based on the connection map. + tokio::spawn(async move { + let mut refresh_interval = interval(Duration::from_millis(3_000)); + refresh_interval.tick().await; + loop { + refresh_interval.tick().await; + let data = connections_clone.data(); + let telemetry_per_minute = telemetry_clone.get_udp_connections_per_minute(); + let telemetry_current = telemetry_clone.get_current_udp_connections(); + let data = data + .into_iter() + .map(|(port, addresses)| { + let connections_per_minute = + telemetry_per_minute.get(&port).copied().unwrap_or_default(); + let current_connections = + telemetry_current.get(&port).copied().unwrap_or_default(); + ( + port, + (addresses, connections_per_minute, current_connections), + ) + }) + .collect(); + *data_clone.lock().expect("not poisoned") = data; + } + }); + } let alias_data = Arc::new(Mutex::default()); if !config.disable_aliasing { let data_clone = Arc::clone(&alias_data); @@ -696,12 +748,14 @@ pub async fn entrypoint(config: ApplicationConfig) -> color_eyre::Result<()> { sni: sni_connections, ssh: ssh_connections, tcp: tcp_connections, + udp: udp_connections, admin_alias: admin_alias_connections, alias: alias_connections, http_data, sni_data, ssh_data, tcp_data, + udp_data, alias_data, system_data, admin_notifications, @@ -711,6 +765,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> color_eyre::Result<()> { api_login, address_delegator: addressing, tcp_handler, + udp_handler, domain: config.mode.domain, http_port: config.http_port.into(), https_port: config.https_port.into(), @@ -721,6 +776,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> color_eyre::Result<()> { disable_https: config.disable_http || config.disable_https, disable_sni: config.disable_http || config.disable_https || config.disable_sni, disable_tcp: config.disable_tcp, + disable_udp: config.disable_udp, disable_aliasing: config.disable_aliasing, buffer_size, pool_size: config.pool_size, @@ -745,7 +801,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> color_eyre::Result<()> { sandhole .admin_alias .insert( - TcpAlias("prometheus.sandhole".into(), ADMIN_ALIAS_PORT), + SockAddrAlias("prometheus.sandhole".into(), ADMIN_ALIAS_PORT), SocketAddr::from(([0, 0, 0, 0], 0)), TokenHolder::System, Arc::new(AdminAliasHandler { diff --git a/src/http/mod.rs b/src/http/mod.rs index 99865d7..fdbf8d8 100644 --- a/src/http/mod.rs +++ b/src/http/mod.rs @@ -20,8 +20,8 @@ use crate::{ error::ServerError, http::{http2::handle_http2_request, http11::handle_http11_request}, keepalive::KeepaliveAlias, + sock_addr_alias::BorrowedSockAddrAlias, ssh::ServerHandlerSender, - tcp_alias::BorrowedTcpAlias, telemetry::{ TELEMETRY_COUNTER_ALIAS_CONNECTIONS, TELEMETRY_COUNTER_HTTP_REQUESTS, TELEMETRY_HISTOGRAM_HTTP_ELAPSED_TIME, TELEMETRY_KEY_ALIAS, TELEMETRY_KEY_HOSTNAME, @@ -641,7 +641,7 @@ where append_to_header(headers, &X_FORWARDED_PORT, port.to_string().as_bytes()); // Add this request to the telemetry for the host if http_data.as_ref().is_some_and(|data| data.is_aliasing) { - counter!(TELEMETRY_COUNTER_ALIAS_CONNECTIONS, TELEMETRY_KEY_ALIAS => BorrowedTcpAlias(&host, &port).to_string()) + counter!(TELEMETRY_COUNTER_ALIAS_CONNECTIONS, TELEMETRY_KEY_ALIAS => BorrowedSockAddrAlias(&host, &port).to_string()) .increment(1); } else { counter!(TELEMETRY_COUNTER_HTTP_REQUESTS, TELEMETRY_KEY_HOSTNAME => host.clone()) diff --git a/src/lib.rs b/src/lib.rs index 49bac72..c714af0 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -25,10 +25,11 @@ use crate::{ fingerprints::FingerprintsValidator, http::ProxyData, quota::TokenHolderUser, - reactor::{AliasReactor, HttpReactor, SniReactor, SshReactor, TcpReactor}, + reactor::{AliasReactor, HttpReactor, SniReactor, SshReactor, TcpReactor, UdpReactor}, + sock_addr_alias::SockAddrAlias, ssh::connection_handler::{SshChannel, SshTunnelHandler}, tcp::TcpHandler, - tcp_alias::TcpAlias, + udp::UdpHandler, }; #[doc(hidden)] @@ -60,12 +61,14 @@ mod keepalive; mod login; mod quota; mod reactor; +mod sock_addr_alias; mod ssh; mod tcp; -mod tcp_alias; mod tcp_listener; mod telemetry; mod tls; +mod udp; +mod udp_listener; // Data collected from the system and displayed on the admin interface. #[derive(Default, Clone)] @@ -112,10 +115,12 @@ pub(crate) struct SandholeServer { pub(crate) sni: Arc, SniReactor>>, // The map for forwarded TCP connections. pub(crate) tcp: Arc, TcpReactor>>, + // The map for forwarded UDP connections. + pub(crate) udp: Arc, UdpReactor>>, // The map for admin aliased connections. - pub(crate) admin_alias: Arc, AliasReactor>>, + pub(crate) admin_alias: Arc, AliasReactor>>, // The map for forwarded aliased connections. - pub(crate) alias: Arc, AliasReactor>>, + pub(crate) alias: Arc, AliasReactor>>, // Data related to the SSH forwardings for the admin interface. pub(crate) ssh_data: DataTable, u64, f64)>, // Data related to the HTTP forwardings for the admin interface. @@ -124,8 +129,11 @@ pub(crate) struct SandholeServer { pub(crate) sni_data: DataTable, u64, f64)>, // Data related to the TCP forwardings for the admin interface. pub(crate) tcp_data: DataTable, u64, f64)>, + // Data related to the UDP forwardings for the admin interface. + pub(crate) udp_data: DataTable, u64, f64)>, // Data related to the alias forwardings for the admin interface. - pub(crate) alias_data: DataTable, u64, f64)>, + pub(crate) alias_data: + DataTable, u64, f64)>, // System data for the admin interface. pub(crate) system_data: Arc>, // Warnings for the admin interface. @@ -141,6 +149,8 @@ pub(crate) struct SandholeServer { pub(crate) address_delegator: Arc>, // Service for handling opening and closing TCP sockets for non-aliased services. pub(crate) tcp_handler: Arc, + // Service for handling opening and closing UDP sockets. + pub(crate) udp_handler: Arc, // The base domain of Sandhole. pub(crate) domain: Option, // Which port Sandhole listens to for HTTP connections. @@ -161,6 +171,8 @@ pub(crate) struct SandholeServer { pub(crate) disable_sni: bool, // If true, TCP is disabled for all ports except for HTTP. pub(crate) disable_tcp: bool, + // If true, UDP is disabled for all ports. + pub(crate) disable_udp: bool, // If true, aliasing is disabled, including SSH and all local forwarding connections. pub(crate) disable_aliasing: bool, // Buffer size for bidirectional copying. diff --git a/src/reactor.rs b/src/reactor.rs index 1c8441d..fe8cc43 100644 --- a/src/reactor.rs +++ b/src/reactor.rs @@ -2,9 +2,10 @@ use std::sync::Arc; use crate::{ certificates::CertificateResolver, - tcp::{PortHandler, TcpHandler}, - tcp_alias::TcpAlias, + sock_addr_alias::SockAddrAlias, + tcp::{TcpHandler, TcpPortHandler}, telemetry::Telemetry, + udp::{UdpHandler, UdpPortHandler}, }; #[cfg_attr(test, mockall::automock)] @@ -62,10 +63,22 @@ impl ConnectionMapReactor for TcpReactor { } } +pub(crate) struct UdpReactor { + pub(crate) handler: Arc, + pub(crate) telemetry: Arc, +} + +impl ConnectionMapReactor for UdpReactor { + fn call(&self, identifiers: Vec) { + self.handler.update_ports(identifiers.clone()); + self.telemetry.udp_reactor(identifiers); + } +} + pub(crate) struct AliasReactor(pub(crate) Arc); -impl ConnectionMapReactor for AliasReactor { - fn call(&self, identifiers: Vec) { +impl ConnectionMapReactor for AliasReactor { + fn call(&self, identifiers: Vec) { self.0.alias_reactor(identifiers); } } diff --git a/src/tcp_alias.rs b/src/sock_addr_alias.rs similarity index 56% rename from src/tcp_alias.rs rename to src/sock_addr_alias.rs index 1ea063c..010b68e 100644 --- a/src/tcp_alias.rs +++ b/src/sock_addr_alias.rs @@ -11,64 +11,64 @@ use color_eyre::eyre::OptionExt; // A TCP alias, with an address and a port. #[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] -pub(crate) struct TcpAlias(pub(crate) String, pub(crate) u16); +pub(crate) struct SockAddrAlias(pub(crate) String, pub(crate) u16); -impl Display for TcpAlias { +impl Display for SockAddrAlias { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { write!(f, "{}:{}", self.0, self.1) } } -impl FromStr for TcpAlias { +impl FromStr for SockAddrAlias { type Err = color_eyre::Report; fn from_str(s: &str) -> Result { let (left, right) = s.rsplit_once(':').ok_or_eyre("Missing : separator")?; - Ok(TcpAlias(left.to_string(), right.parse()?)) + Ok(SockAddrAlias(left.to_string(), right.parse()?)) } } // A borrowed TCP alias, with references to an address and a port. Useful for accessing the TCP connection map. #[derive(Copy, Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] -pub(crate) struct BorrowedTcpAlias<'a>(pub(crate) &'a str, pub(crate) &'a u16); +pub(crate) struct BorrowedSockAddrAlias<'a>(pub(crate) &'a str, pub(crate) &'a u16); -impl Display for BorrowedTcpAlias<'_> { +impl Display for BorrowedSockAddrAlias<'_> { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { write!(f, "{}:{}", self.0, self.1) } } -impl<'a> Borrow for TcpAlias { - fn borrow(&self) -> &(dyn TcpAliasKey + 'a) { +impl<'a> Borrow for SockAddrAlias { + fn borrow(&self) -> &(dyn SockAddrAliasKey + 'a) { self } } -pub(crate) trait TcpAliasKey { - fn key(&self) -> BorrowedTcpAlias<'_>; +pub(crate) trait SockAddrAliasKey { + fn key(&self) -> BorrowedSockAddrAlias<'_>; } -impl TcpAliasKey for TcpAlias { - fn key(&self) -> BorrowedTcpAlias<'_> { - BorrowedTcpAlias(self.0.as_str(), &self.1) +impl SockAddrAliasKey for SockAddrAlias { + fn key(&self) -> BorrowedSockAddrAlias<'_> { + BorrowedSockAddrAlias(self.0.as_str(), &self.1) } } -impl TcpAliasKey for BorrowedTcpAlias<'_> { - fn key(&self) -> BorrowedTcpAlias<'_> { +impl SockAddrAliasKey for BorrowedSockAddrAlias<'_> { + fn key(&self) -> BorrowedSockAddrAlias<'_> { *self } } -impl PartialEq for dyn TcpAliasKey + '_ { +impl PartialEq for dyn SockAddrAliasKey + '_ { fn eq(&self, other: &Self) -> bool { self.key().eq(&other.key()) } } -impl Eq for dyn TcpAliasKey + '_ {} +impl Eq for dyn SockAddrAliasKey + '_ {} -impl Hash for dyn TcpAliasKey + '_ { +impl Hash for dyn SockAddrAliasKey + '_ { fn hash(&self, state: &mut H) { self.key().hash(state) } diff --git a/src/ssh/auth.rs b/src/ssh/auth.rs index 5a692da..2bd6c89 100644 --- a/src/ssh/auth.rs +++ b/src/ssh/auth.rs @@ -17,7 +17,7 @@ use tokio_util::sync::CancellationToken; use crate::{ admin::interface::AdminInterface, connection_handler::ConnectionHttpData, droppable_handle::DroppableHandle, ip::IpFilter, quota::TokenHolder, ssh::FingerprintFn, - tcp_alias::TcpAlias, + sock_addr_alias::SockAddrAlias, }; pub(crate) struct ProxyAutoCancellation { @@ -88,11 +88,12 @@ pub(crate) struct UserData { // Identifier for the user, used for creating quota tokens. pub(crate) quota_key: TokenHolder, // Map to keep track of opened host-based connections (HTTP and SSH), to clean up when the forwarding is canceled. - pub(crate) host_addressing: HashMap), RandomState>, + pub(crate) host_addressing: HashMap), RandomState>, // Map to keep track of opened port-based connections (TCP), to clean up when the forwarding is canceled. - pub(crate) port_addressing: HashMap), RandomState>, + pub(crate) port_addressing: HashMap), RandomState>, // Map to keep track of opened alias-based connections (aliases), to clean up when the forwarding is canceled. - pub(crate) alias_addressing: HashMap), RandomState>, + pub(crate) alias_addressing: + HashMap), RandomState>, // IPs allowed to connect to this user's services. pub(crate) allowlist: Option>, // IPs disallowed from connecting to this user's services. diff --git a/src/ssh/exec.rs b/src/ssh/exec.rs index 028bbcb..2e9b0ed 100644 --- a/src/ssh/exec.rs +++ b/src/ssh/exec.rs @@ -9,7 +9,7 @@ use crate::{ SandholeServer, admin::interface::AdminInterface, ssh::{AuthenticatedData, ServerHandlerSender, auth::UserSessionRestriction}, - tcp_alias::TcpAlias, + sock_addr_alias::SockAddrAlias, }; #[bitflags] @@ -133,7 +133,7 @@ impl SshCommand for AllowedFingerprintsCommand { } // Insert our handler into the TCP alias connections map. context.server.alias.insert( - TcpAlias(address.clone(), 80), + SockAddrAlias(address.clone(), 80), *context.peer, user_data.quota_key.clone(), handler, @@ -198,7 +198,7 @@ impl SshCommand for TcpAliasCommand { } // Insert our handler into the TCP alias connections map. context.server.alias.insert( - TcpAlias(address.clone(), 80), + SockAddrAlias(address.clone(), 80), *context.peer, user_data.quota_key.clone(), handler, diff --git a/src/ssh/forwarding.rs b/src/ssh/forwarding.rs index d6245bc..3f21c7a 100644 --- a/src/ssh/forwarding.rs +++ b/src/ssh/forwarding.rs @@ -27,20 +27,22 @@ use crate::{ connection_handler::ConnectionHandler, connections::ConnectionGetByHttpHost, http::proxy_handler, + sock_addr_alias::{BorrowedSockAddrAlias, SockAddrAlias, SockAddrAliasKey}, ssh::{ AuthenticatedData, ServerHandlerSender, UserData, auth::UserSessionRestriction, connection_handler::SshTunnelHandler, }, - tcp::PortHandler, - tcp_alias::{BorrowedTcpAlias, TcpAlias, TcpAliasKey}, + tcp::TcpPortHandler, telemetry::{ TELEMETRY_COUNTER_ADMIN_ALIAS_CONNECTIONS, TELEMETRY_COUNTER_ALIAS_CONNECTIONS, TELEMETRY_COUNTER_SNI_CONNECTIONS, TELEMETRY_COUNTER_SSH_CONNECTIONS, - TELEMETRY_COUNTER_TCP_CONNECTIONS, TELEMETRY_GAUGE_ADMIN_ALIAS_CONNECTIONS_CURRENT, - TELEMETRY_GAUGE_ALIAS_CONNECTIONS_CURRENT, TELEMETRY_GAUGE_SNI_CONNECTIONS_CURRENT, - TELEMETRY_GAUGE_SSH_CONNECTIONS_CURRENT, TELEMETRY_GAUGE_TCP_CONNECTIONS_CURRENT, + TELEMETRY_COUNTER_TCP_CONNECTIONS, TELEMETRY_COUNTER_UDP_CONNECTIONS, + TELEMETRY_GAUGE_ADMIN_ALIAS_CONNECTIONS_CURRENT, TELEMETRY_GAUGE_ALIAS_CONNECTIONS_CURRENT, + TELEMETRY_GAUGE_SNI_CONNECTIONS_CURRENT, TELEMETRY_GAUGE_SSH_CONNECTIONS_CURRENT, + TELEMETRY_GAUGE_TCP_CONNECTIONS_CURRENT, TELEMETRY_GAUGE_UDP_CONNECTIONS_CURRENT, TELEMETRY_KEY_ALIAS, TELEMETRY_KEY_HOSTNAME, TELEMETRY_KEY_PORT, }, + udp::UdpPortHandler, }; pub(crate) struct RemoteForwardingContext<'a> { @@ -89,6 +91,8 @@ pub(crate) trait ForwardingHandlerStrategy { pub(crate) struct Forwarder; +pub(crate) const UDP_ADDRESS: &str = "udp.sandhole"; + impl Forwarder { pub(crate) async fn remote_forwarding( context: &mut RemoteForwardingContext<'_>, @@ -111,6 +115,10 @@ impl Forwarder { HttpForwardingHandler .remote_forwarding(context, address, port, handle) .await + } else if address == UDP_ADDRESS { + UdpForwardingHandler + .remote_forwarding(context, address, port, handle) + .await } else if context.server.is_alias(address) { AliasForwardingHandler .remote_forwarding(context, address, port, handle) @@ -145,6 +153,11 @@ impl Forwarder { .cancel_remote_forwarding(context, address, port) .await } + _ if address == UDP_ADDRESS => { + UdpForwardingHandler + .cancel_remote_forwarding(context, address, port) + .await + } _ if context.server.is_alias(address) => { AliasForwardingHandler .cancel_remote_forwarding(context, address, port) @@ -192,6 +205,17 @@ impl Forwarder { channel, ) .await + } else if address == UDP_ADDRESS { + UdpForwardingHandler + .local_forwarding( + context, + address, + port, + originator_address, + originator_port, + channel, + ) + .await } else if context.server.is_alias(address) { AliasForwardingHandler .local_forwarding( @@ -222,6 +246,7 @@ pub(crate) struct SshForwardingHandler; pub(crate) struct HttpForwardingHandler; pub(crate) struct AliasForwardingHandler; pub(crate) struct TcpForwardingHandler; +pub(crate) struct UdpForwardingHandler; impl ForwardingHandlerStrategy for SshForwardingHandler { async fn remote_forwarding( @@ -343,7 +368,7 @@ impl ForwardingHandlerStrategy for SshForwardingHandler { .into_bytes(), ); context.user_data.host_addressing.insert( - TcpAlias(address.to_string(), *port as u16), + SockAddrAlias(address.to_string(), *port as u16), (address.to_string(), semaphore), ); Ok(true) @@ -360,7 +385,7 @@ impl ForwardingHandlerStrategy for SshForwardingHandler { if let Some(assigned_host) = context .user_data .host_addressing - .remove(&BorrowedTcpAlias(address, &port) as &dyn TcpAliasKey) + .remove(&BorrowedSockAddrAlias(address, &port) as &dyn SockAddrAliasKey) { context.server.ssh.remove(&assigned_host.0, context.peer); #[cfg(not(coverage_nightly))] @@ -545,7 +570,7 @@ impl ForwardingHandlerStrategy for HttpForwardingHandler { context.user_data.max_pool_size.load(Ordering::Acquire), )); match context.server.alias.insert( - TcpAlias(address.into(), 80), + SockAddrAlias(address.into(), 80), *context.peer, context.user_data.quota_key.clone(), Arc::new(SshTunnelHandler { @@ -598,8 +623,8 @@ impl ForwardingHandlerStrategy for HttpForwardingHandler { .into_bytes(), ); context.user_data.alias_addressing.insert( - TcpAlias(address.into(), 80), - (TcpAlias(address.into(), 80), semaphore), + SockAddrAlias(address.into(), 80), + (SockAddrAlias(address.into(), 80), semaphore), ); Ok(true) } @@ -678,7 +703,7 @@ impl ForwardingHandlerStrategy for HttpForwardingHandler { .into_bytes(), ); context.user_data.host_addressing.insert( - TcpAlias(address.into(), *port as u16), + SockAddrAlias(address.into(), *port as u16), (assigned_host, semaphore), ); Ok(true) @@ -823,7 +848,7 @@ impl ForwardingHandlerStrategy for HttpForwardingHandler { ); } context.user_data.host_addressing.insert( - TcpAlias(address.to_string(), *port as u16), + SockAddrAlias(address.to_string(), *port as u16), (assigned_host, semaphore), ); Ok(true) @@ -867,9 +892,9 @@ impl ForwardingHandlerStrategy for HttpForwardingHandler { if let Some(assigned_alias) = context .user_data .alias_addressing - .remove(&BorrowedTcpAlias(address, &80) as &dyn TcpAliasKey) + .remove(&BorrowedSockAddrAlias(address, &80) as &dyn SockAddrAliasKey) { - let key: &dyn TcpAliasKey = assigned_alias.0.borrow(); + let key: &dyn SockAddrAliasKey = assigned_alias.0.borrow(); context.server.alias.remove(key, context.peer); #[cfg(not(coverage_nightly))] tracing::info!( @@ -897,7 +922,7 @@ impl ForwardingHandlerStrategy for HttpForwardingHandler { if let Some(assigned_alias) = context .user_data .host_addressing - .remove(&BorrowedTcpAlias(address, &{ port }) as &dyn TcpAliasKey) + .remove(&BorrowedSockAddrAlias(address, &{ port }) as &dyn SockAddrAliasKey) { context.server.sni.remove(&assigned_alias.0, context.peer); #[cfg(not(coverage_nightly))] @@ -923,7 +948,7 @@ impl ForwardingHandlerStrategy for HttpForwardingHandler { if let Some(assigned_host) = context .user_data .host_addressing - .remove(&BorrowedTcpAlias(address, &{ port }) as &dyn TcpAliasKey) + .remove(&BorrowedSockAddrAlias(address, &{ port }) as &dyn SockAddrAliasKey) { context.server.http.remove(&assigned_host.0, context.peer); #[cfg(not(coverage_nightly))] @@ -1146,7 +1171,7 @@ impl ForwardingHandlerStrategy for AliasForwardingHandler { context.user_data.max_pool_size.load(Ordering::Acquire), )); match context.server.alias.insert( - TcpAlias(address.to_string(), assigned_port), + SockAddrAlias(address.to_string(), assigned_port), *context.peer, context.user_data.quota_key.clone(), Arc::new(SshTunnelHandler { @@ -1188,8 +1213,8 @@ impl ForwardingHandlerStrategy for AliasForwardingHandler { _ => { // Adding to connection map succeeded. context.user_data.alias_addressing.insert( - TcpAlias(address.to_string(), *port as u16), - (TcpAlias(address.to_string(), assigned_port), semaphore), + SockAddrAlias(address.to_string(), *port as u16), + (SockAddrAlias(address.to_string(), assigned_port), semaphore), ); #[cfg(not(coverage_nightly))] tracing::info!( @@ -1217,13 +1242,12 @@ impl ForwardingHandlerStrategy for AliasForwardingHandler { address: &str, port: u16, ) -> Result { - if let Some(assigned_alias) = - context - .user_data - .alias_addressing - .remove(&BorrowedTcpAlias(address, &{ port }) as &dyn TcpAliasKey) + if let Some(assigned_alias) = context + .user_data + .alias_addressing + .remove(&BorrowedSockAddrAlias(address, &{ port }) as &dyn SockAddrAliasKey) { - let key: &dyn TcpAliasKey = assigned_alias.0.borrow(); + let key: &dyn SockAddrAliasKey = assigned_alias.0.borrow(); context.server.alias.remove(key, context.peer); #[cfg(not(coverage_nightly))] tracing::info!( @@ -1256,17 +1280,16 @@ impl ForwardingHandlerStrategy for AliasForwardingHandler { channel: Channel, ) -> Result { let ip = context.peer.ip().to_canonical(); - if let Some(handler) = context - .server - .admin_alias - .get(&BorrowedTcpAlias(address, &port) as &dyn TcpAliasKey, ip) - { + if let Some(handler) = context.server.admin_alias.get( + &BorrowedSockAddrAlias(address, &port) as &dyn SockAddrAliasKey, + ip, + ) { if let AuthenticatedData::Admin { .. } = context.auth_data { if let Ok(mut io) = handler .aliasing_channel(ip, context.peer.port(), context.key_fingerprint.as_ref()) .await { - let alias = TcpAlias(address.into(), port); + let alias = SockAddrAlias(address.into(), port); let gauge = gauge!(TELEMETRY_GAUGE_ADMIN_ALIAS_CONNECTIONS_CURRENT, TELEMETRY_KEY_ALIAS => alias.to_string()); gauge.increment(1); counter!(TELEMETRY_COUNTER_ADMIN_ALIAS_CONNECTIONS, TELEMETRY_KEY_ALIAS => alias.to_string()) @@ -1314,15 +1337,14 @@ impl ForwardingHandlerStrategy for AliasForwardingHandler { "Non-admin user attempt to local forward admin alias", ) } - } else if let Some(handler) = context - .server - .alias - .get(&BorrowedTcpAlias(address, &port) as &dyn TcpAliasKey, ip) - && let Ok(mut io) = handler - .aliasing_channel(ip, context.peer.port(), context.key_fingerprint.as_ref()) - .await + } else if let Some(handler) = context.server.alias.get( + &BorrowedSockAddrAlias(address, &port) as &dyn SockAddrAliasKey, + ip, + ) && let Ok(mut io) = handler + .aliasing_channel(ip, context.peer.port(), context.key_fingerprint.as_ref()) + .await { - let alias = TcpAlias(address.into(), port); + let alias = SockAddrAlias(address.into(), port); let gauge = gauge!(TELEMETRY_GAUGE_ALIAS_CONNECTIONS_CURRENT, TELEMETRY_KEY_ALIAS => alias.to_string()); gauge.increment(1); counter!(TELEMETRY_COUNTER_ALIAS_CONNECTIONS, TELEMETRY_KEY_ALIAS => alias.to_string()) @@ -1648,7 +1670,7 @@ impl ForwardingHandlerStrategy for TcpForwardingHandler { _ => { // Adding to connection map succeeded. context.user_data.port_addressing.insert( - TcpAlias(address.to_string(), *port as u16), + SockAddrAlias(address.to_string(), *port as u16), (assigned_port, semaphore), ); #[cfg(not(coverage_nightly))] @@ -1682,11 +1704,10 @@ impl ForwardingHandlerStrategy for TcpForwardingHandler { address: &str, port: u16, ) -> Result { - if let Some(assigned_port) = - context - .user_data - .port_addressing - .remove(&BorrowedTcpAlias(address, &{ port }) as &dyn TcpAliasKey) + if let Some(assigned_port) = context + .user_data + .port_addressing + .remove(&BorrowedSockAddrAlias(address, &{ port }) as &dyn SockAddrAliasKey) { context.server.tcp.remove(&assigned_port.0, context.peer); #[cfg(not(coverage_nightly))] @@ -1832,3 +1853,346 @@ impl ForwardingHandlerStrategy for TcpForwardingHandler { Ok(false) } } + +impl ForwardingHandlerStrategy for UdpForwardingHandler { + async fn remote_forwarding( + &mut self, + context: &mut RemoteForwardingContext<'_>, + address: &str, + port: &mut u32, + handle: Handle, + ) -> Result { + // Forbid binding UDP if disabled + if context.server.disable_udp { + let error = eyre!("UDP is disabled"); + #[cfg(not(coverage_nightly))] + tracing::info!( + peer = %context.peer, %port, %error, + "Failed to bind UDP port.", + ); + let _ = context.tx.send( + format!( + "{} {} Cannot listen to UDP on port {}:{} ({})\r\n", + Utc::now().to_rfc3339().dimmed(), + " Error ".black().on_red().bold(), + context + .server + .domain + .as_deref() + .unwrap_or(""), + port, + error, + ) + .into_bytes(), + ); + Ok(false) + // Forbid binding low UDP ports + } else if (1..1024).contains(port) { + let error = eyre!("port too low"); + #[cfg(not(coverage_nightly))] + tracing::info!( + peer = %context.peer, %port, %error, + "Failed to bind UDP port.", + ); + let _ = context.tx.send( + format!( + "{} {} Cannot listen to UDP on port {}:{} ({})\r\n", + Utc::now().to_rfc3339().dimmed(), + " Error ".black().on_red().bold(), + context + .server + .domain + .as_deref() + .unwrap_or(""), + port, + error, + ) + .into_bytes(), + ); + Ok(false) + } else { + // When port is 0, assign a random one + let assigned_port = if *port == 0 { + let assigned_port = match context.server.udp_handler.get_free_port().await { + Ok(port) => port, + Err(error) => { + #[cfg(not(coverage_nightly))] + tracing::info!( + peer = %context.peer, alias = %address, %error, + "Failed to bind random UDP port for alias.", + ); + let _ = context.tx.send( + format!( + "{} {} Cannot listen to UDP on random port of {} ({})\r\n", + Utc::now().to_rfc3339().dimmed(), + " Error ".black().on_red().bold(), + address, + error, + ) + .into_bytes(), + ); + return Ok(false); + } + }; + // Set port to communicate it back to the client + *port = assigned_port.into(); + assigned_port + // Ignore user-requested port, assign any free one + } else if context.server.force_random_ports { + match context.server.udp_handler.get_free_port().await { + Ok(port) => port, + Err(error) => { + #[cfg(not(coverage_nightly))] + tracing::info!( + peer = %context.peer, alias = %address, %error, + "Failed to bind random UDP port for alias.", + ); + let _ = context.tx.send( + format!( + "{} {} Cannot listen to UDP on random port of {} ({})\r\n", + Utc::now().to_rfc3339().dimmed(), + " Error ".black().on_red().bold(), + port, + error, + ) + .into_bytes(), + ); + return Ok(false); + } + } + // Allow user-requested port when server allows binding on any port + } else { + match context + .server + .udp_handler + .create_port_socket(*port as u16) + .await + { + Ok(_) => (), + Err(error) => { + // Creating port listener failed. + #[cfg(not(coverage_nightly))] + tracing::info!( + peer = %context.peer, %port, %error, + "Rejecting UDP.", + ); + let _ = context.tx.send( + format!( + "{} {} Cannot listen to UDP on {}:{} ({})\r\n", + Utc::now().to_rfc3339().dimmed(), + " Error ".black().on_red().bold(), + context + .server + .domain + .as_deref() + .unwrap_or(""), + &port, + error, + ) + .into_bytes(), + ); + return Ok(false); + } + } + *port as u16 + }; + // Add handler to UDP connection map + let semaphore = Arc::new(Semaphore::new( + context.user_data.max_pool_size.load(Ordering::Acquire), + )); + match context.server.udp.insert( + assigned_port, + *context.peer, + context.user_data.quota_key.clone(), + Arc::new(SshTunnelHandler { + allow_fingerprint: Arc::clone(&context.user_data.allow_fingerprint), + http_data: None, + pool: Arc::clone(&semaphore), + pool_timeout: context.server.pool_timeout, + ip_connections: Arc::default(), + max_connections_per_ip: context.server.max_connections_per_ip, + ip_filter: Arc::clone(&context.user_data.ip_filter), + handle, + tx: context.tx.clone(), + peer: *context.peer, + address: address.to_string(), + port: *port, + limiter: context.user_data.limiter.clone(), + }), + ) { + Err(error) => { + // Adding to connection map failed. + #[cfg(not(coverage_nightly))] + tracing::info!( + peer = %context.peer, port = %assigned_port, %error, + "Rejecting UDP.", + ); + let _ = context.tx.send( + format!( + "{} {} Cannot listen to UDP on {}:{} ({})\r\n", + Utc::now().to_rfc3339().dimmed(), + " Error ".black().on_red().bold(), + context + .server + .domain + .as_deref() + .unwrap_or(""), + &assigned_port, + error, + ) + .into_bytes(), + ); + Ok(false) + } + _ => { + // Adding to connection map succeeded. + context.user_data.port_addressing.insert( + SockAddrAlias(address.to_string(), *port as u16), + (assigned_port, semaphore), + ); + #[cfg(not(coverage_nightly))] + tracing::info!( + peer = %context.peer, port = %assigned_port, + "Serving UDP...", + ); + let _ = context.tx.send( + format!( + "{} {:>14} port on {}:{}\r\n", + Utc::now().to_rfc3339().dimmed(), + "Starting UDP".green().bold(), + context + .server + .domain + .as_deref() + .unwrap_or(""), + &assigned_port, + ) + .into_bytes(), + ); + Ok(true) + } + } + } + } + + async fn cancel_remote_forwarding( + &mut self, + context: &mut RemoteForwardingContext<'_>, + address: &str, + port: u16, + ) -> Result { + if let Some(assigned_port) = context + .user_data + .port_addressing + .remove(&BorrowedSockAddrAlias(address, &{ port }) as &dyn SockAddrAliasKey) + { + context.server.udp.remove(&assigned_port.0, context.peer); + #[cfg(not(coverage_nightly))] + tracing::info!( + peer = %context.peer, port = %port, + "Stopped UDP forwarding.", + ); + let _ = context.tx.send( + format!( + "{} {:>14} port on {}:{}\r\n", + Utc::now().to_rfc3339().dimmed(), + "Stopping UDP".yellow().bold(), + context + .server + .domain + .as_deref() + .unwrap_or(""), + &assigned_port.0, + ) + .into_bytes(), + ); + Ok(true) + } else { + Ok(false) + } + } + + async fn local_forwarding( + &mut self, + context: &mut LocalForwardingContext<'_>, + address: &str, + port: u16, + originator_address: &str, + originator_port: u16, + channel: Channel, + ) -> Result { + let ip = context.peer.ip().to_canonical(); + if let Some(handler) = context.server.udp.get(&port, ip) + && let Ok(mut io) = handler + .aliasing_channel(ip, context.peer.port(), context.key_fingerprint.as_ref()) + .await + { + let gauge = gauge!(TELEMETRY_GAUGE_UDP_CONNECTIONS_CURRENT, TELEMETRY_KEY_PORT => port.to_string()); + gauge.increment(1); + counter!(TELEMETRY_COUNTER_UDP_CONNECTIONS, TELEMETRY_KEY_PORT => port.to_string()) + .increment(1); + let _ = handler.log_channel().send( + format!( + "{} {:>14} - {}:{} => {}:{}\r\n", + Utc::now().to_rfc3339().dimmed(), + "Proxying UDP".blue().bold(), + originator_address, + originator_port, + address, + port, + ) + .into_bytes(), + ); + match context.auth_data { + // Serve UDP for unauthed user, then add disconnection timeout if this is the last proxy connection + AuthenticatedData::None { proxy_data } => { + let guard = proxy_data.clone(); + let buffer_size = context.server.buffer_size; + tokio::spawn(async move { + let mut stream = channel.into_stream(); + let _ = copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await; + drop(guard); + gauge.decrement(1); + }); + } + // Serve UDP normally for authed user + _ => { + let buffer_size = context.server.buffer_size; + tokio::spawn(async move { + let mut stream = channel.into_stream(); + let _ = copy_bidirectional_with_sizes( + &mut stream, + &mut io, + buffer_size, + buffer_size, + ) + .await; + gauge.decrement(1); + }); + } + } + #[cfg(not(coverage_nightly))] + tracing::debug!( + peer = %context.peer, remote = %handler.peer, port = %port, + "Accepted UDP connection.", + ); + return Ok(true); + } + let _ = context.tx.send( + format!( + "{} {} Unknown UDP port '{}'\r\n", + Utc::now().to_rfc3339().dimmed(), + " Error ".black().on_red().bold(), + port + ) + .into_bytes(), + ); + Ok(false) + } +} diff --git a/src/tcp.rs b/src/tcp.rs index 0605b20..35fbf92 100644 --- a/src/tcp.rs +++ b/src/tcp.rs @@ -37,13 +37,13 @@ pub(crate) struct TcpHandler { disable_tcp_logs: bool, } -pub(crate) trait PortHandler { +pub(crate) trait TcpPortHandler { async fn create_port_listener(&self, port: u16) -> color_eyre::Result; async fn get_free_port(&self) -> color_eyre::Result; fn update_ports(&self, ports: Vec); } -impl PortHandler for Arc { +impl TcpPortHandler for Arc { // Create a TCP listener on the given port. async fn create_port_listener(&self, port: u16) -> color_eyre::Result { if self.sockets.contains_key(&port) { @@ -61,14 +61,13 @@ impl PortHandler for Arc { loop { match listener.accept().await { Ok((mut stream, address)) => { - let ip = address.ip(); + let ip = address.ip().to_canonical(); if !clone.ip_filter.is_allowed(ip) { #[cfg(not(coverage_nightly))] tracing::info!(%address, "Rejecting TCP connection: IP not allowed."); continue; } // Get the handler for this port - let ip = address.ip().to_canonical(); if let Some(handler) = clone.conn_manager.get(&port, ip) && let Ok(mut channel) = handler.tunneling_channel(ip, address.port()).await @@ -80,7 +79,7 @@ impl PortHandler for Arc { let _ = handler.log_channel().send( format!( "New connection from {}:{} to TCP port {}\r\n", - address.ip().to_canonical(), + ip, address.port(), port ) diff --git a/src/telemetry.rs b/src/telemetry.rs index 33a2c0c..5beef7b 100644 --- a/src/telemetry.rs +++ b/src/telemetry.rs @@ -23,7 +23,7 @@ use tokio::time::sleep; #[cfg_attr(not(feature = "prometheus"), allow(unused_imports))] use crate::droppable_handle::DroppableHandle; -use crate::tcp_alias::TcpAlias; +use crate::sock_addr_alias::SockAddrAlias; // A value that increases with time. struct SlidingWindowCounter { @@ -162,6 +162,7 @@ pub(crate) const TELEMETRY_COUNTER_ALIAS_CONNECTIONS: &str = "sandhole_alias_con pub(crate) const TELEMETRY_COUNTER_ADMIN_ALIAS_CONNECTIONS: &str = "sandhole_admin_alias_connections"; pub(crate) const TELEMETRY_COUNTER_TCP_CONNECTIONS: &str = "sandhole_tcp_connections"; +pub(crate) const TELEMETRY_COUNTER_UDP_CONNECTIONS: &str = "sandhole_udp_connections"; pub(crate) const TELEMETRY_GAUGE_USED_MEMORY: &str = "system_used_memory"; pub(crate) const TELEMETRY_GAUGE_TOTAL_MEMORY: &str = "system_total_memory"; pub(crate) const TELEMETRY_COUNTER_NETWORK_TX: &str = "system_network_tx"; @@ -174,6 +175,7 @@ pub(crate) const TELEMETRY_GAUGE_ALIAS_CONNECTIONS_CURRENT: &str = pub(crate) const TELEMETRY_GAUGE_ADMIN_ALIAS_CONNECTIONS_CURRENT: &str = "sandhole_admin_alias_connections_current"; pub(crate) const TELEMETRY_GAUGE_TCP_CONNECTIONS_CURRENT: &str = "sandhole_tcp_connections_current"; +pub(crate) const TELEMETRY_GAUGE_UDP_CONNECTIONS_CURRENT: &str = "sandhole_udp_connections_current"; pub(crate) const TELEMETRY_GAUGE_CPU_USAGE: &str = "system_cpu_usage"; pub(crate) const TELEMETRY_HISTOGRAM_HTTP_ELAPSED_TIME: &str = "sandhole_http_elapsed_time"; @@ -187,21 +189,26 @@ pub(crate) struct Telemetry { // Connections per minute for each SNI host. sni_connections_per_minute: DashMap, RandomState>, // Connections per minute for each local-forwarded alias. - alias_connections_per_minute: DashMap, RandomState>, + alias_connections_per_minute: DashMap, RandomState>, // Connections per minute for each admin alias. - admin_alias_connections_per_minute: DashMap, RandomState>, + admin_alias_connections_per_minute: + DashMap, RandomState>, // Connections per minute for each TCP port. tcp_connections_per_minute: DashMap, RandomState>, + // Connections per minute for each UDP port. + udp_connections_per_minute: DashMap, RandomState>, // Current connections for each SSH alias. ssh_connections_current: DashMap, RandomState>, // Current connections for each SNI host. sni_connections_current: DashMap, RandomState>, // Current connections for each local-forwarded alias. - alias_connections_current: DashMap, RandomState>, + alias_connections_current: DashMap, RandomState>, // Current connections for each admin alias. - admin_alias_connections_current: DashMap, RandomState>, + admin_alias_connections_current: DashMap, RandomState>, // Current connections for each TCP port. tcp_connections_current: DashMap, RandomState>, + // Current connections for each UDP port. + udp_connections_current: DashMap, RandomState>, // Recorder for Prometheus metrics export. #[cfg(feature = "prometheus")] prometheus_recorder: Option, @@ -246,11 +253,13 @@ impl Telemetry { alias_connections_per_minute: DashMap::default(), admin_alias_connections_per_minute: DashMap::default(), tcp_connections_per_minute: DashMap::default(), + udp_connections_per_minute: DashMap::default(), ssh_connections_current: DashMap::default(), sni_connections_current: DashMap::default(), alias_connections_current: DashMap::default(), admin_alias_connections_current: DashMap::default(), tcp_connections_current: DashMap::default(), + udp_connections_current: DashMap::default(), #[cfg(feature = "prometheus")] prometheus_recorder, #[cfg(feature = "prometheus")] @@ -283,6 +292,10 @@ impl Telemetry { TELEMETRY_COUNTER_TCP_CONNECTIONS, "Total connections for TCP ports" ); + describe_counter!( + TELEMETRY_COUNTER_UDP_CONNECTIONS, + "Total connections for UDP ports" + ); describe_counter!( TELEMETRY_COUNTER_NETWORK_TX, Unit::Bytes, @@ -310,6 +323,10 @@ impl Telemetry { TELEMETRY_GAUGE_TCP_CONNECTIONS_CURRENT, "Current requests for TCP ports" ); + describe_gauge!( + TELEMETRY_GAUGE_UDP_CONNECTIONS_CURRENT, + "Current requests for UDP ports" + ); describe_gauge!(TELEMETRY_GAUGE_CPU_USAGE, Unit::Percent, "Total CPU usage"); describe_gauge!(TELEMETRY_GAUGE_USED_MEMORY, Unit::Bytes, "Used memory"); describe_gauge!(TELEMETRY_GAUGE_TOTAL_MEMORY, Unit::Bytes, "Total memory"); @@ -349,7 +366,9 @@ impl Telemetry { .collect() } - pub(crate) fn get_alias_connections_per_minute(&self) -> HashMap { + pub(crate) fn get_alias_connections_per_minute( + &self, + ) -> HashMap { self.alias_connections_per_minute .iter() .map(|entry| (entry.key().clone(), entry.value().measure())) @@ -358,7 +377,7 @@ impl Telemetry { pub(crate) fn get_admin_alias_connections_per_minute( &self, - ) -> HashMap { + ) -> HashMap { self.admin_alias_connections_per_minute .iter() .map(|entry| (entry.key().clone(), entry.value().measure())) @@ -372,6 +391,13 @@ impl Telemetry { .collect() } + pub(crate) fn get_udp_connections_per_minute(&self) -> HashMap { + self.udp_connections_per_minute + .iter() + .map(|entry| (*entry.key(), entry.value().measure())) + .collect() + } + pub(crate) fn get_current_ssh_connections(&self) -> HashMap { self.ssh_connections_current .iter() @@ -386,7 +412,7 @@ impl Telemetry { .collect() } - pub(crate) fn get_current_alias_connections(&self) -> HashMap { + pub(crate) fn get_current_alias_connections(&self) -> HashMap { self.alias_connections_current .iter() .map(|entry| (entry.key().clone(), entry.value().measure())) @@ -395,7 +421,7 @@ impl Telemetry { pub(crate) fn get_current_admin_alias_connections( &self, - ) -> HashMap { + ) -> HashMap { self.admin_alias_connections_current .iter() .map(|entry| (entry.key().clone(), entry.value().measure())) @@ -409,6 +435,13 @@ impl Telemetry { .collect() } + pub(crate) fn get_current_udp_connections(&self) -> HashMap { + self.udp_connections_current + .iter() + .map(|entry| (*entry.key(), entry.value().measure())) + .collect() + } + pub(crate) fn ssh_reactor(&self, aliases: Vec) { let aliases: HashSet = aliases.into_iter().collect(); self.ssh_connections_per_minute @@ -431,8 +464,8 @@ impl Telemetry { .retain(|key, _| hostnames.contains(key)); } - pub(crate) fn alias_reactor(&self, aliases: Vec) { - let aliases: HashSet = aliases.into_iter().collect(); + pub(crate) fn alias_reactor(&self, aliases: Vec) { + let aliases: HashSet = aliases.into_iter().collect(); self.alias_connections_per_minute .retain(|key, _| aliases.contains(key)); self.alias_connections_current @@ -446,6 +479,14 @@ impl Telemetry { self.tcp_connections_current .retain(|key, _| ports.contains(key)); } + + pub(crate) fn udp_reactor(&self, ports: Vec) { + let ports: HashSet = ports.into_iter().collect(); + self.udp_connections_per_minute + .retain(|key, _| ports.contains(key)); + self.udp_connections_current + .retain(|key, _| ports.contains(key)); + } } impl Recorder for Telemetry { @@ -557,7 +598,7 @@ impl Recorder for Telemetry { TELEMETRY_COUNTER_ALIAS_CONNECTIONS => { for (key, value) in labels { if key == TELEMETRY_KEY_ALIAS { - match value.parse::() { + match value.parse::() { Ok(alias) => { return metrics::Counter::from_arc(Arc::clone( self.alias_connections_per_minute @@ -580,7 +621,7 @@ impl Recorder for Telemetry { TELEMETRY_COUNTER_ADMIN_ALIAS_CONNECTIONS => { for (key, value) in labels { if key == TELEMETRY_KEY_ALIAS { - match value.parse::() { + match value.parse::() { Ok(alias) => { return metrics::Counter::from_arc(Arc::clone( self.admin_alias_connections_per_minute @@ -623,6 +664,29 @@ impl Recorder for Telemetry { } } } + TELEMETRY_COUNTER_UDP_CONNECTIONS => { + for (key, value) in labels { + if key == TELEMETRY_KEY_PORT { + match value.parse::() { + Ok(port) => { + return metrics::Counter::from_arc(Arc::clone( + self.udp_connections_per_minute + .entry(port) + .or_insert(Arc::new(SlidingWindowCounter::new( + prometheus_counter, + Duration::from_secs(60), + ))) + .value(), + )); + } + Err(error) => { + #[cfg(not(coverage_nightly))] + tracing::warn!(port = value, %error, "Invalid port in telemetry."); + } + } + } + } + } _ => (), } prometheus_counter @@ -676,7 +740,7 @@ impl Recorder for Telemetry { TELEMETRY_GAUGE_ALIAS_CONNECTIONS_CURRENT => { for (key, value) in labels { if key == TELEMETRY_KEY_ALIAS { - match value.parse::() { + match value.parse::() { Ok(alias) => { return metrics::Gauge::from_arc(Arc::clone( self.alias_connections_current @@ -696,7 +760,7 @@ impl Recorder for Telemetry { TELEMETRY_GAUGE_ADMIN_ALIAS_CONNECTIONS_CURRENT => { for (key, value) in labels { if key == TELEMETRY_KEY_ALIAS { - match value.parse::() { + match value.parse::() { Ok(alias) => { return metrics::Gauge::from_arc(Arc::clone( self.admin_alias_connections_current @@ -733,6 +797,26 @@ impl Recorder for Telemetry { } } } + TELEMETRY_GAUGE_UDP_CONNECTIONS_CURRENT => { + for (key, value) in labels { + if key == TELEMETRY_KEY_PORT { + match value.parse::() { + Ok(port) => { + return metrics::Gauge::from_arc(Arc::clone( + self.udp_connections_current + .entry(port) + .or_insert(Arc::new(TelemetryGauge::new(prometheus_gauge))) + .value(), + )); + } + Err(error) => { + #[cfg(not(coverage_nightly))] + tracing::warn!(port = value, %error, "Invalid port in telemetry."); + } + } + } + } + } _ => (), } prometheus_gauge diff --git a/src/udp.rs b/src/udp.rs new file mode 100644 index 0000000..64f046b --- /dev/null +++ b/src/udp.rs @@ -0,0 +1,223 @@ +use std::{collections::HashSet, net::IpAddr, sync::Arc}; + +use crate::{ + connection_handler::ConnectionHandler, + connections::ConnectionMap, + droppable_handle::DroppableHandle, + ip::IpFilter, + reactor::UdpReactor, + ssh::connection_handler::SshTunnelHandler, + telemetry::{TELEMETRY_COUNTER_UDP_CONNECTIONS, TELEMETRY_KEY_PORT}, + udp_listener::get_udp_socket, +}; +use ahash::RandomState; +use bon::Builder; +use color_eyre::eyre::Context; +use dashmap::DashMap; +use futures_util::pin_mut; +use metrics::counter; +use tokio::{ + io::{AsyncReadExt, AsyncWriteExt}, + select, +}; + +pub const MAX_PACKET_SIZE: usize = std::mem::size_of::() + u16::MAX as usize; + +/// Creates and returns a buffer on the heap with enough space to contain any possible +/// UDP datagram. +/// +/// This is put on the heap and in a separate function to avoid the 64k buffer from ending +/// up on the stack and blowing up the size of the futures using it. +#[inline] +fn datagram_buffer() -> Box<[u8; MAX_PACKET_SIZE]> { + Box::new([0u8; MAX_PACKET_SIZE]) +} + +// Service that handles creating UDP sockets for reverse forwarding connections. +#[derive(Builder)] +pub(crate) struct UdpHandler { + // Address to listen to when creating sockets. + listen_address: IpAddr, + // Map containing spawned tasks of connections for each socket. + #[builder(skip = DashMap::default())] + sockets: DashMap, RandomState>, + // Connection map to assign a tunneling service for each incoming connection. + conn_manager: Arc, UdpReactor>>, + // Service that identifies whether to allow or block a given IP address. + ip_filter: Arc, + // Whether to send UDP logs to the SSH handles behind the forwarded connections. + disable_udp_logs: bool, +} + +pub(crate) trait UdpPortHandler { + async fn create_port_socket(&self, port: u16) -> color_eyre::Result; + async fn get_free_port(&self) -> color_eyre::Result; + fn update_ports(&self, ports: Vec); +} + +impl UdpPortHandler for Arc { + // Create a UDP listener on the given port. + async fn create_port_socket(&self, port: u16) -> color_eyre::Result { + if self.sockets.contains_key(&port) { + return Ok(port); + } + // Check if we're able to bind to the given address and port. + let socket = get_udp_socket((self.listen_address, port))?; + let port = socket + .local_addr() + .with_context(|| "Missing local address when binding port")? + .port(); + let clone = Arc::clone(self); + let listen_address = self.listen_address; + // Start task that will listen to incoming connections. + let join_handle = DroppableHandle(tokio::spawn(async move { + let mut socket = socket; + loop { + let mut buf = datagram_buffer(); + match socket.peek_from(buf.as_mut()).await { + Ok((_, address)) => { + let ip = address.ip().to_canonical(); + if !clone.ip_filter.is_allowed(ip) { + #[cfg(not(coverage_nightly))] + tracing::info!(%address, "Rejecting UDP connection: IP not allowed."); + continue; + } + // Get the handler for this port + if let Some(handler) = clone.conn_manager.get(&port, ip) + && let Ok(channel) = handler.tunneling_channel(ip, address.port()).await + { + counter!(TELEMETRY_COUNTER_UDP_CONNECTIONS, TELEMETRY_KEY_PORT => port.to_string()) + .increment(1); + // Log new connection to SSH handler + if !clone.disable_udp_logs { + let _ = handler.log_channel().send( + format!( + "New connection from {}:{} to UDP port {}\r\n", + ip, + address.port(), + port + ) + .into_bytes(), + ); + } + if let Err(error) = socket.connect(address).await { + #[cfg(not(coverage_nightly))] + tracing::error!(%port, %error, "Error connecting UDP socket."); + continue; + } + let connected_socket = socket; + socket = get_udp_socket((listen_address, port)).expect(""); + tokio::spawn(async move { + let udp_read = Arc::new(connected_socket); + let udp_write = udp_read.clone(); + let (mut ssh_read, mut ssh_write) = tokio::io::split(channel); + + let udp2ssh = async move { + let mut buf = datagram_buffer(); + loop { + match udp_read + .recv(&mut buf.as_mut()[std::mem::size_of::()..]) + .await + { + Ok(len) => { + *buf[..std::mem::size_of::()] + .as_mut_array() + .expect("length checked") = + (len as u16).to_be_bytes(); + if let Err(error) = ssh_write + .write_all( + &buf[..std::mem::size_of::() + len], + ) + .await + { + #[cfg(not(coverage_nightly))] + tracing::warn!(%port, %error, "Error writing to SSH channel for UDP."); + break; + }; + } + Err(error) => { + #[cfg(not(coverage_nightly))] + tracing::warn!(%port, %error, "Error reading from UDP socket."); + break; + } + } + } + }; + let ssh2udp = async move { + let mut buf = datagram_buffer(); + loop { + match ssh_read.read_u16().await { + Ok(len) => { + if let Err(error) = ssh_read + .read_exact(&mut buf[..len as usize]) + .await + { + #[cfg(not(coverage_nightly))] + tracing::warn!(%port, %error, "Error reading UDP datagram from SSH channel."); + break; + } else if let Err(error) = + udp_write.send(&mut buf[..len as usize]).await + { + #[cfg(not(coverage_nightly))] + tracing::warn!(%port, %error, "Error reading from SSH channel for UDP."); + break; + } + } + Err(error) => { + #[cfg(not(coverage_nightly))] + tracing::warn!(%port, %error, "Error reading UDP datagram size from SSH channel."); + break; + } + } + } + }; + + pin_mut!(udp2ssh); + pin_mut!(ssh2udp); + + select! { + _ = udp2ssh => {} + _ = ssh2udp => {} + } + }); + } + } + Err(error) => { + #[cfg(not(coverage_nightly))] + tracing::warn!(%port, %error, "Error reading from UDP port.") + } + } + } + })); + self.sockets.insert(port, join_handle); + Ok(port) + } + + // Create a TCP listener on a random open port, returning the port number. + async fn get_free_port(&self) -> color_eyre::Result { + // By passing 0 to create_port_listener, the OS will choose a port for us. + self.create_port_socket(0).await + } + + // Handle changes to the proxy ports, creating/deleting listeners as needed. + fn update_ports(&self, ports: Vec) { + // Find the ports listening to the localhost address + let mut ports: HashSet = ports.into_iter().collect(); + // Remove any socket tasks not in the list of localhost port + self.sockets.retain(|port, _| ports.contains(port)); + // Find the list of new ports + ports.retain(|port| !self.sockets.contains_key(port)); + if !ports.is_empty() { + let clone = Arc::clone(self); + // Create port listeners for the new ports + tokio::spawn(async move { + for port in ports.into_iter() { + if let Err(error) = clone.create_port_socket(port).await { + #[cfg(not(coverage_nightly))] + tracing::error!(%port, %error, "Failed to create listener for TCP port."); + } + } + }); + } + } +} diff --git a/src/udp_listener.rs b/src/udp_listener.rs new file mode 100644 index 0000000..f2e200c --- /dev/null +++ b/src/udp_listener.rs @@ -0,0 +1,40 @@ +use std::{io, net::ToSocketAddrs}; + +use socket2::{Domain, Socket, Type}; +use tokio::net::UdpSocket; + +// Create an async UDP listener. +pub(crate) fn get_udp_socket(addr: A) -> 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::DGRAM, + None, + )?; + + socket.set_nonblocking(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.bind(&addr.into())?; + + UdpSocket::from_std(socket.into()) +} diff --git a/tests/integration/main.rs b/tests/integration/main.rs index 7f8f9b8..fc0263d 100644 --- a/tests/integration/main.rs +++ b/tests/integration/main.rs @@ -99,5 +99,7 @@ mod tcp_rate_limit_upload; mod tcp_reject_low_ports; mod tcp_reject_port_above_max; mod tcp_timeout; +mod udp_allow_requested_ports; +mod udp_bind_random_ports; mod websocket_connection; mod websocket_timeout; diff --git a/tests/integration/udp_allow_requested_ports.rs b/tests/integration/udp_allow_requested_ports.rs new file mode 100644 index 0000000..d7c051e --- /dev/null +++ b/tests/integration/udp_allow_requested_ports.rs @@ -0,0 +1,271 @@ +use std::{sync::Arc, time::Duration}; + +use clap::Parser; +use rand::{RngExt, SeedableRng}; +use rand_chacha::ChaCha20Rng; +use russh::keys::ssh_key::private::Ed25519Keypair; +use russh::keys::{key::PrivateKeyWithHashAlg, load_secret_key}; +use russh::{ + Channel, + client::{Msg, Session}, +}; +use sandhole::{ApplicationConfig, entrypoint}; +use tokio::net::UdpSocket; +use tokio::{ + net::TcpStream, + time::{sleep, timeout}, +}; + +use crate::common::SandholeHandle; + +/// This test ensures that UDP works when binding for all ports is allowed. +#[test_log::test(tokio::test(flavor = "multi_thread"))] +async fn udp_allow_requested_ports() { + // 1. Initialize Sandhole + let config = ApplicationConfig::parse_from([ + "sandhole", + "--domain=foobar.tld", + "--user-keys-directory", + &(format!( + "{}/tests/data/user_keys", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--admin-keys-directory", + &(format!( + "{}/tests/data/admin_keys", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--certificates-directory", + &(format!( + "{}/tests/data/certificates", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--private-key-file", + &(format!( + "{}/tests/data/server_keys/ssh", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--acme-cache-directory", + &(format!( + "{}/tests/data/acme_cache", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--disable-directory-creation", + "--listen-address=127.0.0.1", + "--ssh-port=18022", + "--http-port=18080", + "--https-port=18443", + "--acme-use-staging", + "--bind-hostnames=none", + "--allow-requested-ports", + "--idle-connection-timeout=1s", + "--authentication-request-timeout=5s", + "--http-request-timeout=5s", + ]); + let _sandhole_handle = SandholeHandle(tokio::spawn(async move { entrypoint(config).await })); + if timeout(Duration::from_secs(5), async { + while TcpStream::connect("127.0.0.1:18022").await.is_err() { + sleep(Duration::from_millis(100)).await; + } + }) + .await + .is_err() + { + panic!("Timeout waiting for Sandhole to start.") + }; + + // 2. Start SSH client that will be proxied + let key = load_secret_key( + std::path::PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()) + .join("tests/data/private_keys/key1"), + None, + ) + .expect("Missing file key1"); + let ssh_client = SshClient; + let mut session_one = russh::client::connect(Default::default(), "127.0.0.1:18022", ssh_client) + .await + .expect("Failed to connect to SSH server"); + assert!( + session_one + .authenticate_publickey( + "user", + PrivateKeyWithHashAlg::new( + Arc::new(key), + session_one + .best_supported_rsa_hash() + .await + .unwrap() + .flatten() + ) + ) + .await + .expect("SSH authentication failed") + .success(), + "authentication didn't succeed" + ); + session_one + .tcpip_forward("udp.sandhole", 12345) + .await + .expect("tcpip_forward failed"); + + // 3. Connect to the UDP port of our proxy + let udp_socket = UdpSocket::bind("127.0.0.1:0") + .await + .expect("UDP connection failed"); + udp_socket.connect("127.0.0.1:12345").await.unwrap(); + udp_socket.send(b"Ping").await.unwrap(); + let mut buf = [0u8; 32]; + if timeout(Duration::from_secs(5), async { + assert_eq!(udp_socket.recv(&mut buf).await.unwrap(), 4); + }) + .await + .is_err() + { + panic!("Timeout waiting for UDP socket to reply.") + }; + assert_eq!(&buf[..4], b"Pong"); + + // 4. Local-forward the UDP port for random user + let key = russh::keys::PrivateKey::from(Ed25519Keypair::from_seed( + &ChaCha20Rng::from_rng(&mut rand::rng()).random(), + )); + let ssh_client = SshClient; + let mut session_two = russh::client::connect(Default::default(), "127.0.0.1:18022", ssh_client) + .await + .expect("Failed to connect to SSH server"); + assert!( + session_two + .authenticate_publickey( + "user", + PrivateKeyWithHashAlg::new( + Arc::new(key), + session_two + .best_supported_rsa_hash() + .await + .unwrap() + .flatten() + ) + ) + .await + .expect("SSH authentication failed") + .success(), + "authentication didn't succeed" + ); + let mut channel = session_two + .channel_open_direct_tcpip("udp.sandhole", 12345, "::1", 23456) + .await + .expect("Local forwarding failed"); + if timeout(Duration::from_secs(5), async { + channel + .data(&[0x00, 0x04, b'P', b'i', b'n', b'g'][..]) + .await + .unwrap(); + match &mut channel.wait().await.unwrap() { + russh::ChannelMsg::Data { data } => { + assert_eq!(data.to_vec(), [0x00, 0x04, b'P', b'o', b'n', b'g']); + } + msg => panic!("Unexpected message {msg:?}"), + } + }) + .await + .is_err() + { + panic!("Timeout waiting for proxy server to reply.") + }; + + // 5. Local-forward the UDP port for known user + let key = load_secret_key( + std::path::PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()) + .join("tests/data/private_keys/key2"), + None, + ) + .expect("Missing file key2"); + let ssh_client = SshClient; + let mut session_three = + russh::client::connect(Default::default(), "127.0.0.1:18022", ssh_client) + .await + .expect("Failed to connect to SSH server"); + assert!( + session_three + .authenticate_publickey( + "user", + PrivateKeyWithHashAlg::new( + Arc::new(key), + session_three + .best_supported_rsa_hash() + .await + .unwrap() + .flatten() + ) + ) + .await + .expect("SSH authentication failed") + .success(), + "authentication didn't succeed" + ); + let mut channel = session_three + .channel_open_direct_tcpip("udp.sandhole", 12345, "::1", 23456) + .await + .expect("Local forwarding failed"); + if timeout(Duration::from_secs(5), async { + channel + .data(&[0x00, 0x04, b'P', b'i', b'n', b'g'][..]) + .await + .unwrap(); + match &mut channel.wait().await.unwrap() { + russh::ChannelMsg::Data { data } => { + assert_eq!(data.to_vec(), [0x00, 0x04, b'P', b'o', b'n', b'g']); + } + msg => panic!("Unexpected message {msg:?}"), + } + }) + .await + .is_err() + { + panic!("Timeout waiting for proxy server to reply.") + }; + + // 6. Attempt to close UDP forwarding + session_one + .cancel_tcpip_forward("udp.sandhole", 12345) + .await + .expect("cancel_tcpip_forward failed"); +} + +struct SshClient; + +impl russh::client::Handler for SshClient { + type Error = color_eyre::eyre::Error; + + async fn check_server_key( + &mut self, + _key: &russh::keys::PublicKey, + ) -> Result { + Ok(true) + } + + async fn server_channel_open_forwarded_tcpip( + &mut self, + mut channel: Channel, + _connected_address: &str, + _connected_port: u32, + _originator_address: &str, + _originator_port: u32, + _session: &mut Session, + ) -> Result<(), Self::Error> { + tokio::spawn(async move { + match &mut channel.wait().await.unwrap() { + russh::ChannelMsg::Data { data } => { + assert_eq!(data.to_vec(), [0x00, 0x04, b'P', b'i', b'n', b'g']); + } + msg => panic!("Unexpected message {msg:?}"), + } + channel + .data(&[0x00, 0x04, b'P', b'o', b'n', b'g'][..]) + .await + .unwrap(); + channel.eof().await.unwrap(); + }); + Ok(()) + } +} diff --git a/tests/integration/udp_bind_random_ports.rs b/tests/integration/udp_bind_random_ports.rs new file mode 100644 index 0000000..1e30a9c --- /dev/null +++ b/tests/integration/udp_bind_random_ports.rs @@ -0,0 +1,203 @@ +use std::{sync::Arc, time::Duration}; + +use clap::Parser; +use russh::keys::{key::PrivateKeyWithHashAlg, load_secret_key}; +use russh::{ + Channel, + client::{Msg, Session}, +}; +use sandhole::{ApplicationConfig, entrypoint}; +use tokio::net::UdpSocket; +use tokio::{ + net::TcpStream, + time::{sleep, timeout}, +}; + +use crate::common::SandholeHandle; + +/// This test ensures that random ports work for UDP connections. +#[test_log::test(tokio::test(flavor = "multi_thread"))] +async fn udp_bind_random_ports() { + // 1. Initialize Sandhole + let config = ApplicationConfig::parse_from([ + "sandhole", + "--domain=foobar.tld", + "--user-keys-directory", + &(format!( + "{}/tests/data/user_keys", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--admin-keys-directory", + &(format!( + "{}/tests/data/admin_keys", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--certificates-directory", + &(format!( + "{}/tests/data/certificates", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--private-key-file", + &(format!( + "{}/tests/data/server_keys/ssh", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--acme-cache-directory", + &(format!( + "{}/tests/data/acme_cache", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--disable-directory-creation", + "--listen-address=127.0.0.1", + "--ssh-port=18022", + "--http-port=18080", + "--https-port=18443", + "--acme-use-staging", + "--bind-hostnames=none", + "--idle-connection-timeout=1s", + "--authentication-request-timeout=5s", + "--http-request-timeout=5s", + ]); + let _sandhole_handle = SandholeHandle(tokio::spawn(async move { entrypoint(config).await })); + if timeout(Duration::from_secs(5), async { + while TcpStream::connect("127.0.0.1:18022").await.is_err() { + sleep(Duration::from_millis(100)).await; + } + }) + .await + .is_err() + { + panic!("Timeout waiting for Sandhole to start.") + }; + + // 2. Start SSH client that will be proxied + let key = load_secret_key( + std::path::PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()) + .join("tests/data/private_keys/key1"), + None, + ) + .expect("Missing file key1"); + let ssh_client = SshClient; + let mut session_one = russh::client::connect(Default::default(), "127.0.0.1:18022", ssh_client) + .await + .expect("Failed to connect to SSH server"); + assert!( + session_one + .authenticate_publickey( + "user", + PrivateKeyWithHashAlg::new( + Arc::new(key), + session_one + .best_supported_rsa_hash() + .await + .unwrap() + .flatten() + ) + ) + .await + .expect("SSH authentication failed") + .success(), + "authentication didn't succeed" + ); + let mut channel = session_one + .channel_open_session() + .await + .expect("channel_open_session failed"); + session_one + .tcpip_forward("udp.sandhole", 12345) + .await + .expect("tcpip_forward failed"); + let regex = regex::Regex::new(r"foobar\.tld:(\d+)").expect("Invalid regex"); + let Ok(port) = timeout(Duration::from_secs(3), async move { + while let Some(message) = channel.wait().await { + match message { + russh::ChannelMsg::Data { data } => { + let data = + String::from_utf8(data.to_vec()).expect("Invalid UTF-8 from message"); + if let Some(captures) = regex.captures(&data) { + let port = captures + .get(1) + .expect("Missing port capture group") + .as_str() + .to_string(); + return port; + } + } + message => panic!("Unexpected message {message:?}"), + } + } + panic!("Unexpected end of channel"); + }) + .await + else { + panic!("Timed out waiting for port allocation."); + }; + assert!( + port.parse::().expect("should be a valid port number") >= 1024, + "random port must be greater than or equal to 1024" + ); + + // 3. Connect to the UDP port of our proxy + let udp_socket = UdpSocket::bind("127.0.0.1:0") + .await + .expect("UDP connection failed"); + udp_socket + .connect(format!("127.0.0.1:{port}")) + .await + .unwrap(); + udp_socket.send(b"Ping").await.unwrap(); + let mut buf = [0u8; 32]; + if timeout(Duration::from_secs(5), async { + assert_eq!(udp_socket.recv(&mut buf).await.unwrap(), 4); + }) + .await + .is_err() + { + panic!("Timeout waiting for UDP socket to reply.") + }; + assert_eq!(&buf[..4], b"Pong"); + + // 4. Attempt to close UDP forwarding + session_one + .cancel_tcpip_forward("udp.sandhole", 12345) + .await + .expect("cancel_tcpip_forward failed"); +} + +struct SshClient; + +impl russh::client::Handler for SshClient { + type Error = color_eyre::eyre::Error; + + async fn check_server_key( + &mut self, + _key: &russh::keys::PublicKey, + ) -> Result { + Ok(true) + } + + async fn server_channel_open_forwarded_tcpip( + &mut self, + mut channel: Channel, + _connected_address: &str, + _connected_port: u32, + _originator_address: &str, + _originator_port: u32, + _session: &mut Session, + ) -> Result<(), Self::Error> { + tokio::spawn(async move { + match &mut channel.wait().await.unwrap() { + russh::ChannelMsg::Data { data } => { + assert_eq!(data.to_vec(), [0x00, 0x04, b'P', b'i', b'n', b'g']); + } + msg => panic!("Unexpected message {msg:?}"), + } + channel + .data(&[0x00, 0x04, b'P', b'o', b'n', b'g'][..]) + .await + .unwrap(); + channel.eof().await.unwrap(); + }); + Ok(()) + } +}