diff --git a/src/dev/tls.rs b/src/dev/tls.rs index d73945b..536a670 100644 --- a/src/dev/tls.rs +++ b/src/dev/tls.rs @@ -10,6 +10,7 @@ use bytes::Bytes; use http_body_util::{combinators::BoxBody, BodyExt, Full}; use hyper::body::Incoming; use hyper::header::{HeaderValue, CONNECTION, HOST, UPGRADE}; +use hyper::http::uri::Authority; use hyper::server::conn::http1::Builder as ServerBuilder; use hyper::service::service_fn; use hyper::{Request, Response, StatusCode, Uri}; @@ -177,18 +178,18 @@ async fn proxy_request( mut request: Request, routes: Arc>, ) -> std::result::Result, Infallible> { - let hostname = request + let authority = request .headers() .get(HOST) .and_then(|value| value.to_str().ok()) - .and_then(hostname_without_port) - .map(str::to_ascii_lowercase); - let Some(hostname) = hostname else { + .and_then(parse_host_authority); + let Some(authority) = authority else { return Ok(text_response( StatusCode::BAD_REQUEST, - "missing Host header\n", + "missing or invalid Host header\n", )); }; + let hostname = authority.host().to_ascii_lowercase(); let Some(port) = match_route(&hostname, &routes) else { return Ok(text_response( StatusCode::MISDIRECTED_REQUEST, @@ -225,9 +226,11 @@ async fn proxy_request( request .headers_mut() .insert("x-forwarded-proto", HeaderValue::from_static("https")); - if let Ok(value) = HeaderValue::from_str(&hostname) { - request.headers_mut().insert("x-forwarded-host", value); - } + let forwarded_host = HeaderValue::from_str(authority.as_str()) + .expect("parsed HTTP authority is a valid header value"); + request + .headers_mut() + .insert("x-forwarded-host", forwarded_host); let client: Client = Client::builder(TokioExecutor::new()).build_http(); @@ -270,12 +273,20 @@ fn text_response(status: StatusCode, message: &str) -> Response { Response::builder().status(status).body(body).unwrap() } -fn hostname_without_port(host: &str) -> Option<&str> { - let host = host.trim(); - if host.is_empty() { +fn parse_host_authority(value: &str) -> Option { + if value.contains('@') || value.matches(':').count() > 1 { + return None; + } + if let Some((hostname, port)) = value.rsplit_once(':') { + if hostname.is_empty() || port.parse::().ok().filter(|port| *port > 0).is_none() { + return None; + } + } + let authority = value.parse::().ok()?; + if authority.host().is_empty() { return None; } - Some(host.split_once(':').map_or(host, |(hostname, _)| hostname)) + Some(authority) } #[derive(Clone)] @@ -427,6 +438,15 @@ mod tests { assert_eq!(match_route("sites.test", &routes), None); } + #[test] + fn host_authority_accepts_numeric_ports_and_rejects_userinfo() { + let authority = parse_host_authority("intern.dev:8443").unwrap(); + assert_eq!(authority.host(), "intern.dev"); + assert_eq!(authority.port_u16(), Some(8443)); + assert!(parse_host_authority("intern.dev:443@evil.example").is_none()); + assert!(parse_host_authority("intern.dev:not-a-port").is_none()); + } + #[test] fn dns_validation_probes_apex_open_host_and_a_wildcard_site_hostname() { let tls = DevTlsProxyConfig { diff --git a/tests/dev_tls.rs b/tests/dev_tls.rs index 0d80895..92ea8ce 100644 --- a/tests/dev_tls.rs +++ b/tests/dev_tls.rs @@ -252,10 +252,13 @@ tls_proxy = {{ certificate_hosts = ["intern.dev", "*.local.sites.intern.dev"], o // Observable behavior: exact and suffix hosts cross real TLS and HTTP // sockets, preserve logical Host, and select different upstream services. - let front_response = https_get(tls_port, "intern.dev", "intern.dev", cert_der.clone()); + let front_authority = format!("intern.dev:{tls_port}"); + let front_response = https_get(tls_port, "intern.dev", &front_authority, cert_der.clone()); assert!(front_response.contains("200 OK"), "{front_response}"); assert!( - front_response.contains("frontend host=intern.dev proto=https"), + front_response.contains(&format!( + "frontend host={front_authority} forwarded={front_authority} proto=https" + )), "{front_response}" ); let site_response = https_get( @@ -265,11 +268,24 @@ tls_proxy = {{ certificate_hosts = ["intern.dev", "*.local.sites.intern.dev"], o cert_der.clone(), ); assert!( - site_response.contains("site-gateway host=test.local.sites.intern.dev proto=https"), + site_response.contains( + "site-gateway host=test.local.sites.intern.dev forwarded=test.local.sites.intern.dev proto=https" + ), "{site_response}" ); let unknown = https_get(tls_port, "intern.dev", "unknown.test", cert_der); assert!(unknown.contains("421 Misdirected Request"), "{unknown}"); + let malformed = https_get( + tls_port, + "intern.dev", + "intern.dev:443@evil.example", + CertificateDer::from(certified.cert.der().to_vec()), + ); + assert!(malformed.contains("400 Bad Request"), "{malformed}"); + assert!( + malformed.contains("missing or invalid Host header"), + "{malformed}" + ); assert_tls_upgrade_echo(tls_port, "intern.dev", certified.cert.der().to_vec()); // Lifecycle outcome: authenticated control shutdown stops the group and @@ -299,6 +315,7 @@ fn serve_http(listener: TcpListener, label: &'static str) -> thread::JoinHandle< for stream in listener.incoming().flatten() { let mut reader = BufReader::new(stream.try_clone().unwrap()); let mut host = String::new(); + let mut forwarded_host = String::new(); let mut proto = String::new(); let mut upgrade = false; loop { @@ -316,6 +333,9 @@ fn serve_http(listener: TcpListener, label: &'static str) -> thread::JoinHandle< if lower.starts_with("x-forwarded-proto:") { proto = line[18..].trim().to_string(); } + if lower.starts_with("x-forwarded-host:") { + forwarded_host = line[17..].trim().to_string(); + } } if upgrade { let mut stream = stream; @@ -330,7 +350,7 @@ fn serve_http(listener: TcpListener, label: &'static str) -> thread::JoinHandle< stream.write_all(&payload).unwrap(); continue; } - let body = format!("{label} host={host} proto={proto}"); + let body = format!("{label} host={host} forwarded={forwarded_host} proto={proto}"); let mut stream = stream; write!( stream,