cdk_http_client/ws/
native.rs1use futures::{SinkExt, StreamExt};
4use tokio_tungstenite::tungstenite::client::IntoClientRequest;
5use tokio_tungstenite::tungstenite::Message;
6use tokio_tungstenite::WebSocketStream;
7
8use super::WsError;
9
10pub struct WsSender {
12 inner: Box<
13 dyn futures::Sink<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin + Send,
14 >,
15}
16
17impl std::fmt::Debug for WsSender {
18 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
19 f.debug_struct("WsSender").finish_non_exhaustive()
20 }
21}
22
23pub struct WsReceiver {
25 inner: Box<
26 dyn futures::Stream<Item = Result<Message, tokio_tungstenite::tungstenite::Error>>
27 + Unpin
28 + Send,
29 >,
30}
31
32impl std::fmt::Debug for WsReceiver {
33 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
34 f.debug_struct("WsReceiver").finish_non_exhaustive()
35 }
36}
37
38impl WsSender {
39 pub async fn send(&mut self, text: String) -> Result<(), WsError> {
41 self.inner
42 .send(Message::Text(text.into()))
43 .await
44 .map_err(WsError::from_tungstenite)
45 }
46
47 pub async fn close(&mut self) -> Result<(), WsError> {
49 self.inner
50 .send(Message::Close(None))
51 .await
52 .map_err(WsError::from_tungstenite)
53 }
54}
55
56impl WsReceiver {
57 pub async fn recv(&mut self) -> Option<Result<String, WsError>> {
60 loop {
61 match self.inner.next().await {
62 Some(Ok(Message::Text(text))) => return Some(Ok(text.to_string())),
63 Some(Ok(Message::Close(_))) | None => return None,
64 Some(Ok(_)) => continue, Some(Err(e)) => return Some(Err(WsError::from_tungstenite(e))),
66 }
67 }
68 }
69}
70
71pub fn from_websocket_stream<S>(ws_stream: WebSocketStream<S>) -> (WsSender, WsReceiver)
77where
78 WebSocketStream<S>: futures::Sink<Message, Error = tokio_tungstenite::tungstenite::Error>
79 + futures::Stream<Item = Result<Message, tokio_tungstenite::tungstenite::Error>>
80 + Unpin
81 + Send
82 + 'static,
83{
84 let (sink, stream) = ws_stream.split();
85
86 (
87 WsSender {
88 inner: Box::new(sink),
89 },
90 WsReceiver {
91 inner: Box::new(stream),
92 },
93 )
94}
95
96pub async fn connect(
100 url: &str,
101 headers: &[(&str, &str)],
102) -> Result<(WsSender, WsReceiver), WsError> {
103 let mut request = url
104 .into_client_request()
105 .map_err(|e| WsError::Terminal(e.to_string()))?;
106
107 for &(name, value) in headers {
108 let header_name = name
109 .parse::<tokio_tungstenite::tungstenite::http::header::HeaderName>()
110 .map_err(|error| {
111 WsError::Terminal(format!("invalid WebSocket header name `{name}`: {error}"))
112 })?;
113 let header_value = value
114 .parse::<tokio_tungstenite::tungstenite::http::header::HeaderValue>()
115 .map_err(|error| {
116 WsError::Terminal(format!(
117 "invalid value for WebSocket header `{name}`: {error}"
118 ))
119 })?;
120 request.headers_mut().insert(header_name, header_value);
121 }
122
123 let (ws_stream, _) = tokio_tungstenite::connect_async(request)
124 .await
125 .map_err(WsError::from_tungstenite)?;
126
127 Ok(from_websocket_stream(ws_stream))
128}
129
130#[cfg(feature = "tor")]
132pub(crate) async fn connect_tor(
133 tor_client: arti_client::TorClient<tor_rtcompat::PreferredRuntime>,
134 url: &str,
135 headers: &[(&str, &str)],
136) -> Result<(WsSender, WsReceiver), WsError> {
137 let parsed_url =
138 url::Url::parse(url).map_err(|e| WsError::Terminal(format!("Invalid URL: {e}")))?;
139
140 let host = parsed_url
141 .host_str()
142 .ok_or_else(|| WsError::Terminal("WebSocket URL must include a host".to_string()))?;
143 let port = parsed_url
144 .port_or_known_default()
145 .ok_or_else(|| WsError::Terminal("WebSocket URL must include a port".to_string()))?;
146
147 let mut request = url
148 .into_client_request()
149 .map_err(|e| WsError::Terminal(e.to_string()))?;
150
151 for &(name, value) in headers {
152 let header_name = name
153 .parse::<tokio_tungstenite::tungstenite::http::header::HeaderName>()
154 .map_err(|error| {
155 WsError::Terminal(format!("invalid WebSocket header name `{name}`: {error}"))
156 })?;
157 let header_value = value
158 .parse::<tokio_tungstenite::tungstenite::http::header::HeaderValue>()
159 .map_err(|error| {
160 WsError::Terminal(format!(
161 "invalid value for WebSocket header `{name}`: {error}"
162 ))
163 })?;
164 request.headers_mut().insert(header_name, header_value);
165 }
166
167 let stream = tor_client
168 .connect((host, port))
169 .await
170 .map_err(|e| WsError::Transient(e.to_string()))?;
171
172 let (ws_stream, _) =
173 tokio_tungstenite::client_async_tls_with_config(request, stream, None, None)
174 .await
175 .map_err(WsError::from_tungstenite)?;
176
177 Ok(from_websocket_stream(ws_stream))
178}