#![allow(clippy::unused_async)]
use std::rc::Rc;
use crate::codec::{Decoder, Encoder};
use crate::util::{BytePages, BytesMut};
use crate::{Cfg, io::IoRef, io::Waiter, rt, time::sleep, util::select, ws};
#[derive(Clone, Debug)]
pub struct WsSink(Rc<WsSinkInner>);
#[derive(Debug)]
struct WsSinkInner {
io: IoRef,
codec: ws::Codec,
cfg: Cfg<ws::WsClientConfig>,
}
impl WsSink {
pub(crate) fn new(io: IoRef, codec: ws::Codec, cfg: Cfg<ws::WsClientConfig>) -> Self {
Self(Rc::new(WsSinkInner { io, codec, cfg }))
}
pub fn io(&self) -> &IoRef {
&self.0.io
}
pub(crate) fn codec(&self) -> &ws::Codec {
&self.0.codec
}
pub(crate) fn is_closed(&self) -> bool {
self.0.codec.is_closed()
}
pub(crate) fn start_close_timeout(&self) {
if self.0.cfg.close_timeout.non_zero() {
let io = self.0.io.clone();
let close_timeout = self.0.cfg.close_timeout;
rt::spawn(async move {
select(sleep(close_timeout), io.on_disconnect()).await;
if io.is_active() {
io.close();
}
});
}
}
pub async fn send(&self, item: ws::Message) -> Result<(), ws::error::ProtocolError> {
let close = matches!(item, ws::Message::Close(_));
if matches!(
item,
ws::Message::Text(_) | ws::Message::Binary(_) | ws::Message::Continuation(_)
) {
let _ = self.0.io.write_ready().await;
}
if let Err(e) = self.0.io.encode(item, &self.0.codec) {
Err(e)
} else {
if close {
self.start_close_timeout();
}
Ok(())
}
}
pub fn on_disconnect(&self) -> Waiter<'static> {
self.0.io.on_disconnect()
}
}
impl Encoder for WsSink {
type Item = ws::Message;
type Error = ws::error::ProtocolError;
fn encode(&self, item: ws::Message, dst: &mut BytePages) -> Result<(), Self::Error> {
self.0.codec.encode(item, dst)
}
}
impl Decoder for WsSink {
type Item = ws::Frame;
type Error = ws::error::ProtocolError;
fn decode(&self, src: &mut BytesMut) -> Result<Option<ws::Frame>, Self::Error> {
self.0.codec.decode(src)
}
fn decode_eof(&self, src: &mut BytesMut) -> Result<Option<ws::Frame>, Self::Error> {
self.0.codec.decode_eof(src)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{SharedCfg, io::Io, testing::IoTest, time::Millis, time::timeout};
#[crate::rt_test]
async fn clones_share_codec_state() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(4096);
let io = Io::new(server, SharedCfg::new("WS-TEST"));
let sink = WsSink::new(io.get_ref(), ws::Codec::new(), io.shared().get());
let sink2 = sink.clone();
let mut dst = BytePages::default();
sink2.encode(ws::Message::Close(None), &mut dst).unwrap();
assert!(sink.is_closed());
assert!(matches!(
sink.send(ws::Message::Text("t".into())).await,
Err(ws::error::ProtocolError::Closed)
));
}
#[crate::rt_test]
async fn send_waits_for_write_backpressure() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(0);
let cfg = SharedCfg::new("WS-TEST").add(crate::io::IoConfig::new().set_write_buf(64));
let io = Io::new(server, cfg);
let sink = WsSink::new(io.get_ref(), ws::Codec::new(), io.shared().get());
sink.send(ws::Message::Binary(vec![0; 128].into()))
.await
.unwrap();
assert!(io.is_wr_backpressure());
sink.send(ws::Message::Ping("p".into())).await.unwrap();
let sent = std::rc::Rc::new(std::cell::Cell::new(false));
let (sink2, sent2) = (sink.clone(), sent.clone());
let handle = rt::spawn(async move {
sink2.send(ws::Message::Text("t".into())).await.unwrap();
sent2.set(true);
});
sleep(Millis(50)).await;
assert!(!sent.get());
client.remote_buffer_cap(4096);
let _ = client.read().await;
timeout(Millis(1000), handle)
.await
.expect("send was not released")
.unwrap();
assert!(sent.get());
}
#[crate::rt_test]
async fn send_on_disconnect_does_not_wait() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(0);
let cfg = SharedCfg::new("WS-TEST").add(crate::io::IoConfig::new().set_write_buf(64));
let io = Io::new(server, cfg);
let sink = WsSink::new(io.get_ref(), ws::Codec::new(), io.shared().get());
sink.send(ws::Message::Binary(vec![0; 128].into()))
.await
.unwrap();
assert!(io.is_wr_backpressure());
let (sink2, io2) = (sink.clone(), io.get_ref());
rt::spawn(async move {
sleep(Millis(20)).await;
io2.terminate();
});
timeout(Millis(1000), sink2.send(ws::Message::Text("t".into())))
.await
.expect("send was not released")
.unwrap();
}
#[crate::rt_test]
async fn close_timeout() {
let (client, server) = IoTest::create();
client.remote_buffer_cap(4096);
let cfg =
SharedCfg::new("WS-TEST").add(ws::WsClientConfig::new().set_close_timeout(Millis(50)));
let io = Io::new(server, cfg);
let sink = WsSink::new(io.get_ref(), ws::Codec::new(), io.shared().get());
let start = std::time::Instant::now();
sink.send(ws::Message::Close(None)).await.unwrap();
assert!(!client.is_server_dropped());
assert!(sink.io().is_active());
timeout(Millis(1000), async {
while sink.io().is_active() {
sleep(Millis(10)).await;
}
})
.await
.expect("close timeout did not close the connection");
assert!(start.elapsed() >= std::time::Duration::from_millis(50));
}
}