Skip to main content

websock_wasm/
connection.rs

1//! Browser WebSocket connection management.
2
3use std::cell::RefCell;
4use std::rc::Rc;
5use websock_proto::Bytes;
6use websock_proto::{CloseFrame, ConnectOptions, Error, Message, Result};
7
8use futures_channel::{mpsc, oneshot};
9use futures_util::StreamExt;
10use wasm_bindgen::JsCast;
11use wasm_bindgen::prelude::*;
12
13/// Establish a browser WebSocket connection.
14pub async fn connect(url: &str, opts: ConnectOptions) -> Result<Connection> {
15    opts.limits.validate()?;
16    let max_message_size = opts.limits.max_message_size;
17    let ws = if opts.protocols.is_empty() {
18        web_sys::WebSocket::new(url).map_err(js_err)?
19    } else {
20        let arr = js_sys::Array::new();
21        for p in &opts.protocols {
22            arr.push(&JsValue::from_str(p));
23        }
24        web_sys::WebSocket::new_with_str_sequence(url, &arr).map_err(js_err)?
25    };
26
27    ws.set_binary_type(web_sys::BinaryType::Arraybuffer);
28
29    // Channel used to deliver messages to the consumer.
30    let (tx, rx) = mpsc::channel::<Result<Message>>(64);
31
32    // Handle the connection process.
33    let (open_tx, open_rx) = oneshot::channel::<Result<()>>();
34    let open_tx_cell: Rc<RefCell<Option<oneshot::Sender<Result<()>>>>> =
35        Rc::new(RefCell::new(Some(open_tx)));
36
37    let open_tx_cell_onopen = Rc::clone(&open_tx_cell);
38    let wait_onopen = Closure::<dyn FnMut()>::new(move || {
39        if let Some(tx) = open_tx_cell_onopen.borrow_mut().take() {
40            let _ = tx.send(Ok(()));
41        }
42    });
43    ws.set_onopen(Some(wait_onopen.as_ref().unchecked_ref()));
44
45    let open_tx_cell_onerror = Rc::clone(&open_tx_cell);
46    let wait_onerror = Closure::<dyn FnMut(web_sys::Event)>::new(move |_e: web_sys::Event| {
47        if let Some(tx) = open_tx_cell_onerror.borrow_mut().take() {
48            let _ = tx.send(Err(Error::Other("websocket error (before open)".into())));
49        }
50    });
51    ws.set_onerror(Some(wait_onerror.as_ref().unchecked_ref()));
52
53    let open_tx_cell_onclose = Rc::clone(&open_tx_cell);
54    let wait_onclose =
55        Closure::<dyn FnMut(web_sys::CloseEvent)>::new(move |_e: web_sys::CloseEvent| {
56            if let Some(tx) = open_tx_cell_onclose.borrow_mut().take() {
57                let _ = tx.send(Err(Error::Closed));
58            }
59        });
60    ws.set_onclose(Some(wait_onclose.as_ref().unchecked_ref()));
61
62    // Wait until the connection is opened or fails.
63    let open_res = open_rx.await;
64
65    // Always unset the connection process handlers.
66    ws.set_onopen(None);
67    ws.set_onerror(None);
68    ws.set_onclose(None);
69
70    // Drop closures AFTER unsetting.
71    drop(wait_onopen);
72    drop(wait_onerror);
73    drop(wait_onclose);
74
75    match open_res {
76        Ok(Ok(())) => {}
77        Ok(Err(e)) => return Err(e),
78        Err(_) => return Err(Error::Other("onopen waiter dropped".into())),
79    }
80
81    // Set up message/error/close handlers.
82    let mut tx_msg = tx.clone();
83    let ws_onmessage = ws.clone();
84    let onmessage =
85        Closure::<dyn FnMut(web_sys::MessageEvent)>::new(move |e: web_sys::MessageEvent| {
86            let data = e.data();
87
88            if let Some(s) = data.as_string() {
89                if s.len() > max_message_size {
90                    let _ = tx_msg.try_send(Err(Error::Protocol(
91                        "websocket message exceeds max_message_size".into(),
92                    )));
93                    tx_msg.close_channel();
94                    let _ = ws_onmessage.close();
95                    return;
96                }
97                if tx_msg.try_send(Ok(Message::Text(s))).is_err() {
98                    tx_msg.close_channel();
99                    let _ = ws_onmessage.close();
100                }
101                return;
102            }
103
104            if data.is_instance_of::<js_sys::ArrayBuffer>() {
105                let ab: js_sys::ArrayBuffer = data.unchecked_into();
106                let u8arr = js_sys::Uint8Array::new(&ab);
107                if u8arr.length() as usize > max_message_size {
108                    let _ = tx_msg.try_send(Err(Error::Protocol(
109                        "websocket message exceeds max_message_size".into(),
110                    )));
111                    tx_msg.close_channel();
112                    let _ = ws_onmessage.close();
113                    return;
114                }
115                let mut buf = vec![0u8; u8arr.length() as usize];
116                u8arr.copy_to(&mut buf);
117                if tx_msg
118                    .try_send(Ok(Message::Binary(Bytes::from(buf))))
119                    .is_err()
120                {
121                    tx_msg.close_channel();
122                    let _ = ws_onmessage.close();
123                }
124                return;
125            }
126
127            let _ = tx_msg.try_send(Err(Error::Protocol("unsupported message type".into())));
128            tx_msg.close_channel();
129            let _ = ws_onmessage.close();
130        });
131    ws.set_onmessage(Some(onmessage.as_ref().unchecked_ref()));
132
133    let mut tx_err = tx.clone();
134    let onerror = Closure::<dyn FnMut(web_sys::Event)>::new(move |_e: web_sys::Event| {
135        if tx_err
136            .try_send(Err(Error::Other("websocket error".into())))
137            .is_err()
138        {
139            tx_err.close_channel();
140        }
141    });
142    ws.set_onerror(Some(onerror.as_ref().unchecked_ref()));
143
144    let close_frame = Rc::new(RefCell::new(None));
145    let close_frame_handler = Rc::clone(&close_frame);
146    let mut tx_close = tx;
147    let onclose = Closure::<dyn FnMut(web_sys::CloseEvent)>::new(move |e: web_sys::CloseEvent| {
148        *close_frame_handler.borrow_mut() = Some(CloseFrame {
149            code: e.code(),
150            reason: e.reason(),
151        });
152        let _ = tx_close.try_send(Err(Error::Closed));
153        tx_close.close_channel();
154    });
155    ws.set_onclose(Some(onclose.as_ref().unchecked_ref()));
156    let negotiated_subprotocol = ws.protocol();
157
158    Ok(Connection {
159        ws: Rc::new(ws),
160        rx: Some(rx),
161        max_write_buffer_size: opts.limits.max_write_buffer_size,
162        negotiated_subprotocol,
163        close_frame,
164        _onmessage: Some(onmessage),
165        _onerror: Some(onerror),
166        _onclose: Some(onclose),
167    })
168}
169
170/// WebSocket connection wrapper for browser WebSockets.
171pub struct Connection {
172    pub(crate) ws: Rc<web_sys::WebSocket>,
173    pub(crate) rx: Option<mpsc::Receiver<Result<Message>>>,
174    pub(crate) max_write_buffer_size: usize,
175    pub(crate) negotiated_subprotocol: String,
176    pub(crate) close_frame: Rc<RefCell<Option<CloseFrame>>>,
177
178    pub(crate) _onmessage: Option<Closure<dyn FnMut(web_sys::MessageEvent)>>,
179    pub(crate) _onerror: Option<Closure<dyn FnMut(web_sys::Event)>>,
180    pub(crate) _onclose: Option<Closure<dyn FnMut(web_sys::CloseEvent)>>,
181}
182
183impl Connection {
184    /// Send a text or binary message.
185    pub async fn send(&mut self, msg: Message) -> Result<()> {
186        match msg {
187            Message::Text(s) => {
188                check_send_capacity(&self.ws, s.len(), self.max_write_buffer_size)?;
189                self.ws.send_with_str(&s).map_err(js_err)?;
190            }
191            Message::Binary(b) => {
192                check_send_capacity(&self.ws, b.len(), self.max_write_buffer_size)?;
193                self.ws.send_with_u8_array(b.as_ref()).map_err(js_err)?;
194            }
195        }
196        Ok(())
197    }
198
199    /// Receive the next text or binary message.
200    pub async fn recv(&mut self) -> Result<Message> {
201        let rx = self.rx.as_mut().ok_or(Error::Closed)?;
202        rx.next().await.ok_or(Error::Closed)?
203    }
204
205    /// Close the WebSocket connection and wait for the browser close event.
206    pub async fn close(&mut self) -> Result<()> {
207        if self.ws.ready_state() == web_sys::WebSocket::CLOSED {
208            self.rx = None;
209            return Ok(());
210        }
211        if self.ws.ready_state() != web_sys::WebSocket::CLOSING {
212            self.ws.close().map_err(js_err)?;
213        }
214
215        let result = loop {
216            let next = match self.rx.as_mut() {
217                Some(rx) => rx.next().await,
218                None => return Ok(()),
219            };
220            match next {
221                Some(Ok(_)) => continue,
222                Some(Err(Error::Closed)) | None => break Ok(()),
223                Some(Err(error)) => break Err(error),
224            }
225        };
226        self.rx = None;
227        result
228    }
229
230    /// Return the WebSocket subprotocol selected by the server, if any.
231    pub fn negotiated_subprotocol(&self) -> Option<&str> {
232        (!self.negotiated_subprotocol.is_empty()).then_some(self.negotiated_subprotocol.as_str())
233    }
234
235    /// Return the most recently received close-frame metadata, if any.
236    pub fn close_frame(&self) -> Option<CloseFrame> {
237        self.close_frame.borrow().clone()
238    }
239}
240
241impl websock_proto::WebSocketConnection for Connection {
242    fn send<'a>(&'a mut self, msg: Message) -> websock_proto::LocalBoxFuture<'a, Result<()>> {
243        Box::pin(async move { Connection::send(self, msg).await })
244    }
245
246    fn recv<'a>(&'a mut self) -> websock_proto::LocalBoxFuture<'a, Result<Message>> {
247        Box::pin(async move { Connection::recv(self).await })
248    }
249
250    fn close<'a>(&'a mut self) -> websock_proto::LocalBoxFuture<'a, Result<()>> {
251        Box::pin(async move { Connection::close(self).await })
252    }
253}
254
255impl Drop for Connection {
256    fn drop(&mut self) {
257        // If there are other Rc references, do not close.
258        if Rc::strong_count(&self.ws) != 1 {
259            return;
260        }
261
262        if self._onmessage.is_some() {
263            self.ws.set_onmessage(None);
264        }
265        if self._onerror.is_some() {
266            self.ws.set_onerror(None);
267        }
268        if self._onclose.is_some() {
269            self.ws.set_onclose(None);
270        }
271        self.ws.set_onopen(None);
272
273        let _ = self.ws.close();
274    }
275}
276
277/// Convert a JavaScript error into the shared error type.
278pub(crate) fn js_err(e: JsValue) -> Error {
279    Error::Other(format!("{e:?}"))
280}
281
282pub(crate) fn check_send_capacity(
283    ws: &web_sys::WebSocket,
284    message_len: usize,
285    max_write_buffer_size: usize,
286) -> Result<()> {
287    let buffered = usize::try_from(ws.buffered_amount()).unwrap_or(usize::MAX);
288    if buffered.saturating_add(message_len) > max_write_buffer_size {
289        return Err(Error::Other("websocket write buffer limit exceeded".into()));
290    }
291    Ok(())
292}