#![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,
future::Future,
ops::{Deref, DerefMut},
pin::Pin,
sync::Arc,
task::{Context, Poll},
};
use rama_core::Service;
use rama_core::error::{BoxError, ErrorContext, ErrorExt};
use rama_core::extensions::{Extensions, ExtensionsRef};
use rama_core::futures::{Sink, SinkExt as _, Stream, StreamExt as _};
use rama_core::rt::blocking::Io as BlockingIo;
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::blocking::Client as BlockingHttpClient;
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::{CloseFrame, Message, ProtocolError, Role, WebSocket, 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>,
}
struct PreparedHandshakeRequest {
request: Request,
protocols: Option<SecWebSocketProtocol>,
extensions: Option<SecWebSocketExtensions>,
config: Option<WebSocketConfig>,
key: Option<SecWebSocketKey>,
}
impl PreparedHandshakeRequest {
async fn send<S, Body>(
self,
service: &S,
) -> Result<NegotiatedHandshakeRequest<Body>, HandshakeError>
where
S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
{
let uri = self.request.uri().clone();
let response = service.serve(self.request).await.map_err(|err| {
let err: BoxError = err.into();
HandshakeError::HttpRequestError(
err.context(uri)
.context("send initial websocket handshake request (upgrade)"),
)
})?;
Ok(NegotiatedHandshakeRequest {
protocols: self.protocols,
extensions: self.extensions,
config: self.config,
key: self.key,
response,
})
}
}
pub struct WithService<'a, S, Body, Mode = websocket_builder_mode::Async> {
service: &'a S,
builder: RequestBuilder<'a, S, Response<Body>>,
config: Option<WebSocketConfig>,
is_h2: bool,
mode: Mode,
}
impl<S: fmt::Debug, Body, Mode: fmt::Debug> fmt::Debug for WithService<'_, S, Body, Mode> {
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)
.field("mode", &self.mode)
.finish()
}
}
pub mod websocket_builder_mode {
use std::sync::Arc;
use rama_core::rt::blocking::Runtime;
#[derive(Debug)]
#[non_exhaustive]
pub struct Async;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct Blocking<S> {
pub(crate) runtime: Runtime,
pub(crate) service: Arc<S>,
}
}
pub type BlockingWebSocketRequestBuilder<'a, S, Body> =
WebSocketRequestBuilder<WithService<'a, S, Body, websocket_builder_mode::Blocking<S>>>;
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().is_success() {
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, websocket_builder_mode::Async>>
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_and_mode(
service,
Version::HTTP_11,
uri,
websocket_builder_mode::Async,
)
}
pub fn new_h2_with_service<T>(service: &'a S, uri: T) -> Self
where
T: IntoUrl,
{
Self::new_with_service_and_version_and_mode(
service,
Version::HTTP_2,
uri,
websocket_builder_mode::Async,
)
}
pub fn new_with_service_and_request<RequestBody>(
service: &'a S,
request: Request<RequestBody>,
) -> Self
where
RequestBody: Into<rama_http::Body>,
{
Self::new_with_service_request_and_mode(service, request, websocket_builder_mode::Async)
}
}
impl<'a, S, Body, Mode> WebSocketRequestBuilder<WithService<'a, S, Body, Mode>>
where
S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
{
fn new_with_service_and_version_and_mode<T>(
service: &'a S,
version: Version,
uri: T,
mode: Mode,
) -> Self
where
T: IntoUrl,
{
Self {
inner: WithService {
service,
builder: new_ws_request_builder_from_uri_with_service(service, uri, version),
config: Default::default(),
is_h2: version == Version::HTTP_2,
mode,
},
protocols: Default::default(),
extensions: Default::default(),
key: Default::default(),
}
}
fn new_with_service_request_and_mode<RequestBody>(
service: &'a S,
request: Request<RequestBody>,
mode: Mode,
) -> 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 {
service,
builder: new_ws_request_builder_from_request(service, request),
config: Default::default(),
is_h2,
mode,
},
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
}
}
fn prepare_handshake_inner(
self,
extensions: &Extensions,
) -> Result<PreparedHandshakeRequest, 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 request = builder
.build()
.context("build initial websocket handshake request (upgrade)")
.map_err(HandshakeError::HttpRequestError)?;
Ok(PreparedHandshakeRequest {
request,
protocols: self.protocols,
extensions: self.extensions,
config: self.inner.config,
key,
})
}
async fn initiate_handshake_inner(
self,
extensions: Extensions,
) -> Result<NegotiatedHandshakeRequest<Body>, HandshakeError> {
let service = self.inner.service;
let prepared = self.prepare_handshake_inner(&extensions)?;
prepared.send(service).await
}
}
impl<'a, S, Body> WebSocketRequestBuilder<WithService<'a, S, Body, websocket_builder_mode::Async>>
where
S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
{
pub async fn initiate_handshake(
self,
extensions: Extensions,
) -> Result<NegotiatedHandshakeRequest<Body>, HandshakeError> {
self.initiate_handshake_inner(extensions).await
}
pub async fn handshake(self, extensions: Extensions) -> Result<ClientWebSocket, HandshakeError>
where
Body: Send + 'static,
{
let handshake = self.initiate_handshake(extensions).await?;
handshake.complete().await
}
}
impl<'a, S, Body>
WebSocketRequestBuilder<WithService<'a, S, Body, websocket_builder_mode::Blocking<S>>>
where
S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
Body: Send + 'static,
{
fn new_blocking_with_service<T>(client: &'a BlockingHttpClient<S>, uri: T) -> Self
where
T: IntoUrl,
{
Self::new_with_service_and_version_and_mode(
client.get_ref(),
Version::HTTP_11,
uri,
websocket_builder_mode::Blocking {
runtime: client.runtime().clone(),
service: client.clone_service(),
},
)
}
fn new_blocking_h2_with_service<T>(client: &'a BlockingHttpClient<S>, uri: T) -> Self
where
T: IntoUrl,
{
Self::new_with_service_and_version_and_mode(
client.get_ref(),
Version::HTTP_2,
uri,
websocket_builder_mode::Blocking {
runtime: client.runtime().clone(),
service: client.clone_service(),
},
)
}
fn new_blocking_with_service_and_request<RequestBody>(
client: &'a BlockingHttpClient<S>,
request: Request<RequestBody>,
) -> Self
where
RequestBody: Into<rama_http::Body>,
{
Self::new_with_service_request_and_mode(
client.get_ref(),
request,
websocket_builder_mode::Blocking {
runtime: client.runtime().clone(),
service: client.clone_service(),
},
)
}
pub fn try_handshake(self) -> Result<BlockingClientWebSocket, HandshakeError> {
self.try_handshake_with_extensions(Extensions::new())
}
#[expect(
clippy::needless_pass_by_value,
reason = "matches the async handshake API and transfers the extension set"
)]
pub fn try_handshake_with_extensions(
self,
extensions: Extensions,
) -> Result<BlockingClientWebSocket, HandshakeError> {
let runtime = self.inner.mode.runtime.clone();
let service = Arc::clone(&self.inner.mode.service);
let prepared = self.prepare_handshake_inner(&extensions)?;
let completed = runtime.block_on_task(async move {
let handshake = prepared.send(service.as_ref()).await?;
handshake.complete_upgrade().await
})?;
let socket = WebSocket::from_raw_socket(
runtime.io(completed.stream),
Role::Client,
completed.config,
);
Ok(BlockingClientWebSocket {
socket,
response: completed.response,
accepted_protocol: completed.accepted_protocol,
})
}
}
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>,
}
struct CompletedClientHandshake {
stream: rama_http::io::upgrade::Upgraded,
response: response::Parts,
accepted_protocol: Option<AcceptedWebSocketProtocol>,
config: Option<WebSocketConfig>,
}
impl<Body> NegotiatedHandshakeRequest<Body> {
pub async fn complete(self) -> Result<ClientWebSocket, HandshakeError>
where
Body: Send + 'static,
{
let completed = self.complete_upgrade().await?;
let socket =
AsyncWebSocket::from_raw_socket(completed.stream, Role::Client, completed.config).await;
Ok(ClientWebSocket {
socket,
response: completed.response,
accepted_protocol: completed.accepted_protocol,
})
}
async fn complete_upgrade(self) -> Result<CompletedClientHandshake, HandshakeError>
where
Body: Send + 'static,
{
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",
);
#[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,
))),
));
}
self.config
};
let on_upgrade = rama_http::io::upgrade::handle_upgrade(&self.response);
let (parts, body) = self.response.into_parts();
let stream = on_upgrade
.await
.context("upgrade http connection into a raw web socket")
.map_err(HandshakeError::HttpUpgradeError)?
.with_guard(body);
Ok(CompletedClientHandshake {
stream,
response: parts,
accepted_protocol: accepted_data.protocol,
config: maybe_ws_cfg,
})
}
}
#[derive(Debug)]
pub struct ClientWebSocket<S = AsyncWebSocket> {
pub socket: S,
pub response: response::Parts,
pub accepted_protocol: Option<AcceptedWebSocketProtocol>,
}
impl<S> Deref for ClientWebSocket<S> {
type Target = S;
fn deref(&self) -> &Self::Target {
&self.socket
}
}
impl<S> DerefMut for ClientWebSocket<S> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.socket
}
}
impl<S> Stream for ClientWebSocket<S>
where
S: Stream<Item = Result<Message, ProtocolError>> + Unpin,
{
type Item = Result<Message, ProtocolError>;
fn poll_next(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Stream::poll_next(Pin::new(&mut self.get_mut().socket), ctx)
}
}
impl<S> Sink<Message> for ClientWebSocket<S>
where
S: Sink<Message, Error = ProtocolError> + Unpin,
{
type Error = ProtocolError;
fn poll_ready(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Sink::poll_ready(Pin::new(&mut self.get_mut().socket), ctx)
}
fn start_send(self: Pin<&mut Self>, message: Message) -> Result<(), Self::Error> {
Sink::start_send(Pin::new(&mut self.get_mut().socket), message)
}
fn poll_flush(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Sink::poll_flush(Pin::new(&mut self.get_mut().socket), ctx)
}
fn poll_close(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Sink::poll_close(Pin::new(&mut self.get_mut().socket), ctx)
}
}
impl<S> ExtensionsRef for ClientWebSocket<S>
where
S: ExtensionsRef,
{
fn extensions(&self) -> &Extensions {
self.socket.extensions()
}
}
impl<S> ClientWebSocket<S> {
#[must_use]
pub fn map_socket<T>(self, map: impl FnOnce(S) -> T) -> ClientWebSocket<T> {
ClientWebSocket {
socket: map(self.socket),
response: self.response,
accepted_protocol: self.accepted_protocol,
}
}
pub fn send_message(
&mut self,
message: Message,
) -> impl Future<Output = Result<(), ProtocolError>> + Send + '_
where
S: Sink<Message, Error = ProtocolError> + Send + Unpin,
{
self.socket.send(message)
}
pub async fn recv_message(&mut self) -> Result<Message, ProtocolError>
where
S: Stream<Item = Result<Message, ProtocolError>> + Unpin,
{
self.socket.next().await.ok_or_else(|| {
ProtocolError::Io(std::io::Error::new(
std::io::ErrorKind::ConnectionAborted,
"Connection closed: no messages to receive",
))
})?
}
pub async fn close(&mut self, message: Option<CloseFrame>) -> Result<(), ProtocolError>
where
S: Sink<Message, Error = ProtocolError> + Send + Unpin,
{
self.socket.send(Message::Close(message)).await
}
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) -> S {
self.socket
}
}
pub type BlockingWebSocket = WebSocket<BlockingIo<rama_http::io::upgrade::Upgraded>>;
#[derive(Debug)]
pub struct BlockingClientWebSocket {
pub socket: BlockingWebSocket,
pub response: response::Parts,
pub accepted_protocol: Option<AcceptedWebSocketProtocol>,
}
impl Deref for BlockingClientWebSocket {
type Target = BlockingWebSocket;
fn deref(&self) -> &Self::Target {
&self.socket
}
}
impl DerefMut for BlockingClientWebSocket {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.socket
}
}
impl BlockingClientWebSocket {
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 send_message(&mut self, message: Message) -> Result<(), ProtocolError> {
self.socket.send(message)
}
pub fn recv_message(&mut self) -> Result<Message, ProtocolError> {
self.socket.read()
}
pub fn into_inner(self) -> BlockingWebSocket {
self.socket
}
}
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)
}
}
pub trait BlockingHttpClientWebSocketExt<Body>:
private::BlockingHttpClientWebSocketExtSealed<Body>
{
type AsyncService: Service<Request, Output = Response<Body>, Error: Into<BoxError>>;
fn websocket(
&self,
url: impl IntoUrl,
) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body>;
fn websocket_h2(
&self,
url: impl IntoUrl,
) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body>;
fn websocket_with_request<RequestBody: Into<rama_http::Body>>(
&self,
request: Request<RequestBody>,
) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body>;
}
impl<S, Body> BlockingHttpClientWebSocketExt<Body> for BlockingHttpClient<S>
where
S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
Body: Send + 'static,
{
type AsyncService = S;
fn websocket(
&self,
url: impl IntoUrl,
) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body> {
BlockingWebSocketRequestBuilder::new_blocking_with_service(self, url)
}
fn websocket_h2(
&self,
url: impl IntoUrl,
) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body> {
BlockingWebSocketRequestBuilder::new_blocking_h2_with_service(self, url)
}
fn websocket_with_request<RequestBody: Into<rama_http::Body>>(
&self,
request: Request<RequestBody>,
) -> BlockingWebSocketRequestBuilder<'_, Self::AsyncService, Body> {
BlockingWebSocketRequestBuilder::new_blocking_with_service_and_request(self, request)
}
}
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>>
{
}
pub trait BlockingHttpClientWebSocketExtSealed<Body> {}
impl<S, Body> BlockingHttpClientWebSocketExtSealed<Body> for BlockingHttpClient<S>
where
S: Service<Request, Output = Response<Body>, Error: Into<BoxError>>,
Body: Send + 'static,
{
}
}
#[cfg(test)]
mod tests {
use super::*;
use rama_core::{ServiceInput, bytes::Bytes, service::service_fn};
use rama_http::HeaderMap;
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
struct ResponseLease(Arc<AtomicUsize>);
impl Drop for ResponseLease {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::Release);
}
}
#[test]
fn blocking_client_websocket_roundtrip_and_lifetimes() {
fn assert_send<T: Send>() {}
assert_send::<BlockingClientWebSocket>();
let leases_dropped = Arc::new(AtomicUsize::new(0));
let service_leases_dropped = leases_dropped.clone();
let service = service_fn(move |request: Request| {
let leases_dropped = service_leases_dropped.clone();
async move {
let is_h2 = request.version() == Version::HTTP_2;
let accept = if is_h2 {
assert_eq!(request.method(), Method::CONNECT);
assert!(request.headers().typed_get::<SecWebSocketKey>().is_none());
None
} else {
assert_eq!(request.method(), Method::GET);
let key = request
.headers()
.typed_get::<SecWebSocketKey>()
.expect("HTTP/1.1 client handshake request to contain a key");
Some(
headers::SecWebSocketAccept::try_from(key)
.expect("client handshake key to produce an accept value"),
)
};
if request.uri().path().is_some_and(|path| path == "/custom") {
assert_eq!(
request.headers().get("x-rama-test"),
Some(&rama_http::HeaderValue::from_static("custom")),
);
}
let (client_io, server_io) = tokio::io::duplex(4 * 1024);
let (pending, on_upgrade) = rama_http::io::upgrade::pending();
pending.fulfill(rama_http::io::upgrade::Upgraded::new(
ServiceInput::new(client_io),
Bytes::new(),
));
tokio::spawn(async move {
let mut socket = AsyncWebSocket::from_raw_socket(
ServiceInput::new(server_io),
Role::Server,
None,
)
.await;
let message = socket.recv_message().await.unwrap();
socket.send_message(message).await.unwrap();
});
let mut response = Response::new(ResponseLease(leases_dropped));
if let Some(accept) = accept {
*response.status_mut() = StatusCode::SWITCHING_PROTOCOLS;
*response.version_mut() = Version::HTTP_11;
response
.headers_mut()
.typed_insert(headers::Upgrade::websocket());
response
.headers_mut()
.typed_insert(headers::Connection::upgrade());
response.headers_mut().typed_insert(accept);
} else {
*response.status_mut() = StatusCode::OK;
*response.version_mut() = Version::HTTP_2;
}
response.extensions().insert(on_upgrade);
Ok::<_, BoxError>(response)
}
});
let client = BlockingHttpClient::try_new(service).unwrap();
let client_clone = client.clone();
drop(client);
let config = WebSocketConfig::default().with_read_buffer_size(4 * 1024);
let mut from_url = client_clone
.websocket("ws://example.test/echo")
.with_config(config)
.try_handshake()
.unwrap();
assert_eq!(from_url.response().status, StatusCode::SWITCHING_PROTOCOLS);
assert_eq!(from_url.get_config().read_buffer_size, 4 * 1024);
assert_eq!(
from_url
.extensions()
.get_ref::<StreamTransformed>()
.unwrap()
.by,
"rama-http::Upgraded",
);
let request = Request::builder()
.version(Version::HTTP_11)
.uri("ws://example.test/custom")
.header("x-rama-test", "custom")
.body(Body::empty())
.unwrap();
let mut from_request = client_clone
.websocket_with_request(request)
.try_handshake()
.unwrap();
let mut from_h2 = client_clone
.websocket_h2("wss://example.test/h2")
.try_handshake()
.unwrap();
drop(client_clone);
assert_eq!(leases_dropped.load(Ordering::Acquire), 0);
from_url.send_message("from url".into()).unwrap();
assert_eq!(
from_url
.recv_message()
.unwrap()
.into_text()
.unwrap()
.as_str(),
"from url",
);
from_request.send_message("from request".into()).unwrap();
assert_eq!(
from_request
.recv_message()
.unwrap()
.into_text()
.unwrap()
.as_str(),
"from request",
);
from_h2.send_message("from h2".into()).unwrap();
assert_eq!(
from_h2
.recv_message()
.unwrap()
.into_text()
.unwrap()
.as_str(),
"from h2",
);
let BlockingClientWebSocket {
socket: from_url,
response,
accepted_protocol: protocol,
} = from_url;
assert_eq!(response.status, StatusCode::SWITCHING_PROTOCOLS);
assert!(protocol.is_none());
assert_eq!(leases_dropped.load(Ordering::Acquire), 0);
drop(from_url);
assert_eq!(leases_dropped.load(Ordering::Acquire), 1);
drop(from_request);
assert_eq!(leases_dropped.load(Ordering::Acquire), 2);
drop(from_h2);
assert_eq!(leases_dropped.load(Ordering::Acquire), 3);
}
#[cfg(feature = "dial9")]
#[test]
fn blocking_handshake_runs_inside_dial9_session() {
let temp_dir = tempfile::tempdir().unwrap();
let config = rama_core::telemetry::dial9::Dial9Config::builder()
.enabled(true)
.base_path(temp_dir.path().join("blocking-websocket.bin"))
.max_file_size(1024 * 1024)
.max_total_size(4 * 1024 * 1024)
.build()
.unwrap();
let runtime = rama_core::rt::blocking::Runtime::builder()
.with_dial9_config(config)
.try_build()
.unwrap();
let service = service_fn(|request: Request| async move {
assert!(
rama_core::telemetry::dial9::telemetry::TelemetryHandle::current().is_enabled()
);
let key = request
.headers()
.typed_get::<SecWebSocketKey>()
.expect("handshake request to contain a key");
let (client_io, _server_io) = tokio::io::duplex(1024);
let (pending, on_upgrade) = rama_http::io::upgrade::pending();
pending.fulfill(rama_http::io::upgrade::Upgraded::new(
ServiceInput::new(client_io),
Bytes::new(),
));
let mut response = Response::new(());
*response.status_mut() = StatusCode::SWITCHING_PROTOCOLS;
*response.version_mut() = Version::HTTP_11;
response
.headers_mut()
.typed_insert(headers::Upgrade::websocket());
response
.headers_mut()
.typed_insert(headers::Connection::upgrade());
response.headers_mut().typed_insert(
headers::SecWebSocketAccept::try_from(key)
.expect("client handshake key to produce an accept value"),
);
response.extensions().insert(on_upgrade);
Ok::<_, BoxError>(response)
});
let client = BlockingHttpClient::with_runtime(service, &runtime);
let socket = client
.websocket("wss://example.test/socket")
.try_handshake()
.unwrap();
assert_eq!(socket.response().status, StatusCode::SWITCHING_PROTOCOLS);
}
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
}
#[test]
fn h2_handshake_accepts_any_successful_connect_status() {
let mut response = Response::new(());
*response.version_mut() = Version::HTTP_2;
*response.status_mut() = StatusCode::CREATED;
validate_http_server_response(&response, None, None, None)
.expect("successful CONNECT response");
*response.status_mut() = StatusCode::BAD_REQUEST;
assert!(matches!(
validate_http_server_response(&response, None, None, None),
Err(ResponseValidateError::UnexpectedStatusCode(
StatusCode::BAD_REQUEST
))
));
}
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"),
);
}
}