1use 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#[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 #[cfg(feature = "io-util")]
64 pub(super) read_pending_bytes: Option<Vec<u8>>, }
66
67impl WebSocket {
68 pub fn open(url: &str) -> Result<Self, JsError> {
78 Self::setup(web_sys::WebSocket::new(url))
79 }
80
81 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 #[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 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 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 (None, Some(reason)) => self.ws.close_with_code_and_reason(1005, reason),
225 };
226 result.map_err(js_to_js_error)
227 }
228
229 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 pub fn extensions(&self) -> String {
243 self.ws.extensions()
244 }
245
246 pub fn protocol(&self) -> String {
248 self.ws.protocol()
249 }
250
251 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 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}