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