Skip to main content

moirai_http/
upgrade.rs

1//! HTTP/1.1 WebSocket upgrade validation and response serialization.
2
3use std::io;
4use std::time::Duration;
5
6use moirai_async::io::AsyncWriteExt;
7use moirai_async::timer::timeout;
8use moirai_crypto::{base64_decode, base64_encode, sha1};
9
10use crate::request::{HttpRequestHead, read_request_head};
11use crate::websocket::WebSocketStream;
12
13const WEBSOCKET_GUID: &[u8] = b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11";
14const DEFAULT_MAX_HEADER_BYTES: usize = 16 * 1024;
15const DEFAULT_MAX_HEADER_COUNT: usize = 32;
16const DEFAULT_MAX_MESSAGE_BYTES: usize = 16 * 1024 * 1024;
17const DEFAULT_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
18const DEFAULT_FRAME_TIMEOUT: Duration = Duration::from_secs(30);
19
20/// Resource and deadline policy for one WebSocket connection.
21#[derive(Debug, Clone, Copy, PartialEq, Eq)]
22pub struct WebSocketConfig {
23    /// Maximum bytes read for the HTTP request head.
24    pub max_header_bytes: usize,
25    /// Maximum number of HTTP request headers.
26    pub max_header_count: usize,
27    /// Maximum payload bytes in one binary message.
28    pub max_message_bytes: usize,
29    /// Deadline for reading and answering the HTTP upgrade.
30    pub handshake_timeout: Duration,
31    /// Deadline for each WebSocket frame operation.
32    pub frame_timeout: Duration,
33}
34
35impl Default for WebSocketConfig {
36    fn default() -> Self {
37        Self {
38            max_header_bytes: DEFAULT_MAX_HEADER_BYTES,
39            max_header_count: DEFAULT_MAX_HEADER_COUNT,
40            max_message_bytes: DEFAULT_MAX_MESSAGE_BYTES,
41            handshake_timeout: DEFAULT_HANDSHAKE_TIMEOUT,
42            frame_timeout: DEFAULT_FRAME_TIMEOUT,
43        }
44    }
45}
46
47impl WebSocketConfig {
48    /// Construct a WebSocket policy with explicit bounds and deadlines.
49    #[must_use]
50    pub const fn new(
51        max_header_bytes: usize,
52        max_header_count: usize,
53        max_message_bytes: usize,
54        handshake_timeout: Duration,
55        frame_timeout: Duration,
56    ) -> Self {
57        Self {
58            max_header_bytes,
59            max_header_count,
60            max_message_bytes,
61            handshake_timeout,
62            frame_timeout,
63        }
64    }
65
66    pub(crate) fn validate(self) -> io::Result<()> {
67        if self.max_header_bytes == 0
68            || self.max_header_count == 0
69            || self.max_header_count > crate::request::MAX_HEADER_SLOTS
70            || self.max_message_bytes == 0
71            || self.handshake_timeout.is_zero()
72            || self.frame_timeout.is_zero()
73        {
74            return Err(io::Error::new(
75                io::ErrorKind::InvalidInput,
76                "WebSocket limits and deadlines must be non-zero",
77            ));
78        }
79        Ok(())
80    }
81}
82
83/// Validated request metadata returned by a successful upgrade.
84#[derive(Debug, Clone, PartialEq, Eq)]
85pub struct WebSocketUpgrade {
86    request: HttpRequestHead,
87    origin: Option<String>,
88}
89
90impl WebSocketUpgrade {
91    /// Returns the validated HTTP request head.
92    #[must_use]
93    pub const fn request(&self) -> &HttpRequestHead {
94        &self.request
95    }
96
97    /// Returns the browser origin, when the peer sent an `Origin` header.
98    #[must_use]
99    pub fn origin(&self) -> Option<&str> {
100        self.origin.as_deref()
101    }
102}
103
104/// Perform a bounded HTTP/1.1 WebSocket upgrade on an async byte stream.
105///
106/// The returned stream preserves any bytes read after the HTTP delimiter, so a
107/// peer may pipeline its first WebSocket frame in the same TCP packet.
108///
109/// # Errors
110/// Returns invalid-input/data, timeout, or transport errors when the request
111/// does not satisfy the WebSocket upgrade contract or I/O fails.
112pub async fn accept_websocket<S>(
113    stream: S,
114    config: WebSocketConfig,
115) -> io::Result<(WebSocketStream<S>, WebSocketUpgrade)>
116where
117    S: moirai_async::io::AsyncRead + moirai_async::io::AsyncWrite + Unpin,
118{
119    accept_websocket_with_validator(stream, config, |_| Ok(())).await
120}
121
122/// Perform a bounded WebSocket upgrade after validating the parsed request.
123///
124/// `validator` runs after the RFC 6455 request checks and before any `101
125/// Switching Protocols` bytes are written. Consumers use this hook for
126/// trust-boundary checks such as an exact browser-origin policy; a rejected
127/// request therefore never becomes an acknowledged WebSocket session.
128///
129/// # Errors
130/// Returns the validator's error, invalid-input/data, timeout, or transport
131/// errors when the request does not satisfy the upgrade contract or I/O fails.
132pub async fn accept_websocket_with_validator<S, F>(
133    mut stream: S,
134    config: WebSocketConfig,
135    validator: F,
136) -> io::Result<(WebSocketStream<S>, WebSocketUpgrade)>
137where
138    S: moirai_async::io::AsyncRead + moirai_async::io::AsyncWrite + Unpin,
139    F: FnOnce(&HttpRequestHead) -> io::Result<()>,
140{
141    config.validate()?;
142    let (request, remainder) = timeout(
143        config.handshake_timeout,
144        read_request_head(
145            &mut stream,
146            config.max_header_bytes,
147            config.max_header_count,
148        ),
149    )
150    .await
151    .map_err(|_| timed_out("WebSocket handshake read"))??;
152    let key = validate_upgrade(&request)?;
153    validator(&request)?;
154    let accept = websocket_accept_key(key);
155    let response = format!(
156        "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\n\r\n"
157    );
158    timeout(
159        config.handshake_timeout,
160        stream.write_all(response.as_bytes()),
161    )
162    .await
163    .map_err(|_| timed_out("WebSocket handshake write"))??;
164    timeout(config.handshake_timeout, stream.flush())
165        .await
166        .map_err(|_| timed_out("WebSocket handshake flush"))??;
167
168    let origin = request.origin().map(str::to_owned);
169    let upgrade = WebSocketUpgrade { request, origin };
170    Ok((WebSocketStream::new(stream, config, remainder), upgrade))
171}
172
173fn validate_upgrade(request: &HttpRequestHead) -> io::Result<&str> {
174    if request.method() != "GET" {
175        return Err(invalid_upgrade("WebSocket upgrade requires GET"));
176    }
177    if request.version() != "HTTP/1.1" {
178        return Err(invalid_upgrade("WebSocket upgrade requires HTTP/1.1"));
179    }
180    if request.header("host").is_none_or(str::is_empty) {
181        return Err(invalid_upgrade("WebSocket upgrade requires Host"));
182    }
183    if request
184        .header("upgrade")
185        .is_none_or(|value| !value.eq_ignore_ascii_case("websocket"))
186    {
187        return Err(invalid_upgrade("Upgrade header must be websocket"));
188    }
189    let connection = request
190        .header("connection")
191        .ok_or_else(|| invalid_upgrade("Connection header is required"))?;
192    if !connection
193        .split(',')
194        .any(|token| token.trim().eq_ignore_ascii_case("upgrade"))
195    {
196        return Err(invalid_upgrade("Connection header must include Upgrade"));
197    }
198    if request.header("sec-websocket-version") != Some("13") {
199        return Err(invalid_upgrade("Sec-WebSocket-Version must be 13"));
200    }
201    if request.header("content-length").is_some() || request.header("transfer-encoding").is_some() {
202        return Err(invalid_upgrade(
203            "WebSocket upgrade must not carry an HTTP body",
204        ));
205    }
206    let key = request
207        .header("sec-websocket-key")
208        .ok_or_else(|| invalid_upgrade("Sec-WebSocket-Key is required"))?;
209    let decoded = base64_decode(key.as_bytes())
210        .filter(|bytes| bytes.len() == 16)
211        .ok_or_else(|| invalid_upgrade("Sec-WebSocket-Key must encode 16 bytes"))?;
212    if decoded.len() != 16 {
213        return Err(invalid_upgrade("Sec-WebSocket-Key length is invalid"));
214    }
215    Ok(key)
216}
217
218fn websocket_accept_key(key: &str) -> String {
219    let mut input = Vec::with_capacity(key.len());
220    input.extend_from_slice(key.as_bytes());
221    input.extend_from_slice(WEBSOCKET_GUID);
222    base64_encode(&sha1(&input))
223}
224
225fn invalid_upgrade(message: &str) -> io::Error {
226    io::Error::new(io::ErrorKind::InvalidData, message)
227}
228
229fn timed_out(operation: &str) -> io::Error {
230    io::Error::new(io::ErrorKind::TimedOut, operation)
231}
232
233#[cfg(test)]
234mod tests {
235    use super::*;
236    use crate::websocket::WebSocketStream;
237    use moirai_async::io::{AsyncRead, AsyncWrite};
238    use std::collections::VecDeque;
239    use std::future::Future;
240    use std::pin::Pin;
241    use std::sync::Arc;
242    use std::sync::atomic::{AtomicBool, Ordering};
243    use std::task::Waker;
244    use std::task::{Context, Poll};
245    use std::time::Duration;
246
247    struct MemoryStream {
248        input: VecDeque<u8>,
249        output: Vec<u8>,
250    }
251
252    impl MemoryStream {
253        fn new(input: &[u8]) -> Self {
254            Self {
255                input: input.iter().copied().collect(),
256                output: Vec::new(),
257            }
258        }
259    }
260
261    impl AsyncRead for MemoryStream {
262        fn poll_read(
263            mut self: Pin<&mut Self>,
264            _cx: &mut Context<'_>,
265            output: &mut [u8],
266        ) -> Poll<io::Result<usize>> {
267            let count = output.len().min(self.input.len());
268            for slot in output.iter_mut().take(count) {
269                let Some(byte) = self.input.pop_front() else {
270                    return Poll::Ready(Ok(0));
271                };
272                *slot = byte;
273            }
274            Poll::Ready(Ok(count))
275        }
276    }
277
278    impl AsyncWrite for MemoryStream {
279        fn poll_write(
280            mut self: Pin<&mut Self>,
281            _cx: &mut Context<'_>,
282            input: &[u8],
283        ) -> Poll<io::Result<usize>> {
284            let count = input.len().min(5);
285            let input = input
286                .get(..count)
287                .expect("invariant: test writer count is within input length");
288            self.output.extend_from_slice(input);
289            Poll::Ready(Ok(count))
290        }
291
292        fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
293            Poll::Ready(Ok(()))
294        }
295
296        fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
297            Poll::Ready(Ok(()))
298        }
299    }
300
301    impl crate::websocket::OutputBytes for MemoryStream {
302        fn output_bytes(&self) -> &[u8] {
303            &self.output
304        }
305    }
306
307    struct PendingStream {
308        dropped: Option<Arc<AtomicBool>>,
309    }
310
311    impl AsyncRead for PendingStream {
312        fn poll_read(
313            self: Pin<&mut Self>,
314            _cx: &mut Context<'_>,
315            _output: &mut [u8],
316        ) -> Poll<io::Result<usize>> {
317            let _ = self;
318            Poll::Pending
319        }
320    }
321
322    impl AsyncWrite for PendingStream {
323        fn poll_write(
324            self: Pin<&mut Self>,
325            _cx: &mut Context<'_>,
326            _input: &[u8],
327        ) -> Poll<io::Result<usize>> {
328            let _ = self;
329            Poll::Pending
330        }
331
332        fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
333            let _ = self;
334            Poll::Ready(Ok(()))
335        }
336
337        fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
338            let _ = self;
339            Poll::Ready(Ok(()))
340        }
341    }
342
343    impl Drop for PendingStream {
344        fn drop(&mut self) {
345            if let Some(dropped) = self.dropped.take() {
346                dropped.store(true, Ordering::Relaxed);
347            }
348        }
349    }
350
351    fn request(extra: &str) -> Vec<u8> {
352        let mut request = String::from(
353            "GET /metis HTTP/1.1\r\nHost: localhost\r\nUpgrade: websocket\r\nConnection: keep-alive, Upgrade\r\nSec-WebSocket-Version: 13\r\nSec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\nOrigin: http://127.0.0.1:8765\r\n",
354        );
355        if !extra.is_empty() {
356            request.push_str(extra);
357            request.push_str("\r\n");
358        }
359        request.push_str("\r\n");
360        request.into_bytes()
361    }
362
363    #[test]
364    fn valid_upgrade_emits_rfc_response_and_preserves_origin() {
365        let mut bytes = request("");
366        bytes.extend_from_slice(b"first-frame");
367        let stream = MemoryStream::new(&bytes);
368        let (stream, upgrade): (WebSocketStream<MemoryStream>, WebSocketUpgrade) =
369            moirai::block_on(accept_websocket(stream, WebSocketConfig::default()))
370                .expect("upgrade must succeed");
371        assert_eq!(upgrade.origin(), Some("http://127.0.0.1:8765"));
372        assert_eq!(stream.initial_bytes(), b"first-frame");
373        assert_eq!(
374            stream.output_bytes(),
375            b"HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: s3pPLMBiTxaQ9kYGzzhZRbK+xOo=\r\n\r\n"
376        );
377    }
378
379    #[test]
380    fn validator_rejects_before_switching_protocols_response() {
381        let called = Arc::new(AtomicBool::new(false));
382        let marker = Arc::clone(&called);
383        let result = moirai::block_on(accept_websocket_with_validator(
384            MemoryStream::new(&request("")),
385            WebSocketConfig::default(),
386            move |request| {
387                marker.store(request.origin().is_some(), Ordering::Relaxed);
388                Err(io::Error::new(
389                    io::ErrorKind::PermissionDenied,
390                    "origin denied",
391                ))
392            },
393        ));
394        let error = match result {
395            Ok(_) => panic!("validator rejection must stop before response"),
396            Err(error) => error,
397        };
398        assert_eq!(error.kind(), io::ErrorKind::PermissionDenied);
399        assert!(called.load(Ordering::Relaxed));
400    }
401
402    #[test]
403    fn invalid_upgrade_headers_are_rejected() {
404        for replacement in [
405            "Upgrade: http",
406            "Connection: keep-alive",
407            "Sec-WebSocket-Version: 12",
408            "Sec-WebSocket-Key: bad",
409            "Content-Length: 0",
410        ] {
411            let input = request(replacement);
412            let stream = MemoryStream::new(&input);
413            let error = match moirai::block_on(accept_websocket(stream, WebSocketConfig::default()))
414            {
415                Ok(_) => panic!("invalid upgrade must fail"),
416                Err(error) => error,
417            };
418            assert_eq!(error.kind(), io::ErrorKind::InvalidData);
419        }
420    }
421
422    #[test]
423    fn upgrade_requires_http_host() {
424        let mut input = request("");
425        let host = b"Host: localhost\r\n";
426        let start = input
427            .windows(host.len())
428            .position(|window| window == host)
429            .expect("test request contains Host");
430        let end = start
431            .checked_add(host.len())
432            .expect("test Host range fits request");
433        input.drain(start..end);
434        let error = match moirai::block_on(accept_websocket(
435            MemoryStream::new(&input),
436            WebSocketConfig::default(),
437        )) {
438            Ok(_) => panic!("missing Host must fail"),
439            Err(error) => error,
440        };
441        assert_eq!(error.kind(), io::ErrorKind::InvalidData);
442    }
443
444    #[test]
445    fn handshake_deadline_terminates_a_pending_peer() {
446        let config = WebSocketConfig::new(
447            1024,
448            8,
449            1024,
450            Duration::from_millis(10),
451            Duration::from_millis(10),
452        );
453        let error =
454            match moirai::block_on(accept_websocket(PendingStream { dropped: None }, config)) {
455                Ok(_) => panic!("pending handshake must time out"),
456                Err(error) => error,
457            };
458        assert_eq!(error.kind(), io::ErrorKind::TimedOut);
459    }
460
461    #[test]
462    fn dropping_handshake_future_drops_the_owned_stream() {
463        let dropped = Arc::new(AtomicBool::new(false));
464        let config = WebSocketConfig::new(
465            1024,
466            8,
467            1024,
468            Duration::from_secs(1),
469            Duration::from_secs(1),
470        );
471        let mut future = Box::pin(accept_websocket(
472            PendingStream {
473                dropped: Some(Arc::clone(&dropped)),
474            },
475            config,
476        ));
477        let waker = Waker::noop();
478        let mut context = Context::from_waker(waker);
479        assert!(future.as_mut().poll(&mut context).is_pending());
480        drop(future);
481        assert!(dropped.load(Ordering::Relaxed));
482    }
483}