async_wsocket/native/
mod.rs1#[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
32pub 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#[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}