Skip to main content

cdk_http_client/ws/
native.rs

1//! Native WebSocket implementation using tokio-tungstenite
2
3use futures::{SinkExt, StreamExt};
4use tokio_tungstenite::tungstenite::client::IntoClientRequest;
5use tokio_tungstenite::tungstenite::Message;
6use tokio_tungstenite::WebSocketStream;
7
8use super::WsError;
9
10/// WebSocket sender half
11pub 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
23/// WebSocket receiver half
24pub 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    /// Send a text message over the WebSocket
40    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    /// Send a close frame
48    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    /// Receive the next text message. Returns `None` when the connection is closed.
58    /// Non-text messages are silently skipped.
59    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, // skip binary, ping, pong
65                Some(Err(e)) => return Some(Err(WsError::from_tungstenite(e))),
66            }
67        }
68    }
69}
70
71/// Adapt an established WebSocket stream to CDK's sender and receiver types.
72///
73/// This is useful for transports that perform the HTTP upgrade themselves,
74/// such as an encrypted tunnel, and then construct a WebSocket stream over the
75/// resulting bidirectional byte stream.
76pub 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
96/// Connect to a WebSocket endpoint with optional headers.
97///
98/// `headers` is a slice of `(name, value)` pairs to include in the upgrade request.
99pub 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/// Connect to a WebSocket endpoint through an Arti Tor client.
131#[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}