Skip to main content

rig_tungstenite/
lib.rs

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//! The bundled native websocket backend for Rig: a
12//! [`rig_http::ws_client::WebSocketClientExt`] implementation over tokio-tungstenite. rig-core's `tungstenite` feature
13//! opens Responses WebSocket sessions over it with no backend named.
14//!
15//! Sockets use the current Tokio runtime or a lazy fallback runtime. Off-runtime
16//! callers communicate through channels without polling socket I/O themselves.
17//!
18//! ```
19//! let backend = rig_tungstenite::TungsteniteClient::new();
20//! ```
21
22// The native dependency is target-gated; diagnose unsupported WASM use before
23// unresolved imports obscure the required browser backend.
24#[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/// A stateless websocket backend that takes handshake configuration from each
52/// request.
53#[derive(Clone, Copy, Debug, Default)]
54pub struct TungsteniteClient;
55
56#[cfg(not(target_family = "wasm"))]
57impl TungsteniteClient {
58    /// The bundled backend.
59    #[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"))]
90/// Build a tungstenite handshake request, returning an error for an invalid URI.
91/// Caller headers, including authentication, override generated handshake headers.
92fn 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/// A handshake that did not complete in time.
107#[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"))]
134/// Convert a tungstenite failure to a transport error, preserving the status,
135/// headers, and body of a rejected upgrade.
136/// Other failures become [`Error::Instance`].
137fn 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;