diff --git a/src/entrypoint.rs b/src/entrypoint.rs index 0ec08da..26c4505 100644 --- a/src/entrypoint.rs +++ b/src/entrypoint.rs @@ -56,8 +56,8 @@ use crate::{ error::ServerError, fingerprints::FingerprintsValidator, http::{ - DomainRedirect, Protocol, ProxyData, ProxyType, http2_handler, http11_handler, - proxy_handler, + DomainRedirect, Protocol, ProxyData, ProxyType, http2::https_2_handler, + http11::https_11_handler, proxy_handler, }, ip::{IpFilter, IpFilterConfig}, quota::{DummyQuotaHandler, QuotaHandler, QuotaMap}, @@ -999,7 +999,7 @@ async fn handle_https_connection( // Create a Hyper service and serve over the accepted TLS connection. let io = TokioIo::new(stream); let service = service_fn(move |req: Request| { - http2_handler( + https_2_handler( req, address, Arc::clone(&proxy_data), @@ -1030,7 +1030,7 @@ async fn handle_https_connection( // Create a Hyper service and serve over the accepted TLS connection. let io = TokioIo::new(stream); let service = service_fn(move |req: Request| { - http11_handler( + https_11_handler( req, address, Arc::clone(&proxy_data), diff --git a/src/http/http11.rs b/src/http/http11.rs new file mode 100644 index 0000000..4d22956 --- /dev/null +++ b/src/http/http11.rs @@ -0,0 +1,472 @@ +use std::{ + error::Error, fmt::Debug, net::SocketAddr, pin::pin, str::FromStr, sync::Arc, time::Instant, +}; + +use crate::{ + connection_handler::ConnectionHandler, + connections::ConnectionGetByHttpHost, + http::{ + ArcProxyData, HttpError, HttpLog, Protocol, ProxyData, ProxyResponse, ProxyType, + TimedResponse, X_FORWARDED_FOR, X_FORWARDED_HOST, X_FORWARDED_PORT, X_FORWARDED_PROTO, + append_to_header, http_log, + }, + keepalive::KeepaliveAlias, + telemetry::{TELEMETRY_COUNTER_HTTP_REQUESTS, TELEMETRY_KEY_HOSTNAME}, +}; + +use axum::{body::Body as AxumBody, response::IntoResponse}; +use http::header::COOKIE; +use http::{Uri, Version}; +use hyper::{ + Request, Response, StatusCode, + body::Body, + header::{HOST, UPGRADE}, +}; +use hyper_util::rt::TokioIo; +use metrics::counter; +use tokio::{ + io::{AsyncRead, AsyncWrite, copy_bidirectional_with_sizes}, + time::timeout, +}; + +#[cfg_attr( + not(coverage_nightly), + tracing::instrument(skip(proxy_data, handler), level = "debug") +)] +pub(crate) async fn https_11_handler( + request: Request, + tcp_address: SocketAddr, + proxy_data: Arc>, + handler: Arc, + host: String, +) -> color_eyre::Result> +where + M: ConnectionGetByHttpHost> + Send + Sync + 'static, + H: ConnectionHandler + Send + Sync + 'static, + T: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static, + B: Body + Debug + Send + Unpin + 'static, + ::Data: Send + Sync + 'static, + ::Error: Error + Send + Sync + 'static, +{ + match https_11_handler_inner(request, tcp_address, proxy_data, handler, host).await { + Ok(response) => Ok(match response { + ProxyResponse::Axum(response) => response, + ProxyResponse::Proxy(response) => response.into_response(), + }), + Err(error) => Ok(error.into_response()), + } +} + +#[inline] +async fn https_11_handler_inner( + mut request: Request, + tcp_address: SocketAddr, + proxy_data: Arc>, + handler: Arc, + host: String, +) -> Result +where + M: ConnectionGetByHttpHost> + Send + Sync + 'static, + H: ConnectionHandler + Send + Sync + 'static, + T: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static, + B: Body + Debug + Send + Unpin + 'static, + ::Data: Send + Sync + 'static, + ::Error: Error + Send + Sync + 'static, +{ + let timer = Instant::now(); + let host_header_or_uri = match request.version() { + Version::HTTP_2 => request.uri().host().ok_or(HttpError::MissingUriHost)?, + Version::HTTP_11 => match request.headers().get(HOST) { + Some(header_value) => match header_value.to_str() { + Ok(header) => header + .split(':') + .next() + .ok_or(HttpError::InvalidHostHeader)?, + Err(_) => return Err(HttpError::InvalidHostHeader), + }, + None => return Err(HttpError::MissingHostHeader), + }, + version => return Err(HttpError::InvalidHttpVersion(version)), + }; + if host != host_header_or_uri { + return Err(HttpError::MismatchedHostHeader); + } + + let ip = tcp_address.ip().to_canonical(); + + // Read protocol information for X-Forwarded headers + let Protocol::Https { port } = proxy_data.protocol else { + unreachable!("HTTPS-only"); + }; + + // Add proxied info to the proper headers, but don't overwrite any existing proxy headers + let headers = request.headers_mut(); + append_to_header(headers, &X_FORWARDED_FOR, ip.to_string().as_bytes()); + append_to_header(headers, &X_FORWARDED_HOST, host.as_bytes()); + append_to_header(headers, &X_FORWARDED_PROTO, b"https"); + append_to_header(headers, &X_FORWARDED_PORT, port.to_string().as_bytes()); + + // Add this request to the telemetry for the host + counter!(TELEMETRY_COUNTER_HTTP_REQUESTS, TELEMETRY_KEY_HOSTNAME => host.clone()).increment(1); + + let host_clone = host.clone(); + let http_data = handler.http_data(); + let request_host = http_data + .as_ref() + .and_then(|data| data.host.as_deref()) + .unwrap_or(host_clone.as_str()); + + let key = KeepaliveAlias(host.clone(), ip, None); + handle_http11_request( + request, + tcp_address, + proxy_data, + handler, + request_host, + key, + timer, + ) + .await +} + +pub(crate) async fn handle_http11_request( + mut request: Request, + tcp_address: SocketAddr, + proxy_data: Arc>, + handler: Arc, + request_host: &str, + key: KeepaliveAlias, + timer: Instant, +) -> Result +where + M: ConnectionGetByHttpHost> + Send + Sync + 'static, + H: ConnectionHandler + Send + Sync + 'static, + T: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static, + B: Body + Debug + Send + Unpin + 'static, + ::Data: Send + Sync + 'static, + ::Error: Error + Send + Sync + 'static, +{ + let http_log_builder = HttpLog::builder() + .uri(request.uri().path().to_string()) + .method(request.method().to_string()) + .host(key.0.clone()) + .ip(key.1.to_string()); + // Ensure best-effort compatibility of proxy request with HTTP/1.1 format + // -> Add host header if missing + request + .headers_mut() + .insert(HOST, request_host.try_into().expect("valid host")); + // -> Change URI to only include path and query + *request.uri_mut() = request + .uri() + .path_and_query() + .map(|path| Uri::from_str(path.as_str()).expect("valid URI")) + .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().expect("valid header value")); + } + // Check for an Upgrade header + if let Some(request_upgrade) = request.headers().get(UPGRADE) { + // Get HTTP/1.1 sender and remote log channel with upgrades with a new connection + let io = match proxy_data.proxy_type { + ProxyType::Tunneling => { + handler + .tunneling_channel(tcp_address.ip(), tcp_address.port()) + .await + } + ProxyType::Aliasing => { + handler + .aliasing_channel(tcp_address.ip(), tcp_address.port(), key.2.as_ref()) + .await + } + }; + let (mut sender, conn) = hyper::client::conn::http1::handshake(TokioIo::new(io?)).await?; + tokio::spawn(Box::pin(async move { + if let Err(error) = conn.with_upgrades().await { + #[cfg(not(coverage_nightly))] + tracing::warn!(%error, "HTTP/1.1 connection failed."); + } + })); + let tx = handler.log_channel(); + + // If there is an Upgrade header, make sure that it's a valid Websocket upgrade + 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? + } else { + let elapsed_time = timer.elapsed(); + http_log( + http_log_builder + .status(StatusCode::GATEWAY_TIMEOUT.as_u16()) + .elapsed_time(elapsed_time) + .build(), + Some(tx), + proxy_data.disable_http_logs, + ); + return Err(HttpError::RequestTimeout); + } + } + None => sender.send_request(request).await?, + }; + // Check if the underlying server accepts the Upgrade request + match response.status() { + StatusCode::SWITCHING_PROTOCOLS => { + if request_type + == response + .headers() + .get(UPGRADE) + .ok_or(HttpError::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; + let buffer_size = proxy_data.buffer_size; + // Start a task to copy data between the two Upgraded parts + tokio::spawn(async move { + let mut upgraded_request = + TokioIo::new(upgraded_request.await.expect("upgradable request")); + 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_with_sizes( + &mut upgraded_response, + &mut upgraded_request, + buffer_size, + buffer_size, + ) + .await + }) + .await; + } + // If there isn't a Websocket timeout, copy data between both sides unconditionally. + None => { + let _ = copy_bidirectional_with_sizes( + &mut upgraded_response, + &mut upgraded_request, + buffer_size, + buffer_size, + ) + .await; + } + } + }); + } + // Return the response to the client + let http_log_builder = http_log_builder.status(response.status().as_u16()); + Ok(ProxyResponse::Proxy(TimedResponse { + response, + on_drop: Some(Box::new(move || { + http_log( + http_log_builder.elapsed_time(timer.elapsed()).build(), + Some(tx), + proxy_data.disable_http_logs, + ) + })), + })) + } + _ => { + let http_log_builder = http_log_builder.status(response.status().as_u16()); + Ok(ProxyResponse::Proxy(TimedResponse { + response, + on_drop: Some(Box::new(move || { + http_log( + http_log_builder.elapsed_time(timer.elapsed()).build(), + Some(tx), + proxy_data.disable_http_logs, + ) + })), + })) + } + } + } else { + loop { + // If Upgrade header is not present, get HTTP/1.1 sender and remote log channel + let (mut sender, tx, _guard) = if proxy_data.has_pool_queue + && let Some(guard) = proxy_data.get_http11_pool_guard(key.clone()) + { + // Race between new channel and HTTP pool + let mut recv = guard.pool.recv(); + let mut channel_future = pin!(async { + match proxy_data.proxy_type { + ProxyType::Tunneling => { + handler + .tunneling_channel(tcp_address.ip(), tcp_address.port()) + .await + } + ProxyType::Aliasing => { + handler + .aliasing_channel( + tcp_address.ip(), + tcp_address.port(), + key.2.as_ref(), + ) + .await + } + } + }); + loop { + tokio::select! { + result = recv => { + match result { + // Return pool item + Ok(tuple) if tuple.0.is_ready() => break (tuple.0, tuple.1, guard), + // Connection is closed; discard pool item + Ok(_) => { + recv = guard.pool.recv(); + continue; + }, + Err(_) => return Err(HttpError::PoolClosed), + } + } + result = &mut channel_future => { + let (sender, conn) = hyper::client::conn::http1::handshake( + TokioIo::new(result?), + ) + .await?; + tokio::spawn(Box::pin(async move { + if let Err(error) = conn.await { + #[cfg(not(coverage_nightly))] + tracing::warn!(%error, "HTTP/1.1 connection failed."); + } + })); + break (sender, handler.log_channel(), guard); + } + } + } + } else { + // No pool timeout - get recycled sender or create new one + 'sender: loop { + match proxy_data.get_http11_pool_guard(key.clone()) { + Some(guard) => { + while let Ok(sender) = guard.pool.try_recv() { + if sender.0.is_ready() { + break 'sender (sender.0, sender.1, guard); + } + } + } + None => { + let io = match proxy_data.proxy_type { + ProxyType::Tunneling => { + handler + .tunneling_channel(tcp_address.ip(), tcp_address.port()) + .await + } + ProxyType::Aliasing => { + handler + .aliasing_channel( + tcp_address.ip(), + tcp_address.port(), + key.2.as_ref(), + ) + .await + } + }; + let (sender, conn) = + hyper::client::conn::http1::handshake(TokioIo::new(io?)).await?; + tokio::spawn(Box::pin(async move { + if let Err(error) = conn.await { + #[cfg(not(coverage_nightly))] + tracing::warn!(%error, "HTTP/1.1 connection failed."); + } + })); + let guard = proxy_data.create_http11_pool_guard(key.clone()); + break (sender, handler.log_channel(), guard); + } + } + } + }; + + // Create entry for pool + let key_clone = key.clone(); + let pool = { + let pool_ref = proxy_data + .keepalive_http11_pool_map + .entry(key_clone) + .or_insert_with(|| Arc::new(async_channel::unbounded())) + .downgrade(); + pool_ref.0.clone() + }; + + 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.try_send_request(request)).await + { + response + } else { + let elapsed_time = timer.elapsed(); + http_log( + http_log_builder + .status(StatusCode::GATEWAY_TIMEOUT.as_u16()) + .elapsed_time(elapsed_time) + .build(), + Some(tx), + proxy_data.disable_http_logs, + ); + return Err(HttpError::RequestTimeout); + } + } + None => sender.try_send_request(request).await, + }; + let response = match response { + Ok(response) => response, + Err(mut error) => { + if let Some(recovered) = error.take_message() { + #[cfg(not(coverage_nightly))] + tracing::debug!(error = %error.error(), "Recovering HTTP/1.1 request to try again."); + request = recovered; + continue; + } else { + let elapsed_time = timer.elapsed(); + http_log( + http_log_builder + .status(StatusCode::INTERNAL_SERVER_ERROR.as_u16()) + .elapsed_time(elapsed_time) + .build(), + Some(tx), + proxy_data.disable_http_logs, + ); + return Err(error.into_error().into()); + } + } + }; + + // Return the received response to the client + let http_log_builder = http_log_builder.status(response.status().as_u16()); + return Ok(ProxyResponse::Proxy(TimedResponse { + response, + on_drop: Some(Box::new(move || { + // Log HTTP request + http_log( + http_log_builder.elapsed_time(timer.elapsed()).build(), + Some(tx.clone()), + proxy_data.disable_http_logs, + ); + // Return sender to pool + tokio::spawn(async move { + let _ = pool.send((sender, tx)).await; + }); + })), + })); + } + } +} diff --git a/src/http/http2.rs b/src/http/http2.rs new file mode 100644 index 0000000..88e1ee5 --- /dev/null +++ b/src/http/http2.rs @@ -0,0 +1,320 @@ +use std::{error::Error, fmt::Debug, net::SocketAddr, pin::pin, sync::Arc, time::Instant}; + +use crate::{ + connection_handler::ConnectionHandler, + connections::ConnectionGetByHttpHost, + http::{ + ArcProxyData, HttpError, HttpLog, Protocol, ProxyData, ProxyResponse, ProxyType, + TimedResponse, X_FORWARDED_FOR, X_FORWARDED_HOST, X_FORWARDED_PORT, X_FORWARDED_PROTO, + append_to_header, http_log, + }, + keepalive::KeepaliveAlias, + telemetry::{TELEMETRY_COUNTER_HTTP_REQUESTS, TELEMETRY_KEY_HOSTNAME}, +}; + +use axum::{body::Body as AxumBody, response::IntoResponse}; +use http::{Uri, header::CONNECTION, uri::Authority}; +use hyper::{Request, Response, StatusCode, body::Body}; +use hyper_util::rt::{TokioExecutor, TokioIo}; +use metrics::counter; +use tokio::{ + io::{AsyncRead, AsyncWrite}, + time::timeout, +}; + +#[cfg_attr( + not(coverage_nightly), + tracing::instrument(skip(proxy_data, handler), level = "debug") +)] +pub(crate) async fn https_2_handler( + request: Request, + tcp_address: SocketAddr, + proxy_data: Arc>, + handler: Arc, + host: String, +) -> color_eyre::Result> +where + M: ConnectionGetByHttpHost> + Send + Sync + 'static, + H: ConnectionHandler + Send + Sync + 'static, + T: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static, + B: Body + Debug + Send + Unpin + 'static, + ::Data: Send + Sync + 'static, + ::Error: Error + Send + Sync + 'static, +{ + match http2_handler_inner(request, tcp_address, proxy_data, handler, host).await { + Ok(response) => Ok(match response { + ProxyResponse::Axum(response) => response, + ProxyResponse::Proxy(response) => response.into_response(), + }), + Err(error) => Ok(error.into_response()), + } +} + +#[inline] +async fn http2_handler_inner( + mut request: Request, + tcp_address: SocketAddr, + proxy_data: Arc>, + handler: Arc, + host: String, +) -> Result +where + M: ConnectionGetByHttpHost> + Send + Sync + 'static, + H: ConnectionHandler + Send + Sync + 'static, + T: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static, + B: Body + Debug + Send + Unpin + 'static, + ::Data: Send + Sync + 'static, + ::Error: Error + Send + Sync + 'static, +{ + let timer = Instant::now(); + let Some(host_uri) = request.uri().host() else { + return Err(HttpError::MissingHostHeader); + }; + if host != host_uri { + return Err(HttpError::MismatchedHostHeader); + }; + + let ip = tcp_address.ip().to_canonical(); + + // Read protocol information for X-Forwarded headers + let Protocol::Https { port } = proxy_data.protocol else { + unreachable!("HTTPS-only"); + }; + + // Add proxied info to the proper headers, but don't overwrite any existing proxy headers + let headers = request.headers_mut(); + append_to_header(headers, &X_FORWARDED_FOR, ip.to_string().as_bytes()); + append_to_header(headers, &X_FORWARDED_HOST, host.as_bytes()); + append_to_header(headers, &X_FORWARDED_PROTO, b"https"); + append_to_header(headers, &X_FORWARDED_PORT, port.to_string().as_bytes()); + + // Add this request to the telemetry for the host + counter!(TELEMETRY_COUNTER_HTTP_REQUESTS, TELEMETRY_KEY_HOSTNAME => host.clone()).increment(1); + + let host_clone = host.clone(); + let http_data = handler.http_data(); + let request_host = http_data + .as_ref() + .and_then(|data| data.host.as_deref()) + .unwrap_or(host_clone.as_str()); + + let key = KeepaliveAlias(host.clone(), ip, None); + handle_http2_request( + request, + tcp_address, + proxy_data, + handler, + request_host, + key, + timer, + ) + .await +} + +pub(crate) async fn handle_http2_request( + mut request: Request, + tcp_address: SocketAddr, + proxy_data: Arc>, + handler: Arc, + request_host: &str, + key: KeepaliveAlias, + timer: Instant, +) -> Result +where + M: ConnectionGetByHttpHost> + Send + Sync + 'static, + H: ConnectionHandler + Send + Sync + 'static, + T: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static, + B: Body + Debug + Send + Unpin + 'static, + ::Data: Send + Sync + 'static, + ::Error: Error + Send + Sync + 'static, +{ + let http_log_builder = HttpLog::builder() + .uri(request.uri().path().to_string()) + .method(request.method().to_string()) + .host(key.0.clone()) + .ip(key.1.to_string()); + + loop { + // Get HTTP/2 sender and remote log channel + let (mut sender, tx, _guard) = if proxy_data.has_pool_queue + && let Some(guard) = proxy_data.get_http2_pool_guard(key.clone()) + { + // Race between new channel and HTTP pool + let mut recv = guard.pool.recv(); + let mut channel_future = pin!(async { + match proxy_data.proxy_type { + ProxyType::Tunneling => { + handler + .tunneling_channel(tcp_address.ip(), tcp_address.port()) + .await + } + ProxyType::Aliasing => { + handler + .aliasing_channel(tcp_address.ip(), tcp_address.port(), key.2.as_ref()) + .await + } + } + }); + loop { + tokio::select! { + result = recv => { + match result { + // Return pool item + Ok(tuple) if tuple.0.is_ready() => break (tuple.0, tuple.1, guard), + // Connection is closed; discard pool item + Ok(_) => { + recv = guard.pool.recv(); + continue; + }, + Err(_) => return Err(HttpError::PoolClosed), + } + } + result = &mut channel_future => { + let (sender, conn) = hyper::client::conn::http2::handshake( + TokioExecutor::new(), + TokioIo::new(result?), + ) + .await?; + tokio::spawn(Box::pin(async move { + if let Err(error) = conn.await { + #[cfg(not(coverage_nightly))] + tracing::warn!(%error, "HTTP/2 connection failed."); + } + })); + break (sender, handler.log_channel(), guard); + } + } + } + } else { + // No pool timeout - get recycled sender or create new one + 'sender: loop { + match proxy_data.get_http2_pool_guard(key.clone()) { + Some(guard) => { + while let Ok(sender) = guard.pool.try_recv() { + if sender.0.is_ready() { + break 'sender (sender.0, sender.1, guard); + } + } + } + None => { + let io = match proxy_data.proxy_type { + ProxyType::Tunneling => { + handler + .tunneling_channel(tcp_address.ip(), tcp_address.port()) + .await + } + ProxyType::Aliasing => { + handler + .aliasing_channel( + tcp_address.ip(), + tcp_address.port(), + key.2.as_ref(), + ) + .await + } + }; + let (sender, conn) = hyper::client::conn::http2::handshake( + TokioExecutor::new(), + TokioIo::new(io?), + ) + .await?; + tokio::spawn(Box::pin(async move { + if let Err(error) = conn.await { + #[cfg(not(coverage_nightly))] + tracing::warn!(%error, "HTTP/2 connection failed."); + } + })); + let guard = proxy_data.create_http2_pool_guard(key.clone()); + break (sender, handler.log_channel(), guard); + } + } + } + }; + + // Create entry for pool + let key_clone = key.clone(); + let pool = { + let pool_ref = proxy_data + .keepalive_http2_pool_map + .entry(key_clone) + .or_insert_with(|| Arc::new(async_channel::unbounded())) + .downgrade(); + pool_ref.0.clone() + }; + + // Create an HTTP/2 handshake over the selected channel + let mut uri_parts = request.uri().clone().into_parts(); + let authority = uri_parts.authority.as_mut().expect("Host has been checked"); + *authority = Authority::from_maybe_shared( + authority + .as_str() + .replace(authority.host(), request_host) + .into_bytes(), + )?; + *request.uri_mut() = Uri::from_parts(uri_parts)?; + request.headers_mut().insert( + CONNECTION, + "keepalive".try_into().expect("valid HeaderValue"), + ); + + 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.try_send_request(request)).await { + response + } else { + let elapsed_time = timer.elapsed(); + http_log( + http_log_builder + .status(StatusCode::GATEWAY_TIMEOUT.as_u16()) + .elapsed_time(elapsed_time) + .build(), + Some(tx), + proxy_data.disable_http_logs, + ); + return Err(HttpError::RequestTimeout); + } + } + None => sender.try_send_request(request).await, + }; + let response = match response { + Ok(response) => response, + Err(mut error) => { + if let Some(recovered) = error.take_message() { + #[cfg(not(coverage_nightly))] + tracing::debug!(error = %error.error(), "Recovering HTTP/2 request to try again."); + request = recovered; + continue; + } else { + let elapsed_time = timer.elapsed(); + http_log( + http_log_builder + .status(StatusCode::INTERNAL_SERVER_ERROR.as_u16()) + .elapsed_time(elapsed_time) + .build(), + Some(tx), + proxy_data.disable_http_logs, + ); + return Err(error.into_error().into()); + } + } + }; + + let http_log_builder = http_log_builder.status(response.status().as_u16()); + return Ok(ProxyResponse::Proxy(TimedResponse { + response, + on_drop: Some(Box::new(move || { + // Log HTTP request + http_log( + http_log_builder.elapsed_time(timer.elapsed()).build(), + Some(tx.clone()), + proxy_data.disable_http_logs, + ); + // Send sender to pool + tokio::spawn(async move { + let _ = pool.send((sender, tx)).await; + }); + })), + })); + } +} diff --git a/src/http.rs b/src/http/mod.rs similarity index 74% rename from src/http.rs rename to src/http/mod.rs index d20204a..952e352 100644 --- a/src/http.rs +++ b/src/http/mod.rs @@ -3,17 +3,21 @@ use std::{ fmt::Debug, marker::PhantomData, net::SocketAddr, - pin::{Pin, pin}, + pin::Pin, str::FromStr, sync::{Arc, LazyLock}, task::{Context, Poll}, time::{Duration, Instant}, }; +pub(crate) mod http11; +pub(crate) mod http2; + use crate::{ connection_handler::ConnectionHandler, connections::ConnectionGetByHttpHost, error::ServerError, + http::{http2::handle_http2_request, http11::handle_http11_request}, keepalive::KeepaliveAlias, ssh::ServerHandlerSender, tcp_alias::BorrowedTcpAlias, @@ -31,26 +35,19 @@ use axum::{ }; use bon::Builder; use dashmap::DashMap; -use http::{ - HeaderMap, HeaderName, HeaderValue, Uri, Version, - header::CONNECTION, - uri::{Authority, InvalidUri}, -}; -use http::{header::COOKIE, uri::InvalidUriParts}; +use http::uri::InvalidUriParts; +use http::{HeaderMap, HeaderName, HeaderValue, Version, uri::InvalidUri}; use hyper::{ Request, Response, StatusCode, body::{Body, Incoming}, - client::conn::{http1, http2}, - header::{HOST, UPGRADE}, + client::conn::http1::SendRequest as Http11SendRequest, + client::conn::http2::SendRequest as Http2SendRequest, + header::HOST, }; -use hyper_util::rt::{TokioExecutor, TokioIo}; use metrics::{counter, histogram}; use owo_colors::{OwoColorize, Style}; use russh::keys::ssh_key::Fingerprint; -use tokio::{ - io::{AsyncRead, AsyncWrite, copy_bidirectional_with_sizes}, - time::timeout, -}; +use tokio::io::{AsyncRead, AsyncWrite}; static X_FORWARDED_FOR: LazyLock = LazyLock::new(|| HeaderName::from_str("X-Forwarded-For").expect("valid header name")); @@ -61,12 +58,12 @@ static X_FORWARDED_PROTO: LazyLock = static X_FORWARDED_PORT: LazyLock = LazyLock::new(|| HeaderName::from_str("X-Forwarded-Port").expect("valid header name")); -enum ProxyResponse { +pub(crate) enum ProxyResponse { Axum(Response), Proxy(TimedResponse), } -struct TimedResponse { +pub(crate) struct TimedResponse { response: Response, on_drop: Option>, } @@ -185,7 +182,11 @@ fn http_log(data: HttpLog, tx: Option, disable_http_logs: b } // Append the bytes to the given comma-separated entry of HeaderMap -fn append_to_header(headers: &mut HeaderMap, header_name: &HeaderName, new_value: &[u8]) { +pub(crate) fn append_to_header( + headers: &mut HeaderMap, + header_name: &HeaderName, + new_value: &[u8], +) { match headers.entry(header_name) { http::header::Entry::Vacant(entry) => { entry.insert(HeaderValue::from_bytes(new_value).expect("valid header value")); @@ -291,6 +292,8 @@ impl IntoResponse for HttpError { } type KeepalivePool

= Arc<(Sender

, Receiver

)>; +type Http11KeepalivePool = KeepalivePool<(Http11SendRequest, ServerHandlerSender)>; +type Http2KeepalivePool = KeepalivePool<(Http2SendRequest, ServerHandlerSender)>; // Data commonly reused between HTTP proxy requests. #[derive(Builder)] @@ -305,22 +308,10 @@ where { // Keep-alive pool for opened HTTP/1.1 connections. #[builder(default = Arc::default())] - keepalive_http11_pool_map: Arc< - DashMap< - KeepaliveAlias, - KeepalivePool<(http1::SendRequest, ServerHandlerSender)>, - RandomState, - >, - >, + keepalive_http11_pool_map: Arc, RandomState>>, // Keep-alive pool for opened HTTP/2 connections. #[builder(default = Arc::default())] - keepalive_http2_pool_map: Arc< - DashMap< - KeepaliveAlias, - KeepalivePool<(http2::SendRequest, ServerHandlerSender)>, - RandomState, - >, - >, + keepalive_http2_pool_map: Arc, RandomState>>, // Whether the server supports a queue pool for handlers. has_pool_queue: bool, // Connection manager to get handlers from. @@ -358,35 +349,23 @@ where } pub(crate) struct Http11PoolGuard { - pool: Receiver<(http1::SendRequest, ServerHandlerSender)>, + pool: Receiver<(Http11SendRequest, ServerHandlerSender)>, key: KeepaliveAlias, - map: Arc< - DashMap< - KeepaliveAlias, - KeepalivePool<(http1::SendRequest, ServerHandlerSender)>, - RandomState, - >, - >, + map: Arc, RandomState>>, } impl Drop for Http11PoolGuard { fn drop(&mut self) { self.map.remove_if(&self.key, |_, value| { - Arc::strong_count(&value) == 1 && value.1.is_empty() + Arc::strong_count(value) == 1 && value.1.is_empty() }); } } pub(crate) struct Http2PoolGuard { - pool: Receiver<(http2::SendRequest, ServerHandlerSender)>, + pool: Receiver<(Http2SendRequest, ServerHandlerSender)>, key: KeepaliveAlias, - map: Arc< - DashMap< - KeepaliveAlias, - KeepalivePool<(http2::SendRequest, ServerHandlerSender)>, - RandomState, - >, - >, + map: Arc, RandomState>>, } pub(crate) trait ArcProxyData { @@ -399,7 +378,7 @@ pub(crate) trait ArcProxyData { impl Drop for Http2PoolGuard { fn drop(&mut self) { self.map.remove_if(&self.key, |_, value| { - Arc::strong_count(&value) == 1 && value.1.is_empty() + Arc::strong_count(value) == 1 && value.1.is_empty() }); } } @@ -655,712 +634,7 @@ where ) .await } - version => return Err(HttpError::InvalidHttpVersion(version)), - } -} - -#[cfg_attr( - not(coverage_nightly), - tracing::instrument(skip(proxy_data, handler), level = "debug") -)] -pub(crate) async fn http11_handler( - request: Request, - tcp_address: SocketAddr, - proxy_data: Arc>, - handler: Arc, - host: String, -) -> color_eyre::Result> -where - M: ConnectionGetByHttpHost> + Send + Sync + 'static, - H: ConnectionHandler + Send + Sync + 'static, - T: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static, - B: Body + Debug + Send + Unpin + 'static, - ::Data: Send + Sync + 'static, - ::Error: Error + Send + Sync + 'static, -{ - match http11_handler_inner(request, tcp_address, proxy_data, handler, host).await { - Ok(response) => Ok(match response { - ProxyResponse::Axum(response) => response, - ProxyResponse::Proxy(response) => response.into_response(), - }), - Err(error) => Ok(error.into_response()), - } -} - -#[inline] -async fn http11_handler_inner( - request: Request, - tcp_address: SocketAddr, - proxy_data: Arc>, - handler: Arc, - host: String, -) -> Result -where - M: ConnectionGetByHttpHost> + Send + Sync + 'static, - H: ConnectionHandler + Send + Sync + 'static, - T: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static, - B: Body + Debug + Send + Unpin + 'static, - ::Data: Send + Sync + 'static, - ::Error: Error + Send + Sync + 'static, -{ - let timer = Instant::now(); - let host_header_or_uri = match request.version() { - Version::HTTP_2 => request.uri().host().ok_or(HttpError::MissingUriHost)?, - Version::HTTP_11 => match request.headers().get(HOST) { - Some(header_value) => match header_value.to_str() { - Ok(header) => header - .split(':') - .next() - .ok_or(HttpError::InvalidHostHeader)?, - Err(_) => return Err(HttpError::InvalidHostHeader), - }, - None => return Err(HttpError::MissingHostHeader), - }, - version => return Err(HttpError::InvalidHttpVersion(version)), - }; - if host != host_header_or_uri { - return Err(HttpError::MismatchedHostHeader); - } - let ip = tcp_address.ip().to_canonical(); - let host_clone = host.clone(); - let http_data = handler.http_data(); - let request_host = http_data - .as_ref() - .and_then(|data| data.host.as_deref()) - .unwrap_or(host_clone.as_str()); - - let key = KeepaliveAlias(host.clone(), ip, None); - handle_http11_request( - request, - tcp_address, - proxy_data, - handler, - request_host, - key, - timer, - ) - .await -} - -async fn handle_http11_request( - mut request: Request, - tcp_address: SocketAddr, - proxy_data: Arc>, - handler: Arc, - request_host: &str, - key: KeepaliveAlias, - timer: Instant, -) -> Result -where - M: ConnectionGetByHttpHost> + Send + Sync + 'static, - H: ConnectionHandler + Send + Sync + 'static, - T: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static, - B: Body + Debug + Send + Unpin + 'static, - ::Data: Send + Sync + 'static, - ::Error: Error + Send + Sync + 'static, -{ - let http_log_builder = HttpLog::builder() - .uri(request.uri().path().to_string()) - .method(request.method().to_string()) - .host(key.0.clone()) - .ip(key.1.to_string()); - // Ensure best-effort compatibility of proxy request with HTTP/1.1 format - // -> Add host header if missing - request - .headers_mut() - .insert(HOST, request_host.try_into().expect("valid host")); - // -> Change URI to only include path and query - *request.uri_mut() = request - .uri() - .path_and_query() - .map(|path| Uri::from_str(path.as_str()).expect("valid URI")) - .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().expect("valid header value")); - } - // Check for an Upgrade header - if let Some(request_upgrade) = request.headers().get(UPGRADE) { - // Get HTTP/1.1 sender and remote log channel with upgrades with a new connection - let io = match proxy_data.proxy_type { - ProxyType::Tunneling => { - handler - .tunneling_channel(tcp_address.ip(), tcp_address.port()) - .await - } - ProxyType::Aliasing => { - handler - .aliasing_channel(tcp_address.ip(), tcp_address.port(), key.2.as_ref()) - .await - } - }; - let (mut sender, conn) = hyper::client::conn::http1::handshake(TokioIo::new(io?)).await?; - tokio::spawn(Box::pin(async move { - if let Err(error) = conn.with_upgrades().await { - #[cfg(not(coverage_nightly))] - tracing::warn!(%error, "HTTP/1.1 connection failed."); - } - })); - let tx = handler.log_channel(); - - // If there is an Upgrade header, make sure that it's a valid Websocket upgrade - 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? - } else { - let elapsed_time = timer.elapsed(); - http_log( - http_log_builder - .status(StatusCode::GATEWAY_TIMEOUT.as_u16()) - .elapsed_time(elapsed_time) - .build(), - Some(tx), - proxy_data.disable_http_logs, - ); - return Err(HttpError::RequestTimeout); - } - } - None => sender.send_request(request).await?, - }; - // Check if the underlying server accepts the Upgrade request - match response.status() { - StatusCode::SWITCHING_PROTOCOLS => { - if request_type - == response - .headers() - .get(UPGRADE) - .ok_or(HttpError::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; - let buffer_size = proxy_data.buffer_size; - // Start a task to copy data between the two Upgraded parts - tokio::spawn(async move { - let mut upgraded_request = - TokioIo::new(upgraded_request.await.expect("upgradable request")); - 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_with_sizes( - &mut upgraded_response, - &mut upgraded_request, - buffer_size, - buffer_size, - ) - .await - }) - .await; - } - // If there isn't a Websocket timeout, copy data between both sides unconditionally. - None => { - let _ = copy_bidirectional_with_sizes( - &mut upgraded_response, - &mut upgraded_request, - buffer_size, - buffer_size, - ) - .await; - } - } - }); - } - // Return the response to the client - let http_log_builder = http_log_builder.status(response.status().as_u16()); - return Ok(ProxyResponse::Proxy(TimedResponse { - response, - on_drop: Some(Box::new(move || { - http_log( - http_log_builder.elapsed_time(timer.elapsed()).build(), - Some(tx), - proxy_data.disable_http_logs, - ) - })), - })); - } - _ => { - let http_log_builder = http_log_builder.status(response.status().as_u16()); - return Ok(ProxyResponse::Proxy(TimedResponse { - response, - on_drop: Some(Box::new(move || { - http_log( - http_log_builder.elapsed_time(timer.elapsed()).build(), - Some(tx), - proxy_data.disable_http_logs, - ) - })), - })); - } - } - } else { - loop { - // If Upgrade header is not present, get HTTP/1.1 sender and remote log channel - let (mut sender, tx, _guard) = if proxy_data.has_pool_queue - && let Some(guard) = proxy_data.get_http11_pool_guard(key.clone()) - { - // Race between new channel and HTTP pool - let mut recv = guard.pool.recv(); - let mut channel_future = pin!(async { - match proxy_data.proxy_type { - ProxyType::Tunneling => { - handler - .tunneling_channel(tcp_address.ip(), tcp_address.port()) - .await - } - ProxyType::Aliasing => { - handler - .aliasing_channel( - tcp_address.ip(), - tcp_address.port(), - key.2.as_ref(), - ) - .await - } - } - }); - loop { - tokio::select! { - result = recv => { - match result { - // Return pool item - Ok(tuple) if tuple.0.is_ready() => break (tuple.0, tuple.1, guard), - // Connection is closed; discard pool item - Ok(_) => { - recv = guard.pool.recv(); - continue; - }, - Err(_) => return Err(HttpError::PoolClosed), - } - } - result = &mut channel_future => { - let (sender, conn) = hyper::client::conn::http1::handshake( - TokioIo::new(result?), - ) - .await?; - tokio::spawn(Box::pin(async move { - if let Err(error) = conn.await { - #[cfg(not(coverage_nightly))] - tracing::warn!(%error, "HTTP/1.1 connection failed."); - } - })); - break (sender, handler.log_channel(), guard); - } - } - } - } else { - // No pool timeout - get recycled sender or create new one - 'sender: loop { - match proxy_data.get_http11_pool_guard(key.clone()) { - Some(guard) => { - while let Ok(sender) = guard.pool.try_recv() { - if sender.0.is_ready() { - break 'sender (sender.0, sender.1, guard); - } - } - } - None => { - let io = match proxy_data.proxy_type { - ProxyType::Tunneling => { - handler - .tunneling_channel(tcp_address.ip(), tcp_address.port()) - .await - } - ProxyType::Aliasing => { - handler - .aliasing_channel( - tcp_address.ip(), - tcp_address.port(), - key.2.as_ref(), - ) - .await - } - }; - let (sender, conn) = - hyper::client::conn::http1::handshake(TokioIo::new(io?)).await?; - tokio::spawn(Box::pin(async move { - if let Err(error) = conn.await { - #[cfg(not(coverage_nightly))] - tracing::warn!(%error, "HTTP/1.1 connection failed."); - } - })); - let guard = proxy_data.create_http11_pool_guard(key.clone()); - break (sender, handler.log_channel(), guard); - } - } - } - }; - - // Create entry for pool - let key_clone = key.clone(); - let pool = { - let pool_ref = proxy_data - .keepalive_http11_pool_map - .entry(key_clone) - .or_insert_with(|| Arc::new(async_channel::unbounded())) - .downgrade(); - pool_ref.0.clone() - }; - - 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.try_send_request(request)).await - { - response - } else { - let elapsed_time = timer.elapsed(); - http_log( - http_log_builder - .status(StatusCode::GATEWAY_TIMEOUT.as_u16()) - .elapsed_time(elapsed_time) - .build(), - Some(tx), - proxy_data.disable_http_logs, - ); - return Err(HttpError::RequestTimeout); - } - } - None => sender.try_send_request(request).await, - }; - let response = match response { - Ok(response) => response, - Err(mut error) => { - if let Some(recovered) = error.take_message() { - #[cfg(not(coverage_nightly))] - tracing::debug!(error = %error.error(), "Recovering HTTP/1.1 request to try again."); - request = recovered; - continue; - } else { - let elapsed_time = timer.elapsed(); - http_log( - http_log_builder - .status(StatusCode::INTERNAL_SERVER_ERROR.as_u16()) - .elapsed_time(elapsed_time) - .build(), - Some(tx), - proxy_data.disable_http_logs, - ); - return Err(error.into_error().into()); - } - } - }; - - // Return the received response to the client - let http_log_builder = http_log_builder.status(response.status().as_u16()); - return Ok(ProxyResponse::Proxy(TimedResponse { - response, - on_drop: Some(Box::new(move || { - // Log HTTP request - http_log( - http_log_builder.elapsed_time(timer.elapsed()).build(), - Some(tx.clone()), - proxy_data.disable_http_logs, - ); - // Return sender to pool - tokio::spawn(async move { - let _ = pool.send((sender, tx)).await; - }); - })), - })); - } - } -} - -#[cfg_attr( - not(coverage_nightly), - tracing::instrument(skip(proxy_data, handler), level = "debug") -)] -pub(crate) async fn http2_handler( - request: Request, - tcp_address: SocketAddr, - proxy_data: Arc>, - handler: Arc, - host: String, -) -> color_eyre::Result> -where - M: ConnectionGetByHttpHost> + Send + Sync + 'static, - H: ConnectionHandler + Send + Sync + 'static, - T: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static, - B: Body + Debug + Send + Unpin + 'static, - ::Data: Send + Sync + 'static, - ::Error: Error + Send + Sync + 'static, -{ - match http2_handler_inner(request, tcp_address, proxy_data, handler, host).await { - Ok(response) => Ok(match response { - ProxyResponse::Axum(response) => response, - ProxyResponse::Proxy(response) => response.into_response(), - }), - Err(error) => Ok(error.into_response()), - } -} - -#[inline] -async fn http2_handler_inner( - request: Request, - tcp_address: SocketAddr, - proxy_data: Arc>, - handler: Arc, - host: String, -) -> Result -where - M: ConnectionGetByHttpHost> + Send + Sync + 'static, - H: ConnectionHandler + Send + Sync + 'static, - T: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static, - B: Body + Debug + Send + Unpin + 'static, - ::Data: Send + Sync + 'static, - ::Error: Error + Send + Sync + 'static, -{ - let timer = Instant::now(); - let Some(host_uri) = request.uri().host() else { - return Err(HttpError::MissingHostHeader); - }; - if host != host_uri { - return Err(HttpError::MismatchedHostHeader); - }; - let ip = tcp_address.ip().to_canonical(); - let host_clone = host.clone(); - let http_data = handler.http_data(); - let request_host = http_data - .as_ref() - .and_then(|data| data.host.as_deref()) - .unwrap_or(host_clone.as_str()); - - let key = KeepaliveAlias(host.clone(), ip, None); - handle_http2_request( - request, - tcp_address, - proxy_data, - handler, - request_host, - key, - timer, - ) - .await -} - -async fn handle_http2_request( - mut request: Request, - tcp_address: SocketAddr, - proxy_data: Arc>, - handler: Arc, - request_host: &str, - key: KeepaliveAlias, - timer: Instant, -) -> Result -where - M: ConnectionGetByHttpHost> + Send + Sync + 'static, - H: ConnectionHandler + Send + Sync + 'static, - T: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static, - B: Body + Debug + Send + Unpin + 'static, - ::Data: Send + Sync + 'static, - ::Error: Error + Send + Sync + 'static, -{ - let http_log_builder = HttpLog::builder() - .uri(request.uri().path().to_string()) - .method(request.method().to_string()) - .host(key.0.clone()) - .ip(key.1.to_string()); - - loop { - // Get HTTP/2 sender and remote log channel - let (mut sender, tx, _guard) = if proxy_data.has_pool_queue - && let Some(guard) = proxy_data.get_http2_pool_guard(key.clone()) - { - // Race between new channel and HTTP pool - let mut recv = guard.pool.recv(); - let mut channel_future = pin!(async { - match proxy_data.proxy_type { - ProxyType::Tunneling => { - handler - .tunneling_channel(tcp_address.ip(), tcp_address.port()) - .await - } - ProxyType::Aliasing => { - handler - .aliasing_channel(tcp_address.ip(), tcp_address.port(), key.2.as_ref()) - .await - } - } - }); - loop { - tokio::select! { - result = recv => { - match result { - // Return pool item - Ok(tuple) if tuple.0.is_ready() => break (tuple.0, tuple.1, guard), - // Connection is closed; discard pool item - Ok(_) => { - recv = guard.pool.recv(); - continue; - }, - Err(_) => return Err(HttpError::PoolClosed), - } - } - result = &mut channel_future => { - let (sender, conn) = hyper::client::conn::http2::handshake( - TokioExecutor::new(), - TokioIo::new(result?), - ) - .await?; - tokio::spawn(Box::pin(async move { - if let Err(error) = conn.await { - #[cfg(not(coverage_nightly))] - tracing::warn!(%error, "HTTP/2 connection failed."); - } - })); - break (sender, handler.log_channel(), guard); - } - } - } - } else { - // No pool timeout - get recycled sender or create new one - 'sender: loop { - match proxy_data.get_http2_pool_guard(key.clone()) { - Some(guard) => { - while let Ok(sender) = guard.pool.try_recv() { - if sender.0.is_ready() { - break 'sender (sender.0, sender.1, guard); - } - } - } - None => { - let io = match proxy_data.proxy_type { - ProxyType::Tunneling => { - handler - .tunneling_channel(tcp_address.ip(), tcp_address.port()) - .await - } - ProxyType::Aliasing => { - handler - .aliasing_channel( - tcp_address.ip(), - tcp_address.port(), - key.2.as_ref(), - ) - .await - } - }; - let (sender, conn) = hyper::client::conn::http2::handshake( - TokioExecutor::new(), - TokioIo::new(io?), - ) - .await?; - tokio::spawn(Box::pin(async move { - if let Err(error) = conn.await { - #[cfg(not(coverage_nightly))] - tracing::warn!(%error, "HTTP/2 connection failed."); - } - })); - let guard = proxy_data.create_http2_pool_guard(key.clone()); - break (sender, handler.log_channel(), guard); - } - } - } - }; - - // Create entry for pool - let key_clone = key.clone(); - let pool = { - let pool_ref = proxy_data - .keepalive_http2_pool_map - .entry(key_clone) - .or_insert_with(|| Arc::new(async_channel::unbounded())) - .downgrade(); - pool_ref.0.clone() - }; - - // Create an HTTP/2 handshake over the selected channel - let mut uri_parts = request.uri().clone().into_parts(); - let authority = uri_parts.authority.as_mut().expect("Host has been checked"); - *authority = Authority::from_maybe_shared( - authority - .as_str() - .replace(authority.host(), request_host) - .into_bytes(), - )?; - *request.uri_mut() = Uri::from_parts(uri_parts)?; - request.headers_mut().insert( - CONNECTION, - "keepalive".try_into().expect("valid HeaderValue"), - ); - - 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.try_send_request(request)).await { - response - } else { - let elapsed_time = timer.elapsed(); - http_log( - http_log_builder - .status(StatusCode::GATEWAY_TIMEOUT.as_u16()) - .elapsed_time(elapsed_time) - .build(), - Some(tx), - proxy_data.disable_http_logs, - ); - return Err(HttpError::RequestTimeout); - } - } - None => sender.try_send_request(request).await, - }; - let response = match response { - Ok(response) => response, - Err(mut error) => { - if let Some(recovered) = error.take_message() { - #[cfg(not(coverage_nightly))] - tracing::debug!(error = %error.error(), "Recovering HTTP/2 request to try again."); - request = recovered; - continue; - } else { - let elapsed_time = timer.elapsed(); - http_log( - http_log_builder - .status(StatusCode::INTERNAL_SERVER_ERROR.as_u16()) - .elapsed_time(elapsed_time) - .build(), - Some(tx), - proxy_data.disable_http_logs, - ); - return Err(error.into_error().into()); - } - } - }; - - let http_log_builder = http_log_builder.status(response.status().as_u16()); - return Ok(ProxyResponse::Proxy(TimedResponse { - response, - on_drop: Some(Box::new(move || { - // Log HTTP request - http_log( - http_log_builder.elapsed_time(timer.elapsed()).build(), - Some(tx.clone()), - proxy_data.disable_http_logs, - ); - // Send sender to pool - tokio::spawn(async move { - let _ = pool.send((sender, tx)).await; - }); - })), - })); + version => Err(HttpError::InvalidHttpVersion(version)), } } diff --git a/src/keepalive.rs b/src/keepalive.rs index 45b2822..1a89926 100644 --- a/src/keepalive.rs +++ b/src/keepalive.rs @@ -1,7 +1,6 @@ // Shoutout to https://github.com/sunshowers-code/borrow-complex-key-example/blob/main/src/lib.rs use std::{ - borrow::Borrow, hash::{Hash, Hasher}, net::IpAddr, }; @@ -26,64 +25,3 @@ impl Hash for KeepaliveAlias { .hash(state); } } - -// A borrowed TCP alias, with references to an address and a port. Useful for accessing the TCP connection map. -#[derive(Copy, Clone, Debug, Eq, Ord, PartialEq, PartialOrd)] -pub(crate) struct BorrowedKeepaliveAlias<'a>( - pub(crate) &'a str, - pub(crate) &'a IpAddr, - pub(crate) &'a Option, -); - -impl BorrowedKeepaliveAlias<'_> { - pub(crate) fn as_owned(&self) -> KeepaliveAlias { - KeepaliveAlias(self.0.to_string(), *self.1, *self.2) - } -} - -impl Hash for BorrowedKeepaliveAlias<'_> { - fn hash(&self, state: &mut H) { - self.0.hash(state); - self.1.hash(state); - self.2 - .as_ref() - .map(|fingerprint| fingerprint.as_bytes()) - .hash(state); - } -} - -impl<'a> Borrow for KeepaliveAlias { - fn borrow(&self) -> &(dyn KeepaliveAliasKey + 'a) { - self - } -} - -pub(crate) trait KeepaliveAliasKey { - fn key(&self) -> BorrowedKeepaliveAlias<'_>; -} - -impl KeepaliveAliasKey for KeepaliveAlias { - fn key(&self) -> BorrowedKeepaliveAlias<'_> { - BorrowedKeepaliveAlias(self.0.as_str(), &self.1, &self.2) - } -} - -impl KeepaliveAliasKey for BorrowedKeepaliveAlias<'_> { - fn key(&self) -> BorrowedKeepaliveAlias<'_> { - *self - } -} - -impl PartialEq for dyn KeepaliveAliasKey + '_ { - fn eq(&self, other: &Self) -> bool { - self.key().eq(&other.key()) - } -} - -impl Eq for dyn KeepaliveAliasKey + '_ {} - -impl Hash for dyn KeepaliveAliasKey + '_ { - fn hash(&self, state: &mut H) { - self.key().hash(state) - } -} diff --git a/src/ssh/connection_handler.rs b/src/ssh/connection_handler.rs index b2d9a88..8b845e6 100644 --- a/src/ssh/connection_handler.rs +++ b/src/ssh/connection_handler.rs @@ -246,7 +246,7 @@ impl SshTunnelHandler { .ip_connections .entry(ip) .or_insert(Arc::new(Semaphore::new(self.max_connections_per_ip))); - Arc::clone(&entry.value()) + Arc::clone(entry.value()) }; let Ok(_permit) = semaphore.try_acquire_owned() else { return Err(ServerError::IpConnectionLimitReached); diff --git a/src/ssh/forwarding.rs b/src/ssh/forwarding.rs index 08618b3..d6245bc 100644 --- a/src/ssh/forwarding.rs +++ b/src/ssh/forwarding.rs @@ -126,14 +126,21 @@ impl Forwarder { context: &mut RemoteForwardingContext<'_>, address: &str, port: u16, + ssh_port: u16, + http_port: Option, + https_port: Option, ) -> Result { match port { - 22 => { + port if port == 22 || port == ssh_port => { SshForwardingHandler .cancel_remote_forwarding(context, address, port) .await } - 80 | 443 => { + port if port == 80 + || port == 443 + || http_port.is_some_and(|http_port| port == http_port) + || https_port.is_some_and(|https_port| port == https_port) => + { HttpForwardingHandler .cancel_remote_forwarding(context, address, port) .await @@ -159,7 +166,7 @@ impl Forwarder { originator_port: u16, channel: Channel, ) -> Result { - if port == context.server.ssh_port { + if port == 22 || port == context.server.ssh_port { SshForwardingHandler .local_forwarding( context, @@ -170,7 +177,11 @@ impl Forwarder { channel, ) .await - } else if port == context.server.http_port || port == context.server.https_port { + } else if port == 80 + || port == 443 + || port == context.server.http_port + || port == context.server.https_port + { HttpForwardingHandler .local_forwarding( context, diff --git a/src/ssh/mod.rs b/src/ssh/mod.rs index 6f1fa0d..178cdc5 100644 --- a/src/ssh/mod.rs +++ b/src/ssh/mod.rs @@ -887,6 +887,17 @@ impl Handler for ServerHandler { | AuthenticatedData::Admin { user_data, .. } => user_data, AuthenticatedData::None { .. } => return Err(russh::Error::Disconnect), }; + let ssh_port = self.server.ssh_port; + let http_port = if self.server.disable_http { + None + } else { + Some(self.server.http_port) + }; + let https_port = if self.server.disable_https { + None + } else { + Some(self.server.https_port) + }; Forwarder::cancel_remote_forwarding( &mut RemoteForwardingContext { server: &mut self.server, @@ -898,6 +909,9 @@ impl Handler for ServerHandler { }, address.trim(), port as u16, + ssh_port, + http_port, + https_port, ) .await } diff --git a/tests/integration/alias_http2.rs b/tests/integration/alias_http2.rs new file mode 100644 index 0000000..94906d9 --- /dev/null +++ b/tests/integration/alias_http2.rs @@ -0,0 +1,259 @@ +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 sandhole::{ApplicationConfig, entrypoint}; +use tokio::{ + net::TcpStream, + sync::oneshot, + time::{sleep, timeout}, +}; +use tower::Service; + +use crate::common::SandholeHandle; + +/// This test ensures that local forwarding works for alias-only HTTP/2 services. +#[test_log::test(tokio::test(flavor = "multi_thread"))] +async fn alias_http2() { + // 1. Initialize Sandhole + let config = ApplicationConfig::parse_from([ + "sandhole", + "--domain=foobar.tld", + "--user-keys-directory", + &(format!( + "{}/tests/data/user_keys", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--admin-keys-directory", + &(format!( + "{}/tests/data/admin_keys", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--certificates-directory", + &(format!( + "{}/tests/data/certificates", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--private-key-file", + &(format!( + "{}/tests/data/server_keys/ssh", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--acme-cache-directory", + &(format!( + "{}/tests/data/acme_cache", + std::env::var("CARGO_MANIFEST_DIR").unwrap() + )), + "--disable-directory-creation", + "--listen-address=127.0.0.1", + "--ssh-port=18022", + "--http-port=18080", + "--https-port=18443", + "--acme-use-staging", + "--bind-hostnames=all", + "--idle-connection-timeout=1s", + "--authentication-request-timeout=5s", + "--http-request-timeout=5s", + ]); + let _sandhole_handle = SandholeHandle(tokio::spawn(async move { entrypoint(config).await })); + if timeout(Duration::from_secs(5), async { + while TcpStream::connect("127.0.0.1:18022").await.is_err() { + sleep(Duration::from_millis(100)).await; + } + }) + .await + .is_err() + { + panic!("Timeout waiting for Sandhole to start.") + }; + + // 2. Start SSH client that will be proxied via HTTP/2 + let key = load_secret_key( + std::path::PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()) + .join("tests/data/private_keys/key1"), + None, + ) + .expect("Missing file key1"); + let (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 alias of our proxy + let key = load_secret_key( + std::path::PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()) + .join("tests/data/private_keys/key1"), + None, + ) + .expect("Missing file key1"); + let ssh_client = SshAliasClient; + 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 channel = session + .channel_open_direct_tcpip("http2.foobar.tld", 80, "my.hostname", 12345) + .await + .expect("channel_open_direct_tcpip failed"); + let (mut sender, conn) = hyper::client::conn::http2::handshake( + TokioExecutor::new(), + TokioIo::new(channel.into_stream()), + ) + .await + .expect("HTTP handshake failed"); + let jh = tokio::spawn(async move { + if let Err(error) = conn.await { + eprintln!("Connection failed: {error:?}"); + } + }); + 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!"); + jh.abort(); +} + +struct SshClient(Option>); + +impl russh::client::Handler for SshClient { + type Error = color_eyre::eyre::Error; + + async fn check_server_key( + &mut self, + _key: &russh::keys::PublicKey, + ) -> Result { + Ok(true) + } + + async fn server_channel_open_forwarded_tcpip( + &mut self, + 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(()) + } +} + +struct SshAliasClient; + +impl russh::client::Handler for SshAliasClient { + type Error = color_eyre::eyre::Error; + + async fn check_server_key( + &mut self, + _key: &russh::keys::PublicKey, + ) -> Result { + Ok(true) + } +} diff --git a/tests/integration/main.rs b/tests/integration/main.rs index d278cd7..7b080a4 100644 --- a/tests/integration/main.rs +++ b/tests/integration/main.rs @@ -10,6 +10,7 @@ mod admin_restricted_aliases; mod admin_window_change; mod alias_aliasing_tunnel; mod alias_cannot_be_localhost; +mod alias_http2; mod alias_http_aliases; mod alias_ip_connections_limit; mod alias_local_forward_existing_http;