From ce26dfee8d37555cfef2b0f46ee7741c7e59f66a Mon Sep 17 00:00:00 2001 From: sasicodes Date: Sat, 6 Jun 2026 23:41:17 +0530 Subject: [PATCH] Validate protected websocket origins --- crates/peek-relay/src/handler.rs | 105 +++++++++++++++++++++++- crates/peek-relay/tests/integration.rs | 107 +++++++++++++++++++++++++ 2 files changed, 211 insertions(+), 1 deletion(-) diff --git a/crates/peek-relay/src/handler.rs b/crates/peek-relay/src/handler.rs index 77f6626..e6cf92b 100644 --- a/crates/peek-relay/src/handler.rs +++ b/crates/peek-relay/src/handler.rs @@ -8,7 +8,7 @@ use axum::{ ConnectInfo, Request, State, ws::{Message, WebSocket, WebSocketUpgrade, rejection::WebSocketUpgradeRejection}, }, - http::HeaderMap, + http::{HeaderMap, Uri, uri::Authority}, response::Response, }; use futures_util::{SinkExt, StreamExt}; @@ -287,6 +287,7 @@ pub async fn public_handler( return not_found_page(); }; + let password_protected = conn.password.is_some(); if let Some(ref tunnel_password) = conn.password { request = match password_gate(request, &subdomain, tunnel_password).await { Ok(request) => request, @@ -295,6 +296,9 @@ pub async fn public_handler( } if let Ok(ws) = ws { + if password_protected && !websocket_origin_matches_host(request.headers(), &host) { + return forbidden_page(); + } return handle_public_ws_upgrade(ws, request, conn, registry.max_body_size).await; } @@ -434,6 +438,70 @@ fn requested_ws_protocols(headers: &HeaderMap) -> Vec { .unwrap_or_default() } +fn websocket_origin_matches_host(headers: &HeaderMap, request_host: &str) -> bool { + let mut origins = headers.get_all("origin").iter(); + let Some(origin) = origins.next() else { + return true; + }; + if origins.next().is_some() { + return false; + } + + let Ok(origin) = origin.to_str() else { + return false; + }; + origin_matches_host(origin, request_host) +} + +fn origin_matches_host(origin: &str, request_host: &str) -> bool { + let Ok(origin_uri) = origin.parse::() else { + return false; + }; + let Some(scheme) = origin_uri.scheme_str() else { + return false; + }; + if !scheme.eq_ignore_ascii_case("http") && !scheme.eq_ignore_ascii_case("https") { + return false; + } + let Some(origin_authority) = origin_uri.authority() else { + return false; + }; + let Ok(request_authority) = request_host.parse::() else { + return false; + }; + + if !origin_authority + .host() + .eq_ignore_ascii_case(request_authority.host()) + { + return false; + } + + match request_authority.port_u16() { + Some(request_port) => origin_effective_port(scheme, origin_authority) == Some(request_port), + None => match origin_authority.port_u16() { + Some(port) => Some(port) == default_port_for_scheme(scheme), + None => true, + }, + } +} + +fn origin_effective_port(scheme: &str, authority: &Authority) -> Option { + authority + .port_u16() + .or_else(|| default_port_for_scheme(scheme)) +} + +fn default_port_for_scheme(scheme: &str) -> Option { + if scheme.eq_ignore_ascii_case("http") { + Some(80) + } else if scheme.eq_ignore_ascii_case("https") { + Some(443) + } else { + None + } +} + async fn bridge_public_ws( socket: WebSocket, conn: Arc, @@ -670,6 +738,12 @@ fn not_found_page() -> Response { resp } +fn forbidden_page() -> Response { + let mut resp = status_page("Forbidden"); + *resp.status_mut() = axum::http::StatusCode::FORBIDDEN; + resp +} + fn bad_gateway_page() -> Response { let mut resp = status_page("Bad gateway"); *resp.status_mut() = axum::http::StatusCode::BAD_GATEWAY; @@ -754,6 +828,35 @@ mod tests { ); } + #[test] + fn test_origin_matches_host() { + assert!(origin_matches_host( + "https://a8f3k2.example.com", + "a8f3k2.example.com" + )); + assert!(origin_matches_host( + "https://a8f3k2.example.com:443", + "a8f3k2.example.com" + )); + assert!(origin_matches_host( + "http://a8f3k2.example.com:8080", + "a8f3k2.example.com:8080" + )); + assert!(origin_matches_host( + "https://A8F3K2.example.com", + "a8f3k2.example.com" + )); + assert!(!origin_matches_host( + "https://attacker.example.com", + "a8f3k2.example.com" + )); + assert!(!origin_matches_host( + "https://a8f3k2.example.com:444", + "a8f3k2.example.com" + )); + assert!(!origin_matches_host("null", "a8f3k2.example.com")); + } + #[test] fn test_normalize_subdomain() { assert_eq!(normalize_subdomain("My-App"), Some("my-app".into())); diff --git a/crates/peek-relay/tests/integration.rs b/crates/peek-relay/tests/integration.rs index e243766..7d8b2be 100644 --- a/crates/peek-relay/tests/integration.rs +++ b/crates/peek-relay/tests/integration.rs @@ -137,6 +137,31 @@ async fn start_local_server() -> (SocketAddr, &'static str) { (addr, expected_body) } +async fn login_cookie(relay_addr: SocketAddr, host: &str, password: &str) -> String { + let http_client = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .unwrap(); + let resp = http_client + .post(format!("http://{relay_addr}/__peek_auth")) + .header("host", host) + .header("content-type", "application/x-www-form-urlencoded") + .body(format!("password={password}")) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 303); + resp.headers() + .get("set-cookie") + .and_then(|value| value.to_str().ok()) + .unwrap() + .split(';') + .next() + .unwrap() + .to_string() +} + #[tokio::test] async fn test_tunnel_end_to_end() { let (local_addr, expected_body) = start_local_server().await; @@ -296,6 +321,88 @@ async fn test_websocket_tunnel_forwards_headers() { handle.close().await; } +#[tokio::test] +async fn test_password_protected_websocket_rejects_cross_origin() { + let (local_addr, _) = start_local_server().await; + let relay_addr = start_relay("test-ws-origin.local").await; + let host = "victim.test-ws-origin.local"; + + let client = peek_client::TunnelClient::new(&format!("ws://{relay_addr}/tunnel")) + .unwrap() + .with_password("s3cret".into()); + let handle = client + .connect_with_subdomain(local_addr.port(), Some("victim".into())) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(100)).await; + + let cookie = login_cookie(relay_addr, host, "s3cret").await; + let mut request = format!("ws://{relay_addr}/ws") + .into_client_request() + .unwrap(); + request.headers_mut().insert("host", host.parse().unwrap()); + request.headers_mut().insert( + "origin", + "https://attacker.test-ws-origin.local".parse().unwrap(), + ); + request + .headers_mut() + .insert("cookie", cookie.parse().unwrap()); + + let err = match tokio_tungstenite::connect_async(request).await { + Ok(_) => panic!("expected HTTP 403 handshake failure"), + Err(err) => err, + }; + match err { + tokio_tungstenite::tungstenite::Error::Http(response) => { + assert_eq!(response.status().as_u16(), 403); + } + other => panic!("expected HTTP 403 handshake failure, got {other}"), + } + + handle.close().await; +} + +#[tokio::test] +async fn test_password_protected_websocket_allows_same_origin() { + let (local_addr, _) = start_local_server().await; + let relay_addr = start_relay("test-ws-same-origin.local").await; + let host = "victim.test-ws-same-origin.local"; + + let client = peek_client::TunnelClient::new(&format!("ws://{relay_addr}/tunnel")) + .unwrap() + .with_password("s3cret".into()); + let handle = client + .connect_with_subdomain(local_addr.port(), Some("victim".into())) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(100)).await; + + let cookie = login_cookie(relay_addr, host, "s3cret").await; + let mut request = format!("ws://{relay_addr}/ws") + .into_client_request() + .unwrap(); + request.headers_mut().insert("host", host.parse().unwrap()); + request + .headers_mut() + .insert("origin", format!("https://{host}").parse().unwrap()); + request + .headers_mut() + .insert("cookie", cookie.parse().unwrap()); + + let (mut socket, _) = tokio_tungstenite::connect_async(request).await.unwrap(); + socket + .send(tokio_tungstenite::tungstenite::Message::Text( + "hello".into(), + )) + .await + .unwrap(); + let msg = socket.next().await.unwrap().unwrap(); + assert_eq!(msg.into_text().unwrap(), "hello"); + + handle.close().await; +} + #[tokio::test] async fn test_tunnel_not_found() { let relay_addr = start_relay("test2.local").await;