use crate::http::error::Error;
use crate::http::response::Body;
use hyper::header::{self, HeaderMap, HeaderName, HeaderValue};
use hyper::http::request::Parts;
use hyper::{Method, Response, StatusCode};
use std::borrow::Cow;
use std::future::Future;
use tokio_tungstenite::WebSocketStream;
use tokio_tungstenite::tungstenite::handshake::derive_accept_key;
use tokio_tungstenite::tungstenite::protocol;
pub use tokio_tungstenite::tungstenite::Message;
pub use tokio_tungstenite::tungstenite::protocol::{
CloseFrame, WebSocketConfig, frame::coding::CloseCode,
};
#[must_use]
pub struct WebSocketUpgrade<F = DefaultOnFailedUpgrade> {
config: WebSocketConfig,
protocol: Option<HeaderValue>,
sec_websocket_key: HeaderValue,
on_upgrade: hyper::upgrade::OnUpgrade,
on_failed_upgrade: F,
sec_websocket_protocol: Vec<HeaderValue>,
origin: Option<HeaderValue>,
}
impl<F> std::fmt::Debug for WebSocketUpgrade<F> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WebSocketUpgrade")
.field("protocol", &self.protocol)
.field("sec_websocket_protocol", &self.sec_websocket_protocol)
.finish_non_exhaustive()
}
}
impl<F> WebSocketUpgrade<F> {
pub const fn read_buffer_size(mut self, size: usize) -> Self {
self.config.read_buffer_size = size;
self
}
pub const fn write_buffer_size(mut self, size: usize) -> Self {
self.config.write_buffer_size = size;
self
}
pub const fn max_write_buffer_size(mut self, max: usize) -> Self {
self.config.max_write_buffer_size = max;
self
}
pub const fn max_message_size(mut self, max: usize) -> Self {
self.config.max_message_size = Some(max);
self
}
pub const fn max_frame_size(mut self, max: usize) -> Self {
self.config.max_frame_size = Some(max);
self
}
pub const fn accept_unmasked_frames(mut self, accept: bool) -> Self {
self.config.accept_unmasked_frames = accept;
self
}
pub fn protocols<I>(mut self, protocols: I) -> Self
where
I: IntoIterator,
I::Item: Into<Cow<'static, str>>,
{
self.protocol = protocols.into_iter().map(Into::into).find_map(|proto| {
let value = match proto {
Cow::Owned(s) => HeaderValue::from_str(&s).ok()?,
Cow::Borrowed(s) => HeaderValue::from_static(s),
};
self.sec_websocket_protocol
.contains(&value)
.then_some(value)
});
self
}
pub fn requested_protocols(&self) -> impl Iterator<Item = &HeaderValue> {
self.sec_websocket_protocol.iter()
}
#[must_use]
pub const fn selected_protocol(&self) -> Option<&HeaderValue> {
self.protocol.as_ref()
}
#[must_use]
pub const fn origin(&self) -> Option<&HeaderValue> {
self.origin.as_ref()
}
pub fn on_failed_upgrade<C>(self, callback: C) -> WebSocketUpgrade<C>
where
C: OnFailedUpgrade,
{
WebSocketUpgrade {
config: self.config,
protocol: self.protocol,
sec_websocket_key: self.sec_websocket_key,
on_upgrade: self.on_upgrade,
on_failed_upgrade: callback,
sec_websocket_protocol: self.sec_websocket_protocol,
origin: self.origin,
}
}
#[must_use = "the response from `on_upgrade` must be returned from the handler"]
pub fn on_upgrade<C, Fut>(self, callback: C) -> Response<Body>
where
C: FnOnce(WebSocket) -> Fut + Send + 'static,
Fut: Future<Output = ()> + Send + 'static,
F: OnFailedUpgrade,
{
let on_upgrade = self.on_upgrade;
let config = self.config;
let on_failed_upgrade = self.on_failed_upgrade;
let protocol = self.protocol.clone();
tokio::spawn(async move {
let upgraded = match on_upgrade.await {
Ok(upgraded) => upgraded,
Err(err) => {
on_failed_upgrade.call(Error::Internal(err.to_string()));
return;
}
};
let upgraded = hyper_util::rt::TokioIo::new(upgraded);
let inner =
WebSocketStream::from_raw_socket(upgraded, protocol::Role::Server, Some(config))
.await;
callback(WebSocket { inner, protocol }).await;
});
let response = Response::builder()
.status(StatusCode::SWITCHING_PROTOCOLS)
.header(header::CONNECTION, HeaderValue::from_static("upgrade"))
.header(header::UPGRADE, HeaderValue::from_static("websocket"))
.header(
header::SEC_WEBSOCKET_ACCEPT,
derive_accept_key(self.sec_websocket_key.as_bytes()),
)
.body(Body::empty())
.unwrap_or_else(|_| Response::new(Body::empty()));
let mut response = response;
if let Some(protocol) = self.protocol {
response
.headers_mut()
.insert(header::SEC_WEBSOCKET_PROTOCOL, protocol);
}
response
}
}
pub trait OnFailedUpgrade: Send + 'static {
fn call(self, error: Error);
}
impl<F> OnFailedUpgrade for F
where
F: FnOnce(Error) + Send + 'static,
{
fn call(self, error: Error) {
self(error);
}
}
#[non_exhaustive]
#[derive(Debug)]
pub struct DefaultOnFailedUpgrade;
impl OnFailedUpgrade for DefaultOnFailedUpgrade {
fn call(self, _error: Error) {}
}
fn header_eq(headers: &HeaderMap, key: &HeaderName, value: &'static str) -> bool {
headers
.get(key)
.is_some_and(|h| h.as_bytes().eq_ignore_ascii_case(value.as_bytes()))
}
fn header_contains(headers: &HeaderMap, key: &HeaderName, value: &'static str) -> bool {
let Some(header) = headers.get(key) else {
return false;
};
std::str::from_utf8(header.as_bytes()).is_ok_and(|h| h.to_ascii_lowercase().contains(value))
}
impl<S> crate::routing::extract::FromRequest<S> for WebSocketUpgrade<DefaultOnFailedUpgrade>
where
S: Send + Sync,
{
type Rejection = Error;
async fn from_request(req: hyper::Request<Body>, state: &S) -> Result<Self, Self::Rejection> {
let (mut parts, _body) = req.into_parts();
Self::from_request_parts(&mut parts, state)
}
}
impl WebSocketUpgrade<DefaultOnFailedUpgrade> {
pub fn from_request_parts<S>(parts: &mut Parts, _state: &S) -> Result<Self, Error> {
if parts.version > hyper::Version::HTTP_11 {
return Err(Error::Rejection {
status: StatusCode::UPGRADE_REQUIRED,
message: "WebSocket upgrades require HTTP/1.1".to_string(),
});
}
if parts.method != Method::GET {
return Err(Error::Rejection {
status: StatusCode::METHOD_NOT_ALLOWED,
message: "Request method must be `GET`".to_string(),
});
}
if !header_contains(&parts.headers, &header::CONNECTION, "upgrade") {
return Err(Error::Rejection {
status: StatusCode::BAD_REQUEST,
message: "`Connection` header did not include 'upgrade'".to_string(),
});
}
if !header_eq(&parts.headers, &header::UPGRADE, "websocket") {
return Err(Error::Rejection {
status: StatusCode::BAD_REQUEST,
message: "`Upgrade` header did not include 'websocket'".to_string(),
});
}
if !header_eq(&parts.headers, &header::SEC_WEBSOCKET_VERSION, "13") {
return Err(Error::Rejection {
status: StatusCode::BAD_REQUEST,
message: "`Sec-WebSocket-Version` header did not include '13'".to_string(),
});
}
let sec_websocket_key = parts
.headers
.get(header::SEC_WEBSOCKET_KEY)
.cloned()
.ok_or_else(|| Error::Rejection {
status: StatusCode::BAD_REQUEST,
message: "`Sec-WebSocket-Key` header missing".to_string(),
})?;
let on_upgrade = parts
.extensions
.remove::<hyper::upgrade::OnUpgrade>()
.ok_or_else(|| Error::Rejection {
status: StatusCode::UPGRADE_REQUIRED,
message: "Request couldn't be upgraded: no upgrade state was present".to_string(),
})?;
let sec_websocket_protocol = parts
.headers
.get_all(header::SEC_WEBSOCKET_PROTOCOL)
.iter()
.flat_map(|val| val.as_bytes().split(|&b| b == b','))
.filter_map(|proto| HeaderValue::from_bytes(proto.trim_ascii()).ok())
.collect();
let origin = parts.headers.get(header::ORIGIN).cloned();
Ok(Self {
config: WebSocketConfig::default(),
protocol: None,
sec_websocket_key,
on_upgrade,
on_failed_upgrade: DefaultOnFailedUpgrade,
sec_websocket_protocol,
origin,
})
}
}
pub struct WebSocket {
inner: WebSocketStream<hyper_util::rt::TokioIo<hyper::upgrade::Upgraded>>,
protocol: Option<HeaderValue>,
}
impl std::fmt::Debug for WebSocket {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WebSocket").finish_non_exhaustive()
}
}
impl WebSocket {
pub async fn recv(&mut self) -> Option<Result<Message, Error>> {
use futures_util::StreamExt;
self.inner
.next()
.await
.map(|r| r.map_err(|e| Error::Internal(e.to_string())))
}
pub async fn send(&mut self, msg: Message) -> Result<(), Error> {
use futures_util::SinkExt;
self.inner
.send(msg)
.await
.map_err(|e| Error::Internal(e.to_string()))
}
pub async fn flush(&mut self) -> Result<(), Error> {
use futures_util::SinkExt;
self.inner
.flush()
.await
.map_err(|e| Error::Internal(e.to_string()))
}
pub async fn close(mut self) -> Result<(), Error> {
self.inner
.close(None)
.await
.map_err(|e| Error::Internal(e.to_string()))
}
#[must_use]
pub const fn protocol(&self) -> Option<&HeaderValue> {
self.protocol.as_ref()
}
pub fn split(
self,
) -> (
impl futures_util::Sink<Message, Error = Error> + Send,
impl futures_util::Stream<Item = Result<Message, Error>> + Send,
) {
use futures_util::{SinkExt, StreamExt};
let (sink, stream) = self.inner.split();
(
sink.sink_map_err(|e| Error::Internal(e.to_string())),
stream.map(|r| r.map_err(|e| Error::Internal(e.to_string()))),
)
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::unwrap_used)]
use super::*;
use hyper::Request;
fn make_ws_parts() -> Parts {
let mut req = Request::builder()
.method(Method::GET)
.uri("/ws")
.header(header::CONNECTION, "upgrade")
.header(header::UPGRADE, "websocket")
.header(header::SEC_WEBSOCKET_VERSION, "13")
.header(header::SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
.body(())
.unwrap();
let on_upgrade = hyper::upgrade::on(&mut req);
let (mut parts, ()) = req.into_parts();
let _ = parts.extensions.insert(on_upgrade);
parts
}
#[test]
fn rejects_http2_and_above() {
let mut parts = make_ws_parts();
parts.version = hyper::Version::HTTP_2;
let err = WebSocketUpgrade::from_request_parts(&mut parts, &()).unwrap_err();
assert!(matches!(
err,
Error::Rejection {
status: StatusCode::UPGRADE_REQUIRED,
..
}
));
}
#[test]
fn rejects_non_get_method() {
let mut parts = make_ws_parts();
parts.method = Method::POST;
let err = WebSocketUpgrade::from_request_parts(&mut parts, &()).unwrap_err();
assert!(matches!(
err,
Error::Rejection {
status: StatusCode::METHOD_NOT_ALLOWED,
..
}
));
}
#[test]
fn rejects_missing_connection_upgrade_token() {
let mut parts = make_ws_parts();
let _ = parts
.headers
.insert(header::CONNECTION, HeaderValue::from_static("keep-alive"));
let err = WebSocketUpgrade::from_request_parts(&mut parts, &()).unwrap_err();
assert!(matches!(
err,
Error::Rejection {
status: StatusCode::BAD_REQUEST,
..
}
));
}
#[test]
fn rejects_wrong_upgrade_header_value() {
let mut parts = make_ws_parts();
let _ = parts
.headers
.insert(header::UPGRADE, HeaderValue::from_static("h2c"));
let err = WebSocketUpgrade::from_request_parts(&mut parts, &()).unwrap_err();
assert!(matches!(
err,
Error::Rejection {
status: StatusCode::BAD_REQUEST,
..
}
));
}
#[test]
fn rejects_wrong_sec_websocket_version() {
let mut parts = make_ws_parts();
let _ = parts
.headers
.insert(header::SEC_WEBSOCKET_VERSION, HeaderValue::from_static("8"));
let err = WebSocketUpgrade::from_request_parts(&mut parts, &()).unwrap_err();
assert!(matches!(
err,
Error::Rejection {
status: StatusCode::BAD_REQUEST,
..
}
));
}
#[test]
fn rejects_missing_sec_websocket_key() {
let mut parts = make_ws_parts();
let _ = parts.headers.remove(header::SEC_WEBSOCKET_KEY);
let err = WebSocketUpgrade::from_request_parts(&mut parts, &()).unwrap_err();
assert!(matches!(
err,
Error::Rejection {
status: StatusCode::BAD_REQUEST,
..
}
));
}
#[test]
fn rejects_when_no_upgrade_state_is_present() {
let req = Request::builder()
.method(Method::GET)
.uri("/ws")
.header(header::CONNECTION, "upgrade")
.header(header::UPGRADE, "websocket")
.header(header::SEC_WEBSOCKET_VERSION, "13")
.header(header::SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
.body(())
.unwrap();
let (mut parts, ()) = req.into_parts();
let err = WebSocketUpgrade::from_request_parts(&mut parts, &()).unwrap_err();
assert!(matches!(
err,
Error::Rejection {
status: StatusCode::UPGRADE_REQUIRED,
..
}
));
}
#[test]
fn builder_methods_configure_the_underlying_websocket_config() {
let mut parts = make_ws_parts();
let upgrade = WebSocketUpgrade::from_request_parts(&mut parts, &())
.unwrap()
.read_buffer_size(1024)
.write_buffer_size(2048)
.max_write_buffer_size(4096)
.max_message_size(8192)
.max_frame_size(16384)
.accept_unmasked_frames(true);
assert_eq!(upgrade.config.read_buffer_size, 1024);
assert_eq!(upgrade.config.write_buffer_size, 2048);
assert_eq!(upgrade.config.max_write_buffer_size, 4096);
assert_eq!(upgrade.config.max_message_size, Some(8192));
assert_eq!(upgrade.config.max_frame_size, Some(16384));
assert!(upgrade.config.accept_unmasked_frames);
}
#[test]
fn protocols_selects_a_requested_subprotocol_the_server_also_supports() {
let mut req = Request::builder()
.method(Method::GET)
.uri("/ws")
.header(header::CONNECTION, "upgrade")
.header(header::UPGRADE, "websocket")
.header(header::SEC_WEBSOCKET_VERSION, "13")
.header(header::SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
.header(header::SEC_WEBSOCKET_PROTOCOL, "chat, superchat")
.body(())
.unwrap();
let on_upgrade = hyper::upgrade::on(&mut req);
let (mut parts, ()) = req.into_parts();
let _ = parts.extensions.insert(on_upgrade);
let upgrade = WebSocketUpgrade::from_request_parts(&mut parts, &()).unwrap();
let requested: Vec<_> = upgrade
.requested_protocols()
.map(|v| v.to_str().unwrap().to_string())
.collect();
assert_eq!(requested, vec!["chat", "superchat"]);
let upgrade = upgrade.protocols(["superchat"]);
assert_eq!(
upgrade.selected_protocol().unwrap().to_str().unwrap(),
"superchat"
);
}
#[test]
fn protocols_selects_none_when_nothing_matches() {
let mut parts = make_ws_parts();
let upgrade = WebSocketUpgrade::from_request_parts(&mut parts, &())
.unwrap()
.protocols(["some-protocol-the-client-never-asked-for"]);
assert!(upgrade.selected_protocol().is_none());
}
#[test]
fn origin_getter_reflects_the_request_header() {
let mut parts = make_ws_parts();
let _ = parts.headers.insert(
header::ORIGIN,
HeaderValue::from_static("https://example.com"),
);
let upgrade = WebSocketUpgrade::from_request_parts(&mut parts, &()).unwrap();
assert_eq!(
upgrade.origin().unwrap().to_str().unwrap(),
"https://example.com"
);
}
#[test]
fn origin_getter_is_none_when_absent() {
let mut parts = make_ws_parts();
let upgrade = WebSocketUpgrade::from_request_parts(&mut parts, &()).unwrap();
assert!(upgrade.origin().is_none());
}
#[test]
fn websocket_upgrade_debug_does_not_panic() {
let mut parts = make_ws_parts();
let upgrade = WebSocketUpgrade::from_request_parts(&mut parts, &()).unwrap();
assert!(format!("{upgrade:?}").contains("WebSocketUpgrade"));
}
#[test]
fn on_failed_upgrade_swaps_the_callback_type_and_preserves_config() {
let mut parts = make_ws_parts();
let upgrade = WebSocketUpgrade::from_request_parts(&mut parts, &())
.unwrap()
.max_message_size(1234)
.on_failed_upgrade(|_err: Error| {});
assert_eq!(upgrade.config.max_message_size, Some(1234));
}
#[test]
fn default_on_failed_upgrade_silently_ignores_the_error() {
DefaultOnFailedUpgrade.call(Error::Internal("boom".to_string()));
}
}