use bytes::Bytes;
use http_body_util::Empty;
use hyper::{
ext::Protocol, header, header::HeaderValue, http::Extensions, Method, Request, Response,
StatusCode,
};
use hyper_util::rt::{TokioExecutor, TokioIo};
use tokio::io::{AsyncRead, AsyncWrite};
use url::Url;
use super::{HttpRequestBuilder, HttpStream, HttpWebSocket, Negotiation, Options, Role};
use crate::{compression::WebSocketExtensions, Result, WebSocketError};
pub(super) const WEBSOCKET_PROTOCOL: &str = "websocket";
pub(super) fn is_extended_connect(method: &Method, extensions: &Extensions) -> bool {
method == Method::CONNECT
&& extensions
.get::<Protocol>()
.is_some_and(|protocol| protocol.as_str() == WEBSOCKET_PROTOCOL)
}
pub(super) fn build_request(
url: &Url,
options: &Options,
builder: HttpRequestBuilder,
) -> Result<Request<Empty<Bytes>>> {
let scheme = match url.scheme() {
"ws" | "http" => "http",
"wss" | "https" => "https",
_ => return Err(WebSocketError::InvalidHttpScheme),
};
let host = url.host().expect("hostname").to_string();
let authority = if let Some(port) = url.port() {
format!("{host}:{port}")
} else {
host
};
let path = &url[url::Position::BeforePath..];
let uri = format!("{scheme}://{authority}{path}");
let mut request = builder
.method(Method::CONNECT)
.uri(uri)
.body(Empty::<Bytes>::new())
.expect("request build");
request.headers_mut().insert(
header::SEC_WEBSOCKET_VERSION,
HeaderValue::from_static("13"),
);
request
.extensions_mut()
.insert(Protocol::from_static(WEBSOCKET_PROTOCOL));
if let Some(compression) = options.compression.as_ref() {
let extensions = WebSocketExtensions::from(compression);
let header_value = extensions.to_string().parse().expect("extensions header");
request
.headers_mut()
.insert(header::SEC_WEBSOCKET_EXTENSIONS, header_value);
}
Ok(request)
}
pub(super) fn verify<B>(response: &Response<B>, options: Options) -> Result<Negotiation> {
if response.status() != StatusCode::OK {
return Err(WebSocketError::InvalidStatusCode(
response.status().as_u16(),
));
}
let extensions = WebSocketExtensions::from_headers(response.headers());
Negotiation::new(extensions, &options, Role::Client)
}
pub async fn handshake<S>(
url: Url,
io: S,
options: Options,
builder: HttpRequestBuilder,
) -> Result<HttpWebSocket>
where
S: AsyncRead + AsyncWrite + Send + Unpin + 'static,
{
let request = build_request(&url, &options, builder)?;
let (mut sender, conn) =
hyper::client::conn::http2::handshake(TokioExecutor::new(), TokioIo::new(io)).await?;
super::spawn_connection(async move {
if let Err(err) = conn.await {
log::debug!("http2 connection closed: {err:?}");
}
});
let mut response = sender.send_request(request).await.map_err(connect_error)?;
let negotiated = verify(&response, options)?;
let upgraded = hyper::upgrade::on(&mut response)
.await
.map_err(connect_error)?;
Ok(HttpWebSocket::new(
Role::Client,
HttpStream::from(TokioIo::new(upgraded)),
Bytes::new(),
negotiated,
))
}
fn connect_error(err: hyper::Error) -> WebSocketError {
use std::error::Error;
let mut source: Option<&(dyn Error + 'static)> = Some(&err);
while let Some(cause) = source {
if let Some(h2_err) = cause.downcast_ref::<h2::Error>() {
if h2_err.reason() == Some(h2::Reason::PROTOCOL_ERROR) {
return WebSocketError::ExtendedConnectNotSupported;
}
}
source = cause.source();
}
WebSocketError::from(err)
}
#[cfg(test)]
mod tests {
use super::*;
fn request(url: &str, options: Options) -> Request<Empty<Bytes>> {
build_request(&url.parse().unwrap(), &options, hyper::Request::builder()).unwrap()
}
#[test]
fn request_uses_extended_connect() {
let req = request("wss://example.com/chat", Options::default());
assert_eq!(req.method(), Method::CONNECT);
assert_eq!(
req.extensions().get::<Protocol>().map(Protocol::as_str),
Some(WEBSOCKET_PROTOCOL)
);
assert_eq!(req.uri().scheme_str(), Some("https"));
assert_eq!(req.uri().authority().unwrap().as_str(), "example.com");
assert_eq!(req.uri().path(), "/chat");
assert_eq!(
req.headers().get(header::SEC_WEBSOCKET_VERSION).unwrap(),
"13"
);
}
#[test]
fn request_omits_http1_handshake_headers() {
let req = request("wss://example.com/chat", Options::default());
assert!(req.headers().get(header::SEC_WEBSOCKET_KEY).is_none());
assert!(req.headers().get(header::UPGRADE).is_none());
assert!(req.headers().get(header::CONNECTION).is_none());
}
#[test]
fn plaintext_url_maps_to_http_scheme() {
let req = request("ws://example.com:8080/chat", Options::default());
assert_eq!(req.uri().scheme_str(), Some("http"));
assert_eq!(req.uri().authority().unwrap().as_str(), "example.com:8080");
}
#[test]
fn caller_cannot_override_the_websocket_version() {
let req = build_request(
&"wss://example.com/chat".parse().unwrap(),
&Options::default(),
hyper::Request::builder().header(header::SEC_WEBSOCKET_VERSION, "8"),
)
.unwrap();
let versions: Vec<_> = req
.headers()
.get_all(header::SEC_WEBSOCKET_VERSION)
.iter()
.collect();
assert_eq!(versions, vec!["13"]);
}
#[test]
fn caller_headers_are_preserved() {
let req = build_request(
&"wss://example.com/chat".parse().unwrap(),
&Options::default(),
hyper::Request::builder().header("authorization", "Bearer token"),
)
.unwrap();
assert_eq!(req.headers().get("authorization").unwrap(), "Bearer token");
}
#[test]
fn request_offers_compression_when_enabled() {
let req = request(
"wss://example.com/chat",
Options::default().with_balanced_compression(),
);
let offer = req
.headers()
.get(header::SEC_WEBSOCKET_EXTENSIONS)
.unwrap()
.to_str()
.unwrap();
assert!(offer.contains("permessage-deflate"));
}
#[test]
fn rejects_non_websocket_scheme() {
let err = build_request(
&"ftp://example.com/chat".parse().unwrap(),
&Options::default(),
hyper::Request::builder(),
)
.unwrap_err();
assert!(matches!(err, WebSocketError::InvalidHttpScheme));
}
#[test]
fn verify_accepts_200_not_101() {
let ok = Response::builder().status(200).body(()).unwrap();
assert!(verify(&ok, Options::default()).is_ok());
let switching = Response::builder().status(101).body(()).unwrap();
let err = verify(&switching, Options::default()).unwrap_err();
assert!(matches!(err, WebSocketError::InvalidStatusCode(101)));
}
#[test]
fn is_extended_connect_ignores_plain_connect() {
let mut tunnel = Request::builder()
.method(Method::CONNECT)
.uri("https://example.com/")
.body(())
.unwrap();
assert!(!is_extended_connect(tunnel.method(), tunnel.extensions()));
tunnel
.extensions_mut()
.insert(Protocol::from_static(WEBSOCKET_PROTOCOL));
assert!(is_extended_connect(tunnel.method(), tunnel.extensions()));
}
#[test]
fn is_extended_connect_ignores_http1_upgrade() {
let req = Request::builder()
.method(Method::GET)
.uri("/chat")
.header(header::UPGRADE, "websocket")
.body(())
.unwrap();
assert!(!is_extended_connect(req.method(), req.extensions()));
}
#[test]
fn upgrade_rejects_wrong_version() {
let mut req = Request::builder()
.method(Method::CONNECT)
.uri("https://example.com/chat")
.header(header::SEC_WEBSOCKET_VERSION, "8")
.body(())
.unwrap();
req.extensions_mut()
.insert(Protocol::from_static(WEBSOCKET_PROTOCOL));
let err = crate::WebSocket::upgrade_with_options(&mut req, Options::default()).unwrap_err();
assert!(matches!(err, WebSocketError::InvalidSecWebsocketVersion));
}
#[test]
fn upgrade_answers_200_without_accept_header() {
let mut req = Request::builder()
.method(Method::CONNECT)
.uri("https://example.com/chat")
.header(header::SEC_WEBSOCKET_VERSION, "13")
.body(())
.unwrap();
req.extensions_mut()
.insert(Protocol::from_static(WEBSOCKET_PROTOCOL));
let (response, _fut) =
crate::WebSocket::upgrade_with_options(&mut req, Options::default()).unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert!(response
.headers()
.get(header::SEC_WEBSOCKET_ACCEPT)
.is_none());
assert!(response.headers().get(header::UPGRADE).is_none());
assert!(response.headers().get(header::CONNECTION).is_none());
}
#[test]
fn upgrade_negotiates_compression() {
let mut req = Request::builder()
.method(Method::CONNECT)
.uri("https://example.com/chat")
.header(header::SEC_WEBSOCKET_VERSION, "13")
.header(header::SEC_WEBSOCKET_EXTENSIONS, "permessage-deflate")
.body(())
.unwrap();
req.extensions_mut()
.insert(Protocol::from_static(WEBSOCKET_PROTOCOL));
let (response, _fut) = crate::WebSocket::upgrade_with_options(
&mut req,
Options::default().with_balanced_compression(),
)
.unwrap();
let agreed = response
.headers()
.get(header::SEC_WEBSOCKET_EXTENSIONS)
.unwrap()
.to_str()
.unwrap();
assert!(agreed.contains("permessage-deflate"));
}
}