websock_wasm/
connection.rs1use 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
13pub 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 let (tx, rx) = mpsc::channel::<Result<Message>>(64);
31
32 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 let open_res = open_rx.await;
64
65 ws.set_onopen(None);
67 ws.set_onerror(None);
68 ws.set_onclose(None);
69
70 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 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
170pub 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 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 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 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 pub fn negotiated_subprotocol(&self) -> Option<&str> {
232 (!self.negotiated_subprotocol.is_empty()).then_some(self.negotiated_subprotocol.as_str())
233 }
234
235 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 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
277pub(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}