1use std::fmt;
9use std::future::Future;
10
11use base64::Engine as _;
12use bytes::Bytes;
13use http::header::{
14 HeaderName, HeaderValue, CONNECTION, SEC_WEBSOCKET_ACCEPT, SEC_WEBSOCKET_KEY,
15 SEC_WEBSOCKET_PROTOCOL, UPGRADE,
16};
17use http_body_util::Empty;
18use hyper::upgrade::Upgraded;
19use hyper::{Request, Response, StatusCode};
20use hyper_util::rt::TokioIo;
21use sha1::{Digest, Sha1};
22
23const WS_GUID: &[u8] = b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11";
25
26pub const DEFAULT_SUBPROTOCOL: &str = "h2ts";
29
30#[derive(Debug, Clone, Copy, Default)]
33pub struct AcceptOptions {
34 pub allow_implicit_codec: bool,
45}
46
47pub type UpgradedIo = TokioIo<Upgraded>;
51
52#[derive(Debug)]
54pub enum WebSocketError {
55 NotUpgradeRequest,
58 UnsupportedSubprotocol,
63 Upgrade(hyper::Error),
65}
66
67impl WebSocketError {
68 pub fn rejection_response(&self) -> Response<Empty<Bytes>> {
82 let status = match self {
83 WebSocketError::NotUpgradeRequest => StatusCode::UPGRADE_REQUIRED,
84 WebSocketError::UnsupportedSubprotocol => StatusCode::BAD_REQUEST,
85 WebSocketError::Upgrade(_) => StatusCode::INTERNAL_SERVER_ERROR,
86 };
87 Response::builder()
88 .status(status)
89 .body(Empty::new())
90 .expect("static rejection response is well-formed")
91 }
92}
93
94impl fmt::Display for WebSocketError {
95 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
96 match self {
97 WebSocketError::NotUpgradeRequest => f.write_str("not a WebSocket upgrade request"),
98 WebSocketError::UnsupportedSubprotocol => {
99 f.write_str("client offered no supported subprotocol")
100 }
101 WebSocketError::Upgrade(e) => write!(f, "WebSocket upgrade failed: {e}"),
102 }
103 }
104}
105
106impl std::error::Error for WebSocketError {
107 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
108 match self {
109 WebSocketError::Upgrade(e) => Some(e),
110 WebSocketError::NotUpgradeRequest | WebSocketError::UnsupportedSubprotocol => None,
111 }
112 }
113}
114
115fn header_lists<B>(request: &Request<B>, header: HeaderName, needle: &str) -> bool {
118 request
119 .headers()
120 .get(header)
121 .and_then(|v| v.to_str().ok())
122 .map(|list| list.split(',').any(|t| t.trim().eq_ignore_ascii_case(needle)))
123 .unwrap_or(false)
124}
125
126pub fn is_upgrade_request<B>(request: &Request<B>) -> bool {
129 header_lists(request, UPGRADE, "websocket")
130 && header_lists(request, CONNECTION, "upgrade")
131 && request.headers().contains_key(SEC_WEBSOCKET_KEY)
132}
133
134pub fn offered_protocols<B>(request: &Request<B>) -> Vec<&str> {
138 request
139 .headers()
140 .get(SEC_WEBSOCKET_PROTOCOL)
141 .and_then(|v| v.to_str().ok())
142 .map(|list| {
143 list.split(',')
144 .map(str::trim)
145 .filter(|s| !s.is_empty())
146 .collect()
147 })
148 .unwrap_or_default()
149}
150
151#[derive(Debug, PartialEq, Eq)]
154enum Fallback {
155 Accept(Option<String>),
157 Reject,
160}
161
162fn fallback_subprotocol(offered: &[&str], allow_implicit_codec: bool) -> Fallback {
169 if let Some(p) = offered
170 .iter()
171 .find(|p| p.eq_ignore_ascii_case(DEFAULT_SUBPROTOCOL))
172 {
173 return Fallback::Accept(Some(p.to_string()));
174 }
175 if allow_implicit_codec {
176 return Fallback::Accept(offered.first().map(|p| p.to_string()));
177 }
178 Fallback::Reject
179}
180
181fn accept_key(key: &[u8]) -> String {
183 let mut hasher = Sha1::new();
184 hasher.update(key);
185 hasher.update(WS_GUID);
186 base64::engine::general_purpose::STANDARD.encode(hasher.finalize())
187}
188
189#[allow(clippy::type_complexity)]
210pub fn accept_with_options<B, F>(
211 request: &mut Request<B>,
212 select: F,
213 options: AcceptOptions,
214) -> Result<
215 (
216 Response<Empty<Bytes>>,
217 impl Future<Output = Result<UpgradedIo, WebSocketError>>,
218 ),
219 WebSocketError,
220>
221where
222 F: FnOnce(&[&str]) -> Option<String>,
223{
224 if !is_upgrade_request(request) {
225 return Err(WebSocketError::NotUpgradeRequest);
226 }
227 let key = request
229 .headers()
230 .get(SEC_WEBSOCKET_KEY)
231 .ok_or(WebSocketError::NotUpgradeRequest)?;
232 let accept_value = accept_key(key.as_bytes());
233
234 let offered = offered_protocols(request);
238 let decision = match select(&offered) {
239 Some(proto) => Fallback::Accept(Some(proto)),
240 None => fallback_subprotocol(&offered, options.allow_implicit_codec),
241 };
242 drop(offered); let chosen = match decision {
244 Fallback::Accept(proto) => proto,
245 Fallback::Reject => return Err(WebSocketError::UnsupportedSubprotocol),
246 };
247
248 let accept_header =
250 HeaderValue::from_str(&accept_value).expect("base64 is valid header ASCII");
251 let mut response = Response::builder()
252 .status(StatusCode::SWITCHING_PROTOCOLS)
253 .header(CONNECTION, HeaderValue::from_static("Upgrade"))
254 .header(UPGRADE, HeaderValue::from_static("websocket"))
255 .header(SEC_WEBSOCKET_ACCEPT, accept_header)
256 .body(Empty::<Bytes>::new())
257 .expect("static 101 response is well-formed");
258 if let Some(proto) = chosen {
259 if let Ok(value) = HeaderValue::from_str(&proto) {
261 response
262 .headers_mut()
263 .insert(SEC_WEBSOCKET_PROTOCOL, value);
264 }
265 }
266
267 let on_upgrade = hyper::upgrade::on(&mut *request);
270 let fut = async move {
271 let upgraded = on_upgrade.await.map_err(WebSocketError::Upgrade)?;
272 Ok(TokioIo::new(upgraded))
273 };
274 Ok((response, fut))
275}
276
277#[allow(clippy::type_complexity)]
292pub fn accept_with<B, F>(
293 request: &mut Request<B>,
294 select: F,
295) -> Result<
296 (
297 Response<Empty<Bytes>>,
298 impl Future<Output = Result<UpgradedIo, WebSocketError>>,
299 ),
300 WebSocketError,
301>
302where
303 F: FnOnce(&[&str]) -> Option<String>,
304{
305 accept_with_options(request, select, AcceptOptions::default())
306}
307
308#[allow(clippy::type_complexity)]
326pub fn accept<B>(
327 request: &mut Request<B>,
328) -> Result<
329 (
330 Response<Empty<Bytes>>,
331 impl Future<Output = Result<UpgradedIo, WebSocketError>>,
332 ),
333 WebSocketError,
334> {
335 accept_with(request, |_offered| None)
336}
337
338#[cfg(test)]
339mod tests {
340 use super::{
341 accept, accept_with, accept_with_options, fallback_subprotocol, is_upgrade_request,
342 offered_protocols, AcceptOptions, WebSocketError,
343 };
344 use http::header::{CONNECTION, SEC_WEBSOCKET_KEY, SEC_WEBSOCKET_PROTOCOL, UPGRADE};
345 use hyper::{Request, StatusCode};
346
347 fn with_protocol(protocol: Option<&str>) -> Request<()> {
348 let mut b = Request::builder();
349 if let Some(p) = protocol {
350 b = b.header(SEC_WEBSOCKET_PROTOCOL, p);
351 }
352 b.body(()).unwrap()
353 }
354
355 fn upgrade_request(offer: Option<&str>) -> Request<()> {
357 let mut b = Request::builder()
358 .header(UPGRADE, "websocket")
359 .header(CONNECTION, "Upgrade")
360 .header(SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==");
361 if let Some(o) = offer {
362 b = b.header(SEC_WEBSOCKET_PROTOCOL, o);
363 }
364 b.body(()).unwrap()
365 }
366
367 #[test]
368 fn offered_protocols_parses_the_list_in_order() {
369 assert_eq!(offered_protocols(&with_protocol(Some("h2ts"))), ["h2ts"]);
370 assert_eq!(
371 offered_protocols(&with_protocol(Some("chat, h2ts, binary"))),
372 ["chat", "h2ts", "binary"]
373 );
374 assert_eq!(offered_protocols(&with_protocol(Some(" h2ts "))), ["h2ts"]);
375 assert!(offered_protocols(&with_protocol(Some(""))).is_empty());
376 assert!(offered_protocols(&with_protocol(None)).is_empty());
377 }
378
379 #[test]
380 fn is_upgrade_request_requires_all_three_signals() {
381 let full = Request::builder()
382 .header(UPGRADE, "websocket")
383 .header(CONNECTION, "Upgrade")
384 .header(SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
385 .body(())
386 .unwrap();
387 assert!(is_upgrade_request(&full));
388
389 let listed = Request::builder()
391 .header(UPGRADE, "websocket")
392 .header(CONNECTION, "keep-alive, Upgrade")
393 .header(SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
394 .body(())
395 .unwrap();
396 assert!(is_upgrade_request(&listed));
397
398 let no_key = Request::builder()
400 .header(UPGRADE, "websocket")
401 .header(CONNECTION, "Upgrade")
402 .body(())
403 .unwrap();
404 assert!(!is_upgrade_request(&no_key));
405
406 assert!(!is_upgrade_request(&Request::builder().body(()).unwrap()));
408 }
409
410 #[test]
411 fn fallback_prefers_h2ts_then_optionally_the_first_offered() {
412 use super::Fallback::{Accept, Reject};
413
414 assert_eq!(fallback_subprotocol(&["h2ts"], false), Accept(Some("h2ts".into())));
416 assert_eq!(
417 fallback_subprotocol(&["chat", "h2ts"], false),
418 Accept(Some("h2ts".into())),
419 "h2ts is preferred even when it isn't first"
420 );
421 assert_eq!(
422 fallback_subprotocol(&["H2TS"], false),
423 Accept(Some("H2TS".into())),
424 "matched case-insensitively but echoed in the offered casing"
425 );
426
427 assert_eq!(fallback_subprotocol(&["chat"], false), Reject);
429 assert_eq!(fallback_subprotocol(&["chat", "binary"], false), Reject);
430 assert_eq!(fallback_subprotocol(&[], false), Reject);
431
432 assert_eq!(
434 fallback_subprotocol(&["chat", "binary"], true),
435 Accept(Some("chat".into()))
436 );
437 assert_eq!(fallback_subprotocol(&[], true), Accept(None));
439 }
440
441 #[test]
442 fn accept_echoes_h2ts_when_offered() {
443 let mut req = upgrade_request(Some("h2ts"));
444 let (resp, _fut) = accept(&mut req).unwrap();
445 assert_eq!(resp.status(), StatusCode::SWITCHING_PROTOCOLS);
446 assert_eq!(resp.headers().get(SEC_WEBSOCKET_PROTOCOL).unwrap(), "h2ts");
447
448 let mut req = upgrade_request(Some("chat, h2ts"));
450 let (resp, _fut) = accept(&mut req).unwrap();
451 assert_eq!(resp.headers().get(SEC_WEBSOCKET_PROTOCOL).unwrap(), "h2ts");
452 }
453
454 #[test]
455 fn accept_rejects_when_h2ts_absent_by_default() {
456 for offer in [Some("mystery"), Some("chat, binary"), None] {
459 let mut req = upgrade_request(offer);
460 let err = accept(&mut req).err().expect("should reject");
461 assert!(
462 matches!(&err, WebSocketError::UnsupportedSubprotocol),
463 "offer {offer:?}"
464 );
465 assert_eq!(err.rejection_response().status(), StatusCode::BAD_REQUEST);
466 }
467 }
468
469 #[test]
470 fn accept_with_honors_a_selection_even_without_h2ts() {
471 let mut req = upgrade_request(Some("chat, binary"));
473 let (resp, _fut) = accept_with(&mut req, |offered| {
474 offered.iter().find(|p| **p == "binary").map(|p| p.to_string())
475 })
476 .unwrap();
477 assert_eq!(resp.status(), StatusCode::SWITCHING_PROTOCOLS);
478 assert_eq!(resp.headers().get(SEC_WEBSOCKET_PROTOCOL).unwrap(), "binary");
479 }
480
481 #[test]
482 fn allow_implicit_codec_accepts_the_first_offered_codec() {
483 let opts = AcceptOptions {
484 allow_implicit_codec: true,
485 };
486 let mut req = upgrade_request(Some("mystery, other"));
488 let (resp, _fut) = accept_with_options(&mut req, |_| None, opts).unwrap();
489 assert_eq!(resp.status(), StatusCode::SWITCHING_PROTOCOLS);
490 assert_eq!(resp.headers().get(SEC_WEBSOCKET_PROTOCOL).unwrap(), "mystery");
491
492 let mut req = upgrade_request(None);
494 let (resp, _fut) = accept_with_options(&mut req, |_| None, opts).unwrap();
495 assert_eq!(resp.status(), StatusCode::SWITCHING_PROTOCOLS);
496 assert!(resp.headers().get(SEC_WEBSOCKET_PROTOCOL).is_none());
497 }
498
499 #[test]
500 fn accept_key_matches_rfc_6455_example() {
501 assert_eq!(
504 super::accept_key(b"dGhlIHNhbXBsZSBub25jZQ=="),
505 "s3pPLMBiTxaQ9kYGzzhZRbK+xOo="
506 );
507 }
508}