1use 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
19pub 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 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 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 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;