Skip to main content

ntex/ws/
transport.rs

1//! Binary-stream adaptation for WebSocket connections.
2use 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)]
20/// I/O filter that exposes binary WebSocket messages as a byte stream.
21///
22/// Incoming binary messages and fragments are forwarded as bytes. Text
23/// messages are rejected, ping frames receive automatic pong replies, and
24/// close frames close the underlying I/O stream.
25pub struct WsTransport {
26    codec: Codec,
27    flags: Cell<Flags>,
28    peer_code: Cell<Option<CloseCode>>,
29}
30
31impl WsTransport {
32    /// Adds a binary WebSocket transport filter to `io`.
33    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                    // an echoed code this side cannot send, for example 1010
54                    // from a server
55                    let reason = CloseReason::from(CloseCode::Normal);
56                    let _ = self.codec.encode(Message::Close(Some(reason)), w_dst);
57                }
58            });
59        }
60    }
61
62    /// Fails the connection: sends a close frame with `code` and starts a
63    /// graceful shutdown, so the frame is delivered before the connection
64    /// is closed with `err`.
65    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        // echo the peer's close code
77        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        // Wait for the peer's close frame. It cannot arrive after read eof,
85        // and a failed connection is not required to wait for it.
86        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                        // the codec enforces fragment ordering
109                        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                // nothing can be sent after the close frame
156                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)]
169/// Service that adds a [`WsTransport`] filter to an I/O stream.
170pub struct WsTransportService {
171    codec: Codec,
172}
173
174impl WsTransportService {
175    /// Creates a transport service using `codec`.
176    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        // an unmasked frame from a client
271        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        // servers cannot send 1010
355        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        // ping gets an automatic pong
391        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        // outgoing data is sent as binary frames
399        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}