1use std::{cell::Cell, io, task::Poll};
3
4use crate::codec::{Decoder, Encoder};
5use crate::io::{Filter, FilterBuf, FilterLayer, Io, Layer};
6use crate::service::{Ctx, Service};
7
8use super::{CloseCode, CloseReason, Codec, Frame, Item, Message};
9
10bitflags::bitflags! {
11 #[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
12 struct Flags: u8 {
13 const CLOSED = 0b0001;
14 const PEER_CLOSED = 0b0010;
15 const PROTO_ERR = 0b0100;
16 }
17}
18
19#[derive(Clone, Debug)]
20pub struct WsTransport {
26 codec: Codec,
27 flags: Cell<Flags>,
28 peer_code: Cell<Option<CloseCode>>,
29}
30
31impl WsTransport {
32 pub fn create<F: Filter>(io: Io<F>, codec: Codec) -> Io<Layer<WsTransport, F>> {
34 io.add_filter(WsTransport {
35 codec,
36 flags: Cell::new(Flags::empty()),
37 peer_code: Cell::new(None),
38 })
39 }
40
41 fn insert_flags(&self, flags: Flags) {
42 let mut f = self.flags.get();
43 f.insert(flags);
44 self.flags.set(f);
45 }
46
47 fn send_close(&self, buf: &FilterBuf<'_>, code: Option<CloseCode>) {
48 if !self.flags.get().contains(Flags::CLOSED) {
49 self.insert_flags(Flags::CLOSED);
50 buf.with_write_buffers(|_, w_dst| {
51 let reason = code.map(CloseReason::from);
52 if self.codec.encode(Message::Close(reason), w_dst).is_err() {
53 let reason = CloseReason::from(CloseCode::Normal);
56 let _ = self.codec.encode(Message::Close(Some(reason)), w_dst);
57 }
58 });
59 }
60 }
61
62 fn fail(&self, buf: &FilterBuf<'_>, code: CloseCode, err: io::Error) -> io::Error {
66 self.insert_flags(Flags::PROTO_ERR);
67 self.send_close(buf, Some(code));
68 buf.io().close();
69 err
70 }
71}
72
73impl FilterLayer for WsTransport {
74 #[inline]
75 fn shutdown(&self, buf: &FilterBuf<'_>) -> io::Result<Poll<()>> {
76 let code = if self.flags.get().contains(Flags::PEER_CLOSED) {
78 self.peer_code.get()
79 } else {
80 Some(CloseCode::Normal)
81 };
82 self.send_close(buf, code);
83
84 let flags = self.flags.get();
87 if flags.intersects(Flags::PEER_CLOSED | Flags::PROTO_ERR) || buf.io().is_read_eof() {
88 Ok(Poll::Ready(()))
89 } else {
90 Ok(Poll::Pending)
91 }
92 }
93
94 fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
95 buf.with_read_buffers(|r_src, dst| {
96 if let Some(src) = r_src {
97 loop {
98 let Some(frame) = self.codec.decode(src).map_err(|e| {
99 log::trace!("Failed to decode ws codec frames: {e:?}");
100 let err = io::Error::new(io::ErrorKind::InvalidData, e);
101 self.fail(buf, CloseCode::Protocol, err)
102 })?
103 else {
104 break;
105 };
106
107 match frame {
108 Frame::Binary(bin)
110 | Frame::Continuation(
111 Item::FirstBinary(bin) | Item::Continue(bin) | Item::Last(bin),
112 ) => dst.extend_from_slice(&bin),
113 Frame::Continuation(Item::FirstText(_)) => {
114 return Err(self.fail(
115 buf,
116 CloseCode::Unsupported,
117 io::Error::new(
118 io::ErrorKind::InvalidData,
119 "WebSocket Text continuation frames are not supported",
120 ),
121 ));
122 }
123 Frame::Text(_) => {
124 return Err(self.fail(
125 buf,
126 CloseCode::Unsupported,
127 io::Error::new(
128 io::ErrorKind::InvalidData,
129 "WebSockets Text frames are not supported",
130 ),
131 ));
132 }
133 Frame::Ping(msg) => {
134 buf.with_write_buffers(|_, w_dst| {
135 let _ = self.codec.encode(Message::Pong(msg), w_dst);
136 });
137 }
138 Frame::Pong(_) => (),
139 Frame::Close(reason) => {
140 self.peer_code.set(reason.map(|r| r.code));
141 self.insert_flags(Flags::PEER_CLOSED);
142 buf.io().close();
143 break;
144 }
145 }
146 }
147 }
148 Ok(())
149 })
150 }
151
152 fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
153 buf.with_write_buffers(|w_src, w_dst| -> Result<(), super::error::ProtocolError> {
154 if self.flags.get().contains(Flags::CLOSED) {
155 w_src.clear();
157 return Ok(());
158 }
159 while let Some(page) = w_src.take() {
160 self.codec.encode_page(page, w_dst)?;
161 }
162 Ok(())
163 })
164 .map_err(|err| io::Error::new(io::ErrorKind::InvalidData, err))
165 }
166}
167
168#[derive(Clone, Debug)]
169pub struct WsTransportService {
171 codec: Codec,
172}
173
174impl WsTransportService {
175 pub fn new(codec: Codec) -> Self {
177 Self { codec }
178 }
179}
180
181impl<F: Filter> Service<(), Io<F>> for WsTransportService {
182 type Res = Io<Layer<WsTransport, F>>;
183 type Error = io::Error;
184
185 async fn call(&self, io: Io<F>, _: Ctx<'_, Self, ()>) -> Result<Self::Res, Self::Error> {
186 Ok(WsTransport::create(io, self.codec.clone()))
187 }
188}
189
190#[cfg(test)]
191mod tests {
192 use std::{cell::Cell, rc::Rc};
193
194 use super::*;
195 use crate::io::testing::IoTest;
196 use crate::time::{Millis, sleep};
197 use crate::util::{BytePages, Bytes, BytesMut};
198
199 #[derive(Debug)]
200 struct Passthrough;
201
202 impl FilterLayer for Passthrough {
203 fn process_read_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
204 buf.with_read_buffers(|src, dst| {
205 if let Some(src) = src {
206 dst.extend_from_slice(src);
207 src.clear();
208 }
209 });
210 Ok(())
211 }
212
213 fn process_write_buf(&self, buf: &FilterBuf<'_>) -> io::Result<()> {
214 buf.with_write_buffers(BytePages::move_to);
215 Ok(())
216 }
217 }
218
219 fn peer_close(code: Option<CloseCode>) -> Bytes {
220 let mut dst = BytePages::default();
221 Codec::new()
222 .set_client_mode()
223 .encode(Message::Close(code.map(CloseReason::from)), &mut dst)
224 .unwrap();
225 Bytes::from(dst)
226 }
227
228 fn start_shutdown<F: Filter>(io: Io<F>) -> Rc<Cell<bool>> {
229 let done = Rc::new(Cell::new(false));
230 let done2 = done.clone();
231 crate::rt::spawn(async move {
232 let _ = io.shutdown().await;
233 done2.set(true);
234 });
235 done
236 }
237
238 fn assert_close_sent(client: &IoTest) {
239 assert_close_code(client, Some(CloseCode::Normal));
240 }
241
242 fn assert_close_code(client: &IoTest, code: Option<CloseCode>) {
243 let mut data = BytesMut::from(&client.read_any()[..]);
244 assert_eq!(
245 Codec::new().set_client_mode().decode(&mut data).unwrap(),
246 Some(Frame::Close(code.map(CloseReason::from)))
247 );
248 }
249
250 async fn error_sends_close<F: Filter>(
251 client: IoTest,
252 io: Io<F>,
253 input: Bytes,
254 code: CloseCode,
255 ) {
256 client.remote_buffer_cap(1024);
257 let io = WsTransport::create(io, Codec::new());
258
259 client.write(input);
260 let err = io.recv(&crate::codec::BytesCodec).await.unwrap_err();
261 assert_eq!(err.into_inner().kind(), io::ErrorKind::InvalidData);
262 sleep(Millis(50)).await;
263
264 assert_close_code(&client, Some(code));
265 assert!(io.is_closed());
266 }
267
268 #[crate::rt_test]
269 async fn invalid_frame_sends_protocol_close() {
270 let (client, server) = IoTest::create();
272 let input = Bytes::from_static(&[0x82, 0x01, 0x00]);
273 error_sends_close(client, Io::from(server), input, CloseCode::Protocol).await;
274 }
275
276 #[crate::rt_test]
277 async fn invalid_frame_sends_protocol_close_over_inner_filter() {
278 let (client, server) = IoTest::create();
279 let io = Io::from(server).add_filter(Passthrough);
280 let input = Bytes::from_static(&[0x82, 0x01, 0x00]);
281 error_sends_close(client, io, input, CloseCode::Protocol).await;
282 }
283
284 #[crate::rt_test]
285 async fn text_frame_sends_unsupported_close() {
286 let (client, server) = IoTest::create();
287 let mut input = BytePages::default();
288 Codec::new()
289 .set_client_mode()
290 .encode(Message::Text("text".into()), &mut input)
291 .unwrap();
292 let input = Bytes::from(input);
293 error_sends_close(client, Io::from(server), input, CloseCode::Unsupported).await;
294 }
295
296 async fn shutdown_waits_for_peer_close<F: Filter>(client: IoTest, io: Io<F>) {
297 client.remote_buffer_cap(1024);
298 let io = WsTransport::create(io, Codec::new());
299 let done = start_shutdown(io);
300 sleep(Millis(50)).await;
301
302 assert_close_sent(&client);
303 assert!(!done.get());
304
305 client.write(peer_close(Some(CloseCode::Normal)));
306 sleep(Millis(50)).await;
307 assert!(done.get());
308 }
309
310 #[crate::rt_test]
311 async fn shutdown_waits_for_close_reply() {
312 let (client, server) = IoTest::create();
313 shutdown_waits_for_peer_close(client, Io::from(server)).await;
314 }
315
316 #[crate::rt_test]
317 async fn shutdown_waits_for_close_reply_over_inner_filter() {
318 let (client, server) = IoTest::create();
319 shutdown_waits_for_peer_close(client, Io::from(server).add_filter(Passthrough)).await;
320 }
321
322 #[crate::rt_test]
323 async fn shutdown_completes_on_read_eof() {
324 let (client, server) = IoTest::create();
325 client.remote_buffer_cap(1024);
326 let done = start_shutdown(WsTransport::create(Io::from(server), Codec::new()));
327 sleep(Millis(50)).await;
328 assert_close_sent(&client);
329 assert!(!done.get());
330
331 client.close().await;
332 sleep(Millis(50)).await;
333 assert!(done.get());
334 }
335
336 async fn peer_close_is_echoed(code: Option<CloseCode>, reply: Option<CloseCode>) {
337 let (client, server) = IoTest::create();
338 client.remote_buffer_cap(1024);
339 let io = WsTransport::create(Io::from(server), Codec::new());
340
341 client.write(peer_close(code));
342 assert!(io.recv(&crate::codec::BytesCodec).await.unwrap().is_none());
343 sleep(Millis(50)).await;
344
345 assert_close_code(&client, reply);
346 assert!(io.is_closed());
347 }
348
349 #[crate::rt_test]
350 async fn peer_close_code_is_echoed() {
351 peer_close_is_echoed(Some(CloseCode::Away), Some(CloseCode::Away)).await;
352 peer_close_is_echoed(None, None).await;
353
354 peer_close_is_echoed(Some(CloseCode::Extension), Some(CloseCode::Normal)).await;
356 }
357
358 fn peer_frames(msgs: Vec<Message>) -> Bytes {
359 let codec = Codec::new().set_client_mode();
360 let mut dst = BytePages::default();
361 for msg in msgs {
362 codec.encode(msg, &mut dst).unwrap();
363 }
364 Bytes::from(dst)
365 }
366
367 #[crate::rt_test]
368 async fn binary_frames_are_forwarded() {
369 let (client, server) = IoTest::create();
370 client.remote_buffer_cap(1024);
371 let io = crate::service::Pipeline::new((), WsTransportService::new(Codec::new()))
372 .call(Io::from(server))
373 .await
374 .unwrap();
375
376 client.write(peer_frames(vec![
377 Message::Binary(Bytes::from_static(b"one")),
378 Message::Pong(Bytes::from_static(b"ignored")),
379 Message::Continuation(Item::FirstBinary(Bytes::from_static(b"-two"))),
380 Message::Ping(Bytes::from_static(b"ping")),
381 Message::Continuation(Item::Continue(Bytes::from_static(b"-three"))),
382 Message::Continuation(Item::Last(Bytes::from_static(b"-four"))),
383 ]));
384 let mut data = BytesMut::new();
385 while data.len() < 18 {
386 data.extend_from_slice(&io.recv(&crate::codec::BytesCodec).await.unwrap().unwrap());
387 }
388 assert_eq!(&data[..], b"one-two-three-four");
389
390 sleep(Millis(50)).await;
392 let mut data = BytesMut::from(&client.read_any()[..]);
393 assert_eq!(
394 Codec::new().set_client_mode().decode(&mut data).unwrap(),
395 Some(Frame::Pong(Bytes::from_static(b"ping")))
396 );
397
398 io.send(Bytes::from_static(b"out"), &crate::codec::BytesCodec)
400 .await
401 .unwrap();
402 sleep(Millis(50)).await;
403 let mut data = BytesMut::from(&client.read_any()[..]);
404 assert_eq!(
405 Codec::new().set_client_mode().decode(&mut data).unwrap(),
406 Some(Frame::Binary(Bytes::from_static(b"out")))
407 );
408 }
409
410 #[crate::rt_test]
411 async fn text_continuation_sends_unsupported_close() {
412 let (client, server) = IoTest::create();
413 let input = peer_frames(vec![Message::Continuation(Item::FirstText(
414 Bytes::from_static(b"text"),
415 ))]);
416 error_sends_close(client, Io::from(server), input, CloseCode::Unsupported).await;
417 }
418}