Skip to main content

tauri_plugin_websocket/
lib.rs

1// Copyright 2019-2023 Tauri Programme within The Commons Conservancy
2// SPDX-License-Identifier: Apache-2.0
3// SPDX-License-Identifier: MIT
4
5//! Open a WebSocket connection using a Rust client in JS.
6//!
7//! ## Cargo features
8//!
9//! - **rustls-tls** *(enabled by default)*: Enables TLS functionality provided by `rustls` with WebPKI roots.
10//! - **rustls-tls-native-roots**: Enables TLS functionality provided by `rustls` with the platform's native certificate roots.
11//! - **native-tls**: Enables TLS functionality provided by `native-tls`.
12//! - **native-tls-vendored**: Enables the `vendored` feature of `native-tls`.
13//!
14//! At least one TLS feature is required for `wss://`; plain `ws://` works without one.
15
16#![doc(
17    html_logo_url = "https://github.com/tauri-apps/tauri/raw/dev/app-icon.png",
18    html_favicon_url = "https://github.com/tauri-apps/tauri/raw/dev/app-icon.png"
19)]
20
21use futures_util::{stream::SplitSink, SinkExt, StreamExt};
22use http::header::{HeaderName, HeaderValue};
23use serde::{ser::Serializer, Deserialize, Serialize};
24use tauri::{
25    ipc::Channel,
26    plugin::{Builder as PluginBuilder, TauriPlugin},
27    Manager, Runtime, State, Window,
28};
29use tokio::{net::TcpStream, sync::Mutex};
30#[cfg(any(
31    feature = "rustls-tls",
32    feature = "rustls-tls-native-roots",
33    feature = "native-tls"
34))]
35use tokio_tungstenite::connect_async_tls_with_config;
36#[cfg(not(any(
37    feature = "rustls-tls",
38    feature = "rustls-tls-native-roots",
39    feature = "native-tls"
40)))]
41use tokio_tungstenite::connect_async_with_config;
42use tokio_tungstenite::{
43    tungstenite::{
44        client::IntoClientRequest,
45        protocol::{CloseFrame as ProtocolCloseFrame, WebSocketConfig},
46        Message,
47    },
48    Connector, MaybeTlsStream, WebSocketStream,
49};
50
51use std::collections::HashMap;
52use std::str::FromStr;
53
54type Id = u32;
55type WebSocket = WebSocketStream<MaybeTlsStream<TcpStream>>;
56type WebSocketWriter = SplitSink<WebSocket, Message>;
57type Result<T> = std::result::Result<T, Error>;
58
59#[derive(Debug, thiserror::Error)]
60enum Error {
61    #[error(transparent)]
62    Websocket(#[from] tokio_tungstenite::tungstenite::Error),
63    #[error("connection not found for the given id: {0}")]
64    ConnectionNotFound(Id),
65    #[error(transparent)]
66    InvalidHeaderValue(#[from] tokio_tungstenite::tungstenite::http::header::InvalidHeaderValue),
67    #[error(transparent)]
68    InvalidHeaderName(#[from] tokio_tungstenite::tungstenite::http::header::InvalidHeaderName),
69}
70
71impl Serialize for Error {
72    fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
73    where
74        S: Serializer,
75    {
76        serializer.serialize_str(self.to_string().as_str())
77    }
78}
79
80#[derive(Default)]
81struct ConnectionManager(Mutex<HashMap<Id, WebSocketWriter>>);
82
83#[cfg(any(
84    feature = "rustls-tls",
85    feature = "rustls-tls-native-roots",
86    feature = "native-tls"
87))]
88struct TlsConnector(Mutex<Option<Connector>>);
89
90#[derive(Deserialize)]
91#[serde(untagged, rename_all = "camelCase")]
92enum Max {
93    None,
94    Number(usize),
95}
96
97#[derive(Deserialize)]
98#[serde(rename_all = "camelCase")]
99pub(crate) struct ConnectionConfig {
100    pub read_buffer_size: Option<usize>,
101    pub write_buffer_size: Option<usize>,
102    pub max_write_buffer_size: Option<usize>,
103    pub max_message_size: Option<Max>,
104    pub max_frame_size: Option<Max>,
105    #[serde(default)]
106    pub accept_unmasked_frames: bool,
107    pub headers: Option<Vec<(String, String)>>,
108}
109
110impl From<ConnectionConfig> for WebSocketConfig {
111    fn from(config: ConnectionConfig) -> Self {
112        let mut builder =
113            WebSocketConfig::default().accept_unmasked_frames(config.accept_unmasked_frames);
114
115        if let Some(read_buffer_size) = config.read_buffer_size {
116            builder = builder.read_buffer_size(read_buffer_size)
117        }
118
119        if let Some(write_buffer_size) = config.write_buffer_size {
120            builder = builder.write_buffer_size(write_buffer_size)
121        }
122
123        if let Some(max_write_buffer_size) = config.max_write_buffer_size {
124            builder = builder.max_write_buffer_size(max_write_buffer_size)
125        }
126
127        if let Some(max_message_size) = config.max_message_size {
128            let max_size = match max_message_size {
129                Max::None => Option::None,
130                Max::Number(n) => Some(n),
131            };
132            builder = builder.max_message_size(max_size);
133        }
134
135        if let Some(max_frame_size) = config.max_frame_size {
136            let max_size = match max_frame_size {
137                Max::None => Option::None,
138                Max::Number(n) => Some(n),
139            };
140            builder = builder.max_frame_size(max_size);
141        }
142
143        builder
144    }
145}
146
147#[derive(Deserialize, Serialize)]
148struct CloseFrame {
149    pub code: u16,
150    pub reason: String,
151}
152
153#[derive(Deserialize, Serialize)]
154#[serde(tag = "type", content = "data")]
155enum WebSocketMessage {
156    Text(String),
157    Binary(Vec<u8>),
158    Ping(Vec<u8>),
159    Pong(Vec<u8>),
160    Close(Option<CloseFrame>),
161}
162
163#[tauri::command]
164async fn connect<R: Runtime>(
165    window: Window<R>,
166    url: String,
167    on_message: Channel<serde_json::Value>,
168    config: Option<ConnectionConfig>,
169) -> Result<Id> {
170    let id = rand::random();
171    let mut request = url.into_client_request()?;
172
173    if let Some(headers) = config.as_ref().and_then(|c| c.headers.as_ref()) {
174        for (k, v) in headers {
175            let header_name = HeaderName::from_str(k.as_str())?;
176            let header_value = HeaderValue::from_str(v.as_str())?;
177            request.headers_mut().insert(header_name, header_value);
178        }
179    }
180
181    #[cfg(any(
182        feature = "rustls-tls",
183        feature = "rustls-tls-native-roots",
184        feature = "native-tls"
185    ))]
186    let tls_connector = match window.try_state::<TlsConnector>() {
187        Some(tls_connector) => tls_connector.0.lock().await.clone(),
188        None => None,
189    };
190
191    #[cfg(any(
192        feature = "rustls-tls",
193        feature = "rustls-tls-native-roots",
194        feature = "native-tls"
195    ))]
196    let (ws_stream, _) =
197        connect_async_tls_with_config(request, config.map(Into::into), false, tls_connector)
198            .await?;
199    #[cfg(not(any(
200        feature = "rustls-tls",
201        feature = "rustls-tls-native-roots",
202        feature = "native-tls"
203    )))]
204    let (ws_stream, _) = connect_async_with_config(request, config.map(Into::into), false).await?;
205
206    tauri::async_runtime::spawn(async move {
207        let (write, read) = ws_stream.split();
208        let manager = window.state::<ConnectionManager>();
209        manager.0.lock().await.insert(id, write);
210        read.for_each(move |message| {
211            let window_ = window.clone();
212            let on_message_ = on_message.clone();
213            async move {
214                if let Ok(Message::Close(_)) = message {
215                    let manager = window_.state::<ConnectionManager>();
216                    manager.0.lock().await.remove(&id);
217                }
218
219                let response = match message {
220                    Ok(Message::Text(t)) => {
221                        serde_json::to_value(WebSocketMessage::Text(t.to_string())).unwrap()
222                    }
223                    Ok(Message::Binary(t)) => {
224                        serde_json::to_value(WebSocketMessage::Binary(t.to_vec())).unwrap()
225                    }
226                    Ok(Message::Ping(t)) => {
227                        serde_json::to_value(WebSocketMessage::Ping(t.to_vec())).unwrap()
228                    }
229                    Ok(Message::Pong(t)) => {
230                        serde_json::to_value(WebSocketMessage::Pong(t.to_vec())).unwrap()
231                    }
232                    Ok(Message::Close(t)) => {
233                        serde_json::to_value(WebSocketMessage::Close(t.map(|v| CloseFrame {
234                            code: v.code.into(),
235                            reason: v.reason.to_string(),
236                        })))
237                        .unwrap()
238                    }
239                    Ok(Message::Frame(_)) => serde_json::Value::Null, // This value can't be recieved.
240                    Err(e) => serde_json::to_value(Error::from(e)).unwrap(),
241                };
242
243                let _ = on_message_.send(response);
244            }
245        })
246        .await;
247    });
248
249    Ok(id)
250}
251
252#[tauri::command]
253async fn send(
254    manager: State<'_, ConnectionManager>,
255    id: Id,
256    message: WebSocketMessage,
257) -> Result<()> {
258    if let Some(write) = manager.0.lock().await.get_mut(&id) {
259        write
260            .send(match message {
261                WebSocketMessage::Text(t) => Message::Text(t.into()),
262                WebSocketMessage::Binary(t) => Message::Binary(t.into()),
263                WebSocketMessage::Ping(t) => Message::Ping(t.into()),
264                WebSocketMessage::Pong(t) => Message::Pong(t.into()),
265                WebSocketMessage::Close(t) => Message::Close(t.map(|v| ProtocolCloseFrame {
266                    code: v.code.into(),
267                    reason: v.reason.into(),
268                })),
269            })
270            .await?;
271        Ok(())
272    } else {
273        Err(Error::ConnectionNotFound(id))
274    }
275}
276
277pub fn init<R: Runtime>() -> TauriPlugin<R> {
278    Builder::default().build()
279}
280
281#[derive(Default)]
282pub struct Builder {
283    tls_connector: Option<Connector>,
284}
285
286impl Builder {
287    pub fn new() -> Self {
288        Self {
289            tls_connector: None,
290        }
291    }
292
293    pub fn tls_connector(mut self, connector: Connector) -> Self {
294        self.tls_connector.replace(connector);
295        self
296    }
297
298    pub fn build<R: Runtime>(self) -> TauriPlugin<R> {
299        PluginBuilder::new("websocket")
300            .invoke_handler(tauri::generate_handler![connect, send])
301            .setup(|app, _api| {
302                #[cfg(any(feature = "rustls-tls", feature = "rustls-tls-native-roots"))]
303                if (self.tls_connector.is_none()
304                    || matches!(self.tls_connector, Some(Connector::Plain)))
305                    && rustls::crypto::CryptoProvider::get_default().is_none()
306                {
307                    // This can only fail if there is already a default provider which we checked for already.
308                    let _ = rustls::crypto::ring::default_provider().install_default();
309                }
310
311                app.manage(ConnectionManager::default());
312                #[cfg(any(
313                    feature = "rustls-tls",
314                    feature = "rustls-tls-native-roots",
315                    feature = "native-tls"
316                ))]
317                app.manage(TlsConnector(Mutex::new(self.tls_connector)));
318                Ok(())
319            })
320            .build()
321    }
322}