1#![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, 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 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}