use std::fmt;
use std::future::Future;
use base64::Engine as _;
use bytes::Bytes;
use http::header::{
HeaderName, HeaderValue, CONNECTION, SEC_WEBSOCKET_ACCEPT, SEC_WEBSOCKET_KEY,
SEC_WEBSOCKET_PROTOCOL, UPGRADE,
};
use http_body_util::Empty;
use hyper::upgrade::Upgraded;
use hyper::{Request, Response, StatusCode};
use hyper_util::rt::TokioIo;
use sha1::{Digest, Sha1};
const WS_GUID: &[u8] = b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11";
pub const DEFAULT_SUBPROTOCOL: &str = "h2ts";
#[derive(Debug, Clone, Copy, Default)]
pub struct AcceptOptions {
pub allow_implicit_codec: bool,
}
pub type UpgradedIo = TokioIo<Upgraded>;
#[derive(Debug)]
pub enum WebSocketError {
NotUpgradeRequest,
UnsupportedSubprotocol,
Upgrade(hyper::Error),
}
impl WebSocketError {
pub fn rejection_response(&self) -> Response<Empty<Bytes>> {
let status = match self {
WebSocketError::NotUpgradeRequest => StatusCode::UPGRADE_REQUIRED,
WebSocketError::UnsupportedSubprotocol => StatusCode::BAD_REQUEST,
WebSocketError::Upgrade(_) => StatusCode::INTERNAL_SERVER_ERROR,
};
Response::builder()
.status(status)
.body(Empty::new())
.expect("static rejection response is well-formed")
}
}
impl fmt::Display for WebSocketError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
WebSocketError::NotUpgradeRequest => f.write_str("not a WebSocket upgrade request"),
WebSocketError::UnsupportedSubprotocol => {
f.write_str("client offered no supported subprotocol")
}
WebSocketError::Upgrade(e) => write!(f, "WebSocket upgrade failed: {e}"),
}
}
}
impl std::error::Error for WebSocketError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
WebSocketError::Upgrade(e) => Some(e),
WebSocketError::NotUpgradeRequest | WebSocketError::UnsupportedSubprotocol => None,
}
}
}
fn header_lists<B>(request: &Request<B>, header: HeaderName, needle: &str) -> bool {
request
.headers()
.get(header)
.and_then(|v| v.to_str().ok())
.map(|list| list.split(',').any(|t| t.trim().eq_ignore_ascii_case(needle)))
.unwrap_or(false)
}
pub fn is_upgrade_request<B>(request: &Request<B>) -> bool {
header_lists(request, UPGRADE, "websocket")
&& header_lists(request, CONNECTION, "upgrade")
&& request.headers().contains_key(SEC_WEBSOCKET_KEY)
}
pub fn offered_protocols<B>(request: &Request<B>) -> Vec<&str> {
request
.headers()
.get(SEC_WEBSOCKET_PROTOCOL)
.and_then(|v| v.to_str().ok())
.map(|list| {
list.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
.collect()
})
.unwrap_or_default()
}
#[derive(Debug, PartialEq, Eq)]
enum Fallback {
Accept(Option<String>),
Reject,
}
fn fallback_subprotocol(offered: &[&str], allow_implicit_codec: bool) -> Fallback {
if let Some(p) = offered
.iter()
.find(|p| p.eq_ignore_ascii_case(DEFAULT_SUBPROTOCOL))
{
return Fallback::Accept(Some(p.to_string()));
}
if allow_implicit_codec {
return Fallback::Accept(offered.first().map(|p| p.to_string()));
}
Fallback::Reject
}
fn accept_key(key: &[u8]) -> String {
let mut hasher = Sha1::new();
hasher.update(key);
hasher.update(WS_GUID);
base64::engine::general_purpose::STANDARD.encode(hasher.finalize())
}
#[allow(clippy::type_complexity)]
pub fn accept_with_options<B, F>(
request: &mut Request<B>,
select: F,
options: AcceptOptions,
) -> Result<
(
Response<Empty<Bytes>>,
impl Future<Output = Result<UpgradedIo, WebSocketError>>,
),
WebSocketError,
>
where
F: FnOnce(&[&str]) -> Option<String>,
{
if !is_upgrade_request(request) {
return Err(WebSocketError::NotUpgradeRequest);
}
let key = request
.headers()
.get(SEC_WEBSOCKET_KEY)
.ok_or(WebSocketError::NotUpgradeRequest)?;
let accept_value = accept_key(key.as_bytes());
let offered = offered_protocols(request);
let decision = match select(&offered) {
Some(proto) => Fallback::Accept(Some(proto)),
None => fallback_subprotocol(&offered, options.allow_implicit_codec),
};
drop(offered); let chosen = match decision {
Fallback::Accept(proto) => proto,
Fallback::Reject => return Err(WebSocketError::UnsupportedSubprotocol),
};
let accept_header =
HeaderValue::from_str(&accept_value).expect("base64 is valid header ASCII");
let mut response = Response::builder()
.status(StatusCode::SWITCHING_PROTOCOLS)
.header(CONNECTION, HeaderValue::from_static("Upgrade"))
.header(UPGRADE, HeaderValue::from_static("websocket"))
.header(SEC_WEBSOCKET_ACCEPT, accept_header)
.body(Empty::<Bytes>::new())
.expect("static 101 response is well-formed");
if let Some(proto) = chosen {
if let Ok(value) = HeaderValue::from_str(&proto) {
response
.headers_mut()
.insert(SEC_WEBSOCKET_PROTOCOL, value);
}
}
let on_upgrade = hyper::upgrade::on(&mut *request);
let fut = async move {
let upgraded = on_upgrade.await.map_err(WebSocketError::Upgrade)?;
Ok(TokioIo::new(upgraded))
};
Ok((response, fut))
}
#[allow(clippy::type_complexity)]
pub fn accept_with<B, F>(
request: &mut Request<B>,
select: F,
) -> Result<
(
Response<Empty<Bytes>>,
impl Future<Output = Result<UpgradedIo, WebSocketError>>,
),
WebSocketError,
>
where
F: FnOnce(&[&str]) -> Option<String>,
{
accept_with_options(request, select, AcceptOptions::default())
}
#[allow(clippy::type_complexity)]
pub fn accept<B>(
request: &mut Request<B>,
) -> Result<
(
Response<Empty<Bytes>>,
impl Future<Output = Result<UpgradedIo, WebSocketError>>,
),
WebSocketError,
> {
accept_with(request, |_offered| None)
}
#[cfg(test)]
mod tests {
use super::{
accept, accept_with, accept_with_options, fallback_subprotocol, is_upgrade_request,
offered_protocols, AcceptOptions, WebSocketError,
};
use http::header::{CONNECTION, SEC_WEBSOCKET_KEY, SEC_WEBSOCKET_PROTOCOL, UPGRADE};
use hyper::{Request, StatusCode};
fn with_protocol(protocol: Option<&str>) -> Request<()> {
let mut b = Request::builder();
if let Some(p) = protocol {
b = b.header(SEC_WEBSOCKET_PROTOCOL, p);
}
b.body(()).unwrap()
}
fn upgrade_request(offer: Option<&str>) -> Request<()> {
let mut b = Request::builder()
.header(UPGRADE, "websocket")
.header(CONNECTION, "Upgrade")
.header(SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==");
if let Some(o) = offer {
b = b.header(SEC_WEBSOCKET_PROTOCOL, o);
}
b.body(()).unwrap()
}
#[test]
fn offered_protocols_parses_the_list_in_order() {
assert_eq!(offered_protocols(&with_protocol(Some("h2ts"))), ["h2ts"]);
assert_eq!(
offered_protocols(&with_protocol(Some("chat, h2ts, binary"))),
["chat", "h2ts", "binary"]
);
assert_eq!(offered_protocols(&with_protocol(Some(" h2ts "))), ["h2ts"]);
assert!(offered_protocols(&with_protocol(Some(""))).is_empty());
assert!(offered_protocols(&with_protocol(None)).is_empty());
}
#[test]
fn is_upgrade_request_requires_all_three_signals() {
let full = Request::builder()
.header(UPGRADE, "websocket")
.header(CONNECTION, "Upgrade")
.header(SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
.body(())
.unwrap();
assert!(is_upgrade_request(&full));
let listed = Request::builder()
.header(UPGRADE, "websocket")
.header(CONNECTION, "keep-alive, Upgrade")
.header(SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
.body(())
.unwrap();
assert!(is_upgrade_request(&listed));
let no_key = Request::builder()
.header(UPGRADE, "websocket")
.header(CONNECTION, "Upgrade")
.body(())
.unwrap();
assert!(!is_upgrade_request(&no_key));
assert!(!is_upgrade_request(&Request::builder().body(()).unwrap()));
}
#[test]
fn fallback_prefers_h2ts_then_optionally_the_first_offered() {
use super::Fallback::{Accept, Reject};
assert_eq!(fallback_subprotocol(&["h2ts"], false), Accept(Some("h2ts".into())));
assert_eq!(
fallback_subprotocol(&["chat", "h2ts"], false),
Accept(Some("h2ts".into())),
"h2ts is preferred even when it isn't first"
);
assert_eq!(
fallback_subprotocol(&["H2TS"], false),
Accept(Some("H2TS".into())),
"matched case-insensitively but echoed in the offered casing"
);
assert_eq!(fallback_subprotocol(&["chat"], false), Reject);
assert_eq!(fallback_subprotocol(&["chat", "binary"], false), Reject);
assert_eq!(fallback_subprotocol(&[], false), Reject);
assert_eq!(
fallback_subprotocol(&["chat", "binary"], true),
Accept(Some("chat".into()))
);
assert_eq!(fallback_subprotocol(&[], true), Accept(None));
}
#[test]
fn accept_echoes_h2ts_when_offered() {
let mut req = upgrade_request(Some("h2ts"));
let (resp, _fut) = accept(&mut req).unwrap();
assert_eq!(resp.status(), StatusCode::SWITCHING_PROTOCOLS);
assert_eq!(resp.headers().get(SEC_WEBSOCKET_PROTOCOL).unwrap(), "h2ts");
let mut req = upgrade_request(Some("chat, h2ts"));
let (resp, _fut) = accept(&mut req).unwrap();
assert_eq!(resp.headers().get(SEC_WEBSOCKET_PROTOCOL).unwrap(), "h2ts");
}
#[test]
fn accept_rejects_when_h2ts_absent_by_default() {
for offer in [Some("mystery"), Some("chat, binary"), None] {
let mut req = upgrade_request(offer);
let err = accept(&mut req).err().expect("should reject");
assert!(
matches!(&err, WebSocketError::UnsupportedSubprotocol),
"offer {offer:?}"
);
assert_eq!(err.rejection_response().status(), StatusCode::BAD_REQUEST);
}
}
#[test]
fn accept_with_honors_a_selection_even_without_h2ts() {
let mut req = upgrade_request(Some("chat, binary"));
let (resp, _fut) = accept_with(&mut req, |offered| {
offered.iter().find(|p| **p == "binary").map(|p| p.to_string())
})
.unwrap();
assert_eq!(resp.status(), StatusCode::SWITCHING_PROTOCOLS);
assert_eq!(resp.headers().get(SEC_WEBSOCKET_PROTOCOL).unwrap(), "binary");
}
#[test]
fn allow_implicit_codec_accepts_the_first_offered_codec() {
let opts = AcceptOptions {
allow_implicit_codec: true,
};
let mut req = upgrade_request(Some("mystery, other"));
let (resp, _fut) = accept_with_options(&mut req, |_| None, opts).unwrap();
assert_eq!(resp.status(), StatusCode::SWITCHING_PROTOCOLS);
assert_eq!(resp.headers().get(SEC_WEBSOCKET_PROTOCOL).unwrap(), "mystery");
let mut req = upgrade_request(None);
let (resp, _fut) = accept_with_options(&mut req, |_| None, opts).unwrap();
assert_eq!(resp.status(), StatusCode::SWITCHING_PROTOCOLS);
assert!(resp.headers().get(SEC_WEBSOCKET_PROTOCOL).is_none());
}
#[test]
fn accept_key_matches_rfc_6455_example() {
assert_eq!(
super::accept_key(b"dGhlIHNhbXBsZSBub25jZQ=="),
"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="
);
}
}