Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 13 additions & 0 deletions src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -58,3 +60,14 @@ impl ConnectionMode {
pub async fn connect(url: &Url, mode: &ConnectionMode) -> Result<WebSocket, Error> {
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, Error> {
WebSocket::connect_with_headers(url, mode, headers).await
}
98 changes: 88 additions & 10 deletions src/native/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -25,14 +27,23 @@ use crate::socket::WebSocket;
use crate::ConnectionMode;

pub async fn connect(url: &Url, mode: &ConnectionMode) -> Result<WebSocket, Error> {
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<WebSocket, Error> {
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<WebSocket, Error> {
async fn connect_direct(url: &Url, headers: HeaderMap) -> Result<WebSocket, Error> {
let host: &str = url.host_str().ok_or_else(Error::empty_host)?;
let port: u16 = url
.port_or_known_default()
Expand All @@ -42,23 +53,31 @@ async fn connect_direct(url: &Url) -> Result<WebSocket, Error> {

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<WebSocket, Error> {
async fn connect_proxy(
url: &Url,
proxy: SocketAddr,
headers: HeaderMap,
) -> Result<WebSocket, Error> {
let host: &str = url.host_str().ok_or_else(Error::empty_host)?;
let port: u16 = url
.port_or_known_default()
.ok_or_else(Error::invalid_port)?;
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<WebSocket, Error> {
let stream = client_async(url, stream).await?;
async fn connect_stream(
url: &Url,
stream: TcpStream,
headers: HeaderMap,
) -> Result<WebSocket, Error> {
let stream = client_async(url, stream, headers).await?;
Ok(WebSocket::tokio(Box::new(stream)))
}

Expand All @@ -73,8 +92,10 @@ async fn connect_stream(url: &Url, stream: TcpStream) -> Result<WebSocket, Error
async fn client_async(
url: &Url,
stream: TcpStream,
headers: HeaderMap,
) -> Result<WebSocketStream<MaybeTlsStream<TcpStream>>, 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)
}

Expand All @@ -87,6 +108,7 @@ async fn client_async(
async fn client_async(
url: &Url,
stream: TcpStream,
headers: HeaderMap,
) -> Result<WebSocketStream<MaybeTlsStream<TcpStream>>, Error> {
if url.scheme() == "wss" {
return Err(tokio_tungstenite::tungstenite::Error::Url(
Expand All @@ -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<tokio_tungstenite::tungstenite::handshake::client::Request, Error> {
let mut request = url.as_str().into_client_request()?;
request.headers_mut().extend(headers);
Ok(request)
}

#[inline]
pub async fn accept<S>(raw_stream: S) -> Result<WebSocketStream<S>, Error>
where
Expand All @@ -121,3 +153,49 @@ where
{
WebSocketStream::from_raw_socket(raw_stream, Role::Server, None).await
}

#[cfg(test)]
mod tests {
use tokio::net::TcpListener;
use tokio_tungstenite::tungstenite::handshake::server::Request;

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");
}

#[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();
}
}
10 changes: 10 additions & 0 deletions src/socket.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<Self, Error> {
crate::native::connect_with_headers(url, mode, headers).await
}
}

impl Sink<Message> for WebSocket {
Expand Down
Loading