use futures::{
channel::mpsc::{unbounded, UnboundedReceiver},
stream::StreamExt,
};
use send_wrapper::SendWrapper;
use std::{
pin::Pin,
str::FromStr,
task::{ready, Context, Poll},
};
use url::Url;
use wasm_bindgen::prelude::*;
use web_sys::{CloseEvent, Event, MessageEvent};
use crate::{
frame::{Frame, OpCode},
Result, WebSocketError,
};
type EventClosure = Closure<dyn FnMut(Event)>;
pub struct WebSocket {
stream: web_sys::WebSocket,
receiver: UnboundedReceiver<Result<Frame>>,
_handlers: SendWrapper<[EventClosure; 4]>,
}
impl WebSocket {
pub async fn connect(url: Url) -> Result<Self> {
let (socket, mut outcome) = Self::open(&url)?;
match outcome.next().await {
Some(Ok(())) => Ok(socket),
_ => Err(WebSocketError::ConnectionClosed),
}
}
fn open(url: &Url) -> Result<(Self, UnboundedReceiver<Result<()>>)> {
let stream = web_sys::WebSocket::new(url.as_str()).map_err(WebSocketError::Js)?;
stream.set_binary_type(web_sys::BinaryType::Arraybuffer);
let (tx, rx) = unbounded();
let (outcome_tx, outcome_rx) = unbounded();
let onopen = {
let outcome_tx = outcome_tx.clone();
EventClosure::new(move |_: Event| {
let _ = outcome_tx.unbounded_send(Ok(()));
})
};
let onerror = {
let outcome_tx = outcome_tx.clone();
EventClosure::new(move |_: Event| {
let _ = outcome_tx.unbounded_send(Err(WebSocketError::ConnectionClosed));
})
};
let onmessage = {
let tx = tx.clone();
EventClosure::new(move |event: Event| {
let data = event.unchecked_into::<MessageEvent>().data();
let maybe_fv = if data.has_type::<js_sys::JsString>() {
let str_value = data.unchecked_into::<js_sys::JsString>();
Some(Frame::text(String::from(str_value)))
} else if data.has_type::<js_sys::ArrayBuffer>() {
let buffer_value =
js_sys::Uint8Array::new(&data.unchecked_into::<js_sys::ArrayBuffer>())
.to_vec();
Some(Frame::binary(buffer_value))
} else {
None
};
if let Some(fv) = maybe_fv {
let _ = tx.unbounded_send(Ok(fv));
}
})
};
let onclose = EventClosure::new(move |event: Event| {
let close_event = event.unchecked_into::<CloseEvent>();
if !close_event.was_clean() {
web_sys::console::warn_1(
&js_sys::JsString::from_str("WebSocket CloseEvent wasClean() == false")
.unwrap(), );
}
let close_frame = Frame::close(close_event.code().into(), close_event.reason());
let _ = tx.unbounded_send(Ok(close_frame));
let _ = tx.unbounded_send(Err(WebSocketError::ConnectionClosed));
let _ = outcome_tx.unbounded_send(Err(WebSocketError::ConnectionClosed));
});
stream.set_onopen(Some(onopen.as_ref().unchecked_ref()));
stream.set_onerror(Some(onerror.as_ref().unchecked_ref()));
stream.set_onmessage(Some(onmessage.as_ref().unchecked_ref()));
stream.set_onclose(Some(onclose.as_ref().unchecked_ref()));
let socket = Self {
stream,
receiver: rx,
_handlers: SendWrapper::new([onopen, onerror, onmessage, onclose]),
};
Ok((socket, outcome_rx))
}
pub async fn next_frame(&mut self) -> Result<Frame> {
use futures::StreamExt;
match self.next().await {
Some(res) => res,
None => Err(WebSocketError::ConnectionClosed),
}
}
}
impl Drop for WebSocket {
fn drop(&mut self) {
self.stream.set_onopen(None);
self.stream.set_onerror(None);
self.stream.set_onmessage(None);
self.stream.set_onclose(None);
let _ = self.stream.close();
}
}
impl futures::Sink<Frame> for WebSocket {
type Error = WebSocketError;
fn poll_ready(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<()>> {
Poll::Ready(Ok(()))
}
fn start_send(self: Pin<&mut Self>, frame: Frame) -> Result<()> {
match frame.opcode() {
OpCode::Text => self
.stream
.send_with_str(frame.as_str())
.map_err(|_| WebSocketError::ConnectionClosed),
OpCode::Binary => self
.stream
.send_with_js_u8_array(&js_sys::Uint8Array::from(frame.payload().as_ref()))
.map_err(|_| WebSocketError::ConnectionClosed),
OpCode::Close => {
let code = frame.close_code().ok_or(WebSocketError::ConnectionClosed)?;
match frame.close_reason() {
Ok(Some(reason)) => self.stream.close_with_code_and_reason(code.into(), reason),
Ok(None) => self.stream.close_with_code(code.into()),
Err(err) => return Err(err),
}
.map_err(|_| WebSocketError::ConnectionClosed)
}
_ => Ok(()),
}
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_close(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<()>> {
let ret = self.stream.close().map_err(WebSocketError::Js);
Poll::Ready(ret)
}
}
impl futures::Stream for WebSocket {
type Item = Result<Frame>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
match ready!(self.receiver.poll_next_unpin(cx)) {
Some(Ok(message)) => Poll::Ready(Some(Ok(message))),
Some(Err(e)) => {
if matches!(e, WebSocketError::ConnectionClosed) {
Poll::Ready(None)
} else {
Poll::Ready(Some(Err(e)))
}
}
None => Poll::Ready(None),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[wasm_bindgen_test::wasm_bindgen_test]
async fn connect_to_closed_port_fails() {
let url = Url::parse("ws://127.0.0.1:1/").unwrap();
assert!(WebSocket::connect(url).await.is_err());
}
}