1#![cfg_attr(docsrs, feature(doc_cfg))]
2#![cfg_attr(
3 test,
4 allow(
5 clippy::expect_used,
6 clippy::indexing_slicing,
7 clippy::panic,
8 clippy::unwrap_used
9 )
10)]
11#[cfg(target_family = "wasm")]
25compile_error!(
26 "rig-tungstenite is a native websocket backend (tokio-tungstenite). On wasm, implement \
27 `rig_http::ws_client::WebSocketClientExt` over `web_sys::WebSocket` and open sessions with \
28 `connect_with(..)`."
29);
30
31#[cfg(not(target_family = "wasm"))]
32pub use tokio_tungstenite;
33
34#[cfg(not(target_family = "wasm"))]
35mod connection;
36#[cfg(not(target_family = "wasm"))]
37mod runtime;
38
39#[cfg(not(target_family = "wasm"))]
40use connection::{DirectConnection, ForwardedConnection};
41#[cfg(not(target_family = "wasm"))]
42use rig_http::http_client::{Error, NoBody, Request, Result};
43#[cfg(not(target_family = "wasm"))]
44use rig_http::ws_client::{BoxedWebSocketConnection, ConnectOptions, WebSocketClientExt};
45#[cfg(not(target_family = "wasm"))]
46use std::time::Duration;
47#[cfg(not(target_family = "wasm"))]
48use tokio_tungstenite::tungstenite::{self, client::IntoClientRequest};
49
50#[cfg(not(target_family = "wasm"))]
51#[derive(Clone, Copy, Debug, Default)]
54pub struct TungsteniteClient;
55
56#[cfg(not(target_family = "wasm"))]
57impl TungsteniteClient {
58 #[must_use]
60 pub fn new() -> Self {
61 Self
62 }
63}
64
65#[cfg(not(target_family = "wasm"))]
66impl WebSocketClientExt for TungsteniteClient {
67 async fn connect(
68 &self,
69 request: Request<NoBody>,
70 options: ConnectOptions,
71 ) -> Result<BoxedWebSocketConnection> {
72 let request = client_request(request)?;
73
74 #[cfg(not(target_family = "wasm"))]
75 if !runtime::in_tokio() {
76 let timeout = options.timeout;
77 return runtime::run_off_runtime(async move {
78 let socket = handshake(request, timeout).await?;
79 ForwardedConnection::spawn(socket)
80 })
81 .await?;
82 }
83
84 let socket = handshake(request, options.timeout).await?;
85 Ok(Box::new(DirectConnection::new(socket)))
86 }
87}
88
89#[cfg(not(target_family = "wasm"))]
90fn client_request(request: Request<NoBody>) -> Result<tungstenite::handshake::client::Request> {
93 let (parts, _) = request.into_parts();
94 let mut request = parts
95 .uri
96 .to_string()
97 .into_client_request()
98 .map_err(from_tungstenite)?;
99 for (name, value) in &parts.headers {
100 request.headers_mut().insert(name, value.clone());
101 }
102 Ok(request)
103}
104
105#[cfg(not(target_family = "wasm"))]
106#[derive(Debug, thiserror::Error)]
108#[error("timed out connecting the websocket after {0:?}")]
109struct ConnectTimeout(Duration);
110
111#[cfg(not(target_family = "wasm"))]
112async fn handshake(
113 request: tungstenite::handshake::client::Request,
114 timeout: Option<Duration>,
115) -> Result<connection::Socket> {
116 let connect = async {
117 tokio_tungstenite::connect_async(request)
118 .await
119 .map(|(socket, _)| socket)
120 .map_err(from_tungstenite)
121 };
122
123 let Some(timeout) = timeout else {
124 return connect.await;
125 };
126
127 match rig_http::wasm_compat::timeout(timeout, connect).await {
128 Ok(result) => result,
129 Err(_) => Err(Error::instance(ConnectTimeout(timeout))),
130 }
131}
132
133#[cfg(not(target_family = "wasm"))]
134fn from_tungstenite(error: tungstenite::Error) -> Error {
138 let tungstenite::Error::Http(response) = error else {
139 return Error::instance(error);
140 };
141
142 let (parts, body) = (*response).into_parts();
143 let body = body
144 .map(|body| String::from_utf8_lossy(&body).into_owned())
145 .unwrap_or_default();
146
147 Error::non_success_with_details(parts.status, parts.headers, body)
148}
149
150#[cfg(all(test, not(target_family = "wasm")))]
151mod tests;