From 89bc4cb86f6586e01e817a67fafdd198b8aafded Mon Sep 17 00:00:00 2001 From: Eric Rodrigues Pires Date: Mon, 30 Dec 2024 16:34:31 -0300 Subject: [PATCH] Add IP allowlist and blocklist Closes #30 --- CHANGELOG.md | 2 + Cargo.lock | 22 ++++ Cargo.toml | 6 +- book/src/cli.md | 8 ++ src/addressing.rs | 28 +++-- src/config.rs | 46 +++----- src/error.rs | 4 + src/ip.rs | 186 +++++++++++++++++++++++++++++++ src/lib.rs | 166 ++++++++++++++++----------- src/tcp.rs | 17 ++- tests/config_disable_aliasing.rs | 6 +- tests/config_disable_http.rs | 4 + tests/config_disable_tcp.rs | 11 ++ 13 files changed, 399 insertions(+), 107 deletions(-) create mode 100644 src/ip.rs diff --git a/CHANGELOG.md b/CHANGELOG.md index 4be576f..c7dc725 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,8 @@ - Add `--disable-tcp` CLI flag. - Add `--random-subdomain-filter-profanities` CLI flag. - Add `--requested-domain-filter-profanities` CLI flag. +- Add `--ip-allowlist` CLI flag. +- Add `--ip-blocklist` CLI flag. ### Changed diff --git a/Cargo.lock b/Cargo.lock index 45ca903..dd0a56e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2052,6 +2052,16 @@ version = "2.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ddc24109865250148c2e0f3d25d4f0f479571723792d3802153c60922a4fb708" +[[package]] +name = "ipnet-trie" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fd839300aeeb010677522802bd17878847da94e4871ebe2653aa8087c5385af" +dependencies = [ + "ipnet", + "prefix-trie", +] + [[package]] name = "is_terminal_polyfill" version = "1.70.1" @@ -2816,6 +2826,16 @@ dependencies = [ "termtree", ] +[[package]] +name = "prefix-trie" +version = "0.2.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85fe48f29e6e6fcf123d0d03d63028dbe4c4a738023d35d525df4882f4929418" +dependencies = [ + "ipnet", + "num-traits", +] + [[package]] name = "pretty-duration" version = "0.1.1" @@ -3438,6 +3458,8 @@ dependencies = [ "hyper", "hyper-util", "insta", + "ipnet", + "ipnet-trie", "itertools 0.13.0", "log", "mockall", diff --git a/Cargo.toml b/Cargo.toml index aedb1f5..bb71f34 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -5,11 +5,11 @@ edition = "2021" rust-version = "1.82.0" description = "Expose HTTP/SSH/TCP services through SSH port forwarding." repository = "https://github.com/EpicEric/sandhole" -homepage = "https://epiceric.github.io/sandhole/" +homepage = "https://sandhole.eric.dev.br/" license = "MIT" authors = ["Eric Rodrigues Pires "] readme = "README.md" -keywords = ["ssh", "proxy"] +keywords = ["ssh", "proxy", "tunnel", "hole-punching"] categories = ["network-programming", "web-programming", "authentication"] exclude = [".github", "book", "docker-compose-example", "tests/data"] @@ -31,6 +31,8 @@ human_bytes = { version = "0.4.3", default-features = false } humantime = "2.1.0" hyper = { version = "1.5.0", features = ["full"] } hyper-util = { version = "0.1", features = ["full"] } +ipnet = "2.10.1" +ipnet-trie = "0.2.0" itertools = "0.13.0" log = "0.4.22" notify = "7.0.0" diff --git a/book/src/cli.md b/book/src/cli.md index 5a5ad60..1674101 100644 --- a/book/src/cli.md +++ b/book/src/cli.md @@ -210,6 +210,14 @@ Expose HTTP/SSH/TCP services through SSH port forwarding. Beware that this can lead to false positives being blocked! + --ip-allowlist <CIDR> + Comma-separated list of IP networks to allow. Setting this will block + unknown IPs from connecting + + --ip-blocklist <CIDR> + Comma-separated list of IP networks to block. Setting this will allow + unknown IPs to connect, unless --ip-allowlist is set + --idle-connection-timeout <DURATION> Grace period for dangling/unauthenticated SSH connections before they are forcefully disconnected. diff --git a/src/addressing.rs b/src/addressing.rs index fcb5d40..31ed3d8 100644 --- a/src/addressing.rs +++ b/src/addressing.rs @@ -365,14 +365,15 @@ impl AddressDelegator { #[cfg(test)] mod address_delegator_tests { + use std::{collections::HashSet, net::SocketAddr}; + use mockall::predicate::*; use rand::rngs::OsRng; use regex::Regex; use ssh_key::HashAlg; - use std::net::SocketAddr; use webpki::types::DnsName; - use crate::{config::BindHostnames, RandomSubdomainSeed}; + use crate::config::{BindHostnames, RandomSubdomainSeed}; use super::{AddressDelegator, AddressDelegatorData, MockResolver}; @@ -841,10 +842,12 @@ mod address_delegator_tests { random_subdomain_filter_profanities: false, requested_domain_filter: None, }); - let mut set = std::collections::HashSet::with_capacity(200_000); - let regex = Regex::new(r"^[0-9a-z]{6}\.root\.tld$").unwrap(); // 99.99% chance of collision with naïve implementation - for _ in 0..200_000 { + static SIZE: usize = 200_000; + let mut set = HashSet::with_capacity(SIZE); + let regex = Regex::new(r"^[0-9a-z]{6}\.root\.tld$").unwrap(); + let initial_block_rng = *delegator.block_rng.lock().unwrap(); + for _ in 0..SIZE { let address = delegator .get_http_address( "some.address", @@ -866,6 +869,8 @@ mod address_delegator_tests { ); set.insert(address); } + let final_block_rng = *delegator.block_rng.lock().unwrap(); + assert_eq!((final_block_rng - initial_block_rng) as usize, SIZE); } #[tokio::test] @@ -887,10 +892,12 @@ mod address_delegator_tests { random_subdomain_filter_profanities: true, requested_domain_filter: None, }); - let mut set = std::collections::HashSet::with_capacity(5_600); + // 99.99999...% chance of collision with naïve implementation + static SIZE: usize = 10_000; + let mut set = HashSet::with_capacity(SIZE); let regex = Regex::new(r"^[0-9a-z]{4}\.root\.tld$").unwrap(); - // 99.99% chance of collision with naïve implementation - for _ in 0..5_600 { + let initial_block_rng = *delegator.block_rng.lock().unwrap(); + for _ in 0..SIZE { let address = delegator .get_http_address( "some.address", @@ -912,6 +919,11 @@ mod address_delegator_tests { ); set.insert(address); } + let final_block_rng = *delegator.block_rng.lock().unwrap(); + assert!( + (final_block_rng - initial_block_rng) as usize > SIZE, + "expected at least one word to get filtered" + ); } #[tokio::test] diff --git a/src/config.rs b/src/config.rs index a831823..cd31f50 100644 --- a/src/config.rs +++ b/src/config.rs @@ -2,6 +2,7 @@ use std::{num::NonZero, path::PathBuf}; use clap::{command, Parser, ValueEnum}; use humantime::Duration; +use ipnet::IpNet; use webpki::types::DnsName; // Which value to seed with when generating random subdomains, for determinism. @@ -126,31 +127,16 @@ pub struct ApplicationConfig { pub listen_address: String, /// Port to listen for SSH connections. - #[arg( - long, - default_value_t = 2222, - value_parser = validate_port, - value_name = "PORT" - )] - pub ssh_port: u16, + #[arg(long, default_value_t = NonZero::new(2222).unwrap(), value_name = "PORT")] + pub ssh_port: NonZero, /// Port to listen for HTTP connections. - #[arg( - long, - default_value_t = 80, - value_parser = validate_port, - value_name = "PORT" - )] - pub http_port: u16, + #[arg(long, default_value_t = NonZero::new(80).unwrap(), value_name = "PORT")] + pub http_port: NonZero, /// Port to listen for HTTPS connections. - #[arg( - long, - default_value_t = 443, - value_parser = validate_port, - value_name = "PORT" - )] - pub https_port: u16, + #[arg(long, default_value_t = NonZero::new(443).unwrap(), value_name = "PORT")] + pub https_port: NonZero, /// Allow connecting to SSH via the HTTPS port as well. /// This can be useful in networks that block binding to other ports. @@ -279,6 +265,16 @@ pub struct ApplicationConfig { #[arg(long, default_value_t = false)] pub requested_domain_filter_profanities: bool, + /// Comma-separated list of IP networks to allow. + /// Setting this will block unknown IPs from connecting. + #[arg(long, value_delimiter = ',', value_name = "CIDR")] + pub ip_allowlist: Option>, + + /// Comma-separated list of IP networks to block. + /// Setting this will allow unknown IPs to connect, unless --ip-allowlist is set. + #[arg(long, value_delimiter = ',', value_name = "CIDR")] + pub ip_blocklist: Option>, + /// 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. @@ -322,11 +318,3 @@ fn validate_txt_record_prefix(value: &str) -> Result { Ok(value.to_string()) } } - -fn validate_port(value: &str) -> Result { - match value.parse::() { - Err(err) => Err(err.to_string()), - Ok(0) => Err("port cannot be zero".into()), - Ok(port) => Ok(port), - } -} diff --git a/src/error.rs b/src/error.rs index 1cc8bc9..6dc4ab3 100644 --- a/src/error.rs +++ b/src/error.rs @@ -1,5 +1,7 @@ use std::path::PathBuf; +use ipnet::IpNet; + #[derive(thiserror::Error, Debug)] pub(crate) enum ServerError { #[error("Invalid config: {0}")] @@ -26,4 +28,6 @@ pub(crate) enum ServerError { UnknownHttpScheme, #[error("Missing directory {0}")] MissingDirectory(PathBuf), + #[error("Duplicate network CIDR {0}")] + DuplicateNetworkCidr(IpNet), } diff --git a/src/ip.rs b/src/ip.rs new file mode 100644 index 0000000..65f576f --- /dev/null +++ b/src/ip.rs @@ -0,0 +1,186 @@ +use std::net::IpAddr; + +use ipnet::IpNet; +use ipnet_trie::IpnetTrie; + +use crate::error::ServerError; + +// Connection policy applied to an IP range. +#[derive(PartialEq, Eq, Clone, Copy)] +enum IpPolicy { + Allow, + Deny, +} + +// Service that identifies whether to allow or block a given IP address. +pub(crate) struct IpFilter { + // Which policy to apply for IPs not found in the trie. + default_policy: IpPolicy, + // Trie for efficient lookup of IPs by the network prefix. + data: IpnetTrie, +} + +pub(crate) struct IpFilterConfig { + pub(crate) allowlist: Option>, + pub(crate) blocklist: Option>, +} + +impl IpFilter { + pub(crate) fn new(config: IpFilterConfig) -> anyhow::Result { + let IpFilterConfig { + allowlist, + blocklist, + } = config; + let mut data = IpnetTrie::new(); + let mut default_policy = IpPolicy::Allow; + if let Some(allowlist) = allowlist { + if !allowlist.is_empty() { + default_policy = IpPolicy::Deny; + } + for network in allowlist { + if data.insert(network, IpPolicy::Allow).is_some() { + return Err(ServerError::DuplicateNetworkCidr(network).into()); + } + } + } + if let Some(blocklist) = blocklist { + for network in blocklist { + if data.insert(network, IpPolicy::Deny).is_some() { + return Err(ServerError::DuplicateNetworkCidr(network).into()); + } + } + } + Ok(IpFilter { + default_policy, + data, + }) + } + + pub(crate) fn is_allowed(&self, address: IpAddr) -> bool { + self.data + .longest_match(&IpNet::from(address.to_canonical())) + .map(|(_, policy)| *policy) + .unwrap_or_else(|| self.default_policy) + == IpPolicy::Allow + } +} + +#[cfg(test)] +mod ip_filter_tests { + use std::{net::IpAddr, str::FromStr}; + + use ipnet::IpNet; + + use super::{IpFilter, IpFilterConfig}; + + #[test] + fn should_allow_anyone_if_no_lists() { + let filter = IpFilter::new(IpFilterConfig { + allowlist: None, + blocklist: None, + }) + .unwrap(); + assert!(filter.is_allowed(IpAddr::from_str("127.0.0.1").unwrap())); + assert!(filter.is_allowed(IpAddr::from_str("10.0.2.127").unwrap())); + assert!(filter.is_allowed(IpAddr::from_str("1234:dead:beef::154").unwrap())); + assert!(filter.is_allowed(IpAddr::from_str("1234:0db8:502e::3c").unwrap())); + } + + #[test] + fn should_allow_anyone_if_empty_lists() { + let filter = IpFilter::new(IpFilterConfig { + allowlist: Some(vec![]), + blocklist: Some(vec![]), + }) + .unwrap(); + assert!(filter.is_allowed(IpAddr::from_str("127.0.0.1").unwrap())); + assert!(filter.is_allowed(IpAddr::from_str("10.0.2.127").unwrap())); + assert!(filter.is_allowed(IpAddr::from_str("1234:dead:beef::154").unwrap())); + assert!(filter.is_allowed(IpAddr::from_str("1234:0db8:502e::3c").unwrap())); + } + + #[test] + fn should_allow_addresses_not_in_blocklist() { + let filter = IpFilter::new(IpFilterConfig { + allowlist: None, + blocklist: Some(vec![ + IpNet::from_str("10.0.0.0/20").unwrap(), + IpNet::from_str("1234:dead::/32").unwrap(), + ]), + }) + .unwrap(); + assert!(filter.is_allowed(IpAddr::from_str("127.0.0.1").unwrap())); + assert!(!filter.is_allowed(IpAddr::from_str("10.0.2.127").unwrap())); + assert!(!filter.is_allowed(IpAddr::from_str("1234:dead:beef::154").unwrap())); + assert!(filter.is_allowed(IpAddr::from_str("1234:0db8:502e::3c").unwrap())); + } + + #[test] + fn should_reject_addresses_not_in_allowlist() { + let filter = IpFilter::new(IpFilterConfig { + allowlist: Some(vec![ + IpNet::from_str("127.0.0.0/24").unwrap(), + IpNet::from_str("10.0.0.0/18").unwrap(), + ]), + blocklist: None, + }) + .unwrap(); + assert!(filter.is_allowed(IpAddr::from_str("127.0.0.1").unwrap())); + assert!(filter.is_allowed(IpAddr::from_str("10.0.2.127").unwrap())); + assert!(!filter.is_allowed(IpAddr::from_str("1234:dead:beef::154").unwrap())); + assert!(!filter.is_allowed(IpAddr::from_str("1234:0db8:502e::3c").unwrap())); + } + + #[test] + fn should_only_accept_allowlist() { + let filter = IpFilter::new(IpFilterConfig { + allowlist: Some(vec![ + IpNet::from_str("127.0.0.0/24").unwrap(), + IpNet::from_str("10.0.0.0/18").unwrap(), + ]), + blocklist: Some(vec![ + IpNet::from_str("10.0.0.0/20").unwrap(), + IpNet::from_str("1234:dead::/32").unwrap(), + ]), + }) + .unwrap(); + assert!(filter.is_allowed(IpAddr::from_str("127.0.0.1").unwrap())); + assert!(!filter.is_allowed(IpAddr::from_str("10.0.2.127").unwrap())); + assert!(!filter.is_allowed(IpAddr::from_str("1234:dead:beef::154").unwrap())); + assert!(!filter.is_allowed(IpAddr::from_str("1234:0db8:502e::3c").unwrap())); + } + + #[test] + fn should_fail_if_duplicated_network() { + assert!( + IpFilter::new(IpFilterConfig { + allowlist: Some(vec![IpNet::from_str("127.0.0.0/24").unwrap()]), + blocklist: Some(vec![IpNet::from_str("127.0.0.0/24").unwrap()]), + }) + .is_err(), + "shouldn't allow same network in both allowlist and blocklist" + ); + assert!( + IpFilter::new(IpFilterConfig { + allowlist: Some(vec![ + IpNet::from_str("127.0.0.0/24").unwrap(), + IpNet::from_str("127.0.0.0/24").unwrap() + ]), + blocklist: None, + }) + .is_err(), + "shouldn't allow same network in allowlist twice" + ); + assert!( + IpFilter::new(IpFilterConfig { + allowlist: None, + blocklist: Some(vec![ + IpNet::from_str("127.0.0.0/24").unwrap(), + IpNet::from_str("127.0.0.0/24").unwrap() + ]), + }) + .is_err(), + "shouldn't allow same network in blocklist twice" + ); + } +} diff --git a/src/lib.rs b/src/lib.rs index a436306..5d41b13 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -7,6 +7,7 @@ use std::{ future, marker::PhantomData, net::{IpAddr, SocketAddr}, + num::NonZero, sync::{atomic::AtomicUsize, Arc, Mutex, RwLock}, time::Duration, }; @@ -20,6 +21,7 @@ use hyper_util::{ rt::{TokioExecutor, TokioIo}, server::conn::auto, }; +use ip::{IpFilter, IpFilterConfig}; use log::{debug, error, info, warn}; use login::ApiLogin; use quota::{DummyQuotaHandler, QuotaHandler, QuotaMap}; @@ -73,6 +75,7 @@ mod droppable_handle; mod error; mod fingerprints; mod http; +mod ip; mod login; mod quota; mod ssh; @@ -255,11 +258,13 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { .with_context(|| "Error intializing login API")?; // Initialize the ACME ALPN service if a contact email has been provided. let alpn_resolver: Box = match config.acme_contact_email { - Some(contact) if config.https_port == 443 => Box::new(AcmeResolver::new( - config.acme_cache_directory, - contact, - config.acme_use_staging, - )), + Some(contact) if config.https_port == NonZero::new(443).unwrap() => { + Box::new(AcmeResolver::new( + config.acme_cache_directory, + contact, + config.acme_use_staging, + )) + } Some(_) => { warn!( "ACME challenges are only supported on HTTPS port 443 (currently {}). Disabling.", @@ -283,6 +288,11 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { .await .with_context(|| "Error setting up certificates watcher")?, ); + // Initialize the IP address allowlist/blocklist service. + let ip_filter = Arc::new(IpFilter::new(IpFilterConfig { + allowlist: config.ip_allowlist, + blocklist: config.ip_blocklist, + })?); let telemetry = Arc::new(Telemetry::new()); let quota_handler: Arc> = match config.quota_per_user { Some(max_quota) => Arc::new(Box::new(Arc::new(QuotaMap::new(max_quota.into())))), @@ -314,6 +324,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { let tcp_handler: Arc = Arc::new(TcpHandler::new( config.listen_address, Arc::clone(&tcp_connections), + Arc::clone(&ip_filter), config.tcp_connection_timeout.map(Into::into), config.disable_tcp_logs, )); @@ -353,58 +364,66 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { // Telemetry tasks let http_data = Arc::new(RwLock::default()); - let data_clone = Arc::clone(&http_data); - let connections_clone = Arc::clone(&http_connections); - let telemetry_clone = Arc::clone(&telemetry); - // Periodically update HTTP data, based on the connection map and the telemetry counters. - tokio::spawn(async move { - loop { - sleep(Duration::from_millis(3_000)).await; - let data = connections_clone.data(); - let telemetry = telemetry_clone.get_http_requests_per_minute(); - let data = data - .into_iter() - .map(|(hostname, addresses)| { - let requests_per_minute = *telemetry.get(&hostname).unwrap_or(&0f64); - (hostname, (addresses, requests_per_minute)) - }) - .collect(); - *data_clone.write().unwrap() = data; - } - }); + if !config.disable_http { + let data_clone = Arc::clone(&http_data); + let connections_clone = Arc::clone(&http_connections); + let telemetry_clone = Arc::clone(&telemetry); + // Periodically update HTTP data, based on the connection map and the telemetry counters. + tokio::spawn(async move { + loop { + sleep(Duration::from_millis(3_000)).await; + let data = connections_clone.data(); + let telemetry = telemetry_clone.get_http_requests_per_minute(); + let data = data + .into_iter() + .map(|(hostname, addresses)| { + let requests_per_minute = *telemetry.get(&hostname).unwrap_or(&0f64); + (hostname, (addresses, requests_per_minute)) + }) + .collect(); + *data_clone.write().unwrap() = data; + } + }); + } let ssh_data = Arc::new(RwLock::default()); - let data_clone = Arc::clone(&ssh_data); - let connections_clone = Arc::clone(&ssh_connections); - // Periodically update SSH data, based on the connection map. - tokio::spawn(async move { - loop { - sleep(Duration::from_millis(3_000)).await; - let data = connections_clone.data(); - *data_clone.write().unwrap() = data; - } - }); + if !config.disable_aliasing { + let data_clone = Arc::clone(&ssh_data); + let connections_clone = Arc::clone(&ssh_connections); + // Periodically update SSH data, based on the connection map. + tokio::spawn(async move { + loop { + sleep(Duration::from_millis(3_000)).await; + let data = connections_clone.data(); + *data_clone.write().unwrap() = data; + } + }); + } let tcp_data = Arc::new(RwLock::default()); - let data_clone = Arc::clone(&tcp_data); - let connections_clone = Arc::clone(&tcp_connections); - // Periodically update TCP data, based on the connection map. - tokio::spawn(async move { - loop { - sleep(Duration::from_millis(3_000)).await; - let data = connections_clone.data(); - *data_clone.write().unwrap() = data; - } - }); + if !config.disable_tcp { + let data_clone = Arc::clone(&tcp_data); + let connections_clone = Arc::clone(&tcp_connections); + // Periodically update TCP data, based on the connection map. + tokio::spawn(async move { + loop { + sleep(Duration::from_millis(3_000)).await; + let data = connections_clone.data(); + *data_clone.write().unwrap() = data; + } + }); + } let alias_data = Arc::new(RwLock::default()); - let data_clone = Arc::clone(&alias_data); - let connections_clone = Arc::clone(&alias_connections); - // Periodically update alias data, based on the connection map. - tokio::spawn(async move { - loop { - sleep(Duration::from_millis(3_000)).await; - let data = connections_clone.data(); - *data_clone.write().unwrap() = data; - } - }); + if !config.disable_aliasing { + let data_clone = Arc::clone(&alias_data); + let connections_clone = Arc::clone(&alias_connections); + // Periodically update alias data, based on the connection map. + tokio::spawn(async move { + loop { + sleep(Duration::from_millis(3_000)).await; + let data = connections_clone.data(); + *data_clone.write().unwrap() = data; + } + }); + } let system_data = Arc::new(RwLock::default()); let data_clone = Arc::clone(&system_data); // Periodically update system data (every second, as to keep network TX/RX rates accurate). @@ -455,7 +474,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { domain_redirect: Arc::clone(&domain_redirect), // HTTP only. protocol: Protocol::Http { - port: config.http_port, + port: config.http_port.into(), }, // Always use aliasing channels instead of tunneling channels. proxy_type: ProxyType::Aliasing, @@ -483,9 +502,9 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { address_delegator: addressing, tcp_handler, domain: config.domain, - http_port: config.http_port, - https_port: config.https_port, - ssh_port: config.ssh_port, + http_port: config.http_port.into(), + https_port: config.https_port.into(), + ssh_port: config.ssh_port.into(), force_random_ports: !config.allow_requested_ports, disable_http: config.disable_http, disable_tcp: config.disable_tcp, @@ -503,13 +522,14 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { let mut join_handle_http = if config.disable_http { tokio::spawn(future::pending()) } else { - let http_listener = TcpListener::bind((listen_address, config.http_port)) + let http_listener = TcpListener::bind((listen_address, config.http_port.into())) .await .with_context(|| "Error listening to HTTP port")?; info!( "Listening for HTTP connections on port {}.", config.http_port ); + let ip_filter_clone = Arc::clone(&ip_filter); let http_proxy_data = Arc::new(ProxyData { conn_manager: Arc::clone(&http_connections), telemetry: Arc::clone(&telemetry), @@ -517,12 +537,12 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { // Use TLS redirect if --force-https is set, otherwise allow HTTP. protocol: if config.force_https { Protocol::TlsRedirect { - from: config.http_port, - to: config.https_port, + from: config.http_port.into(), + to: config.https_port.into(), } } else { Protocol::Http { - port: config.http_port, + port: config.http_port.into(), } }, // Always use tunneling channels. @@ -542,6 +562,11 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { break; } }; + let ip = address.ip(); + if !ip_filter_clone.is_allowed(ip) { + info!("Rejecting HTTP connection for {}: not allowed", ip); + continue; + } // Create a Hyper service and serve over the accepted TCP connection. let service = service_fn(move |req: Request| { proxy_handler(req, address, None, Arc::clone(&proxy_data)) @@ -560,7 +585,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { let mut join_handle_https = if config.disable_http { tokio::spawn(future::pending()) } else { - let https_listener = TcpListener::bind((listen_address, config.https_port)) + let https_listener = TcpListener::bind((listen_address, config.https_port.into())) .await .with_context(|| "Error listening to HTTPS port")?; info!( @@ -573,12 +598,13 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { .with_no_client_auth() .with_cert_resolver(certificates), ); + let ip_filter_clone = Arc::clone(&ip_filter); let https_proxy_data = Arc::new(ProxyData { conn_manager: http_connections, telemetry: Arc::clone(&telemetry), domain_redirect: Arc::clone(&domain_redirect), protocol: Protocol::Https { - port: config.https_port, + port: config.https_port.into(), }, // Always use tunneling channels. proxy_type: ProxyType::Tunneling, @@ -598,6 +624,11 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { break; } }; + let ip = address.ip(); + if !ip_filter_clone.is_allowed(ip) { + info!("Rejecting HTTPS connection for {}: not allowed", ip); + continue; + } if config.connect_ssh_on_https_port { // Check if this is an SSH-2.0 handshake. let mut buf = [0u8; 8]; @@ -669,7 +700,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { }; // Start Sandhole on SSH port - let ssh_listener = TcpListener::bind((listen_address, config.ssh_port)) + let ssh_listener = TcpListener::bind((listen_address, config.ssh_port.into())) .await .with_context(|| "Error listening to SSH port")?; info!("Listening for SSH connections on port {}.", config.ssh_port); @@ -687,6 +718,11 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { break; }, }; + let ip = address.ip(); + if !ip_filter.is_allowed(ip) { + info!("Rejecting SSH connection for {}: not allowed", ip); + continue; + } handle_ssh_connection(stream, address, &ssh_config, &mut sandhole).await; } _ = &mut signal_handler => { diff --git a/src/tcp.rs b/src/tcp.rs index ee80725..4500900 100644 --- a/src/tcp.rs +++ b/src/tcp.rs @@ -4,12 +4,13 @@ use crate::{ connection_handler::ConnectionHandler, connections::{ConnectionMap, ConnectionMapReactor}, droppable_handle::DroppableHandle, + ip::IpFilter, ssh::SshTunnelHandler, }; use anyhow::Context; use async_trait::async_trait; use dashmap::DashMap; -use log::error; +use log::{error, info}; use tokio::{io::copy_bidirectional, net::TcpListener, time::timeout}; // Service that handles creating TCP sockets for reverse forwarding connections. @@ -20,6 +21,8 @@ pub(crate) struct TcpHandler { sockets: DashMap>, // Connection map to assign a tunneling service for each incoming connection. conn_manager: Arc, Arc>>, + // Service that identifies whether to allow or block a given IP address. + ip_filter: Arc, // Optional duration to time out TCP connections. tcp_connection_timeout: Option, // Whether to send TCP logs to the SSH handles behind the forwarded connections. @@ -30,6 +33,7 @@ impl TcpHandler { pub(crate) fn new( listen_address: String, conn_manager: Arc, Arc>>, + ip_filter: Arc, tcp_connection_timeout: Option, disable_tcp_logs: bool, ) -> Self { @@ -37,6 +41,7 @@ impl TcpHandler { listen_address, sockets: DashMap::new(), conn_manager, + ip_filter, tcp_connection_timeout, disable_tcp_logs, } @@ -65,11 +70,17 @@ impl PortHandler for Arc { let clone = Arc::clone(self); let tcp_connection_timeout = self.tcp_connection_timeout; let disable_tcp_logs = self.disable_tcp_logs; + let ip_filter_clone = Arc::clone(&self.ip_filter); // Start task that will listen to incoming connections. let join_handle = DroppableHandle(tokio::spawn(async move { loop { match listener.accept().await { Ok((mut stream, address)) => { + let ip = address.ip(); + if !ip_filter_clone.is_allowed(ip) { + info!("Rejecting TCP connection for {}: not allowed", ip); + continue; + } // Get the handler for this port if let Some(handler) = clone.conn_manager.get(&port) { if let Ok(mut channel) = handler @@ -137,7 +148,9 @@ impl ConnectionMapReactor for Arc { // Create port listeners for the new ports tokio::spawn(async move { for port in ports.into_iter() { - let _ = clone.create_port_listener(port).await; + if let Err(err) = clone.create_port_listener(port).await { + error!("Failed to create listener for port {}: {}", port, err); + } } }); } diff --git a/tests/config_disable_aliasing.rs b/tests/config_disable_aliasing.rs index 53a6d52..395b613 100644 --- a/tests/config_disable_aliasing.rs +++ b/tests/config_disable_aliasing.rs @@ -112,7 +112,11 @@ async fn config_disable_aliasing() { .tcpip_forward("test.foobar.tld", 80) .await .is_ok(), - "shouldn't have failed to bind regular HTTP" + "shouldn't have failed to bind HTTP" + ); + assert!( + session_one.tcpip_forward("localhost", 12345).await.is_ok(), + "shouldn't have failed to bind TCP" ); // 3. Start SSH proxy that will fail to local forward diff --git a/tests/config_disable_http.rs b/tests/config_disable_http.rs index 1994ff8..3cdac4f 100644 --- a/tests/config_disable_http.rs +++ b/tests/config_disable_http.rs @@ -107,6 +107,10 @@ async fn config_disable_http() { session_one.tcpip_forward("some.address", 80).await.is_ok(), "shouldn't have failed to create HTTP alias" ); + assert!( + session_one.tcpip_forward("some.proxy", 90).await.is_ok(), + "shouldn't have failed to bind alias" + ); } struct SshClient; diff --git a/tests/config_disable_tcp.rs b/tests/config_disable_tcp.rs index 6544249..fc049bf 100644 --- a/tests/config_disable_tcp.rs +++ b/tests/config_disable_tcp.rs @@ -90,6 +90,17 @@ async fn config_disable_tcp() { TcpStream::connect("127.0.0.1:12345").await.is_err(), "shouldn't listen on TCP port" ); + assert!( + session_one + .tcpip_forward("test.foobar.tld", 80) + .await + .is_ok(), + "shouldn't have failed to bind HTTP" + ); + assert!( + session_one.tcpip_forward("some.alias", 12345).await.is_ok(), + "shouldn't have failed to bind alias" + ); } struct SshClient; -- 2.51.2