1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
22pub struct WebSocketConfig {
23 pub max_header_bytes: usize,
25 pub max_header_count: usize,
27 pub max_message_bytes: usize,
29 pub handshake_timeout: Duration,
31 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 #[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#[derive(Debug, Clone, PartialEq, Eq)]
85pub struct WebSocketUpgrade {
86 request: HttpRequestHead,
87 origin: Option<String>,
88}
89
90impl WebSocketUpgrade {
91 #[must_use]
93 pub const fn request(&self) -> &HttpRequestHead {
94 &self.request
95 }
96
97 #[must_use]
99 pub fn origin(&self) -> Option<&str> {
100 self.origin.as_deref()
101 }
102}
103
104pub 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
122pub 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}