Skip to main content

gloo_net/websocket/
futures.rs

1//! The wrapper around `WebSocket` API using the Futures API to be used in async rust
2//!
3//! # Example
4//!
5//! ```rust
6//! use gloo_net::websocket::{Message, futures::WebSocket};
7//! use wasm_bindgen_futures::spawn_local;
8//! use futures::{SinkExt, StreamExt};
9//!
10//! # macro_rules! console_log {
11//! #    ($($expr:expr),*) => {{}};
12//! # }
13//! # fn no_run() {
14//! let mut ws = WebSocket::open("wss://echo.websocket.org").unwrap();
15//! let (mut write, mut read) = ws.split();
16//!
17//! spawn_local(async move {
18//!     write.send(Message::Text(String::from("test"))).await.unwrap();
19//!     write.send(Message::Text(String::from("test 2"))).await.unwrap();
20//! });
21//!
22//! spawn_local(async move {
23//!     while let Some(msg) = read.next().await {
24//!         console_log!(format!("1. {:?}", msg))
25//!     }
26//!     console_log!("WebSocket Closed")
27//! })
28//! # }
29//! ```
30use crate::js_to_js_error;
31use crate::websocket::{events::CloseEvent, Message, State, WebSocketError};
32use futures_channel::mpsc;
33use futures_core::{ready, Stream};
34use futures_sink::Sink;
35use gloo_utils::errors::JsError;
36use pin_project::{pin_project, pinned_drop};
37use std::cell::RefCell;
38use std::pin::Pin;
39use std::rc::Rc;
40use std::task::{Context, Poll, Waker};
41use wasm_bindgen::prelude::*;
42use wasm_bindgen::JsCast;
43use web_sys::{BinaryType, MessageEvent};
44
45/// Wrapper around browser's WebSocket API.
46#[allow(missing_debug_implementations)]
47#[pin_project(PinnedDrop)]
48pub struct WebSocket {
49    ws: web_sys::WebSocket,
50    sink_waker: Rc<RefCell<Option<Waker>>>,
51    #[pin]
52    message_receiver: mpsc::UnboundedReceiver<StreamMessage>,
53    #[allow(clippy::type_complexity)]
54    closures: (
55        Closure<dyn FnMut()>,
56        Closure<dyn FnMut(MessageEvent)>,
57        Closure<dyn FnMut(web_sys::Event)>,
58        Closure<dyn FnMut(web_sys::CloseEvent)>,
59    ),
60    /// Leftover bytes when using `AsyncRead`.
61    ///
62    /// These bytes are drained and returned in subsequent calls to `poll_read`.
63    #[cfg(feature = "io-util")]
64    pub(super) read_pending_bytes: Option<Vec<u8>>, // Same size as `Vec<u8>` alone thanks to niche optimization
65}
66
67impl WebSocket {
68    /// Establish a WebSocket connection.
69    ///
70    /// This function may error in the following cases:
71    /// - The port to which the connection is being attempted is being blocked.
72    /// - The URL is invalid.
73    ///
74    /// The error returned is [`JsError`]. See the
75    /// [MDN Documentation](https://developer.mozilla.org/en-US/docs/Web/API/WebSocket/WebSocket#exceptions_thrown)
76    /// to learn more.
77    pub fn open(url: &str) -> Result<Self, JsError> {
78        Self::setup(web_sys::WebSocket::new(url))
79    }
80
81    /// Establish a WebSocket connection.
82    ///
83    /// This function may error in the following cases:
84    /// - The port to which the connection is being attempted is being blocked.
85    /// - The URL is invalid.
86    /// - The specified protocol is not supported
87    ///
88    /// The error returned is [`JsError`]. See the
89    /// [MDN Documentation](https://developer.mozilla.org/en-US/docs/Web/API/WebSocket/WebSocket#exceptions_thrown)
90    /// to learn more.
91    pub fn open_with_protocol(url: &str, protocol: &str) -> Result<Self, JsError> {
92        Self::setup(web_sys::WebSocket::new_with_str(url, protocol))
93    }
94
95    /// Establish a WebSocket connection.
96    ///
97    /// This function may error in the following cases:
98    /// - The port to which the connection is being attempted is being blocked.
99    /// - The URL is invalid.
100    /// - The specified protocols are not supported
101    /// - The protocols cannot be converted to a JSON string list
102    ///
103    /// The error returned is [`JsError`]. See the
104    /// [MDN Documentation](https://developer.mozilla.org/en-US/docs/Web/API/WebSocket/WebSocket#exceptions_thrown)
105    /// to learn more.
106    ///
107    /// This function requires `json` features because protocols are parsed by `serde` into `JsValue`.
108    #[cfg_attr(docsrs, doc(cfg(feature = "json")))]
109    #[cfg(feature = "json")]
110    pub fn open_with_protocols<S: AsRef<str> + serde::Serialize>(
111        url: &str,
112        protocols: &[S],
113    ) -> Result<Self, JsError> {
114        let json = <JsValue as gloo_utils::format::JsValueSerdeExt>::from_serde(protocols)
115            .map_err(|err| {
116                js_sys::Error::new(&format!(
117                    "Failed to convert protocols to Javascript value: {err}"
118                ))
119            })?;
120        Self::setup(web_sys::WebSocket::new_with_str_sequence(url, &json))
121    }
122
123    fn setup(ws: Result<web_sys::WebSocket, JsValue>) -> Result<Self, JsError> {
124        let waker: Rc<RefCell<Option<Waker>>> = Rc::new(RefCell::new(None));
125        let ws = ws.map_err(js_to_js_error)?;
126
127        // We rely on this because the other type Blob can be converted to Vec<u8> only through a
128        // promise which makes it awkward to use in our event callbacks where we want to guarantee
129        // the order of the events stays the same.
130        ws.set_binary_type(BinaryType::Arraybuffer);
131
132        let (sender, receiver) = mpsc::unbounded();
133
134        let open_callback: Closure<dyn FnMut()> = {
135            let waker = Rc::clone(&waker);
136            Closure::wrap(Box::new(move || {
137                if let Some(waker) = waker.borrow_mut().take() {
138                    waker.wake();
139                }
140            }) as Box<dyn FnMut()>)
141        };
142
143        let add_event_listener_options = web_sys::AddEventListenerOptions::new();
144        add_event_listener_options.set_once(true);
145        ws.add_event_listener_with_callback_and_add_event_listener_options(
146            "open",
147            open_callback.as_ref().unchecked_ref(),
148            &add_event_listener_options,
149        )
150        .map_err(js_to_js_error)?;
151
152        let message_callback: Closure<dyn FnMut(MessageEvent)> = {
153            let sender = sender.clone();
154            Closure::wrap(Box::new(move |e: MessageEvent| {
155                let msg = parse_message(e);
156                let _ = sender.unbounded_send(StreamMessage::Message(msg));
157            }) as Box<dyn FnMut(MessageEvent)>)
158        };
159
160        ws.add_event_listener_with_callback("message", message_callback.as_ref().unchecked_ref())
161            .map_err(js_to_js_error)?;
162
163        let error_callback: Closure<dyn FnMut(web_sys::Event)> = {
164            let sender = sender.clone();
165            let waker = Rc::clone(&waker);
166            Closure::wrap(Box::new(move |_e: web_sys::Event| {
167                if let Some(waker) = waker.borrow_mut().take() {
168                    waker.wake();
169                }
170                let _ = sender.unbounded_send(StreamMessage::ErrorEvent);
171            }) as Box<dyn FnMut(web_sys::Event)>)
172        };
173
174        ws.add_event_listener_with_callback("error", error_callback.as_ref().unchecked_ref())
175            .map_err(js_to_js_error)?;
176
177        let close_callback: Closure<dyn FnMut(web_sys::CloseEvent)> = {
178            Closure::wrap(Box::new(move |e: web_sys::CloseEvent| {
179                let close_event = CloseEvent {
180                    code: e.code(),
181                    reason: e.reason(),
182                    was_clean: e.was_clean(),
183                };
184                let _ = sender.unbounded_send(StreamMessage::CloseEvent(close_event));
185                let _ = sender.unbounded_send(StreamMessage::ConnectionClose);
186            }) as Box<dyn FnMut(web_sys::CloseEvent)>)
187        };
188
189        let add_event_listener_options = web_sys::AddEventListenerOptions::new();
190        add_event_listener_options.set_once(true);
191        ws.add_event_listener_with_callback_and_add_event_listener_options(
192            "close",
193            close_callback.as_ref().unchecked_ref(),
194            &add_event_listener_options,
195        )
196        .map_err(js_to_js_error)?;
197
198        Ok(Self {
199            ws,
200            sink_waker: waker,
201            message_receiver: receiver,
202            closures: (
203                open_callback,
204                message_callback,
205                error_callback,
206                close_callback,
207            ),
208            #[cfg(feature = "io-util")]
209            read_pending_bytes: None,
210        })
211    }
212
213    /// Closes the websocket.
214    ///
215    /// See the [MDN Documentation](https://developer.mozilla.org/en-US/docs/Web/API/WebSocket/close#parameters)
216    /// to learn about parameters passed to this function and when it can return an `Err(_)`
217    pub fn close(self, code: Option<u16>, reason: Option<&str>) -> Result<(), JsError> {
218        let result = match (code, reason) {
219            (None, None) => self.ws.close(),
220            (Some(code), None) => self.ws.close_with_code(code),
221            (Some(code), Some(reason)) => self.ws.close_with_code_and_reason(code, reason),
222            // default code is 1005 so we use it,
223            // see: https://developer.mozilla.org/en-US/docs/Web/API/WebSocket/close#parameters
224            (None, Some(reason)) => self.ws.close_with_code_and_reason(1005, reason),
225        };
226        result.map_err(js_to_js_error)
227    }
228
229    /// The current state of the websocket.
230    pub fn state(&self) -> State {
231        let ready_state = self.ws.ready_state();
232        match ready_state {
233            0 => State::Connecting,
234            1 => State::Open,
235            2 => State::Closing,
236            3 => State::Closed,
237            _ => unreachable!(),
238        }
239    }
240
241    /// The extensions in use.
242    pub fn extensions(&self) -> String {
243        self.ws.extensions()
244    }
245
246    /// The sub-protocol in use.
247    pub fn protocol(&self) -> String {
248        self.ws.protocol()
249    }
250
251    /// Number of pending, unsent bytes queued, but not transmitted to the network.
252    pub fn buffered_amount(&self) -> u32 {
253        self.ws.buffered_amount()
254    }
255}
256
257impl TryFrom<web_sys::WebSocket> for WebSocket {
258    type Error = JsError;
259
260    fn try_from(ws: web_sys::WebSocket) -> Result<Self, Self::Error> {
261        Self::setup(Ok(ws))
262    }
263}
264
265#[derive(Clone)]
266enum StreamMessage {
267    ErrorEvent,
268    CloseEvent(CloseEvent),
269    Message(Message),
270    ConnectionClose,
271}
272
273fn parse_message(event: MessageEvent) -> Message {
274    if let Ok(array_buffer) = event.data().dyn_into::<js_sys::ArrayBuffer>() {
275        let array = js_sys::Uint8Array::new(&array_buffer);
276        Message::Bytes(array.to_vec())
277    } else if let Ok(txt) = event.data().dyn_into::<js_sys::JsString>() {
278        Message::Text(String::from(&txt))
279    } else {
280        unreachable!("message event, received Unknown: {:?}", event.data());
281    }
282}
283
284impl Sink<Message> for WebSocket {
285    type Error = WebSocketError;
286
287    fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
288        let ready_state = self.ws.ready_state();
289        if ready_state == 0 {
290            *self.sink_waker.borrow_mut() = Some(cx.waker().clone());
291            Poll::Pending
292        } else {
293            Poll::Ready(Ok(()))
294        }
295    }
296
297    fn start_send(self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
298        let result = match item {
299            Message::Bytes(bytes) => self.ws.send_with_blob(
300                &web_sys::Blob::new_with_u8_array_sequence(&js_sys::Array::of1(
301                    &js_sys::Uint8Array::from(bytes.as_slice()),
302                ))
303                .map_err(|e| WebSocketError::MessageSendError(js_to_js_error(e)))?,
304            ),
305            Message::Text(message) => self.ws.send_with_str(&message),
306        };
307        match result {
308            Ok(_) => Ok(()),
309            Err(e) => Err(WebSocketError::MessageSendError(js_to_js_error(e))),
310        }
311    }
312
313    fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
314        Poll::Ready(Ok(()))
315    }
316
317    fn poll_close(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
318        Poll::Ready(Ok(()))
319    }
320}
321
322impl Stream for WebSocket {
323    type Item = Result<Message, WebSocketError>;
324
325    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
326        let msg = ready!(self.project().message_receiver.poll_next(cx));
327        match msg {
328            Some(StreamMessage::Message(msg)) => Poll::Ready(Some(Ok(msg))),
329            Some(StreamMessage::ErrorEvent) => {
330                Poll::Ready(Some(Err(WebSocketError::ConnectionError)))
331            }
332            Some(StreamMessage::CloseEvent(e)) => {
333                Poll::Ready(Some(Err(WebSocketError::ConnectionClose(e))))
334            }
335            Some(StreamMessage::ConnectionClose) => Poll::Ready(None),
336            None => Poll::Ready(None),
337        }
338    }
339}
340
341#[pinned_drop]
342impl PinnedDrop for WebSocket {
343    fn drop(self: Pin<&mut Self>) {
344        self.ws.close().unwrap();
345
346        for (ty, cb) in [
347            ("open", self.closures.0.as_ref()),
348            ("message", self.closures.1.as_ref()),
349            ("error", self.closures.2.as_ref()),
350        ] {
351            let _ = self
352                .ws
353                .remove_event_listener_with_callback(ty, cb.unchecked_ref());
354        }
355
356        let close_event_init = web_sys::CloseEventInit::new();
357        close_event_init.set_code(1000);
358        close_event_init.set_reason("client dropped");
359        if let Ok(close_event) =
360            web_sys::CloseEvent::new_with_event_init_dict("close", &close_event_init)
361        {
362            let _ = self.ws.dispatch_event(&close_event);
363        }
364    }
365}
366
367#[cfg(test)]
368mod tests {
369    use crate::http::Request;
370
371    use super::*;
372    use futures::{SinkExt, StreamExt};
373    use wasm_bindgen_test::*;
374
375    wasm_bindgen_test_configure!(run_in_browser);
376
377    #[wasm_bindgen_test]
378    async fn can_build_url_with_parameters_and_ampersand_is_not_added() {
379        let url = "http://something.com/get?param1=value1";
380        let request = Request::get(url).build().unwrap();
381        assert_eq!(
382            request.url(),
383            url,
384            "url of the built request should be equal to the parameter provided {url}"
385        );
386    }
387
388    #[wasm_bindgen_test]
389    async fn websocket_works() {
390        let ws_echo_server_url =
391            option_env!("WS_ECHO_SERVER_URL").expect("Did you set WS_ECHO_SERVER_URL?");
392
393        let ws = WebSocket::open(ws_echo_server_url).unwrap();
394        let (mut sender, mut receiver) = ws.split();
395
396        sender
397            .send(Message::Bytes(String::from("test 1").as_bytes().to_vec()))
398            .await
399            .unwrap();
400        sender
401            .send(Message::Text(String::from("test 2")))
402            .await
403            .unwrap();
404
405        // ignore first message
406        // the echo-server uses it to send it's info in the first message
407        let _ = receiver.next().await;
408
409        assert_eq!(
410            receiver.next().await.unwrap().unwrap(),
411            Message::Bytes("test 1".to_string().as_bytes().to_vec())
412        );
413        assert_eq!(
414            receiver.next().await.unwrap().unwrap(),
415            Message::Text("test 2".to_string())
416        );
417    }
418}