#![expect(
clippy::unreachable,
reason = "vendored from upstream `tungstenite-rs`: arms gated on caller-validated WebSocket protocol state that the type system can't enforce"
)]
use std::fmt;
use std::ops::{Deref, DerefMut};
use rama_core::Service;
use rama_core::error::{BoxError, ErrorContext};
use rama_core::extensions::{Extensions, ExtensionsRef};
use rama_core::telemetry::tracing;
use rama_http::conn::TargetHttpVersion;
use rama_http::headers::sec_websocket_extensions::{Extension, PerMessageDeflateConfig};
use rama_http::headers::sec_websocket_protocol::AcceptedWebSocketProtocol;
use rama_http::headers::{
HeaderMapExt, HttpRequestBuilderExt as _, SecWebSocketExtensions, SecWebSocketKey,
SecWebSocketProtocol,
};
use rama_http::proto::h2::ext::Protocol;
use rama_http::service::client::ext::{IntoHeaderName, IntoHeaderValue};
use rama_http::service::client::{HttpClientExt, IntoUrl, RequestBuilder};
use rama_http::{Body, Method, Request, Response, StatusCode, Version, header, headers};
use rama_http::{request, response};
use rama_net::extensions::StreamTransformed;
use rama_utils::str::NonEmptyStr;
use crate::protocol::{Role, WebSocketConfig};
use crate::runtime::AsyncWebSocket;
#[derive(Debug, Clone)]
pub struct WebSocketRequestBuilder<B> {
inner: B,
protocols: Option<SecWebSocketProtocol>,
extensions: Option<SecWebSocketExtensions>,
key: Option<SecWebSocketKey>,
}
#[derive(Debug)]
pub struct HandshakeRequest {
pub request: Request,
pub protocols: Option<SecWebSocketProtocol>,
pub extensions: Option<SecWebSocketExtensions>,
pub key: Option<SecWebSocketKey>,
}
pub struct WithService<'a, S, Body> {
builder: RequestBuilder<'a, S, Response<Body>>,
config: Option<WebSocketConfig>,
is_h2: bool,
}
impl<S: fmt::Debug, Body> fmt::Debug for WithService<'_, S, Body> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WithService")
.field("builder", &self.builder)
.field("config", &self.config)
.field("is_h2", &self.is_h2)
.finish()
}
}
fn new_ws_request_builder_from_uri<T>(uri: T, version: Version) -> request::Builder
where
T: TryInto<rama_net::uri::Uri, Error: Into<rama_http::HttpError>>,
{
let builder = Request::builder()
.version(version)
.uri(uri)
.typed_header(headers::SecWebSocketVersion::V13);
match version {
version @ (Version::HTTP_10 | Version::HTTP_11) => builder
.method(Method::GET)
.version(version)
.typed_header(headers::Upgrade::websocket())
.typed_header(headers::Connection::upgrade()),
Version::HTTP_2 => builder.method(Method::CONNECT).version(Version::HTTP_2),
_ => unreachable!("bug"),
}
}
fn new_ws_request_builder_from_uri_with_service<'a, S, Body, T>(
service: &'a S,
uri: T,
version: Version,
) -> RequestBuilder<'a, S, Response<Body>>
where
S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
T: IntoUrl,
{
let builder = match version {
version @ (Version::HTTP_10 | Version::HTTP_11) => service
.get(uri)
.version(version)
.typed_header(headers::Upgrade::websocket())
.typed_header(headers::Connection::upgrade()),
Version::HTTP_2 => service.connect(uri).version(Version::HTTP_2),
_ => unreachable!("bug"),
};
builder.typed_header(headers::SecWebSocketVersion::V13)
}
fn new_ws_request_builder_from_request<'a, S, Body, RequestBody>(
service: &'a S,
mut request: Request<RequestBody>,
) -> RequestBuilder<'a, S, Response<Body>>
where
S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
RequestBody: Into<rama_http::Body>,
{
if !request
.headers()
.contains_key(header::SEC_WEBSOCKET_VERSION)
{
request
.headers_mut()
.typed_insert(headers::SecWebSocketVersion::V13);
}
match request.version() {
Version::HTTP_10 | Version::HTTP_11 => {
if request.headers().get(header::UPGRADE).is_none() {
request
.headers_mut()
.typed_insert(headers::Upgrade::websocket());
}
if request.headers().get(header::CONNECTION).is_none() {
request
.headers_mut()
.typed_insert(headers::Connection::upgrade());
}
}
_ => (),
}
service.build_from_request(request)
}
#[derive(Debug)]
pub enum ResponseValidateError {
UnexpectedStatusCode(StatusCode),
UnexpectedHttpVersion(Version),
MissingUpgradeWebSocketHeader,
MissingConnectionUpgradeHeader,
SecWebSocketAcceptKeyMismatch,
ProtocolMismatch(Option<NonEmptyStr>),
ExtensionMismatch(Option<Extension>),
}
#[derive(Debug)]
pub enum HandshakeError {
ValidationError(ResponseValidateError),
HttpRequestError(BoxError),
HttpUpgradeError(BoxError),
}
impl fmt::Display for ResponseValidateError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::UnexpectedStatusCode(status_code) => {
write!(f, "unexpected HTTP status code: {status_code}")
}
Self::UnexpectedHttpVersion(version) => {
write!(f, "unexpected HTTP version: {version:?}")
}
Self::MissingUpgradeWebSocketHeader => {
write!(f, "missing upgrade WebSocket header")
}
Self::MissingConnectionUpgradeHeader => {
write!(f, "missing connection upgrade header")
}
Self::SecWebSocketAcceptKeyMismatch => {
write!(f, "key mismatch for sec-websocket-accept header")
}
Self::ProtocolMismatch(protocol) => {
write!(f, "protocol mismatch: {protocol:?}")
}
Self::ExtensionMismatch(extension) => {
write!(f, "extension mismatch: {extension:?}")
}
}
}
}
impl std::error::Error for ResponseValidateError {}
impl fmt::Display for HandshakeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::ValidationError(error) => {
write!(f, "response validation failed: {error}")
}
Self::HttpRequestError(error) => {
write!(f, "http request error: {error}")
}
Self::HttpUpgradeError(error) => {
write!(f, "http upgrade error: {error}")
}
}
}
}
impl std::error::Error for HandshakeError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::ValidationError(error) => Some(error as &dyn std::error::Error),
Self::HttpRequestError(error) | Self::HttpUpgradeError(error) => error.source(),
}
}
}
#[derive(Default, Debug)]
pub struct AcceptedWebSocketData {
pub protocol: Option<AcceptedWebSocketProtocol>,
pub extension: Option<Extension>,
}
pub fn validate_http_server_response<Body>(
response: &Response<Body>,
key: Option<headers::SecWebSocketKey>,
protocols: Option<SecWebSocketProtocol>,
extensions: Option<SecWebSocketExtensions>,
) -> Result<AcceptedWebSocketData, ResponseValidateError> {
tracing::trace!(
http.version = ?response.version(),
http.response.status = ?response.status(),
ws.protocols = ?protocols,
ws.extensions = ?extensions,
"validate http server response"
);
match response.version() {
Version::HTTP_10 | Version::HTTP_11 => {
let response_status = response.status();
if response_status != StatusCode::SWITCHING_PROTOCOLS {
return Err(ResponseValidateError::UnexpectedStatusCode(response_status));
}
if !response
.headers()
.typed_get::<headers::Upgrade>()
.map(|u| u.is_websocket())
.unwrap_or_default()
{
return Err(ResponseValidateError::MissingUpgradeWebSocketHeader);
}
if !response
.headers()
.typed_get::<headers::Connection>()
.map(|c| c.contains_upgrade())
.unwrap_or_default()
{
return Err(ResponseValidateError::MissingConnectionUpgradeHeader);
}
if let Some(key) = key {
let sec_websocket_accept_header = response
.headers()
.typed_get::<headers::SecWebSocketAccept>();
let expected_accept =
headers::SecWebSocketAccept::try_from(key).map_err(|err| {
tracing::debug!("failed to create WS accept header from key: {err}");
ResponseValidateError::SecWebSocketAcceptKeyMismatch
})?;
if sec_websocket_accept_header != Some(expected_accept) {
tracing::trace!(
"unexpected websocket accept key: {sec_websocket_accept_header:?}"
);
return Err(ResponseValidateError::SecWebSocketAcceptKeyMismatch);
}
}
}
Version::HTTP_2 => {
let response_status = response.status();
if response.status() != StatusCode::OK {
return Err(ResponseValidateError::UnexpectedStatusCode(response_status));
}
}
version => {
return Err(ResponseValidateError::UnexpectedHttpVersion(version));
}
}
let mut accepted_extension = None;
match (
response
.headers()
.typed_get::<SecWebSocketExtensions>()
.map(|ext| ext.0.head),
extensions,
) {
(None, Some(allowed_extensions)) => {
tracing::trace!(
ws.extensions = ?allowed_extensions,
"server selected no WS extensions despite client supporting some (valid, move on without)",
);
}
(Some(Extension::PerMessageDeflate(server_cfg)), Some(client_extensions)) => {
accepted_extension = client_extensions
.0.iter()
.find_map(|client_ext| {
if let Extension::PerMessageDeflate(client_cfg) = client_ext {
return Some(Ok(Extension::PerMessageDeflate(PerMessageDeflateConfig {
client_max_window_bits: match (
server_cfg.client_max_window_bits,
client_cfg.client_max_window_bits,
) {
(None, None | Some(_)) => None,
(Some(srv), maybe_offered) => {
if !(8..=15).contains(&srv) || maybe_offered.map(|offered| offered != 0 && srv > offered).unwrap_or_default() {
tracing::debug!("server offered invalid client_max_window_bits (pmd)... ext mismatch!");
return Some(Err(
ResponseValidateError::ExtensionMismatch(Some(
Extension::PerMessageDeflate(server_cfg.clone()),
)),
));
}
Some(srv)
}
},
server_max_window_bits: match (
server_cfg.server_max_window_bits,
client_cfg.server_max_window_bits,
) {
(None, None | Some(_)) => None,
(Some(their_bits), maybe_our_bits) => {
if !(8..=15).contains(&their_bits)
|| maybe_our_bits
.map(|our_bits| our_bits != 0 && their_bits > our_bits)
.unwrap_or_default()
{
tracing::debug!("server offered invalid server_max_window_bits (pmd)... ext mismatch!");
return Some(Err(
ResponseValidateError::ExtensionMismatch(Some(
Extension::PerMessageDeflate(server_cfg.clone()),
)),
));
}
Some(their_bits)
}
},
server_no_context_takeover: server_cfg.server_no_context_takeover,
client_no_context_takeover: client_cfg.client_no_context_takeover,
identifier: server_cfg.identifier.clone(),
})));
}
None
})
.transpose()?;
}
(Some(server_ext), _) => {
tracing::debug!("server offered ext, but client (we) not!");
return Err(ResponseValidateError::ExtensionMismatch(Some(server_ext)));
}
(None, None) => (),
}
let mut accepted_protocol = None;
match (
response
.headers()
.typed_get::<SecWebSocketProtocol>()
.map(|h| h.accept_first_protocol()),
protocols,
) {
(None, None) => (),
(None, Some(allowed_protocols)) => {
tracing::trace!(
ws.protocols = ?allowed_protocols,
"server selected no WS subprotocol despite client proposing some (valid, proceed without)",
);
}
(Some(header), None) => {
return Err(ResponseValidateError::ProtocolMismatch(Some(header.0)));
}
(Some(protocol_header), Some(sub_protocols)) => {
match sub_protocols.contains(&protocol_header.0) {
Some(protocol) => accepted_protocol = Some(protocol),
None => {
return Err(ResponseValidateError::ProtocolMismatch(Some(
protocol_header.0,
)));
}
};
}
}
Ok(AcceptedWebSocketData {
protocol: accepted_protocol,
extension: accepted_extension,
})
}
impl WebSocketRequestBuilder<request::Builder> {
pub fn new<T>(uri: T) -> Self
where
T: TryInto<rama_net::uri::Uri, Error: Into<rama_http::HttpError>>,
{
Self::new_with_version(uri, Version::HTTP_11)
}
pub fn new_h2<T>(uri: T) -> Self
where
T: TryInto<rama_net::uri::Uri, Error: Into<rama_http::HttpError>>,
{
Self::new_with_version(uri, Version::HTTP_2)
}
fn new_with_version<T>(uri: T, version: Version) -> Self
where
T: TryInto<rama_net::uri::Uri, Error: Into<rama_http::HttpError>>,
{
Self {
inner: new_ws_request_builder_from_uri(uri, version),
protocols: Default::default(),
extensions: Default::default(),
key: Default::default(),
}
}
#[must_use]
pub fn with_header<K, V>(self, name: K, value: V) -> Self
where
K: TryInto<rama_http::HeaderName, Error: Into<rama_http::HttpError>>,
V: TryInto<rama_http::HeaderValue, Error: Into<rama_http::HttpError>>,
{
Self {
inner: self.inner.header(name, value),
protocols: self.protocols,
extensions: self.extensions,
key: self.key,
}
}
#[must_use]
pub fn with_typed_header<H>(self, header: H) -> Self
where
H: headers::HeaderEncode,
{
Self {
inner: self.inner.typed_header(header),
protocols: self.protocols,
extensions: self.extensions,
key: self.key,
}
}
pub fn build_handshake(self) -> Result<HandshakeRequest, BoxError> {
let builder = match self.protocols.as_ref() {
Some(protocols) => self.inner.typed_header(protocols),
None => self.inner,
};
let builder = match self.extensions.as_ref() {
Some(extensions) => builder.typed_header(extensions),
None => builder,
};
let mut request = builder
.body(Body::empty())
.context("request failed to build (invalid custom header?)")?;
let mut key = None;
if request.version() != Version::HTTP_2 {
let k = self.key.unwrap_or_else(headers::SecWebSocketKey::random);
request.headers_mut().typed_insert(&k);
key = Some(k);
}
request
.extensions()
.insert(Protocol::from_static("websocket"));
Ok(HandshakeRequest {
request,
protocols: self.protocols,
extensions: self.extensions,
key,
})
}
}
impl<'a, S, Body> WebSocketRequestBuilder<WithService<'a, S, Body>>
where
S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
{
pub fn new_with_service<T>(service: &'a S, uri: T) -> Self
where
T: IntoUrl,
{
Self::new_with_service_and_version(service, Version::HTTP_11, uri)
}
pub fn new_h2_with_service<T>(service: &'a S, uri: T) -> Self
where
T: IntoUrl,
{
Self::new_with_service_and_version(service, Version::HTTP_2, uri)
}
fn new_with_service_and_version<T>(service: &'a S, version: Version, uri: T) -> Self
where
T: IntoUrl,
{
Self {
inner: WithService {
builder: new_ws_request_builder_from_uri_with_service(service, uri, version),
config: Default::default(),
is_h2: version == Version::HTTP_2,
},
protocols: Default::default(),
extensions: Default::default(),
key: Default::default(),
}
}
pub fn new_with_service_and_request<RequestBody>(
service: &'a S,
request: Request<RequestBody>,
) -> Self
where
RequestBody: Into<rama_http::Body>,
{
let key = request.headers().typed_get();
let is_h2 = request.version() == Version::HTTP_2;
let protocols = request.headers().typed_get();
let extensions = request.headers().typed_get();
Self {
inner: WithService {
builder: new_ws_request_builder_from_request(service, request),
config: Default::default(),
is_h2,
},
protocols,
extensions,
key,
}
}
#[must_use]
pub fn with_header<K, V>(self, name: K, value: V) -> Self
where
K: IntoHeaderName,
V: IntoHeaderValue,
{
Self {
inner: WithService {
builder: self.inner.builder.header(name, value),
..self.inner
},
protocols: self.protocols,
extensions: self.extensions,
key: self.key,
}
}
#[must_use]
pub fn with_header_overwrite<K, V>(self, name: K, value: V) -> Self
where
K: IntoHeaderName,
V: IntoHeaderValue,
{
Self {
inner: WithService {
builder: self.inner.builder.overwrite_header(name, value),
..self.inner
},
protocols: self.protocols,
extensions: self.extensions,
key: self.key,
}
}
#[must_use]
pub fn with_typed_header<H>(self, header: H) -> Self
where
H: headers::HeaderEncode,
{
Self {
inner: WithService {
builder: self.inner.builder.typed_header(header),
..self.inner
},
protocols: self.protocols,
extensions: self.extensions,
key: self.key,
}
}
#[must_use]
pub fn with_typed_header_overwrite<H>(self, header: H) -> Self
where
H: headers::HeaderEncode,
{
Self {
inner: WithService {
builder: self.inner.builder.overwrite_typed_header(header),
..self.inner
},
protocols: self.protocols,
extensions: self.extensions,
key: self.key,
}
}
#[cfg(feature = "compression")]
rama_utils::macros::generate_set_and_with! {
#[must_use]
#[cfg_attr(docsrs, doc(cfg(feature = "compression")))]
pub fn per_message_deflate(mut self) -> Self {
self.extensions = match self.extensions.take() {
Some(ext) => {
Some(ext.with_extra_extension(Extension::PerMessageDeflate(Default::default())))
},
None => Some(SecWebSocketExtensions::per_message_deflate()),
};
self.inner.config = Some(self.inner.config.take().unwrap_or_default().with_per_message_deflate_default());
self
}
}
#[cfg(feature = "compression")]
rama_utils::macros::generate_set_and_with! {
#[must_use]
#[cfg_attr(docsrs, doc(cfg(feature = "compression")))]
pub fn per_message_deflate_overwrite_extensions(mut self) -> Self {
self.extensions = Some(SecWebSocketExtensions::per_message_deflate());
self.inner.config = Some(self.inner.config.take().unwrap_or_default().with_per_message_deflate_default());
self
}
}
#[cfg(feature = "compression")]
rama_utils::macros::generate_set_and_with! {
#[must_use]
#[cfg_attr(docsrs, doc(cfg(feature = "compression")))]
pub fn per_message_deflate_with_config(mut self, config: impl Into<crate::protocol::PerMessageDeflateConfig>) -> Self {
let config = config.into();
self.extensions = match self.extensions.take() {
Some(ext) => {
Some(ext.with_extra_extension(Extension::PerMessageDeflate((&config).into())))
}
None => Some(SecWebSocketExtensions::per_message_deflate_with_config((&config).into())),
};
self.inner.config = Some(
self.inner
.config
.take()
.unwrap_or_default()
.with_per_message_deflate(config),
);
self
}
}
#[cfg(feature = "compression")]
rama_utils::macros::generate_set_and_with! {
#[must_use]
#[cfg_attr(docsrs, doc(cfg(feature = "compression")))]
pub fn per_message_deflate_with_config_overwrite_extensions(mut self, config: impl Into<crate::protocol::PerMessageDeflateConfig>) -> Self {
let config = config.into();
self.extensions = Some(SecWebSocketExtensions::per_message_deflate_with_config((&config).into()));
self.inner.config = Some(
self.inner
.config
.take()
.unwrap_or_default()
.with_per_message_deflate(config),
);
self
}
}
rama_utils::macros::generate_set_and_with! {
pub fn config(mut self, cfg: Option<WebSocketConfig>) -> Self {
self.inner.config = cfg;
self
}
}
pub async fn initiate_handshake(
self,
extensions: Extensions,
) -> Result<NegotiatedHandshakeRequest<Body>, HandshakeError> {
extensions.insert(StreamTransformed {
by: "rama-ws::WebSocketClient",
});
let builder = match self.protocols.as_ref() {
Some(protocols) => self.inner.builder.overwrite_typed_header(protocols),
None => self.inner.builder,
};
let builder = match self.extensions.as_ref() {
Some(extensions) => builder.typed_header(extensions),
None => builder,
};
let mut key = None;
let builder = if !self.inner.is_h2 {
extensions.insert(TargetHttpVersion(Version::HTTP_11));
let k = self.key.unwrap_or_else(headers::SecWebSocketKey::random);
let builder = builder.overwrite_typed_header(&k);
key = Some(k);
builder
} else {
extensions.insert(TargetHttpVersion(Version::HTTP_2));
builder
};
let builder = builder.extension(Protocol::from_static("websocket"));
if let Some(ext) = builder.extensions() {
ext.extend(&extensions);
}
let response = builder
.send()
.await
.context("send initial websocket handshake request (upgrade)")
.map_err(HandshakeError::HttpRequestError)?;
Ok(NegotiatedHandshakeRequest {
protocols: self.protocols,
extensions: self.extensions,
config: self.inner.config,
key,
response,
})
}
pub async fn handshake(
self,
extensions: Extensions,
) -> Result<ClientWebSocket, HandshakeError> {
let handshake = self.initiate_handshake(extensions).await?;
handshake.complete().await
}
}
impl<B> WebSocketRequestBuilder<B> {
rama_utils::macros::generate_set_and_with! {
pub fn protocols(mut self, protocols: Option<SecWebSocketProtocol>) -> Self {
self.protocols = protocols;
self
}
}
rama_utils::macros::generate_set_and_with! {
pub fn key(mut self, key: Option<headers::SecWebSocketKey>) -> Self {
self.key = key;
self
}
}
}
pub fn apply_response_data_to_base_websocket_config<Body>(
base_cfg: Option<WebSocketConfig>,
res: &mut Response<Body>,
) -> Option<WebSocketConfig> {
let accepted_pmd_cfg = res
.headers()
.typed_get::<SecWebSocketExtensions>()
.map(|ext| ext.0.head)
.and_then(|ext| {
if let Extension::PerMessageDeflate(cfg) = ext {
Some(cfg)
} else {
None
}
});
if let Some(accepted_protocol) = res
.headers()
.typed_get::<SecWebSocketProtocol>()
.map(|h| h.accept_first_protocol())
{
res.extensions().insert(accepted_protocol);
}
#[cfg(feature = "compression")]
{
if let Some(pmd_cfg) = accepted_pmd_cfg {
let mut ws_cfg = base_cfg.unwrap_or_default();
ws_cfg.per_message_deflate = Some(pmd_cfg.into());
Some(ws_cfg)
} else if let Some(mut ws_cfg) = base_cfg {
ws_cfg.per_message_deflate = None;
Some(ws_cfg)
} else {
base_cfg
}
}
#[cfg(not(feature = "compression"))]
{
if accepted_pmd_cfg.is_some() {
tracing::error!(
"per-message-deflate is used but compression feature is disabled. Enable it if you wish to use this extension."
);
}
base_cfg
}
}
pub struct NegotiatedHandshakeRequest<Body> {
pub protocols: Option<SecWebSocketProtocol>,
pub extensions: Option<SecWebSocketExtensions>,
pub config: Option<WebSocketConfig>,
pub key: Option<SecWebSocketKey>,
pub response: Response<Body>,
}
impl<Body> NegotiatedHandshakeRequest<Body> {
pub async fn complete(self) -> Result<ClientWebSocket, HandshakeError> {
let accepted_data = validate_http_server_response(
&self.response,
self.key,
self.protocols,
self.extensions,
)
.map_err(HandshakeError::ValidationError)?;
tracing::trace!(
websocket.protocol = ?accepted_data.protocol,
websocket.extension = ?accepted_data.extension,
"websocket handshake http response is valid",
);
let stream = rama_http::io::upgrade::handle_upgrade(&self.response)
.await
.context("upgrade http connection into a raw web socket")
.map_err(HandshakeError::HttpUpgradeError)?;
let (parts, _) = self.response.into_parts();
#[cfg(feature = "compression")]
let maybe_ws_cfg = {
let mut ws_cfg = self.config.unwrap_or_default();
if let Some(Extension::PerMessageDeflate(pmd_cfg)) = accepted_data.extension {
tracing::trace!(
"apply accepted per-message-deflate cfg into WS client config: {pmd_cfg:?}"
);
ws_cfg.per_message_deflate = Some(pmd_cfg.into());
} else {
ws_cfg.per_message_deflate = None;
}
Some(ws_cfg)
};
#[cfg(not(feature = "compression"))]
let maybe_ws_cfg = {
if let Some(Extension::PerMessageDeflate(pmd_cfg)) = accepted_data.extension {
tracing::error!(
"per-message-deflate is used but compression feature is disabled. Enable it if you wish to use this extension."
);
return Err(HandshakeError::ValidationError(
ResponseValidateError::ExtensionMismatch(Some(Extension::PerMessageDeflate(
pmd_cfg,
))),
));
}
None
};
let socket = AsyncWebSocket::from_raw_socket(stream, Role::Client, maybe_ws_cfg).await;
Ok(ClientWebSocket {
socket,
response: parts,
accepted_protocol: accepted_data.protocol,
})
}
}
#[derive(Debug)]
pub struct ClientWebSocket {
socket: AsyncWebSocket,
response: response::Parts,
accepted_protocol: Option<AcceptedWebSocketProtocol>,
}
impl Deref for ClientWebSocket {
type Target = AsyncWebSocket;
fn deref(&self) -> &Self::Target {
&self.socket
}
}
impl DerefMut for ClientWebSocket {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.socket
}
}
impl ClientWebSocket {
pub fn response(&self) -> &response::Parts {
&self.response
}
pub fn accepted_protocol(&self) -> Option<&str> {
self.accepted_protocol.as_ref().map(|p| p.0.as_ref())
}
pub fn into_inner(self) -> AsyncWebSocket {
self.socket
}
pub fn into_parts(
self,
) -> (
AsyncWebSocket,
response::Parts,
Option<AcceptedWebSocketProtocol>,
) {
(self.socket, self.response, self.accepted_protocol)
}
}
pub trait HttpClientWebSocketExt<Body>:
private::HttpClientWebSocketExtSealed<Body> + Sized + Send + Sync + 'static
{
fn websocket(&self, url: impl IntoUrl) -> WebSocketRequestBuilder<WithService<'_, Self, Body>>;
fn websocket_h2(
&self,
url: impl IntoUrl,
) -> WebSocketRequestBuilder<WithService<'_, Self, Body>>;
fn websocket_with_request<RequestBody: Into<rama_http::Body>>(
&self,
req: Request<RequestBody>,
) -> WebSocketRequestBuilder<WithService<'_, Self, Body>>;
}
impl<S, Body> HttpClientWebSocketExt<Body> for S
where
S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
{
fn websocket(&self, url: impl IntoUrl) -> WebSocketRequestBuilder<WithService<'_, Self, Body>> {
WebSocketRequestBuilder::new_with_service(self, url)
}
fn websocket_h2(
&self,
url: impl IntoUrl,
) -> WebSocketRequestBuilder<WithService<'_, Self, Body>> {
WebSocketRequestBuilder::new_h2_with_service(self, url)
}
fn websocket_with_request<RequestBody: Into<rama_http::Body>>(
&self,
req: Request<RequestBody>,
) -> WebSocketRequestBuilder<WithService<'_, Self, Body>> {
WebSocketRequestBuilder::new_with_service_and_request(self, req)
}
}
mod private {
use super::*;
pub trait HttpClientWebSocketExtSealed<Body> {}
impl<S, Body> HttpClientWebSocketExtSealed<Body> for S where
S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>
{
}
}
#[cfg(test)]
mod tests {
use super::*;
use rama_http::HeaderMap;
fn offered_pmd(raw: &str) -> Option<SecWebSocketExtensions> {
let mut headers = HeaderMap::new();
headers.insert(
header::SEC_WEBSOCKET_EXTENSIONS,
raw.parse().expect("valid sec-websocket-extensions header"),
);
headers.typed_get::<SecWebSocketExtensions>()
}
fn h2_response_with_pmd(raw: &str) -> Response<()> {
let mut response = Response::new(());
*response.version_mut() = Version::HTTP_2;
*response.status_mut() = StatusCode::OK;
response.headers_mut().insert(
header::SEC_WEBSOCKET_EXTENSIONS,
raw.parse().expect("valid sec-websocket-extensions header"),
);
response
}
fn validate_pmd(
server_raw: &str,
offered_raw: &str,
) -> Result<Option<u8>, ResponseValidateError> {
let response = h2_response_with_pmd(server_raw);
let accepted =
validate_http_server_response(&response, None, None, offered_pmd(offered_raw))?;
match accepted.extension {
Some(Extension::PerMessageDeflate(cfg)) => Ok(cfg.client_max_window_bits),
other => panic!("expected per-message-deflate extension, got {other:?}"),
}
}
#[test]
fn valueless_client_max_window_bits_accepts_server_choice() {
assert_eq!(
Some(15),
validate_pmd(
"permessage-deflate; client_max_window_bits=15",
"permessage-deflate; client_max_window_bits",
)
.expect("valueless offer should accept the server's window bits"),
);
}
#[test]
fn explicit_client_max_window_bits_rejects_larger_server_choice() {
assert!(matches!(
validate_pmd(
"permessage-deflate; client_max_window_bits=15",
"permessage-deflate; client_max_window_bits=10",
),
Err(ResponseValidateError::ExtensionMismatch(_)),
));
}
#[test]
fn explicit_client_max_window_bits_accepts_smaller_server_choice() {
assert_eq!(
Some(10),
validate_pmd(
"permessage-deflate; client_max_window_bits=10",
"permessage-deflate; client_max_window_bits=12",
)
.expect("server choosing a smaller window should validate"),
);
}
}