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