diff --git a/justfile b/justfile index 3c010ab..a1dbb6c 100644 --- a/justfile +++ b/justfile @@ -14,4 +14,4 @@ cli: to-html --no-prompt "cargo run --quiet -- --help" > cli.html clippy: - cargo clippy --all-targets --fix --allow-dirty && cargo fmt --all + cargo clippy --all-targets --fix --allow-dirty --allow-staged && cargo fmt --all diff --git a/src/connection_handler.rs b/src/connection_handler.rs index 2e526a4..22b18d9 100644 --- a/src/connection_handler.rs +++ b/src/connection_handler.rs @@ -11,6 +11,7 @@ pub(crate) struct ConnectionHttpData { // Port to redirect HTTP requests to. If missing, do not redirect. pub(crate) redirect_http_to_https_port: Option, pub(crate) is_aliasing: bool, + pub(crate) http2: bool, } // Trait for creating tunneling or aliasing channels (via an underlying SSH session). diff --git a/src/error.rs b/src/error.rs index d59f00f..792cdba 100644 --- a/src/error.rs +++ b/src/error.rs @@ -1,15 +1,20 @@ use std::path::PathBuf; +use http::Version; use ipnet::IpNet; #[derive(thiserror::Error, Debug)] pub(crate) enum ServerError { #[error("Invalid config: {0}")] InvalidConfig(String), + #[error("Missing URI host")] + MissingUriHost, #[error("Missing Host header")] MissingHostHeader, #[error("Invalid Host header")] InvalidHostHeader, + #[error("Invalid HTTP version {0:?}")] + InvalidHttpVersion(Version), #[error("Missing Upgrade header")] MissingUpgradeHeader, #[error("Request timed out")] diff --git a/src/http.rs b/src/http.rs index 2e6a102..a3ef957 100644 --- a/src/http.rs +++ b/src/http.rs @@ -1,5 +1,6 @@ use std::error::Error; use std::marker::PhantomData; +use std::str::FromStr; use std::time::{Duration, Instant}; use std::{net::SocketAddr, sync::Arc}; @@ -13,13 +14,15 @@ use axum::{ body::Body as AxumBody, response::{IntoResponse, Redirect}, }; +use http::header::COOKIE; +use http::{Uri, Version}; use hyper::{ Request, Response, StatusCode, body::Body, header::{HOST, UPGRADE}, }; -use hyper_util::rt::TokioIo; -use log::warn; +use hyper_util::rt::{TokioExecutor, TokioIo}; +use log::{debug, warn}; use russh::keys::ssh_key::Fingerprint; use tokio::{ io::{AsyncRead, AsyncWrite, copy_bidirectional}, @@ -141,7 +144,7 @@ where M: ConnectionGetByHttpHost>, H: ConnectionHandler, T: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static, - B: Body + Send + 'static, + B: Body + Send + Unpin + 'static, ::Data: Send + Sync + 'static, ::Error: Error + Send + Sync + 'static, { @@ -152,15 +155,27 @@ where let disable_http_logs = proxy_data.disable_http_logs; let timer = Instant::now(); // Retrieve host from the headers - let host = request - .headers() - .get(HOST) - .ok_or(ServerError::MissingHostHeader)? - .to_str()? - .split(':') - .next() - .ok_or(ServerError::InvalidHostHeader)? - .to_owned(); + let host = match request.version() { + Version::HTTP_2 => request.uri().host().ok_or(ServerError::MissingUriHost), + Version::HTTP_11 => match request.headers().get(HOST) { + Some(header_value) => match header_value.to_str() { + Ok(header) => header + .split(':') + .next() + .ok_or(ServerError::InvalidHostHeader), + Err(_) => Err(ServerError::InvalidHostHeader), + }, + None => Err(ServerError::MissingHostHeader), + }, + version => Err(ServerError::InvalidHttpVersion(version)), + }; + let host = match host { + Ok(host) => host.to_owned(), + Err(err) => { + debug!("Failed to parse host: {:?}", err); + return Ok((StatusCode::BAD_REQUEST, "").into_response()); + } + }; let ip = tcp_address.ip().to_canonical().to_string(); // Find the HTTP handler for the given host let Some(handler) = conn_manager.get_by_http_host(&host) else { @@ -187,9 +202,12 @@ where return Ok((StatusCode::NOT_FOUND, "").into_response()); }; let http_data = handler.http_data().await; - let redirect_http_to_https_port = http_data - .as_ref() - .and_then(|data| data.redirect_http_to_https_port); + let redirect_http_to_https_port = + http_data + .as_ref() + .and_then(|data: &crate::connection_handler::ConnectionHttpData| { + data.redirect_http_to_https_port + }); // Read protocol information for X-Forwarded headers let (proto, port) = match ( protocol, @@ -269,20 +287,20 @@ where return Ok((StatusCode::NOT_FOUND, "").into_response()); }; let tx = handler.log_channel(); - // Create an HTTP handshake over the selected channel - let (mut sender, conn) = hyper::client::conn::http1::handshake(TokioIo::new(io)).await?; - let method = request.method().to_string(); let uri = request.uri().path().to_string(); - // Check for an Upgrade header - match request.headers().get(UPGRADE) { - // If not present, handle the request as usual - None => { - tokio::spawn(async move { + let is_http2 = http_data.as_ref().map(|data| data.http2).unwrap_or(false); + match request.version() { + Version::HTTP_2 if is_http2 => { + // Create an HTTP/2 handshake over the selected channel + let (mut sender, conn) = + hyper::client::conn::http2::handshake(TokioExecutor::new(), TokioIo::new(io)) + .await?; + tokio::spawn(Box::pin(async move { if let Err(err) = conn.await { - warn!("Connection failed: {:?}", err); + warn!("HTTP/2 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)) @@ -303,87 +321,158 @@ where tx, disable_http_logs, ); - // Return the received response to the client Ok(response.into_response()) } + Version::HTTP_11 | Version::HTTP_2 => { + // Ensure best-effort compatibility of proxy request with HTTP/1.1 format + // -> Add host header if missing + request + .headers_mut() + .entry(HOST) + .or_insert_with(|| host.clone().try_into().unwrap()); + // -> Change URI to only include path and query + *request.uri_mut() = request + .uri() + .path_and_query() + .map(|path| Uri::from_str(path.as_str()).unwrap()) + .unwrap_or_default(); + // -> Decompress cookies: https://www.rfc-editor.org/rfc/rfc7540#section-8.1.2.5 + if let http::header::Entry::Occupied(occupied_entry) = + request.headers_mut().entry(COOKIE) + { + let (header, values) = occupied_entry.remove_entry_mult(); + let mut value = vec![]; + for header_value in values { + if !value.is_empty() { + value.extend_from_slice(b"; "); + } + value.extend_from_slice(header_value.as_bytes()); + } + request + .headers_mut() + .insert(header, value.try_into().unwrap()); + } - // 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!("Connection failed: {:?}", err); + // Create an HTTP/1.1 handshake over the selected channel + let (mut sender, conn) = + 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()) } - }); - 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( - &mut upgraded_response, - &mut upgraded_request, - ) - .await; - } + + // 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); + } + }); + 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( + &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.into_response()) } - _ => Ok(response.into_response()), } } + version => { + warn!("Unsupported HTTP version {:?}", version); + Ok((StatusCode::BAD_REQUEST, "").into_response()) + } } } @@ -392,13 +481,14 @@ where mod proxy_handler_tests { use axum::{ Router, - routing::{any, get, post}, + extract::{WebSocketUpgrade, ws}, + routing::{any, get, post, put}, }; use bytes::Bytes; use futures_util::{SinkExt, StreamExt}; - use http_body_util::Empty; + use http_body_util::{BodyExt, Empty}; use hyper::{HeaderMap, Request, StatusCode, body::Incoming, service::service_fn}; - use hyper_util::rt::TokioIo; + use hyper_util::rt::{TokioExecutor, TokioIo}; use std::{marker::PhantomData, sync::Arc, time::Duration}; use tokio::{io::DuplexStream, sync::mpsc, time::sleep}; use tokio_tungstenite::client_async; @@ -417,7 +507,7 @@ mod proxy_handler_tests { use super::{DomainRedirect, Protocol, ProxyData, ProxyType, proxy_handler}; #[tokio::test] - async fn errors_on_missing_host_header() { + async fn returns_bad_request_on_missing_host_header() { let conn_manager: Arc< ConnectionMap< String, @@ -454,7 +544,8 @@ mod proxy_handler_tests { }), ) .await; - assert!(response.is_err(), "should error on missing host header"); + let response = response.expect("should return response when missing host header"); + assert_eq!(response.status(), hyper::StatusCode::BAD_REQUEST); } #[tokio::test] @@ -496,8 +587,7 @@ mod proxy_handler_tests { }), ) .await; - assert!(response.is_ok(), "should return response when not found"); - let response = response.unwrap(); + let response = response.expect("should return response when not found"); assert_eq!(response.status(), hyper::StatusCode::NOT_FOUND); } @@ -540,8 +630,7 @@ mod proxy_handler_tests { }), ) .await; - assert!(response.is_ok(), "should return response when redirect"); - let response = response.unwrap(); + let response = response.expect("should return response when redirect"); assert_eq!(response.status(), hyper::StatusCode::SEE_OTHER); assert_eq!( response.headers().get("location").unwrap(), @@ -569,6 +658,7 @@ mod proxy_handler_tests { Some(ConnectionHttpData { redirect_http_to_https_port: None, is_aliasing: false, + http2: false, }) }); conn_manager @@ -605,11 +695,7 @@ mod proxy_handler_tests { }), ) .await; - assert!( - response.is_ok(), - "should return response when HTTPS redirect" - ); - let response = response.unwrap(); + let response = response.expect("should return response when HTTPS redirect"); assert_eq!(response.status(), hyper::StatusCode::PERMANENT_REDIRECT); assert_eq!( response.headers().get("location").unwrap(), @@ -637,6 +723,7 @@ mod proxy_handler_tests { Some(ConnectionHttpData { redirect_http_to_https_port: None, is_aliasing: false, + http2: false, }) }); conn_manager @@ -673,11 +760,8 @@ mod proxy_handler_tests { }), ) .await; - assert!( - response.is_ok(), - "should return response whyen HTTPS redirect to non-standard port" - ); - let response = response.unwrap(); + let response = + response.expect("should return response whyen HTTPS redirect to non-standard port"); assert_eq!(response.status(), hyper::StatusCode::PERMANENT_REDIRECT); assert_eq!( response.headers().get("location").unwrap(), @@ -705,6 +789,7 @@ mod proxy_handler_tests { Some(ConnectionHttpData { redirect_http_to_https_port: Some(443), is_aliasing: false, + http2: false, }) }); conn_manager @@ -741,11 +826,7 @@ mod proxy_handler_tests { }), ) .await; - assert!( - response.is_ok(), - "should return response when HTTPS redirect" - ); - let response = response.unwrap(); + let response = response.expect("should return response when HTTPS redirect"); assert_eq!(response.status(), hyper::StatusCode::PERMANENT_REDIRECT); assert_eq!( response.headers().get("location").unwrap(), @@ -773,6 +854,7 @@ mod proxy_handler_tests { Some(ConnectionHttpData { redirect_http_to_https_port: Some(8443), is_aliasing: false, + http2: false, }) }); conn_manager @@ -809,11 +891,8 @@ mod proxy_handler_tests { }), ) .await; - assert!( - response.is_ok(), - "should return response whyen HTTPS redirect to non-standard port" - ); - let response = response.unwrap(); + let response = + response.expect("should return response when HTTPS redirect to non-standard port"); assert_eq!(response.status(), hyper::StatusCode::PERMANENT_REDIRECT); assert_eq!( response.headers().get("location").unwrap(), @@ -847,6 +926,7 @@ mod proxy_handler_tests { Some(ConnectionHttpData { redirect_http_to_https_port: None, is_aliasing: false, + http2: false, }) }); conn_manager @@ -940,6 +1020,7 @@ mod proxy_handler_tests { Some(ConnectionHttpData { redirect_http_to_https_port: None, is_aliasing: false, + http2: false, }) }); conn_manager @@ -1001,8 +1082,7 @@ mod proxy_handler_tests { ) .await; assert!(!logging_rx.is_empty(), "should log after proxying request"); - assert!(response.is_ok(), "should return response after proxy"); - let response = response.unwrap(); + let response = response.expect("should return response after proxy"); assert_eq!(response.status(), hyper::StatusCode::OK); let body = response.into_body(); let body = axum::body::to_bytes(body, 32).await.unwrap(); @@ -1036,6 +1116,7 @@ mod proxy_handler_tests { Some(ConnectionHttpData { redirect_http_to_https_port: Some(443), is_aliasing: false, + http2: false, }) }); conn_manager @@ -1097,8 +1178,7 @@ mod proxy_handler_tests { ) .await; assert!(!logging_rx.is_empty(), "should log after proxying request"); - assert!(response.is_ok(), "should return response after proxy"); - let response = response.unwrap(); + let response = response.expect("should return response after proxy"); assert_eq!(response.status(), hyper::StatusCode::OK); let body = response.into_body(); let body = axum::body::to_bytes(body, 32).await.unwrap(); @@ -1132,6 +1212,7 @@ mod proxy_handler_tests { Some(ConnectionHttpData { redirect_http_to_https_port: None, is_aliasing: false, + http2: false, }) }); conn_manager @@ -1193,11 +1274,7 @@ mod proxy_handler_tests { ) .await; assert!(!logging_rx.is_empty(), "should log after proxying request"); - assert!( - response.is_ok(), - "should return response after proxying request" - ); - let response = response.unwrap(); + let response = response.expect("should return response after proxying request"); assert_eq!(response.status(), hyper::StatusCode::OK); let body = response.into_body(); let body = axum::body::to_bytes(body, 32).await.unwrap(); @@ -1231,6 +1308,7 @@ mod proxy_handler_tests { Some(ConnectionHttpData { redirect_http_to_https_port: None, is_aliasing: false, + http2: false, }) }); conn_manager @@ -1244,11 +1322,9 @@ mod proxy_handler_tests { let (socket, stream) = tokio::io::duplex(1024); let router = Router::new().route( "/ws", - any(|ws: axum::extract::WebSocketUpgrade| async move { + any(|ws: WebSocketUpgrade| async move { ws.on_upgrade(|mut socket| async move { - let _ = socket - .send(axum::extract::ws::Message::Text("Success.".into())) - .await; + let _ = socket.send(ws::Message::Text("Success.".into())).await; let _ = socket.close().await; }) }), @@ -1303,4 +1379,125 @@ mod proxy_handler_tests { jh.abort(); jh2.abort(); } + + #[tokio::test] + async fn returns_http2_response_for_existing_handler() { + 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( + "http2.handler".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( + "/http2", + put(|headers: HeaderMap, body: String| async move { + if headers.get("X-Forwarded-For").unwrap() == "127.0.0.1" + && headers.get("X-Forwarded-Host").unwrap() == "http2.handler" + && body == "The future of HTTP!" + { + "Success." + } else { + "Failure." + } + }), + ); + 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 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: None, + 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 (mut sender, conn) = + hyper::client::conn::http2::handshake(TokioExecutor::new(), TokioIo::new(stream)) + .await + .unwrap(); + let jh3 = tokio::spawn(conn); + let request = Request::builder() + .method("PUT") + .uri("https://http2.handler/http2") + .body(String::from("The future of HTTP!")) + .unwrap(); + let response = sender + .send_request(request) + .await + .expect("should return response after proxy"); + assert!(!logging_rx.is_empty(), "should log after proxying request"); + assert_eq!(response.status(), hyper::StatusCode::OK); + let body = String::from_utf8( + response + .into_body() + .collect() + .await + .expect("Error collecting response") + .to_bytes() + .into(), + ) + .unwrap(); + assert_eq!(body, "Success."); + jh.abort(); + jh2.abort(); + jh3.abort(); + } } diff --git a/src/lib.rs b/src/lib.rs index 7e69c1b..7397455 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -16,9 +16,10 @@ use std::{ use addressing::AddressDelegatorData; use anyhow::Context; +use connection_handler::ConnectionHandler; use connections::HttpAliasingConnection; use http::{DomainRedirect, ProxyData, ProxyType}; -use hyper::{Request, body::Incoming, server::conn::http1, service::service_fn}; +use hyper::{Request, body::Incoming, service::service_fn}; use hyper_util::{ rt::{TokioExecutor, TokioIo}, server::conn::auto, @@ -591,8 +592,8 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { }); let io = TokioIo::new(stream); tokio::spawn(async move { - let server = http1::Builder::new(); - let conn = server.serve_connection(io, service).with_upgrades(); + let server = auto::Builder::new(TokioExecutor::new()); + let conn = server.serve_connection_with_upgrades(io, service); match tcp_connection_timeout { Some(duration) => { let _ = timeout(duration, conn).await; @@ -618,14 +619,21 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { config.https_port ); let certificates_clone = Arc::clone(&certificates); - let tls_server_config = Arc::new( - ServerConfig::builder() - .with_no_client_auth() - .with_cert_resolver(certificates), - ); + let mut http11_server_config = ServerConfig::builder() + .with_no_client_auth() + .with_cert_resolver(certificates); + let mut http2_server_config = http11_server_config.clone(); + http11_server_config + .alpn_protocols + .extend_from_slice(&[b"http/1.1".to_vec()]); + let http11_server_config = Arc::new(http11_server_config); + http2_server_config + .alpn_protocols + .extend_from_slice(&[b"h2".to_vec(), b"http/1.1".to_vec()]); + let http2_server_config = Arc::new(http2_server_config); let ip_filter_clone = Arc::clone(&ip_filter); let https_proxy_data = Arc::new(ProxyData { - conn_manager: http_connections, + conn_manager: Arc::clone(&http_connections), telemetry: Arc::clone(&telemetry), domain_redirect: Arc::clone(&domain_redirect), protocol: Protocol::Https { @@ -655,10 +663,12 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { continue; } let proxy_data = Arc::clone(&https_proxy_data); - let server_config = Arc::clone(&tls_server_config); let ssh_config = Arc::clone(&ssh_config_clone); let mut sandhole = Arc::clone(&sandhole_clone); let certificates = Arc::clone(&certificates_clone); + let http_connections = Arc::clone(&http_connections); + let http2_server_config = Arc::clone(&http2_server_config); + let http11_server_config = Arc::clone(&http11_server_config); tokio::spawn(async move { if let Err(err) = stream.set_nodelay(true) { warn!("Error setting nodelay for {}: {}", address, err); @@ -683,7 +693,8 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { tokio::pin!(acceptor); match acceptor.as_mut().await { Ok(handshake) => { - if is_tls_alpn_challenge(&handshake.client_hello()) { + let client_hello = handshake.client_hello(); + if is_tls_alpn_challenge(&client_hello) { // Handle ALPN challenges with the ACME resolver. if let Some(challenge_config) = certificates.challenge_rustls_config() @@ -694,13 +705,26 @@ pub async fn entrypoint(config: ApplicationConfig) -> anyhow::Result<()> { } } else { // Handle regular HTTPS TLS stream. + let is_http2 = match client_hello + .server_name() + .and_then(|host| http_connections.get(host)) + { + Some(conn) => { + conn.http_data().await.is_some_and(|data| data.http2) + } + None => false, + }; + let server_config = if is_http2 { + http2_server_config + } else { + http11_server_config + }; match handshake.into_stream(server_config).await { Ok(stream) => { + let io = TokioIo::new(stream); let server = auto::Builder::new(TokioExecutor::new()); - let conn = server.serve_connection_with_upgrades( - TokioIo::new(stream), - service, - ); + let conn = + server.serve_connection_with_upgrades(io, service); match tcp_connection_timeout { Some(duration) => { let _ = timeout(duration, conn).await; diff --git a/src/login.rs b/src/login.rs index 02fbb35..5bf244f 100644 --- a/src/login.rs +++ b/src/login.rs @@ -86,7 +86,7 @@ impl ApiLogin { .build() .with_context(|| "Unable to build HTTPS client")? } else { - return Err(ServerError::UnknownHttpScheme).with_context(|| "Invalid API login URL")?; + return Err(ServerError::UnknownHttpScheme).with_context(|| "Invalid API login URL"); }; Ok(ApiLogin { configurer: PhantomData, diff --git a/src/ssh.rs b/src/ssh.rs index 8605bba..60f547f 100644 --- a/src/ssh.rs +++ b/src/ssh.rs @@ -192,6 +192,7 @@ impl UserData { http_data: Arc::new(RwLock::new(ConnectionHttpData { redirect_http_to_https_port: None, is_aliasing: false, + http2: false, })), ip_filter: Arc::new(RwLock::new(None)), tcp_alias_only: false, @@ -283,6 +284,7 @@ enum ExecCommand { AllowedFingerprints, TcpAlias, ForceHttps, + Http2, IpAllowlist, IpBlocklist, } @@ -811,6 +813,21 @@ impl Handler for ServerHandler { .redirect_http_to_https_port = Some(self.server.https_port); self.commands.insert(ExecCommand::ForceHttps); } + // - `http2` allows serving HTTP/2 to the HTTP endpoints. + ( + "http2", + AuthenticatedData::User { user_data, .. } + | AuthenticatedData::Admin { user_data, .. }, + ) => { + if self.commands.contains(ExecCommand::Http2) { + self.tx + .send(b"Invalid option \"http2\": duplicated command\r\n".to_vec()); + success = false; + break; + } + user_data.http_data.write().await.http2 = true; + self.commands.insert(ExecCommand::Http2); + } // - `ip-allowlist` requires tunneling/aliasing connections to come from // specific IP ranges. ( diff --git a/tests/https_http11_fallback.rs b/tests/https_http11_fallback.rs new file mode 100644 index 0000000..3eab87f --- /dev/null +++ b/tests/https_http11_fallback.rs @@ -0,0 +1,206 @@ +use std::{sync::Arc, time::Duration}; + +use axum::{Router, extract::Request, routing::get}; +use clap::Parser; +use http::Uri; +use http::header::{COOKIE, HOST}; +use http_body_util::BodyExt; +use hyper::{StatusCode, body::Incoming, server::conn::http1::Builder, service::service_fn}; +use hyper_util::rt::{TokioExecutor, TokioIo}; +use russh::keys::{key::PrivateKeyWithHashAlg, load_secret_key}; +use russh::{ + Channel, + client::{Msg, Session}, +}; +use rustls::{ + RootCertStore, + pki_types::{CertificateDer, pem::PemObject}, +}; +use sandhole::{ApplicationConfig, entrypoint}; +use tokio::{ + net::TcpStream, + time::{sleep, timeout}, +}; +use tokio_rustls::TlsConnector; +use tower::Service; + +#[tokio::test(flavor = "multi_thread")] +async fn https_http11_fallback() { + // 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=all", + "--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 via HTTP/2 + 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("http11.foobar.tld", 443) + .await + .expect("tcpip_forward failed"); + + // 3. Connect to the HTTP/2 port of our proxy + let mut root_store = RootCertStore::empty(); + root_store.add_parsable_certificates( + CertificateDer::pem_file_iter(concat!( + env!("CARGO_MANIFEST_DIR"), + "/tests/data/ca/rootCA.pem" + )) + .and_then(|iter| iter.collect::, _>>()) + .expect("Failed to parse certificates"), + ); + let tls_config = Arc::new( + rustls::ClientConfig::builder() + .with_root_certificates(root_store) + .with_no_client_auth(), + ); + let connector = TlsConnector::from(tls_config); + let tcp_stream = TcpStream::connect("127.0.0.1:18443") + .await + .expect("TCP connection failed"); + let tls_stream = connector + .connect("http11.foobar.tld".try_into().unwrap(), tcp_stream) + .await + .expect("TLS stream failed"); + let (mut sender, conn) = + hyper::client::conn::http2::handshake(TokioExecutor::new(), TokioIo::new(tls_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("https://http11.foobar.tld/hello") + .header(COOKIE, "foo=1") + .header(COOKIE, "bar=2") + .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 http11.foobar.tld!"); +} + +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( + "/hello", + get(async |request: Request| { + if request.headers().get(HOST).unwrap() == "http11.foobar.tld" + && request.headers().get(COOKIE).unwrap() == "foo=1; bar=2" + && request.uri() == &"/hello".parse::().unwrap() + { + dbg!(request.headers()); + "Hello from http11.foobar.tld!" + } else { + "Error" + } + }), + ); + let service = service_fn(move |req: Request| router.clone().call(req)); + tokio::spawn(async move { + Builder::new() + .serve_connection(TokioIo::new(channel.into_stream()), service) + .await + .expect("Invalid request"); + }); + Ok(()) + } +} diff --git a/tests/https_http2.rs b/tests/https_http2.rs new file mode 100644 index 0000000..791ccb4 --- /dev/null +++ b/tests/https_http2.rs @@ -0,0 +1,231 @@ +use std::{sync::Arc, time::Duration}; + +use axum::{Router, extract::Request, routing::get}; +use clap::Parser; +use http::Version; +use http_body_util::BodyExt; +use hyper::{StatusCode, body::Incoming, server::conn::http2::Builder, service::service_fn}; +use hyper_util::rt::{TokioExecutor, TokioIo}; +use russh::{ + Channel, + client::{Msg, Session}, +}; +use russh::{ + ChannelId, + keys::{key::PrivateKeyWithHashAlg, load_secret_key}, +}; +use rustls::{ + RootCertStore, + pki_types::{CertificateDer, pem::PemObject}, +}; +use sandhole::{ApplicationConfig, entrypoint}; +use tokio::{ + net::TcpStream, + sync::oneshot, + time::{sleep, timeout}, +}; +use tokio_rustls::TlsConnector; +use tower::Service; + +#[tokio::test(flavor = "multi_thread")] +async fn https_http2() { + // 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=all", + "--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 via HTTP/2 + let key = load_secret_key( + concat!(env!("CARGO_MANIFEST_DIR"), "/tests/data/private_keys/key1"), + None, + ) + .expect("Missing file key1"); + let (tx, rx) = oneshot::channel(); + let ssh_client = SshClient(Some(tx)); + 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("http2.foobar.tld", 443) + .await + .expect("tcpip_forward failed"); + let channel = session + .channel_open_session() + .await + .expect("channel_open_session failed"); + channel + .exec(true, "http2") + .await + .expect("exec http2 failed"); + let Ok(channel_id) = timeout(Duration::from_secs(2), async { rx.await.unwrap() }).await else { + panic!("Timeout waiting for server to reply."); + }; + assert_eq!(channel_id, channel.id()); + + // 3. Connect to the HTTP/2 port of our proxy + let mut root_store = RootCertStore::empty(); + root_store.add_parsable_certificates( + CertificateDer::pem_file_iter(concat!( + env!("CARGO_MANIFEST_DIR"), + "/tests/data/ca/rootCA.pem" + )) + .and_then(|iter| iter.collect::, _>>()) + .expect("Failed to parse certificates"), + ); + let tls_config = Arc::new( + rustls::ClientConfig::builder() + .with_root_certificates(root_store) + .with_no_client_auth(), + ); + let connector = TlsConnector::from(tls_config); + let tcp_stream = TcpStream::connect("127.0.0.1:18443") + .await + .expect("TCP connection failed"); + let tls_stream = connector + .connect("http2.foobar.tld".try_into().unwrap(), tcp_stream) + .await + .expect("TLS stream failed"); + let (mut sender, conn) = + hyper::client::conn::http2::handshake(TokioExecutor::new(), TokioIo::new(tls_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("https://http2.foobar.tld/") + .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 http2.foobar.tld!"); + session + .cancel_tcpip_forward("http2.foobar.tld", 443) + .await + .expect("cancel_tcpip_forward failed"); +} + +struct SshClient(Option>); + +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 |request: Request| { + if request.version() == Version::HTTP_2 { + "Hello from http2.foobar.tld!" + } else { + "Error" + } + }), + ); + let service = service_fn(move |req: Request| router.clone().call(req)); + tokio::spawn(async move { + Builder::new(TokioExecutor::new()) + .serve_connection(TokioIo::new(channel.into_stream()), service) + .await + .expect("Invalid request"); + }); + Ok(()) + } + + async fn channel_success( + &mut self, + channel: russh::ChannelId, + _session: &mut Session, + ) -> Result<(), Self::Error> { + if let Some(tx) = self.0.take() { + tx.send(channel).unwrap(); + }; + Ok(()) + } +} diff --git a/tests/lib_close_with_unix_signals.rs b/tests/lib_close_with_unix_signals.rs index 98b921e..ab604e2 100644 --- a/tests/lib_close_with_unix_signals.rs +++ b/tests/lib_close_with_unix_signals.rs @@ -15,6 +15,10 @@ use tokio::{ #[tokio::test(flavor = "multi_thread")] async fn lib_configure_from_scratch() { + let _ = env_logger::builder() + .filter_module("sandhole", log::LevelFilter::Debug) + .is_test(true) + .try_init(); for signal in [SIGINT, SIGTERM] { // 1. Initialize Sandhole let config = ApplicationConfig::parse_from([ diff --git a/tests/ssh_invalid_exec_commands.rs b/tests/ssh_invalid_exec_commands.rs index 7347d0c..aeb76ab 100644 --- a/tests/ssh_invalid_exec_commands.rs +++ b/tests/ssh_invalid_exec_commands.rs @@ -184,7 +184,28 @@ async fn ssh_invalid_exec_commands() { }; assert_eq!(channel_id, channel.id()); assert!(rx.is_empty(), "rx shouldn't have any remaining messages"); - // 2h. Fail to run `ip-allowlist` with invalid CIDR + // 2h. Fail to run `http2` twice + channel + .exec(true, "http2 http2") + .await + .expect("exec http2 failed"); + let Ok(channel_id) = timeout(Duration::from_secs(2), async { rx.recv().await.unwrap() }).await + else { + panic!("Timeout waiting for server to reply."); + }; + assert_eq!(channel_id, channel.id()); + assert!(rx.is_empty(), "rx shouldn't have any remaining messages"); + channel + .exec(true, "http2") + .await + .expect("exec http2 failed"); + let Ok(channel_id) = timeout(Duration::from_secs(2), async { rx.recv().await.unwrap() }).await + else { + panic!("Timeout waiting for server to reply."); + }; + assert_eq!(channel_id, channel.id()); + assert!(rx.is_empty(), "rx shouldn't have any remaining messages"); + // 2i. Fail to run `ip-allowlist` with invalid CIDR channel .exec(true, "ip-allowlist=10.0.0") .await @@ -195,7 +216,7 @@ async fn ssh_invalid_exec_commands() { }; assert_eq!(channel_id, channel.id()); assert!(rx.is_empty(), "rx shouldn't have any remaining messages"); - // 2i. Fail to run `ip-allowlist` with no CIDRs + // 2j. Fail to run `ip-allowlist` with no CIDRs channel .exec(true, "ip-allowlist=") .await @@ -206,7 +227,7 @@ async fn ssh_invalid_exec_commands() { }; assert_eq!(channel_id, channel.id()); assert!(rx.is_empty(), "rx shouldn't have any remaining messages"); - // 2j. Fail to run `ip-allowlist` twice + // 2k. Fail to run `ip-allowlist` twice channel .exec( true, @@ -230,7 +251,7 @@ async fn ssh_invalid_exec_commands() { }; assert_eq!(channel_id, channel.id()); assert!(rx.is_empty(), "rx shouldn't have any remaining messages"); - // 2k. Fail to run `ip-blocklist` with invalid CIDR + // 2l. Fail to run `ip-blocklist` with invalid CIDR channel .exec(true, "ip-blocklist=10.0.0") .await @@ -241,7 +262,7 @@ async fn ssh_invalid_exec_commands() { }; assert_eq!(channel_id, channel.id()); assert!(rx.is_empty(), "rx shouldn't have any remaining messages"); - // 2l. Fail to run `ip-blocklist` with no CIDRs + // 2m. Fail to run `ip-blocklist` with no CIDRs channel .exec(true, "ip-blocklist=") .await @@ -252,7 +273,7 @@ async fn ssh_invalid_exec_commands() { }; assert_eq!(channel_id, channel.id()); assert!(rx.is_empty(), "rx shouldn't have any remaining messages"); - // 2m. Fail to run `ip-blocklist` twice + // 2n. Fail to run `ip-blocklist` twice channel .exec( true, @@ -276,7 +297,7 @@ async fn ssh_invalid_exec_commands() { }; assert_eq!(channel_id, channel.id()); assert!(rx.is_empty(), "rx shouldn't have any remaining messages"); - // 2n. Fail to run an unknown command + // 2o. Fail to run an unknown command channel .exec(true, "unknown-command") .await