Skip to main content

async_wsocket/native/
mod.rs

1// Copyright (c) 2022-2024 Yuki Kishimoto
2// Distributed under the MIT software license
3
4//! Native
5
6#[cfg(feature = "socks")]
7use std::net::SocketAddr;
8
9use tokio::net::TcpStream;
10use tokio_tungstenite::tungstenite::client::IntoClientRequest;
11pub use tokio_tungstenite::tungstenite::http::{HeaderMap, HeaderName, HeaderValue};
12pub use tokio_tungstenite::tungstenite::protocol::{Role, WebSocketConfig};
13pub use tokio_tungstenite::tungstenite::Message;
14use tokio_tungstenite::MaybeTlsStream;
15pub use tokio_tungstenite::{accept_async, accept_async_with_config, WebSocketStream};
16use url::Url;
17
18mod error;
19#[cfg(feature = "socks")]
20mod socks;
21
22pub use self::error::Error;
23#[cfg(feature = "socks")]
24use self::socks::TcpSocks5Stream;
25use crate::socket::WebSocket;
26use crate::ConnectionMode;
27
28pub async fn connect(url: &Url, mode: &ConnectionMode) -> Result<WebSocket, Error> {
29    connect_with_headers(url, mode, HeaderMap::new()).await
30}
31
32/// Connect with additional HTTP headers in the WebSocket upgrade request.
33pub async fn connect_with_headers(
34    url: &Url,
35    mode: &ConnectionMode,
36    headers: HeaderMap,
37) -> Result<WebSocket, Error> {
38    match mode {
39        ConnectionMode::Direct => connect_direct(url, headers).await,
40        #[cfg(feature = "socks")]
41        ConnectionMode::Proxy(proxy) => connect_proxy(url, *proxy, headers).await,
42    }
43}
44
45async fn connect_direct(url: &Url, headers: HeaderMap) -> Result<WebSocket, Error> {
46    let host: &str = url.host_str().ok_or_else(Error::empty_host)?;
47    let port: u16 = url
48        .port_or_known_default()
49        .ok_or_else(Error::invalid_port)?;
50
51    let host: String = format!("{}:{}", host, port);
52
53    let tcp_stream: TcpStream = tokio_happy_eyeballs::connect(host).await?;
54
55    connect_stream(url, tcp_stream, headers).await
56}
57
58#[cfg(feature = "socks")]
59async fn connect_proxy(
60    url: &Url,
61    proxy: SocketAddr,
62    headers: HeaderMap,
63) -> Result<WebSocket, Error> {
64    let host: &str = url.host_str().ok_or_else(Error::empty_host)?;
65    let port: u16 = url
66        .port_or_known_default()
67        .ok_or_else(Error::invalid_port)?;
68    let addr: String = format!("{host}:{port}");
69
70    let conn: TcpStream = TcpSocks5Stream::connect(proxy, addr).await?;
71    connect_stream(url, conn, headers).await
72}
73
74async fn connect_stream(
75    url: &Url,
76    stream: TcpStream,
77    headers: HeaderMap,
78) -> Result<WebSocket, Error> {
79    let stream = client_async(url, stream, headers).await?;
80    Ok(WebSocket::tokio(Box::new(stream)))
81}
82
83// NOT REMOVE `Box::pin`!
84// Use `Box::pin` to fix stack overflow on windows targets due to large `Future`
85#[cfg(any(
86    feature = "native-tls",
87    feature = "native-tls-vendored",
88    feature = "rustls-tls-native-roots",
89    feature = "rustls-tls-webpki-roots"
90))]
91async fn client_async(
92    url: &Url,
93    stream: TcpStream,
94    headers: HeaderMap,
95) -> Result<WebSocketStream<MaybeTlsStream<TcpStream>>, Error> {
96    let request = request_with_headers(url, headers)?;
97    let (stream, _) = Box::pin(tokio_tungstenite::client_async_tls(request, stream)).await?;
98    Ok(stream)
99}
100
101#[cfg(not(any(
102    feature = "native-tls",
103    feature = "native-tls-vendored",
104    feature = "rustls-tls-native-roots",
105    feature = "rustls-tls-webpki-roots"
106)))]
107async fn client_async(
108    url: &Url,
109    stream: TcpStream,
110    headers: HeaderMap,
111) -> Result<WebSocketStream<MaybeTlsStream<TcpStream>>, Error> {
112    if url.scheme() == "wss" {
113        return Err(tokio_tungstenite::tungstenite::Error::Url(
114            tokio_tungstenite::tungstenite::error::UrlError::TlsFeatureNotEnabled,
115        )
116        .into());
117    }
118
119    let request = request_with_headers(url, headers)?;
120    let (stream, _) = Box::pin(tokio_tungstenite::client_async(
121        request,
122        MaybeTlsStream::Plain(stream),
123    ))
124    .await?;
125    Ok(stream)
126}
127
128fn request_with_headers(
129    url: &Url,
130    headers: HeaderMap,
131) -> Result<tokio_tungstenite::tungstenite::handshake::client::Request, Error> {
132    let mut request = url.as_str().into_client_request()?;
133    request.headers_mut().extend(headers);
134    Ok(request)
135}
136
137#[cfg(test)]
138mod tests {
139    use tokio::net::TcpListener;
140    use tokio_tungstenite::tungstenite::handshake::server::Request;
141
142    use super::*;
143
144    #[test]
145    fn request_with_headers_adds_headers_to_upgrade_request() {
146        let url = Url::parse("wss://relay.example.com").unwrap();
147        let mut headers = HeaderMap::new();
148        headers.insert("user-agent", HeaderValue::from_static("nostr-sdk"));
149
150        let request = request_with_headers(&url, headers).unwrap();
151
152        assert_eq!(request.headers().get("user-agent").unwrap(), "nostr-sdk");
153        assert_eq!(request.headers().get("host").unwrap(), "relay.example.com");
154    }
155
156    #[tokio::test]
157    #[allow(clippy::result_large_err)]
158    async fn connect_with_headers_sends_headers_in_upgrade_request() {
159        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
160        let address = listener.local_addr().unwrap();
161
162        let server = tokio::spawn(async move {
163            let (stream, _) = listener.accept().await.unwrap();
164            tokio_tungstenite::accept_hdr_async(stream, |request: &Request, response| {
165                assert_eq!(request.headers().get("user-agent").unwrap(), "nostr-sdk");
166                Ok(response)
167            })
168            .await
169            .unwrap();
170        });
171
172        let url = Url::parse(&format!("ws://{address}")).unwrap();
173        let mut headers = HeaderMap::new();
174        headers.insert("user-agent", HeaderValue::from_static("nostr-sdk"));
175
176        connect_with_headers(&url, &ConnectionMode::Direct, headers)
177            .await
178            .unwrap();
179        server.await.unwrap();
180    }
181}