From dedaeb9c709226b147b820d5a5294f81d645bc6c Mon Sep 17 00:00:00 2001 From: Eric Rodrigues Pires Date: Sat, 22 Mar 2025 14:25:42 -0300 Subject: [PATCH] Fix reading messages from all data channels and add more tests --- CHANGELOG.md | 1 + src/error.rs | 2 - src/http.rs | 432 +++++++++++++----- src/lib.rs | 19 +- src/main.rs | 8 +- src/ssh.rs | 69 +-- tests/http_addressing_profanities.rs | 218 +++++++++ tests/https_http11_fallback.rs | 1 - tests/lib_already_bound_ports.rs | 63 +++ ...ssh_ignore_data_on_non_session_channels.rs | 147 ++++++ tests/ssh_single_data_channel.rs | 101 ++++ tests/tcp_allow_requested_ports.rs | 4 +- 12 files changed, 903 insertions(+), 162 deletions(-) create mode 100644 tests/http_addressing_profanities.rs create mode 100644 tests/lib_already_bound_ports.rs create mode 100644 tests/ssh_ignore_data_on_non_session_channels.rs create mode 100644 tests/ssh_single_data_channel.rs diff --git a/CHANGELOG.md b/CHANGELOG.md index 921a86d..af2e8d5 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,6 +4,7 @@ ### Fixed +- Fix reading messages from all data channels. - Improve logic for SSH exec commands. - Flush messages to data channel when closing connection. - Show cursor in admin interface only on shutdown. diff --git a/src/error.rs b/src/error.rs index 792cdba..bc3fe17 100644 --- a/src/error.rs +++ b/src/error.rs @@ -17,8 +17,6 @@ pub(crate) enum ServerError { InvalidHttpVersion(Version), #[error("Missing Upgrade header")] MissingUpgradeHeader, - #[error("Request timed out")] - RequestTimeout, #[error("Invalid file path")] InvalidFilePath, #[error("Already bound by another service")] diff --git a/src/http.rs b/src/http.rs index a3ef957..231cf4b 100644 --- a/src/http.rs +++ b/src/http.rs @@ -74,8 +74,8 @@ fn http_log(data: HttpLog, tx: Option>>, disable_h _ => "44", }; let line = format!( - " \x1b[2m{:19}\x1b[22m \x1b[{}m[{:3}] \x1b[0;1;30;{}m{:^7}\x1b[0m {} => {} \x1b[2m({}) {}\x1b[0m\r\n", - chrono::Local::now().format("%Y-%m-%dT%H:%M:%S"), + " \x1b[2m{}\x1b[22m \x1b[{}m[{}] \x1b[0;1;30;{}m {} \x1b[0m {} => {} \x1b[2m({}) {}\x1b[0m\r\n", + chrono::Local::now().format("%Y-%m-%dT%H:%M:%S%:z"), status_escape_color, status, method_escape_color, @@ -303,10 +303,14 @@ where })); let response = match proxy_data.http_request_timeout { // Await for a response under the given duration. - Some(duration) => timeout(duration, sender.send_request(request)) - .await - .map_err(|_| ServerError::RequestTimeout)??, - None => sender.send_request(request).await?, + Some(duration) => { + if let Ok(response) = timeout(duration, sender.send_request(request)).await { + response?.into_response() + } else { + (StatusCode::REQUEST_TIMEOUT, "").into_response() + } + } + None => sender.send_request(request).await?.into_response(), }; let elapsed_time = timer.elapsed(); http_log( @@ -321,7 +325,7 @@ where tx, disable_http_logs, ); - Ok(response.into_response()) + Ok(response) } Version::HTTP_11 | Version::HTTP_2 => { // Ensure best-effort compatibility of proxy request with HTTP/1.1 format @@ -358,115 +362,121 @@ where hyper::client::conn::http1::handshake(TokioIo::new(io)).await?; // Check for an Upgrade header - match request.headers().get(UPGRADE) { - // If Upgrade header is not present, handle the request as usual - None => { - tokio::spawn(async move { - if let Err(err) = conn.await { - warn!("HTTP/1.1 connection failed: {:?}", err); - } - }); - let response = match proxy_data.http_request_timeout { - // Await for a response under the given duration. - Some(duration) => timeout(duration, sender.send_request(request)) - .await - .map_err(|_| ServerError::RequestTimeout)??, - None => sender.send_request(request).await?, - }; - let elapsed_time = timer.elapsed(); - http_log( - HttpLog { - ip: &ip, - status: response.status().as_u16(), - method: &method, - host: &host, - uri: &uri, - elapsed_time, - }, - tx, - disable_http_logs, - ); - // Return the received response to the client - Ok(response.into_response()) - } - + if let Some(request_upgrade) = request.headers().get(UPGRADE) { // If there is an Upgrade header, make sure that it's a valid Websocket upgrade. - Some(request_upgrade) => { - tokio::spawn(async move { - if let Err(err) = conn.with_upgrades().await { - warn!("HTTP/1.1 connection with upgrades failed: {:?}", err); + tokio::spawn(async move { + if let Err(err) = conn.with_upgrades().await { + warn!("HTTP/1.1 connection with upgrades failed: {:?}", err); + } + }); + let request_type = request_upgrade.to_str()?.to_string(); + // Retrieve the OnUpgrade from the incoming request + let upgraded_request = hyper::upgrade::on(&mut request); + let mut response = match proxy_data.http_request_timeout { + // Await for a response under the given duration. + Some(duration) => { + if let Ok(response) = timeout(duration, sender.send_request(request)).await + { + response?.into_response() + } else { + (StatusCode::REQUEST_TIMEOUT, "").into_response() } - }); - let request_type = request_upgrade.to_str()?.to_string(); - // Retrieve the OnUpgrade from the incoming request - let upgraded_request = hyper::upgrade::on(&mut request); - let mut response = match proxy_data.http_request_timeout { - // Await for a response under the given duration. - Some(duration) => timeout(duration, sender.send_request(request)) - .await - .map_err(|_| ServerError::RequestTimeout)??, - None => sender.send_request(request).await?, - }; - let elapsed_time = timer.elapsed(); - http_log( - HttpLog { - ip: &ip, - status: response.status().as_u16(), - method: &method, - host: &host, - uri: &uri, - elapsed_time, - }, - tx, - disable_http_logs, - ); - // Check if the underlying server accepts the Upgrade request - match response.status() { - StatusCode::SWITCHING_PROTOCOLS => { - if request_type - == response - .headers() - .get(UPGRADE) - .ok_or(ServerError::MissingUpgradeHeader)? - .to_str()? - { - // Retrieve the upgraded connection from the response - let upgraded_response = hyper::upgrade::on(&mut response).await?; - let websocket_timeout = proxy_data.websocket_timeout; - // Start a task to copy data between the two Upgraded parts - tokio::spawn(async move { - let mut upgraded_request = - TokioIo::new(upgraded_request.await.unwrap()); - let mut upgraded_response = TokioIo::new(upgraded_response); - match websocket_timeout { - // If there is a Websocket timeout, copy until the deadline is reached. - Some(duration) => { - let _ = timeout(duration, async { - copy_bidirectional( - &mut upgraded_response, - &mut upgraded_request, - ) - .await - }) - .await; - } - // If there isn't a Websocket timeout, copy data between both sides unconditionally. - None => { - let _ = copy_bidirectional( + } + None => sender.send_request(request).await?.into_response(), + }; + let elapsed_time = timer.elapsed(); + http_log( + HttpLog { + ip: &ip, + status: response.status().as_u16(), + method: &method, + host: &host, + uri: &uri, + elapsed_time, + }, + tx, + disable_http_logs, + ); + // Check if the underlying server accepts the Upgrade request + match response.status() { + StatusCode::SWITCHING_PROTOCOLS => { + if request_type + == response + .headers() + .get(UPGRADE) + .ok_or(ServerError::MissingUpgradeHeader)? + .to_str()? + { + // Retrieve the upgraded connection from the response + let upgraded_response = hyper::upgrade::on(&mut response).await?; + let websocket_timeout = proxy_data.websocket_timeout; + // Start a task to copy data between the two Upgraded parts + tokio::spawn(async move { + let mut upgraded_request = + TokioIo::new(upgraded_request.await.unwrap()); + let mut upgraded_response = TokioIo::new(upgraded_response); + match websocket_timeout { + // If there is a Websocket timeout, copy until the deadline is reached. + Some(duration) => { + let _ = timeout(duration, async { + copy_bidirectional( &mut upgraded_response, &mut upgraded_request, ) - .await; - } + .await + }) + .await; + } + // If there isn't a Websocket timeout, copy data between both sides unconditionally. + None => { + let _ = copy_bidirectional( + &mut upgraded_response, + &mut upgraded_request, + ) + .await; } - }); - } - // Return the response to the client - Ok(response.into_response()) + } + }); } - _ => Ok(response.into_response()), + // Return the response to the client + Ok(response) } + _ => Ok(response), } + } else { + // If Upgrade header is not present, simply handle the request + tokio::spawn(async move { + if let Err(err) = conn.await { + warn!("HTTP/1.1 connection failed: {:?}", err); + } + }); + let response = match proxy_data.http_request_timeout { + // Await for a response under the given duration. + Some(duration) => { + if let Ok(response) = timeout(duration, sender.send_request(request)).await + { + response?.into_response() + } else { + (StatusCode::REQUEST_TIMEOUT, "").into_response() + } + } + None => sender.send_request(request).await?.into_response(), + }; + let elapsed_time = timer.elapsed(); + http_log( + HttpLog { + ip: &ip, + status: response.status().as_u16(), + method: &method, + host: &host, + uri: &uri, + elapsed_time, + }, + tx, + disable_http_logs, + ); + // Return the received response to the client + Ok(response) } } version => { @@ -486,6 +496,7 @@ mod proxy_handler_tests { }; use bytes::Bytes; use futures_util::{SinkExt, StreamExt}; + use http::Version; use http_body_util::{BodyExt, Empty}; use hyper::{HeaderMap, Request, StatusCode, body::Incoming, service::service_fn}; use hyper_util::rt::{TokioExecutor, TokioIo}; @@ -498,7 +509,6 @@ mod proxy_handler_tests { config::LoadBalancing, connection_handler::{ConnectionHttpData, MockConnectionHandler}, connections::ConnectionMap, - http::ServerError, quota::{DummyQuotaHandler, TokenHolder, UserIdentification}, reactor::MockConnectionMapReactor, telemetry::Telemetry, @@ -901,7 +911,7 @@ mod proxy_handler_tests { } #[tokio::test] - async fn returns_error_for_outgoing_request_timeout() { + async fn returns_error_for_outgoing_http11_request_timeout() { let conn_manager: Arc< ConnectionMap< String, @@ -938,6 +948,7 @@ mod proxy_handler_tests { ) .unwrap(); let request = Request::builder() + .version(Version::HTTP_11) .method("GET") .uri("/slow_endpoint") .header("host", "slow.handler") @@ -982,15 +993,204 @@ mod proxy_handler_tests { ) .await; assert!( - logging_rx.is_empty(), - "shouldn't log if failed to proxy request" + !logging_rx.is_empty(), + "should log after timing out request" + ); + let response = response.expect("should return response after proxy"); + assert_eq!(response.status(), hyper::StatusCode::REQUEST_TIMEOUT); + jh.abort(); + } + + #[tokio::test] + async fn returns_error_for_outgoing_websocket_request_timeout() { + let conn_manager: Arc< + ConnectionMap< + String, + Arc>, + MockConnectionMapReactor, + >, + > = Arc::new(ConnectionMap::new( + LoadBalancing::Allow, + Arc::new(Box::new(DummyQuotaHandler)), + None, + )); + let (server, handler) = tokio::io::duplex(1024); + let (logging_tx, logging_rx) = mpsc::unbounded_channel::>(); + let mut mock = MockConnectionHandler::new(); + mock.expect_log_channel() + .once() + .return_once(move || Some(logging_tx)); + mock.expect_tunneling_channel() + .once() + .return_once(move |_, _| Ok(handler)); + mock.expect_http_data().once().return_once(move || { + Some(ConnectionHttpData { + redirect_http_to_https_port: None, + is_aliasing: false, + http2: false, + }) + }); + conn_manager + .insert( + "with.websocket".into(), + "127.0.0.1:12345".parse().unwrap(), + TokenHolder::User(UserIdentification::Username("a".into())), + Arc::new(mock), + ) + .unwrap(); + let (socket, stream) = tokio::io::duplex(1024); + let router = Router::new().route( + "/ws", + any(|ws: WebSocketUpgrade| async move { + sleep(Duration::from_secs(1)).await; + ws.on_upgrade(|mut socket| async move { + let _ = socket.send(ws::Message::Text("Success.".into())).await; + let _ = socket.close().await; + }) + }), + ); + let router_service = service_fn(move |req: Request| router.clone().call(req)); + let jh = tokio::spawn(async move { + hyper_util::server::conn::auto::Builder::new(hyper_util::rt::TokioExecutor::new()) + .serve_connection_with_upgrades(TokioIo::new(server), router_service) + .await + .expect("Invalid request"); + }); + assert!(logging_rx.is_empty(), "shouldn't log before request"); + let proxy_service = service_fn(move |request| { + proxy_handler( + request, + "127.0.0.1:12345".parse().unwrap(), + None, + Arc::new(ProxyData { + conn_manager: Arc::clone(&conn_manager), + telemetry: Arc::new(Telemetry::new()), + domain_redirect: Arc::new(DomainRedirect { + from: "main.domain".into(), + to: "https://example.com".into(), + }), + protocol: Protocol::Https { port: 443 }, + proxy_type: ProxyType::Tunneling, + http_request_timeout: Some(Duration::from_millis(500)), + websocket_timeout: None, + disable_http_logs: false, + _phantom_data: PhantomData, + }), + ) + }); + let jh2 = tokio::spawn(async move { + hyper_util::server::conn::auto::Builder::new(hyper_util::rt::TokioExecutor::new()) + .serve_connection_with_upgrades(TokioIo::new(socket), proxy_service) + .await + .expect("Invalid request"); + }); + let Err(err) = client_async("ws://with.websocket/ws", stream).await else { + panic!("should've errored when establishing Websocket connection"); + }; + match err { + tokio_tungstenite::tungstenite::Error::Http(response) => { + assert!( + response.status() == StatusCode::REQUEST_TIMEOUT, + "should've timed out Websocket request" + ) + } + _ => panic!(), + } + assert!( + !logging_rx.is_empty(), + "should log after upgrade proxying request" ); - assert!(response.is_err(), "should fail if timed out"); - let error = response.unwrap_err(); - assert!(matches!( - error.downcast().unwrap(), - ServerError::RequestTimeout + jh.abort(); + jh2.abort(); + } + + #[tokio::test] + async fn returns_error_for_outgoing_http2_request_timeout() { + let conn_manager: Arc< + ConnectionMap< + String, + Arc>, + MockConnectionMapReactor, + >, + > = Arc::new(ConnectionMap::new( + LoadBalancing::Allow, + Arc::new(Box::new(DummyQuotaHandler)), + None, )); + let (server, handler) = tokio::io::duplex(1024); + let (logging_tx, logging_rx) = mpsc::unbounded_channel::>(); + let mut mock = MockConnectionHandler::new(); + mock.expect_log_channel() + .once() + .return_once(move || Some(logging_tx)); + mock.expect_tunneling_channel() + .once() + .return_once(move |_, _| Ok(handler)); + mock.expect_http_data().once().return_once(move || { + Some(ConnectionHttpData { + redirect_http_to_https_port: None, + is_aliasing: false, + http2: true, + }) + }); + conn_manager + .insert( + "slow.handler".into(), + "127.0.0.1:12345".parse().unwrap(), + TokenHolder::User(UserIdentification::Username("a".into())), + Arc::new(mock), + ) + .unwrap(); + let request = Request::builder() + .version(Version::HTTP_2) + .method("GET") + .uri("https://slow.handler/slow_endpoint") + .body(Empty::::new()) + .unwrap(); + let router = Router::new().route( + "/slow_endpoint", + get(async || { + sleep(Duration::from_secs(1)).await; + "Slow hello." + }), + ); + let router_service = service_fn(move |req: Request| router.clone().call(req)); + let jh = tokio::spawn(async move { + hyper_util::server::conn::auto::Builder::new(hyper_util::rt::TokioExecutor::new()) + .serve_connection(TokioIo::new(server), router_service) + .await + .expect("Invalid request"); + }); + assert!( + logging_rx.is_empty(), + "shouldn't log before handling request" + ); + let response = proxy_handler( + request, + "127.0.0.1:12345".parse().unwrap(), + None, + Arc::new(ProxyData { + conn_manager: Arc::clone(&conn_manager), + telemetry: Arc::new(Telemetry::new()), + domain_redirect: Arc::new(DomainRedirect { + from: "main.domain".into(), + to: "https://example.com".into(), + }), + protocol: Protocol::Https { port: 443 }, + proxy_type: ProxyType::Tunneling, + http_request_timeout: Some(Duration::from_millis(500)), + websocket_timeout: None, + disable_http_logs: false, + _phantom_data: PhantomData, + }), + ) + .await; + assert!( + !logging_rx.is_empty(), + "should log after timing out request" + ); + let response = response.expect("should return response after proxy"); + assert_eq!(response.status(), hyper::StatusCode::REQUEST_TIMEOUT); jh.abort(); } diff --git a/src/lib.rs b/src/lib.rs index 7397455..258530c 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -61,6 +61,7 @@ use crate::{ addressing::{AddressDelegator, DnsResolver}, certificates::{AlpnChallengeResolver, CertificateResolver, DummyAlpnChallengeResolver}, connections::ConnectionMap, + droppable_handle::DroppableHandle, error::ServerError, fingerprints::FingerprintsValidator, http::{Protocol, proxy_handler}, @@ -536,7 +537,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { // HTTP handler let mut join_handle_http = if config.disable_http { - tokio::spawn(future::pending()) + DroppableHandle(tokio::spawn(future::pending())) } else { let http_listener = TcpListener::bind((config.listen_address, config.http_port.into())) .await @@ -568,7 +569,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { disable_http_logs: config.disable_http_logs, _phantom_data: PhantomData, }); - tokio::spawn(async move { + DroppableHandle(tokio::spawn(async move { loop { let proxy_data = Arc::clone(&http_proxy_data); let (stream, address) = match http_listener.accept().await { @@ -604,12 +605,12 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { } }); } - }) + })) }; // HTTPS handler (with optional SSH handling) let mut join_handle_https = if config.disable_http { - tokio::spawn(future::pending()) + DroppableHandle(tokio::spawn(future::pending())) } else { let https_listener = TcpListener::bind((config.listen_address, config.https_port.into())) .await @@ -648,7 +649,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { }); let sandhole_clone = Arc::clone(&sandhole); let ssh_config_clone = Arc::clone(&ssh_config); - tokio::spawn(async move { + DroppableHandle(tokio::spawn(async move { loop { let (stream, address) = match https_listener.accept().await { Ok((stream, address)) => (stream, address), @@ -752,7 +753,7 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { } }); } - }) + })) }; // Start Sandhole on SSH port @@ -787,17 +788,15 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { _ = &mut signal_handler => { break; } - _ = &mut join_handle_http => { + _ = &mut join_handle_http.0 => { break; } - _ = &mut join_handle_https => { + _ = &mut join_handle_https.0 => { break; } } } info!("Sandhole is shutting down."); - join_handle_http.abort(); - join_handle_https.abort(); Ok(()) } diff --git a/src/main.rs b/src/main.rs index 4b39a53..33b5690 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,9 +1,15 @@ use clap::Parser; +use log::error; use sandhole::{ApplicationConfig, entrypoint}; #[tokio::main] async fn main() -> anyhow::Result<()> { env_logger::Builder::from_env(env_logger::Env::default().default_filter_or("info")).init(); let config = ApplicationConfig::parse(); - entrypoint(config).await + if let Err(err) = entrypoint(config).await { + error!("Unable to start Sandhole: {}", err); + Err(err) + } else { + Ok(()) + } } diff --git a/src/ssh.rs b/src/ssh.rs index 60f547f..4ee7d52 100644 --- a/src/ssh.rs +++ b/src/ssh.rs @@ -305,6 +305,8 @@ pub(crate) struct ServerHandler { cancellation_token: CancellationToken, // User-specific data, set after authentication. auth_data: AuthenticatedData, + // ID for the open session channel. + channel_id: Option, // Sender for data session messages, used for sending logs and TUI state to the client. tx: OptionalSender, // Handle for the opened data session task. Initially None. @@ -343,6 +345,7 @@ impl Server for Arc { proxy_count: Arc::new(AtomicUsize::new(0)), }, commands: Default::default(), + channel_id: None, tx: OptionalSender(None), open_session_join_handle: None, server: Arc::clone(self), @@ -366,6 +369,7 @@ impl Handler for ServerHandler { } return Ok(false); }; + self.channel_id = Some(channel.id()); let (tx, mut rx) = mpsc::unbounded_channel::>(); let graceful_cancellation_token = CancellationToken::new(); let graceful_shutdown_rx = graceful_cancellation_token.clone(); @@ -539,42 +543,47 @@ impl Handler for ServerHandler { // Handle data received from the client such as key presses. async fn data( &mut self, - _channel: ChannelId, + channel: ChannelId, data: &[u8], _session: &mut Session, ) -> Result<(), Self::Error> { - debug!("received data {:?}", data); - match &mut self.auth_data { - // Ignore other commands for non-admin users - AuthenticatedData::None { .. } | AuthenticatedData::User { .. } => (), - AuthenticatedData::Admin { admin_data, .. } => { - // Handle the proper key press in the admin TUI. - if let Some(admin_interface) = admin_data.admin_interface.as_mut() { - match data { - // Tab - b"\t" => admin_interface.next_tab(), - // Shift+Tab - b"\x1b[Z" => admin_interface.previous_tab(), - // Up - b"\x1b[A" | b"k" => admin_interface.move_up(), - // Down - b"\x1b[B" | b"j" => admin_interface.move_down(), - // Esc - b"\x1b" => admin_interface.cancel(), - // Enter - b"\r" => admin_interface.enter(), - // Delete - b"\x1b[3~" => admin_interface.delete(), - // Ctrl+C - b"\x03" => admin_interface.disable(), - _ => (), + if self + .channel_id + .is_some_and(|channel_id| channel_id == channel) + { + debug!("received data {:?}", data); + match &mut self.auth_data { + // Ignore other commands for non-admin users + AuthenticatedData::None { .. } | AuthenticatedData::User { .. } => (), + AuthenticatedData::Admin { admin_data, .. } => { + // Handle the proper key press in the admin TUI. + if let Some(admin_interface) = admin_data.admin_interface.as_mut() { + match data { + // Tab + b"\t" => admin_interface.next_tab(), + // Shift+Tab + b"\x1b[Z" => admin_interface.previous_tab(), + // Up + b"\x1b[A" | b"k" => admin_interface.move_up(), + // Down + b"\x1b[B" | b"j" => admin_interface.move_down(), + // Esc + b"\x1b" => admin_interface.cancel(), + // Enter + b"\r" => admin_interface.enter(), + // Delete + b"\x1b[3~" => admin_interface.delete(), + // Ctrl+C + b"\x03" => admin_interface.disable(), + _ => (), + } } } } - } - // Ctrl+C (0x03) ends the session and disconnects the client - if data == b"\x03" { - self.cancellation_token.cancel(); + // Ctrl+C (0x03) ends the session and disconnects the client + if data == b"\x03" { + self.cancellation_token.cancel(); + } } Ok(()) } diff --git a/tests/http_addressing_profanities.rs b/tests/http_addressing_profanities.rs new file mode 100644 index 0000000..2129317 --- /dev/null +++ b/tests/http_addressing_profanities.rs @@ -0,0 +1,218 @@ +use std::{sync::Arc, time::Duration}; + +use axum::{Router, extract::Request, routing::get}; +use clap::Parser; +use http_body_util::BodyExt; +use hyper::{StatusCode, body::Incoming, service::service_fn}; +use hyper_util::{ + rt::{TokioExecutor, TokioIo}, + server::conn::auto::Builder, +}; +use russh::keys::{key::PrivateKeyWithHashAlg, load_secret_key}; +use russh::{ + Channel, + client::{Msg, Session}, +}; +use sandhole::{ApplicationConfig, entrypoint}; +use tokio::{ + net::TcpStream, + time::{sleep, timeout}, +}; +use tower::Service; + +#[tokio::test(flavor = "multi_thread")] +async fn http_addressing_profanities() { + // 1. Initialize Sandhole + let _ = env_logger::builder() + .filter_module("sandhole", log::LevelFilter::Debug) + .is_test(true) + .try_init(); + let config = ApplicationConfig::parse_from([ + "sandhole", + "--domain=fuck.tld", + "--user-keys-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/user_keys"), + "--admin-keys-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/admin_keys"), + "--certificates-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/certificates"), + "--private-key-file", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/server_keys/ssh"), + "--acme-cache-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/acme_cache"), + "--disable-directory-creation", + "--listen-address=127.0.0.1", + "--ssh-port=18022", + "--http-port=18080", + "--https-port=18443", + "--acme-use-staging", + "--bind-hostnames=all", + "--requested-domain-filter-profanities", + "--idle-connection-timeout=1s", + "--authentication-request-timeout=5s", + "--http-request-timeout=5s", + ]); + 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 with a profanity and a non-profanity name + let key = load_secret_key( + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/private_keys/key1"), + None, + ) + .expect("Missing file key1"); + let ssh_client = SshClient; + let mut session = russh::client::connect(Default::default(), "127.0.0.1:18022", ssh_client) + .await + .expect("Failed to connect to SSH server"); + assert!( + session + .authenticate_publickey( + "user", + PrivateKeyWithHashAlg::new( + Arc::new(key), + session.best_supported_rsa_hash().await.unwrap().flatten() + ) + ) + .await + .expect("SSH authentication failed") + .success(), + "authentication didn't succeed" + ); + let mut channel = session + .channel_open_session() + .await + .expect("channel_open_session failed"); + session + .tcpip_forward("shit.fuck.tld", 80) + .await + .expect("tcpip_forward failed"); + let regex = regex::Regex::new(r"http://(\S+)").expect("Invalid regex"); + let Ok((_, hostname)) = 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 address = captures.get(0).unwrap().as_str().to_string(); + let hostname = captures + .get(1) + .expect("Missing hostname matching group") + .as_str() + .split(':') + .next() + .unwrap() + .to_string(); + return (address, hostname); + } + } + message => panic!("Unexpected message {:?}", message), + } + } + panic!("Unexpected end of channel"); + }) + .await + else { + panic!("Timed out waiting for subdomain allocation."); + }; + assert!( + regex::Regex::new(r"^[a-z0-9]+\.fuck\.tld$") + .unwrap() + .is_match(&hostname), + "hostname should've matched regex" + ); + assert!( + !hostname.starts_with("shit."), + "hostname shouldn't start with profanity" + ); + session + .tcpip_forward("valid-as.fuck.tld", 80) + .await + .expect("tcpip_forward failed"); + + // 3. Connect to our HTTP proxies + for host in [&hostname, "valid-as.fuck.tld"] { + let tcp_stream = TcpStream::connect("127.0.0.1:18080") + .await + .expect("TCP connection failed"); + let (mut sender, conn) = hyper::client::conn::http1::handshake(TokioIo::new(tcp_stream)) + .await + .expect("HTTP handshake failed"); + tokio::spawn(async move { + if let Err(err) = conn.await { + eprintln!("Connection failed: {:?}", err); + } + }); + let request = Request::builder() + .method("GET") + .uri("/") + .header("host", host) + .body(http_body_util::Empty::::new()) + .unwrap(); + let Ok(response) = timeout(Duration::from_secs(5), async move { + sender + .send_request(request) + .await + .expect("Error sending HTTP request") + }) + .await + else { + panic!("Timeout waiting for request to finish."); + }; + assert_eq!(response.status(), StatusCode::OK); + let response_body = String::from_utf8( + response + .into_body() + .collect() + .await + .expect("Error collecting response") + .to_bytes() + .into(), + ) + .expect("Invalid response body"); + assert_eq!(response_body, "Hello from a profane place!"); + } +} + +struct SshClient; + +impl russh::client::Handler for SshClient { + type Error = anyhow::Error; + + async fn check_server_key( + &mut self, + _key: &russh::keys::PublicKey, + ) -> Result { + Ok(true) + } + + async fn server_channel_open_forwarded_tcpip( + &mut self, + channel: Channel, + _connected_address: &str, + _connected_port: u32, + _originator_address: &str, + _originator_port: u32, + _session: &mut Session, + ) -> Result<(), Self::Error> { + let router = Router::new().route("/", get(async || "Hello from a profane place!")); + let service = service_fn(move |req: Request| router.clone().call(req)); + tokio::spawn(async move { + Builder::new(TokioExecutor::new()) + .serve_connection_with_upgrades(TokioIo::new(channel.into_stream()), service) + .await + .expect("Invalid request"); + }); + Ok(()) + } +} diff --git a/tests/https_http11_fallback.rs b/tests/https_http11_fallback.rs index 3eab87f..3af10ad 100644 --- a/tests/https_http11_fallback.rs +++ b/tests/https_http11_fallback.rs @@ -187,7 +187,6 @@ impl russh::client::Handler for SshClient { && request.headers().get(COOKIE).unwrap() == "foo=1; bar=2" && request.uri() == &"/hello".parse::().unwrap() { - dbg!(request.headers()); "Hello from http11.foobar.tld!" } else { "Error" diff --git a/tests/lib_already_bound_ports.rs b/tests/lib_already_bound_ports.rs new file mode 100644 index 0000000..9545450 --- /dev/null +++ b/tests/lib_already_bound_ports.rs @@ -0,0 +1,63 @@ +use std::time::Duration; + +use clap::Parser; +use sandhole::{ApplicationConfig, entrypoint}; +use tokio::{ + net::TcpListener, + time::{sleep, timeout}, +}; + +#[tokio::test(flavor = "multi_thread")] +async fn lib_already_bound_ports() { + let _ = env_logger::builder() + .filter_module("sandhole", log::LevelFilter::Debug) + .is_test(true) + .try_init(); + for (i, port) in [18022, 18080, 18443].into_iter().enumerate() { + if i > 0 { + sleep(Duration::from_millis(200)).await; + } + // 1. Bind the specific port before Sandhole does + let listener = TcpListener::bind(("127.0.0.1", port)) + .await + .expect("should be able to bind open port"); + // 2. Fail to initialize Sandhole + if timeout(Duration::from_secs(2), async { + let config = ApplicationConfig::parse_from([ + "sandhole", + "--domain=foobar.tld", + "--user-keys-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/user_keys"), + "--admin-keys-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/admin_keys"), + "--certificates-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/certificates"), + "--private-key-file", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/server_keys/ssh"), + "--acme-cache-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/acme_cache"), + "--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", + ]); + assert!( + entrypoint(config).await.is_err(), + "should've failed to start server" + ) + }) + .await + .is_err() + { + panic!("Timeout waiting for Sandhole to start.") + }; + // 3. Unbind the port + drop(listener); + } +} diff --git a/tests/ssh_ignore_data_on_non_session_channels.rs b/tests/ssh_ignore_data_on_non_session_channels.rs new file mode 100644 index 0000000..7e58a56 --- /dev/null +++ b/tests/ssh_ignore_data_on_non_session_channels.rs @@ -0,0 +1,147 @@ +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::io::AsyncWriteExt; +use tokio::{ + io::AsyncReadExt, + net::TcpStream, + time::{sleep, timeout}, +}; + +#[tokio::test(flavor = "multi_thread")] +async fn ssh_ignore_data_on_non_session_channels() { + // 1. Initialize Sandhole + let _ = env_logger::builder() + .filter_module("sandhole", log::LevelFilter::Debug) + .is_test(true) + .try_init(); + let config = ApplicationConfig::parse_from([ + "sandhole", + "--domain=foobar.tld", + "--user-keys-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/user_keys"), + "--admin-keys-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/admin_keys"), + "--certificates-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/certificates"), + "--private-key-file", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/server_keys/ssh"), + "--acme-cache-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/acme_cache"), + "--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", + ]); + 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( + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/private_keys/key1"), + None, + ) + .expect("Missing file key1"); + let ssh_client = SshClient; + let mut session = russh::client::connect(Default::default(), "127.0.0.1:18022", ssh_client) + .await + .expect("Failed to connect to SSH server"); + assert!( + session + .authenticate_publickey( + "user", + PrivateKeyWithHashAlg::new( + Arc::new(key), + session.best_supported_rsa_hash().await.unwrap().flatten() + ) + ) + .await + .expect("SSH authentication failed") + .success(), + "authentication didn't succeed" + ); + session + .tcpip_forward("foobar.tld", 12345) + .await + .expect("tcpip_forward failed"); + let _channel = session + .channel_open_session() + .await + .expect("channel_open_session_failed"); + + // 3. Send data through the TCP port of our proxy + let mut tcp_stream = TcpStream::connect("127.0.0.1:12345") + .await + .expect("TCP connection failed"); + tcp_stream + .write_all(&[3]) + .await + .expect("should be able to send data"); + let mut buf = [0u8; 1]; + tcp_stream.read_exact(&mut buf).await.unwrap(); + assert_eq!(&buf, &[9]); + sleep(Duration::from_millis(100)).await; + assert!(!session.is_closed(), "session shouldn't have been closed"); +} + +struct SshClient; + +impl russh::client::Handler for SshClient { + type Error = anyhow::Error; + + async fn check_server_key( + &mut self, + _key: &russh::keys::PublicKey, + ) -> Result { + Ok(true) + } + + async fn server_channel_open_forwarded_tcpip( + &mut self, + channel: Channel, + _connected_address: &str, + _connected_port: u32, + _originator_address: &str, + _originator_port: u32, + _session: &mut Session, + ) -> Result<(), Self::Error> { + tokio::spawn(async move { + let mut channel = channel; + match channel.wait().await.unwrap() { + russh::ChannelMsg::Data { data } => { + channel + .data(data.iter().map(|i| i * i).collect::>().as_ref()) + .await + .unwrap(); + } + _ => { + channel.exit_status(1).await.unwrap(); + } + } + channel.eof().await.unwrap(); + }); + Ok(()) + } +} diff --git a/tests/ssh_single_data_channel.rs b/tests/ssh_single_data_channel.rs new file mode 100644 index 0000000..f172ff7 --- /dev/null +++ b/tests/ssh_single_data_channel.rs @@ -0,0 +1,101 @@ +use std::{sync::Arc, time::Duration}; + +use clap::Parser; +use russh::client; +use russh::keys::{key::PrivateKeyWithHashAlg, load_secret_key}; +use sandhole::{ApplicationConfig, entrypoint}; +use tokio::{ + net::TcpStream, + time::{sleep, timeout}, +}; + +#[tokio::test(flavor = "multi_thread")] +async fn ssh_single_data_channel() { + // 1. Initialize Sandhole + let _ = env_logger::builder() + .filter_module("sandhole", log::LevelFilter::Debug) + .is_test(true) + .try_init(); + let config = ApplicationConfig::parse_from([ + "sandhole", + "--domain=foobar.tld", + "--user-keys-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/user_keys"), + "--admin-keys-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/admin_keys"), + "--certificates-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/certificates"), + "--private-key-file", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/server_keys/ssh"), + "--acme-cache-directory", + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/acme_cache"), + "--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=2s", + "--authentication-request-timeout=5s", + "--http-request-timeout=5s", + ]); + 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 fail to alias to localhost + let key = load_secret_key( + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/private_keys/key1"), + None, + ) + .expect("Missing file key1"); + let ssh_client = SshClient; + let mut session = client::connect(Default::default(), "127.0.0.1:18022", ssh_client) + .await + .expect("Failed to connect to SSH server"); + assert!( + session + .authenticate_publickey( + "user", + PrivateKeyWithHashAlg::new( + Arc::new(key), + session.best_supported_rsa_hash().await.unwrap().flatten() + ) + ) + .await + .expect("SSH authentication failed") + .success(), + "authentication didn't succeed" + ); + let _channel = session + .channel_open_session() + .await + .expect("channel_open_session failed"); + assert!( + session.channel_open_session().await.is_err(), + "shouldn't open more than one open session channel" + ); + assert!(!session.is_closed(), "session shouldn't have been closed"); +} + +struct SshClient; + +impl client::Handler for SshClient { + type Error = anyhow::Error; + + async fn check_server_key( + &mut self, + _key: &russh::keys::PublicKey, + ) -> Result { + Ok(true) + } +} diff --git a/tests/tcp_allow_requested_ports.rs b/tests/tcp_allow_requested_ports.rs index 6c94850..10f298e 100644 --- a/tests/tcp_allow_requested_ports.rs +++ b/tests/tcp_allow_requested_ports.rs @@ -145,7 +145,7 @@ async fn tcp_allow_requested_ports() { panic!("Timeout waiting for proxy server to reply.") }; - // 4. Local-forward the TCP port for known user + // 5. Local-forward the TCP port for known user let key = load_secret_key( concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/private_keys/key2"), None, @@ -192,7 +192,7 @@ async fn tcp_allow_requested_ports() { panic!("Timeout waiting for proxy server to reply.") }; - // 5. Attempt to close TCP forwarding + // 6. Attempt to close TCP forwarding session_one .cancel_tcpip_forward("foobar.tld", 12345) .await -- 2.51.2