From cc9c4d26a9eaebd7333c1d588037c0225e6746b2 Mon Sep 17 00:00:00 2001 From: DanConwayDev Date: Thu, 23 Jul 2026 06:05:02 +0100 Subject: [PATCH 1/2] Add HTTP headers to WebSocket connections --- src/lib.rs | 13 +++++++++ src/native/mod.rs | 69 ++++++++++++++++++++++++++++++++++++++++------- src/socket.rs | 10 +++++++ 3 files changed, 82 insertions(+), 10 deletions(-) diff --git a/src/lib.rs b/src/lib.rs index e179942..bb6cde0 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -24,6 +24,8 @@ pub mod wasm; pub use self::message::Message; #[cfg(not(target_arch = "wasm32"))] pub use self::native::Error; +#[cfg(not(target_arch = "wasm32"))] +pub use self::native::{HeaderMap, HeaderName, HeaderValue}; pub use self::socket::WebSocket; #[cfg(target_arch = "wasm32")] pub use self::wasm::Error; @@ -58,3 +60,14 @@ impl ConnectionMode { pub async fn connect(url: &Url, mode: &ConnectionMode) -> Result { WebSocket::connect(url, mode).await } + +/// Connect with additional HTTP headers in the WebSocket upgrade request. +#[cfg(not(target_arch = "wasm32"))] +#[inline] +pub async fn connect_with_headers( + url: &Url, + mode: &ConnectionMode, + headers: HeaderMap, +) -> Result { + WebSocket::connect_with_headers(url, mode, headers).await +} diff --git a/src/native/mod.rs b/src/native/mod.rs index ecc0a90..12e70cc 100644 --- a/src/native/mod.rs +++ b/src/native/mod.rs @@ -8,6 +8,8 @@ use std::net::SocketAddr; use tokio::io::{AsyncRead, AsyncWrite}; use tokio::net::TcpStream; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +pub use tokio_tungstenite::tungstenite::http::{HeaderMap, HeaderName, HeaderValue}; use tokio_tungstenite::tungstenite::protocol::Role; pub use tokio_tungstenite::tungstenite::Message; use tokio_tungstenite::MaybeTlsStream; @@ -25,14 +27,23 @@ use crate::socket::WebSocket; use crate::ConnectionMode; pub async fn connect(url: &Url, mode: &ConnectionMode) -> Result { + connect_with_headers(url, mode, HeaderMap::new()).await +} + +/// Connect with additional HTTP headers in the WebSocket upgrade request. +pub async fn connect_with_headers( + url: &Url, + mode: &ConnectionMode, + headers: HeaderMap, +) -> Result { match mode { - ConnectionMode::Direct => connect_direct(url).await, + ConnectionMode::Direct => connect_direct(url, headers).await, #[cfg(feature = "socks")] - ConnectionMode::Proxy(proxy) => connect_proxy(url, *proxy).await, + ConnectionMode::Proxy(proxy) => connect_proxy(url, *proxy, headers).await, } } -async fn connect_direct(url: &Url) -> Result { +async fn connect_direct(url: &Url, headers: HeaderMap) -> Result { let host: &str = url.host_str().ok_or_else(Error::empty_host)?; let port: u16 = url .port_or_known_default() @@ -42,11 +53,15 @@ async fn connect_direct(url: &Url) -> Result { let tcp_stream: TcpStream = tokio_happy_eyeballs::connect(host).await?; - connect_stream(url, tcp_stream).await + connect_stream(url, tcp_stream, headers).await } #[cfg(feature = "socks")] -async fn connect_proxy(url: &Url, proxy: SocketAddr) -> Result { +async fn connect_proxy( + url: &Url, + proxy: SocketAddr, + headers: HeaderMap, +) -> Result { let host: &str = url.host_str().ok_or_else(Error::empty_host)?; let port: u16 = url .port_or_known_default() @@ -54,11 +69,15 @@ async fn connect_proxy(url: &Url, proxy: SocketAddr) -> Result let addr: String = format!("{host}:{port}"); let conn: TcpStream = TcpSocks5Stream::connect(proxy, addr).await?; - connect_stream(url, conn).await + connect_stream(url, conn, headers).await } -async fn connect_stream(url: &Url, stream: TcpStream) -> Result { - let stream = client_async(url, stream).await?; +async fn connect_stream( + url: &Url, + stream: TcpStream, + headers: HeaderMap, +) -> Result { + let stream = client_async(url, stream, headers).await?; Ok(WebSocket::tokio(Box::new(stream))) } @@ -73,8 +92,10 @@ async fn connect_stream(url: &Url, stream: TcpStream) -> Result Result>, Error> { - let (stream, _) = Box::pin(tokio_tungstenite::client_async_tls(url.as_str(), stream)).await?; + let request = request_with_headers(url, headers)?; + let (stream, _) = Box::pin(tokio_tungstenite::client_async_tls(request, stream)).await?; Ok(stream) } @@ -87,6 +108,7 @@ async fn client_async( async fn client_async( url: &Url, stream: TcpStream, + headers: HeaderMap, ) -> Result>, Error> { if url.scheme() == "wss" { return Err(tokio_tungstenite::tungstenite::Error::Url( @@ -95,14 +117,24 @@ async fn client_async( .into()); } + let request = request_with_headers(url, headers)?; let (stream, _) = Box::pin(tokio_tungstenite::client_async( - url.as_str(), + request, MaybeTlsStream::Plain(stream), )) .await?; Ok(stream) } +fn request_with_headers( + url: &Url, + headers: HeaderMap, +) -> Result { + let mut request = url.as_str().into_client_request()?; + request.headers_mut().extend(headers); + Ok(request) +} + #[inline] pub async fn accept(raw_stream: S) -> Result, Error> where @@ -121,3 +153,20 @@ where { WebSocketStream::from_raw_socket(raw_stream, Role::Server, None).await } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn request_with_headers_adds_headers_to_upgrade_request() { + let url = Url::parse("wss://relay.example.com").unwrap(); + let mut headers = HeaderMap::new(); + headers.insert("user-agent", HeaderValue::from_static("nostr-sdk")); + + let request = request_with_headers(&url, headers).unwrap(); + + assert_eq!(request.headers().get("user-agent").unwrap(), "nostr-sdk"); + assert_eq!(request.headers().get("host").unwrap(), "relay.example.com"); + } +} diff --git a/src/socket.rs b/src/socket.rs index 5b59bf4..6d508e9 100644 --- a/src/socket.rs +++ b/src/socket.rs @@ -55,6 +55,16 @@ impl WebSocket { Ok(socket) } + + /// Connect with additional HTTP headers in the WebSocket upgrade request. + #[cfg(not(target_arch = "wasm32"))] + pub async fn connect_with_headers( + url: &Url, + mode: &ConnectionMode, + headers: crate::native::HeaderMap, + ) -> Result { + crate::native::connect_with_headers(url, mode, headers).await + } } impl Sink for WebSocket { From 5f52ff6a00824a50fe45a4e07788689564a64426 Mon Sep 17 00:00:00 2001 From: DanConwayDev Date: Thu, 23 Jul 2026 06:36:10 +0100 Subject: [PATCH 2/2] Test WebSocket upgrade headers --- src/native/mod.rs | 29 +++++++++++++++++++++++++++++ 1 file changed, 29 insertions(+) diff --git a/src/native/mod.rs b/src/native/mod.rs index 12e70cc..7788b4c 100644 --- a/src/native/mod.rs +++ b/src/native/mod.rs @@ -156,6 +156,9 @@ where #[cfg(test)] mod tests { + use tokio::net::TcpListener; + use tokio_tungstenite::tungstenite::handshake::server::Request; + use super::*; #[test] @@ -169,4 +172,30 @@ mod tests { assert_eq!(request.headers().get("user-agent").unwrap(), "nostr-sdk"); assert_eq!(request.headers().get("host").unwrap(), "relay.example.com"); } + + #[tokio::test] + #[allow(clippy::result_large_err)] + async fn connect_with_headers_sends_headers_in_upgrade_request() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.unwrap(); + tokio_tungstenite::accept_hdr_async(stream, |request: &Request, response| { + assert_eq!(request.headers().get("user-agent").unwrap(), "nostr-sdk"); + Ok(response) + }) + .await + .unwrap(); + }); + + let url = Url::parse(&format!("ws://{address}")).unwrap(); + let mut headers = HeaderMap::new(); + headers.insert("user-agent", HeaderValue::from_static("nostr-sdk")); + + connect_with_headers(&url, &ConnectionMode::Direct, headers) + .await + .unwrap(); + server.await.unwrap(); + } }