1#![allow(clippy::unused_async)]
2use std::rc::Rc;
3
4use crate::codec::{Decoder, Encoder};
5use crate::util::{BytePages, BytesMut};
6use crate::{Cfg, io::IoRef, io::Waiter, rt, time::sleep, util::select, ws};
7
8#[derive(Clone, Debug)]
9pub struct WsSink(Rc<WsSinkInner>);
15
16#[derive(Debug)]
17struct WsSinkInner {
18 io: IoRef,
19 codec: ws::Codec,
20 cfg: Cfg<ws::WsClientConfig>,
21}
22
23impl WsSink {
24 pub(crate) fn new(io: IoRef, codec: ws::Codec, cfg: Cfg<ws::WsClientConfig>) -> Self {
25 Self(Rc::new(WsSinkInner { io, codec, cfg }))
26 }
27
28 pub fn io(&self) -> &IoRef {
30 &self.0.io
31 }
32
33 pub(crate) fn codec(&self) -> &ws::Codec {
35 &self.0.codec
36 }
37
38 pub(crate) fn is_closed(&self) -> bool {
40 self.0.codec.is_closed()
41 }
42
43 pub(crate) fn start_close_timeout(&self) {
44 if self.0.cfg.close_timeout.non_zero() {
45 let io = self.0.io.clone();
46 let close_timeout = self.0.cfg.close_timeout;
47 rt::spawn(async move {
48 select(sleep(close_timeout), io.on_disconnect()).await;
49 if io.is_active() {
50 io.close();
51 }
52 });
53 }
54 }
55
56 pub async fn send(&self, item: ws::Message) -> Result<(), ws::error::ProtocolError> {
66 let close = matches!(item, ws::Message::Close(_));
67
68 if matches!(
69 item,
70 ws::Message::Text(_) | ws::Message::Binary(_) | ws::Message::Continuation(_)
71 ) {
72 let _ = self.0.io.write_ready().await;
74 }
75
76 if let Err(e) = self.0.io.encode(item, &self.0.codec) {
77 Err(e)
78 } else {
79 if close {
80 self.start_close_timeout();
81 }
82 Ok(())
83 }
84 }
85
86 pub fn on_disconnect(&self) -> Waiter<'static> {
88 self.0.io.on_disconnect()
89 }
90}
91
92impl Encoder for WsSink {
93 type Item = ws::Message;
94 type Error = ws::error::ProtocolError;
95
96 fn encode(&self, item: ws::Message, dst: &mut BytePages) -> Result<(), Self::Error> {
97 self.0.codec.encode(item, dst)
98 }
99}
100
101impl Decoder for WsSink {
102 type Item = ws::Frame;
103 type Error = ws::error::ProtocolError;
104
105 fn decode(&self, src: &mut BytesMut) -> Result<Option<ws::Frame>, Self::Error> {
106 self.0.codec.decode(src)
107 }
108
109 fn decode_eof(&self, src: &mut BytesMut) -> Result<Option<ws::Frame>, Self::Error> {
110 self.0.codec.decode_eof(src)
111 }
112}
113
114#[cfg(test)]
115mod tests {
116 use super::*;
117 use crate::{SharedCfg, io::Io, testing::IoTest, time::Millis, time::timeout};
118
119 #[crate::rt_test]
120 async fn clones_share_codec_state() {
121 let (client, server) = IoTest::create();
122 client.remote_buffer_cap(4096);
123 let io = Io::new(server, SharedCfg::new("WS-TEST"));
124 let sink = WsSink::new(io.get_ref(), ws::Codec::new(), io.shared().get());
125 let sink2 = sink.clone();
126
127 let mut dst = BytePages::default();
129 sink2.encode(ws::Message::Close(None), &mut dst).unwrap();
130 assert!(sink.is_closed());
131 assert!(matches!(
132 sink.send(ws::Message::Text("t".into())).await,
133 Err(ws::error::ProtocolError::Closed)
134 ));
135 }
136
137 #[crate::rt_test]
138 async fn send_waits_for_write_backpressure() {
139 let (client, server) = IoTest::create();
140 client.remote_buffer_cap(0);
141 let cfg =
142 SharedCfg::new("WS-TEST").add(crate::io::IoConfig::new().set_write_backpressure(64));
143 let io = Io::new(server, cfg);
144 let sink = WsSink::new(io.get_ref(), ws::Codec::new(), io.shared().get());
145
146 sink.send(ws::Message::Binary(vec![0; 128].into()))
147 .await
148 .unwrap();
149 assert!(io.is_wr_backpressure());
150
151 sink.send(ws::Message::Ping("p".into())).await.unwrap();
153
154 let sent = std::rc::Rc::new(std::cell::Cell::new(false));
155 let (sink2, sent2) = (sink.clone(), sent.clone());
156 let handle = rt::spawn(async move {
157 sink2.send(ws::Message::Text("t".into())).await.unwrap();
158 sent2.set(true);
159 });
160 sleep(Millis(50)).await;
161 assert!(!sent.get());
162
163 client.remote_buffer_cap(4096);
165 let _ = client.read().await;
166 timeout(Millis(1000), handle)
167 .await
168 .expect("send was not released")
169 .unwrap();
170 assert!(sent.get());
171 }
172
173 #[crate::rt_test]
174 async fn send_on_disconnect_does_not_wait() {
175 let (client, server) = IoTest::create();
176 client.remote_buffer_cap(0);
177 let cfg =
178 SharedCfg::new("WS-TEST").add(crate::io::IoConfig::new().set_write_backpressure(64));
179 let io = Io::new(server, cfg);
180 let sink = WsSink::new(io.get_ref(), ws::Codec::new(), io.shared().get());
181
182 sink.send(ws::Message::Binary(vec![0; 128].into()))
183 .await
184 .unwrap();
185 assert!(io.is_wr_backpressure());
186
187 let (sink2, io2) = (sink.clone(), io.get_ref());
188 rt::spawn(async move {
189 sleep(Millis(20)).await;
190 io2.terminate();
191 });
192 timeout(Millis(1000), sink2.send(ws::Message::Text("t".into())))
193 .await
194 .expect("send was not released")
195 .unwrap();
196 }
197
198 #[crate::rt_test]
199 async fn close_timeout() {
200 let (client, server) = IoTest::create();
201 client.remote_buffer_cap(4096);
202 let cfg =
203 SharedCfg::new("WS-TEST").add(ws::WsClientConfig::new().set_close_timeout(Millis(50)));
204 let io = Io::new(server, cfg);
205 let sink = WsSink::new(io.get_ref(), ws::Codec::new(), io.shared().get());
206
207 let start = std::time::Instant::now();
208 sink.send(ws::Message::Close(None)).await.unwrap();
209 assert!(!client.is_server_dropped());
210 assert!(sink.io().is_active());
211
212 timeout(Millis(1000), async {
214 while sink.io().is_active() {
215 sleep(Millis(10)).await;
216 }
217 })
218 .await
219 .expect("close timeout did not close the connection");
220 assert!(start.elapsed() >= std::time::Duration::from_millis(50));
221 }
222}