use std::rc::Rc;
use impulse_utils::prelude::*;
use leptos::prelude::*;
use wasm_bindgen::JsCast;
use wasm_bindgen::prelude::*;
use web_sys::{BinaryType, CloseEvent, Event, MessageEvent, WebSocket};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum WebSocketReadyState {
Connecting,
Open,
Closing,
Closed,
}
#[derive(Clone, Debug)]
pub enum WebSocketMessage {
Text(String),
Binary(Vec<u8>),
}
struct WebSocketInner {
socket: WebSocket,
_on_open: Closure<dyn FnMut(Event)>,
_on_message: Closure<dyn FnMut(MessageEvent)>,
_on_error: Closure<dyn FnMut(Event)>,
_on_close: Closure<dyn FnMut(CloseEvent)>,
}
impl Drop for WebSocketInner {
fn drop(&mut self) {
self.socket.set_onopen(None);
self.socket.set_onmessage(None);
self.socket.set_onerror(None);
self.socket.set_onclose(None);
let _ = self.socket.close();
}
}
#[derive(Clone)]
pub struct WebSocketHandle {
pub state: ReadSignal<WebSocketReadyState>,
pub message: ReadSignal<Option<WebSocketMessage>>,
inner: Rc<WebSocketInner>,
}
impl WebSocketHandle {
pub fn ready_state(&self) -> WebSocketReadyState {
match self.inner.socket.ready_state() {
WebSocket::CONNECTING => WebSocketReadyState::Connecting,
WebSocket::OPEN => WebSocketReadyState::Open,
WebSocket::CLOSING => WebSocketReadyState::Closing,
_ => WebSocketReadyState::Closed,
}
}
pub fn send_text(&self, text: &str) -> CResult<()> {
self
.inner
.socket
.send_with_str(text)
.map_err(|e| ClientError::from_str(format!("WebSocket text send failed: {e:?}")))
}
pub fn send_binary(&self, data: &[u8]) -> CResult<()> {
self
.inner
.socket
.send_with_u8_array(data)
.map_err(|e| ClientError::from_str(format!("WebSocket binary send failed: {e:?}")))
}
pub fn close(&self) -> CResult<()> {
self
.inner
.socket
.close()
.map_err(|e| ClientError::from_str(format!("WebSocket close failed: {e:?}")))
}
pub fn close_with_reason(&self, code: u16, reason: &str) -> CResult<()> {
self
.inner
.socket
.close_with_code_and_reason(code, reason)
.map_err(|e| ClientError::from_str(format!("WebSocket close failed: {e:?}")))
}
pub fn raw(&self) -> &WebSocket {
&self.inner.socket
}
}
pub fn use_websocket(url: impl AsRef<str>) -> CResult<WebSocketHandle> {
use_websocket_inner(url.as_ref(), None)
}
pub fn use_websocket_with_protocols(url: impl AsRef<str>, protocols: &[&str]) -> CResult<WebSocketHandle> {
let arr = js_sys::Array::new();
for p in protocols {
arr.push(&JsValue::from_str(p));
}
use_websocket_inner(url.as_ref(), Some(arr.into()))
}
fn use_websocket_inner(url: &str, protocols: Option<JsValue>) -> CResult<WebSocketHandle> {
let socket = match protocols {
Some(p) => WebSocket::new_with_str_sequence(url, &p),
None => WebSocket::new(url),
}
.map_err(|e| ClientError::from_str(format!("Failed to open WebSocket {url}: {e:?}")))?;
socket.set_binary_type(BinaryType::Arraybuffer);
let (state, set_state) = signal(WebSocketReadyState::Connecting);
let (message, set_message) = signal::<Option<WebSocketMessage>>(None);
let on_open = Closure::<dyn FnMut(Event)>::new(move |_e: Event| {
set_state.set(WebSocketReadyState::Open);
});
socket.set_onopen(Some(on_open.as_ref().unchecked_ref()));
let on_message = Closure::<dyn FnMut(MessageEvent)>::new(move |e: MessageEvent| {
let data = e.data();
if let Some(text) = data.as_string() {
set_message.set(Some(WebSocketMessage::Text(text)));
} else if data.is_instance_of::<js_sys::ArrayBuffer>() {
let arr = js_sys::Uint8Array::new(&data);
let mut buf = vec![0u8; arr.length() as usize];
arr.copy_to(&mut buf);
set_message.set(Some(WebSocketMessage::Binary(buf)));
} else {
log::warn!("WebSocket received unsupported message data type");
}
});
socket.set_onmessage(Some(on_message.as_ref().unchecked_ref()));
let on_error = Closure::<dyn FnMut(Event)>::new(move |e: Event| {
log::error!("WebSocket error: {e:?}");
});
socket.set_onerror(Some(on_error.as_ref().unchecked_ref()));
let on_close = Closure::<dyn FnMut(CloseEvent)>::new(move |_e: CloseEvent| {
set_state.set(WebSocketReadyState::Closed);
});
socket.set_onclose(Some(on_close.as_ref().unchecked_ref()));
Ok(WebSocketHandle {
state,
message,
inner: Rc::new(WebSocketInner {
socket,
_on_open: on_open,
_on_message: on_message,
_on_error: on_error,
_on_close: on_close,
}),
})
}