Skip to main content

moirai_http/
websocket.rs

1//! Bounded RFC 6455 message framing over a Moirai async byte stream.
2
3use std::io;
4
5use moirai_async::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
6use moirai_async::timer::timeout;
7
8use crate::upgrade::WebSocketConfig;
9
10const FIN: u8 = 0x80;
11const RSV_MASK: u8 = 0x70;
12const OPCODE_MASK: u8 = 0x0f;
13const MASK: u8 = 0x80;
14const CLOSE: u8 = 0x8;
15const PING: u8 = 0x9;
16const PONG: u8 = 0xa;
17const BINARY: u8 = 0x2;
18
19/// A bounded, message-oriented RFC 6455 stream.
20pub struct WebSocketStream<S> {
21    stream: S,
22    config: WebSocketConfig,
23    prefix: Vec<u8>,
24    prefix_position: usize,
25    closed: bool,
26}
27
28impl<S> WebSocketStream<S> {
29    pub(crate) fn new(stream: S, config: WebSocketConfig, prefix: Vec<u8>) -> Self {
30        Self {
31            stream,
32            config,
33            prefix,
34            prefix_position: 0,
35            closed: false,
36        }
37    }
38}
39
40impl<S: AsyncRead + AsyncWrite + Unpin> WebSocketStream<S> {
41    /// Receive the next complete binary message.
42    ///
43    /// Ping frames are answered with pong frames and are not returned. Text,
44    /// continuation, fragmented, reserved-bit and unmasked client frames are
45    /// rejected, and non-minimal payload lengths are invalid. A close frame or
46    /// any frame I/O error transitions the stream to a terminal state; a
47    /// partially consumed or written frame is never retried on the same stream.
48    ///
49    /// # Errors
50    /// Returns malformed-frame, timeout, connection, or write failures.
51    pub async fn recv_message(&mut self) -> io::Result<Vec<u8>> {
52        if self.closed {
53            return Err(closed_error());
54        }
55        let timeout_duration = self.config.frame_timeout;
56        match timeout(timeout_duration, self.recv_message_inner()).await {
57            Ok(Ok(message)) => Ok(message),
58            Ok(Err(error)) => {
59                self.closed = true;
60                Err(error)
61            }
62            Err(_) => {
63                self.closed = true;
64                Err(timed_out("WebSocket frame receive"))
65            }
66        }
67    }
68
69    /// Send one unfragmented binary message to the peer.
70    ///
71    /// A frame I/O error or timeout transitions the stream to a terminal state
72    /// because a partial write cannot be retried without duplicating bytes.
73    ///
74    /// # Errors
75    /// Returns a size, timeout, connection, or write failure.
76    pub async fn send_binary(&mut self, payload: &[u8]) -> io::Result<()> {
77        if self.closed {
78            return Err(closed_error());
79        }
80        if payload.len() > self.config.max_message_bytes {
81            return Err(io::Error::new(
82                io::ErrorKind::InvalidData,
83                "WebSocket message exceeds configured byte bound",
84            ));
85        }
86        let timeout_duration = self.config.frame_timeout;
87        match timeout(timeout_duration, self.send_frame(BINARY, payload)).await {
88            Ok(Ok(())) => Ok(()),
89            Ok(Err(error)) => {
90                self.closed = true;
91                Err(error)
92            }
93            Err(_) => {
94                self.closed = true;
95                Err(timed_out("WebSocket binary send"))
96            }
97        }
98    }
99
100    /// Send a close frame and transition the stream to a terminal state.
101    ///
102    /// # Errors
103    /// Returns invalid-input, timeout, connection, or write failures.
104    pub async fn close(&mut self, code: u16, reason: &[u8]) -> io::Result<()> {
105        if self.closed {
106            return Ok(());
107        }
108        if reason.len() > 123 || std::str::from_utf8(reason).is_err() {
109            return Err(io::Error::new(
110                io::ErrorKind::InvalidInput,
111                "WebSocket close reason must be valid UTF-8 and at most 123 bytes",
112            ));
113        }
114        let capacity = reason.len().checked_add(2).ok_or_else(|| {
115            io::Error::new(io::ErrorKind::InvalidInput, "close reason is too large")
116        })?;
117        let mut payload = Vec::with_capacity(capacity);
118        payload.extend_from_slice(&code.to_be_bytes());
119        payload.extend_from_slice(reason);
120        validate_close_payload(&payload)?;
121        let timeout_duration = self.config.frame_timeout;
122        let result = match timeout(timeout_duration, self.send_frame(CLOSE, &payload)).await {
123            Ok(result) => result,
124            Err(_) => Err(timed_out("WebSocket close send")),
125        };
126        self.closed = true;
127        result
128    }
129
130    async fn recv_message_inner(&mut self) -> io::Result<Vec<u8>> {
131        loop {
132            let mut header = [0u8; 2];
133            self.read_exact(&mut header).await?;
134            let first = header[0];
135            let second = header[1];
136            if first & RSV_MASK != 0 {
137                return Err(protocol_error("WebSocket reserved bits are not supported"));
138            }
139            let opcode = first & OPCODE_MASK;
140            let final_frame = first & FIN != 0;
141            let masked = second & MASK != 0;
142            let length_code = second & 0x7f;
143            let payload_length = self.read_length(length_code).await?;
144            if (length_code == 126 && payload_length < 126)
145                || (length_code == 127 && payload_length <= 65_535)
146            {
147                return Err(protocol_error(
148                    "WebSocket payload length is not minimally encoded",
149                ));
150            }
151            let control = opcode >= CLOSE;
152            if control {
153                if !final_frame || payload_length > 125 {
154                    return Err(protocol_error("WebSocket control frame is invalid"));
155                }
156            } else if !final_frame {
157                return Err(protocol_error(
158                    "Fragmented WebSocket messages are not supported",
159                ));
160            }
161            if (!control && opcode != BINARY) || (control && !matches!(opcode, CLOSE | PING | PONG))
162            {
163                return Err(protocol_error(
164                    "WebSocket frame opcode is not a supported message or control frame",
165                ));
166            }
167            if !masked {
168                return Err(protocol_error("Client WebSocket frames must be masked"));
169            }
170            let mask = self.read_mask().await?;
171            if opcode == BINARY {
172                if payload_length > self.config.max_message_bytes {
173                    return Err(io::Error::new(
174                        io::ErrorKind::InvalidData,
175                        "WebSocket message exceeds configured byte bound",
176                    ));
177                }
178                return self.read_payload(payload_length, mask).await;
179            }
180            if opcode == PING {
181                let payload = self.read_payload(payload_length, mask).await?;
182                self.send_frame(PONG, &payload).await?;
183                continue;
184            }
185            if opcode == PONG {
186                let _ = self.read_payload(payload_length, mask).await?;
187                continue;
188            }
189            if opcode == CLOSE {
190                let payload = self.read_payload(payload_length, mask).await?;
191                validate_close_payload(&payload)?;
192                self.closed = true;
193                self.send_frame(CLOSE, &payload).await?;
194                return Err(io::Error::new(
195                    io::ErrorKind::UnexpectedEof,
196                    "WebSocket peer closed the connection",
197                ));
198            }
199            return Err(protocol_error("WebSocket frame opcode is not supported"));
200        }
201    }
202
203    async fn read_length(&mut self, length_code: u8) -> io::Result<usize> {
204        match length_code {
205            0..=125 => Ok(usize::from(length_code)),
206            126 => {
207                let mut bytes = [0u8; 2];
208                self.read_exact(&mut bytes).await?;
209                Ok(usize::from(u16::from_be_bytes(bytes)))
210            }
211            127 => {
212                let mut bytes = [0u8; 8];
213                self.read_exact(&mut bytes).await?;
214                let length = u64::from_be_bytes(bytes);
215                if length & (1u64 << 63) != 0 {
216                    return Err(protocol_error(
217                        "WebSocket payload length has its high bit set",
218                    ));
219                }
220                usize::try_from(length).map_err(|_| {
221                    io::Error::new(
222                        io::ErrorKind::InvalidData,
223                        "WebSocket payload length cannot be represented",
224                    )
225                })
226            }
227            _ => Err(protocol_error("WebSocket payload length code is invalid")),
228        }
229    }
230
231    async fn read_mask(&mut self) -> io::Result<[u8; 4]> {
232        let mut mask = [0u8; 4];
233        self.read_exact(&mut mask).await?;
234        Ok(mask)
235    }
236
237    async fn read_payload(&mut self, length: usize, mask: [u8; 4]) -> io::Result<Vec<u8>> {
238        let mut payload = vec![0u8; length];
239        self.read_exact(&mut payload).await?;
240        for (byte, mask_byte) in payload.iter_mut().zip(mask.iter().cycle()) {
241            *byte ^= *mask_byte;
242        }
243        Ok(payload)
244    }
245
246    async fn read_exact(&mut self, output: &mut [u8]) -> io::Result<()> {
247        let prefix_available = self.prefix.len().saturating_sub(self.prefix_position);
248        let from_prefix = output.len().min(prefix_available);
249        if from_prefix != 0 {
250            let start = self.prefix_position;
251            let end = start.checked_add(from_prefix).ok_or_else(|| {
252                io::Error::other("WebSocket prefix position arithmetic overflowed")
253            })?;
254            let source = self
255                .prefix
256                .get(start..end)
257                .ok_or_else(|| io::Error::other("WebSocket prefix bounds are inconsistent"))?;
258            let destination = output
259                .get_mut(..from_prefix)
260                .ok_or_else(|| io::Error::other("WebSocket output bounds are inconsistent"))?;
261            destination.copy_from_slice(source);
262            self.prefix_position = end;
263        }
264        let mut filled = from_prefix;
265        while filled < output.len() {
266            let destination = output
267                .get_mut(filled..)
268                .ok_or_else(|| io::Error::other("WebSocket output bounds are inconsistent"))?;
269            let count = self.stream.read(destination).await?;
270            if count == 0 {
271                return Err(io::Error::new(
272                    io::ErrorKind::UnexpectedEof,
273                    "connection closed inside a WebSocket frame",
274                ));
275            }
276            filled = filled
277                .checked_add(count)
278                .ok_or_else(|| io::Error::other("WebSocket read length overflow"))?;
279        }
280        Ok(())
281    }
282
283    async fn send_frame(&mut self, opcode: u8, payload: &[u8]) -> io::Result<()> {
284        let capacity = payload
285            .len()
286            .checked_add(10)
287            .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "message is too large"))?;
288        let mut frame = Vec::with_capacity(capacity);
289        frame.push(FIN | opcode);
290        match payload.len() {
291            length @ 0..=125 => frame.push(u8::try_from(length).expect("invariant: length <= 125")),
292            length @ 126..=65_535 => {
293                frame.push(126);
294                let length =
295                    u16::try_from(length).expect("invariant: extended WebSocket length fits u16");
296                frame.extend_from_slice(&length.to_be_bytes());
297            }
298            length => {
299                frame.push(127);
300                let length = u64::try_from(length).map_err(|_| {
301                    io::Error::new(io::ErrorKind::InvalidInput, "message is too large")
302                })?;
303                frame.extend_from_slice(&length.to_be_bytes());
304            }
305        }
306        frame.extend_from_slice(payload);
307        self.stream.write_all(&frame).await?;
308        self.stream.flush().await
309    }
310
311    #[cfg(test)]
312    pub(crate) fn initial_bytes(&self) -> &[u8] {
313        self.prefix
314            .get(self.prefix_position..)
315            .expect("invariant: WebSocket prefix position stays within the prefix")
316    }
317
318    #[cfg(test)]
319    pub(crate) fn output_bytes(&self) -> &[u8]
320    where
321        S: OutputBytes,
322    {
323        self.stream.output_bytes()
324    }
325}
326
327fn validate_close_payload(payload: &[u8]) -> io::Result<()> {
328    if payload.len() == 1 {
329        return Err(protocol_error("WebSocket close payload has one byte"));
330    }
331    if payload.len() >= 2 {
332        let code = payload
333            .get(..2)
334            .and_then(|bytes| bytes.try_into().ok())
335            .map(u16::from_be_bytes)
336            .ok_or_else(|| protocol_error("WebSocket close code is truncated"))?;
337        let valid_range = (1000..=2999).contains(&code) || (3000..=4999).contains(&code);
338        if !valid_range || matches!(code, 1004 | 1005 | 1006 | 1015) {
339            return Err(protocol_error("WebSocket close code is reserved"));
340        }
341        std::str::from_utf8(payload.get(2..).unwrap_or_default())
342            .map_err(|_| protocol_error("WebSocket close reason is not UTF-8"))?;
343    }
344    Ok(())
345}
346
347fn protocol_error(message: &str) -> io::Error {
348    io::Error::new(io::ErrorKind::InvalidData, message)
349}
350
351fn timed_out(operation: &str) -> io::Error {
352    io::Error::new(io::ErrorKind::TimedOut, operation)
353}
354
355fn closed_error() -> io::Error {
356    io::Error::new(io::ErrorKind::BrokenPipe, "WebSocket is closed")
357}
358
359#[cfg(test)]
360pub(crate) trait OutputBytes {
361    fn output_bytes(&self) -> &[u8];
362}
363
364#[cfg(test)]
365#[path = "websocket_tests.rs"]
366mod tests;