use std::{
fmt,
future::Future,
ops::{Deref, DerefMut},
pin::Pin,
task::{Context, Poll},
};
use rama_core::{
Service,
error::{BoxError, BoxErrorExt as _, ErrorContext},
extensions::{Extensions, ExtensionsRef},
futures::{Sink, SinkExt as _, Stream, StreamExt as _},
rt::Executor,
telemetry::tracing::{self, Instrument},
};
#[cfg(feature = "compression")]
use rama_http::headers::sec_websocket_extensions;
use rama_http::{
Method, Request, Response, StatusCode, Version,
headers::{
self, HeaderMapExt,
sec_websocket_extensions::{Extension, PerMessageDeflateConfig},
},
io::upgrade,
layer::upgrade::UpgradeResponse,
proto::h2::ext::Protocol,
request,
service::web::response::{self, Headers, IntoResponse},
};
use rama_net::extensions::StreamTransformed;
use rama_utils::{
collections::non_empty_smallvec,
str::{NonEmptyStr, non_empty_str},
};
use crate::{
Message, ProtocolError, WebSocketIo,
protocol::{CloseFrame, Role, WebSocketConfig},
runtime::AsyncWebSocket,
};
#[derive(Debug)]
pub enum RequestValidateError {
UnexpectedHttpMethod(Method),
UnexpectedHttpVersion(Version),
UnexpectedPseudoProtocolHeader(Option<Protocol>),
MissingUpgradeWebSocketHeader,
MissingConnectionUpgradeHeader,
InvalidSecWebSocketVersionHeader,
InvalidSecWebSocketKeyHeader,
InvalidSecWebSocketProtocolHeader(BoxError),
}
impl fmt::Display for RequestValidateError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::UnexpectedHttpMethod(method) => {
write!(f, "unexpected HTTP method: {method:?}")
}
Self::UnexpectedHttpVersion(version) => {
write!(f, "unexpected HTTP version: {version:?}")
}
Self::UnexpectedPseudoProtocolHeader(maybe_protocol) => {
write!(
f,
"missing or invalid pseudo h2 protocol header: {maybe_protocol:?}"
)
}
Self::MissingUpgradeWebSocketHeader => {
write!(f, "missing upgrade WebSocket header")
}
Self::MissingConnectionUpgradeHeader => {
write!(f, "missing connection upgrade header")
}
Self::InvalidSecWebSocketVersionHeader => {
write!(f, "missing or invalid sec-websocket-version header")
}
Self::InvalidSecWebSocketKeyHeader => {
write!(f, "missing or invalid sec-websocket-key header")
}
Self::InvalidSecWebSocketProtocolHeader(err) => {
write!(f, "invalid sec-websocket-protocol header: {err}")
}
}
}
}
impl std::error::Error for RequestValidateError {}
#[derive(Debug)]
pub struct ClientRequestData {
pub accept_header: Option<headers::SecWebSocketAccept>,
pub protocol: Option<headers::SecWebSocketProtocol>,
pub extensions: Option<headers::SecWebSocketExtensions>,
}
pub fn validate_http_client_request<Body>(
request: &Request<Body>,
) -> Result<ClientRequestData, RequestValidateError> {
tracing::trace!(
http.version = ?request.version(),
"validate http client request"
);
let mut accept_header = None;
match request.version() {
Version::HTTP_10 | Version::HTTP_11 => {
match request.method() {
&Method::GET => (),
method => return Err(RequestValidateError::UnexpectedHttpMethod(method.clone())),
}
if !request
.headers()
.typed_get::<headers::Upgrade>()
.map(|u| u.is_websocket())
.unwrap_or_default()
{
return Err(RequestValidateError::MissingUpgradeWebSocketHeader);
}
if !request
.headers()
.typed_get::<headers::Connection>()
.map(|c| c.contains_upgrade())
.unwrap_or_default()
{
return Err(RequestValidateError::MissingConnectionUpgradeHeader);
}
accept_header = match request.headers().typed_get::<headers::SecWebSocketKey>() {
Some(key) => headers::SecWebSocketAccept::try_from(key)
.inspect_err(|err| {
tracing::debug!(
"failed to create accept typed header from given key: {err}"
)
})
.ok(),
None => return Err(RequestValidateError::InvalidSecWebSocketKeyHeader),
};
}
Version::HTTP_2 => {
match request.method() {
&Method::CONNECT => (),
method => return Err(RequestValidateError::UnexpectedHttpMethod(method.clone())),
}
match request.extensions().get_ref::<Protocol>() {
None => return Err(RequestValidateError::UnexpectedPseudoProtocolHeader(None)),
Some(protocol) => {
if !protocol.as_str().trim().eq_ignore_ascii_case("websocket") {
return Err(RequestValidateError::UnexpectedPseudoProtocolHeader(Some(
protocol.clone(),
)));
}
}
}
}
version => {
return Err(RequestValidateError::UnexpectedHttpVersion(version));
}
}
if request
.headers()
.typed_get::<headers::SecWebSocketVersion>()
.is_none()
{
return Err(RequestValidateError::InvalidSecWebSocketVersionHeader);
}
let protocols_header = request.headers().typed_get();
let extensions_header = request.headers().typed_get();
Ok(ClientRequestData {
accept_header,
protocol: protocols_header,
extensions: extensions_header,
})
}
#[derive(Debug, Clone, Default)]
pub struct WebSocketAcceptor {
protocols: Option<headers::SecWebSocketProtocol>,
protocols_flex: bool,
extensions: Option<headers::SecWebSocketExtensions>,
}
impl WebSocketAcceptor {
#[inline]
#[must_use]
pub fn new() -> Self {
Default::default()
}
rama_utils::macros::generate_set_and_with! {
pub fn protocols_flex(mut self, flexible: bool) -> Self {
self.protocols_flex = flexible;
self
}
}
rama_utils::macros::generate_set_and_with! {
pub fn protocols(mut self, protocols: Option<headers::SecWebSocketProtocol>) -> Self {
self.protocols = protocols;
self
}
}
rama_utils::macros::generate_set_and_with! {
pub fn echo_protocols(mut self) -> Self {
self.protocols = Some(headers::SecWebSocketProtocol(non_empty_smallvec![
ECHO_SERVICE_SUB_PROTOCOL_DEFAULT,
ECHO_SERVICE_SUB_PROTOCOL_UPPER,
ECHO_SERVICE_SUB_PROTOCOL_LOWER,
]));
self
}
}
rama_utils::macros::generate_set_and_with! {
pub fn extensions(mut self, extensions: Option<headers::SecWebSocketExtensions>) -> Self {
self.extensions = extensions;
self
}
}
#[cfg(feature = "compression")]
rama_utils::macros::generate_set_and_with! {
#[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(headers::SecWebSocketExtensions::per_message_deflate()),
};
self
}
}
#[cfg(feature = "compression")]
rama_utils::macros::generate_set_and_with! {
#[cfg_attr(docsrs, doc(cfg(feature = "compression")))]
pub fn per_message_deflate_overwrite_extensions(mut self) -> Self {
self.extensions = Some(headers::SecWebSocketExtensions::per_message_deflate());
self
}
}
#[cfg(feature = "compression")]
rama_utils::macros::generate_set_and_with! {
#[cfg_attr(docsrs, doc(cfg(feature = "compression")))]
pub fn per_message_deflate_with_config(mut self, config: impl Into<sec_websocket_extensions::PerMessageDeflateConfig>) -> Self {
self.extensions = match self.extensions.take() {
Some(ext) => {
Some(ext.with_extra_extension(Extension::PerMessageDeflate(config.into())))
},
None => Some(headers::SecWebSocketExtensions::per_message_deflate_with_config(config.into())),
};
self
}
}
#[cfg(feature = "compression")]
rama_utils::macros::generate_set_and_with! {
#[cfg_attr(docsrs, doc(cfg(feature = "compression")))]
pub fn per_message_deflate_with_config_overwrite_extensions(mut self, config: impl Into<sec_websocket_extensions::PerMessageDeflateConfig>) -> Self {
self.extensions = Some(headers::SecWebSocketExtensions::per_message_deflate_with_config(config.into()));
self
}
}
}
impl WebSocketAcceptor {
pub fn into_service<S>(self, service: S) -> WebSocketAcceptorService<S> {
WebSocketAcceptorService {
acceptor: self,
config: None,
exec: None,
service,
}
}
pub fn into_service_with_executor<S>(
self,
service: S,
exec: Executor,
) -> WebSocketAcceptorService<S> {
WebSocketAcceptorService {
acceptor: self,
config: None,
exec: Some(exec),
service,
}
}
#[must_use]
pub fn into_echo_service(mut self) -> WebSocketAcceptorService<WebSocketEchoService> {
self.norm_echo_flex_protocols_if_protocols_is_none();
WebSocketAcceptorService {
acceptor: self,
config: None,
exec: None,
service: WebSocketEchoService::new(),
}
}
#[must_use]
pub fn into_echo_service_with_executor(
mut self,
exec: Executor,
) -> WebSocketAcceptorService<WebSocketEchoService> {
self.norm_echo_flex_protocols_if_protocols_is_none();
WebSocketAcceptorService {
acceptor: self,
config: None,
exec: Some(exec),
service: WebSocketEchoService::new(),
}
}
fn norm_echo_flex_protocols_if_protocols_is_none(&mut self) {
if self.protocols.is_none() {
self.protocols_flex = true;
self.protocols = Some(headers::SecWebSocketProtocol(non_empty_smallvec![
ECHO_SERVICE_SUB_PROTOCOL_DEFAULT,
ECHO_SERVICE_SUB_PROTOCOL_UPPER,
ECHO_SERVICE_SUB_PROTOCOL_LOWER,
]));
}
}
}
impl<Body> Service<Request<Body>> for WebSocketAcceptor
where
Body: Send + 'static,
{
type Output = UpgradeResponse<Request<Body>, Response>;
type Error = Response;
async fn serve(&self, req: Request<Body>) -> Result<Self::Output, Self::Error> {
let extensions = Extensions::default();
extensions.insert(StreamTransformed {
by: "rama-ws::WebSocketAcceptor",
});
match validate_http_client_request(&req) {
Ok(request_data) => {
let accepted_protocol = match (
self.protocols_flex,
request_data.protocol,
self.protocols.as_ref(),
) {
(false, Some(protocols), None) => {
tracing::debug!(
"WebSocketAcceptor: protocols found while none were expected: {protocols:?}"
);
return Err(StatusCode::BAD_REQUEST.into_response());
}
(false, None, Some(protocols)) => {
tracing::debug!(
"WebSocketAcceptor: no protocols found while one of following was expected: {protocols:?}"
);
return Err(StatusCode::BAD_REQUEST.into_response());
}
(_, None, None) | (true, None, Some(_)) => None,
(true, Some(found_protocols), None) => {
Some(found_protocols.accept_first_protocol())
}
(_, Some(found_protocols), Some(expected_protocols)) => {
if let Some(protocol) =
found_protocols.contains_any(expected_protocols.iter())
{
Some(protocol)
} else {
tracing::debug!(
"WebSocketAcceptor: no protocols from found protocol ({found_protocols:?}) matched for expected protocols: {expected_protocols:?}"
);
return Err(StatusCode::BAD_REQUEST.into_response());
}
}
};
let accepted_extension = match (request_data.extensions, self.extensions.as_ref()) {
(None, _) | (_, None) => None,
(Some(request_extensions), Some(allowed_extensions)) => {
request_extensions.0.iter().find_map(|request_ext| {
for allowed_ext in allowed_extensions.0.iter() {
if let (
Extension::PerMessageDeflate(request_pmd),
Extension::PerMessageDeflate(allowed_pmd),
) = (&request_ext, allowed_ext)
{
let mut resp = PerMessageDeflateConfig {
identifier: allowed_pmd.identifier.clone(),
client_no_context_takeover: request_pmd
.client_no_context_takeover
&& allowed_pmd.client_no_context_takeover,
server_no_context_takeover: allowed_pmd
.server_no_context_takeover,
..Default::default()
};
let srv_cap = allowed_pmd.server_max_window_bits.unwrap_or(15);
let srv_cap = if srv_cap == 0 {
15
} else {
srv_cap.clamp(8, 15)
};
let cli_req_srv = request_pmd
.server_max_window_bits
.map(|v| if v == 0 { 15 } else { v.clamp(8, 15) });
let chosen_srv_bits = match (cli_req_srv, Some(srv_cap)) {
(Some(client_bits), Some(cap)) => {
Some(client_bits.min(cap))
}
(None, Some(cap)) => Some(cap),
_ => None,
};
resp.server_max_window_bits = match chosen_srv_bits {
Some(bits) if bits < 15 || cli_req_srv.is_some() => {
Some(bits)
}
_ => None,
};
resp.client_max_window_bits = request_pmd
.client_max_window_bits
.map(|client_bits_offer| {
let offer = if client_bits_offer == 0 {
15
} else {
client_bits_offer.clamp(8, 15)
};
let cap =
allowed_pmd.client_max_window_bits.unwrap_or(offer);
if cap == 0 {
offer
} else {
offer.min(cap.clamp(8, 15))
}
});
tracing::trace!(
"accept and use ws deflate ext w/ config: {resp:?}"
);
return Some(Extension::PerMessageDeflate(resp));
}
}
None
})
}
};
let protocols_header = match accepted_protocol {
Some(p) => {
tracing::trace!("inject accepted ws protocol in cfg: {p:?}");
extensions.insert(p.clone());
Some(p.into_header())
}
None => None,
};
let extensions_header = match accepted_extension {
Some(ext) => {
tracing::trace!("inject accepted ws extension in cfg: {ext:?}");
extensions.insert(ext.clone());
Some(ext.into_header())
}
None => None,
};
match req.version() {
version @ (Version::HTTP_10 | Version::HTTP_11) => {
let accept_header = request_data.accept_header.ok_or_else(|| {
tracing::debug!("WebSocketAcceptor: missing accept header (no key?)");
StatusCode::BAD_REQUEST.into_response()
})?;
let mut response = (
StatusCode::SWITCHING_PROTOCOLS,
response::Headers((
accept_header,
headers::Upgrade::websocket(),
headers::Connection::upgrade(),
)),
)
.into_response();
*response.version_mut() = version;
if let Some(protocols) = protocols_header {
response.headers_mut().typed_insert(protocols);
}
if let Some(extensions) = extensions_header {
response.headers_mut().typed_insert(extensions);
}
Ok(UpgradeResponse {
response,
request: req,
extensions,
})
}
Version::HTTP_2 => {
let mut response = StatusCode::OK.into_response();
*response.version_mut() = Version::HTTP_2;
if let Some(protocols) = protocols_header {
response.headers_mut().typed_insert(protocols);
}
if let Some(extensions) = extensions_header {
response.headers_mut().typed_insert(extensions);
}
Ok(UpgradeResponse {
response,
request: req,
extensions,
})
}
version => {
tracing::debug!(
http.version = ?version,
"WebSocketAcceptor: http client request has unexpected http version"
);
Err(StatusCode::BAD_REQUEST.into_response())
}
}
}
Err(err) => {
let response =
if matches!(err, RequestValidateError::InvalidSecWebSocketVersionHeader) {
(
Headers::single(headers::SecWebSocketVersion::V13),
StatusCode::BAD_REQUEST,
)
.into_response()
} else {
StatusCode::BAD_REQUEST.into_response()
};
tracing::debug!("WebSocketAcceptor: http client request failed to validate: {err}");
Err(response)
}
}
}
}
#[derive(Debug, Clone)]
pub struct WebSocketAcceptorService<S> {
acceptor: WebSocketAcceptor,
config: Option<WebSocketConfig>,
exec: Option<Executor>,
service: S,
}
impl<S> WebSocketAcceptorService<S> {
rama_utils::macros::generate_set_and_with! {
pub fn config(mut self, cfg: Option<WebSocketConfig>) -> Self {
self.config = cfg;
self
}
}
}
#[derive(Debug)]
pub struct ServerWebSocket<S = AsyncWebSocket> {
pub socket: S,
pub request: request::Parts,
}
impl<S> Deref for ServerWebSocket<S> {
type Target = S;
fn deref(&self) -> &Self::Target {
&self.socket
}
}
impl<S> DerefMut for ServerWebSocket<S> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.socket
}
}
impl<S> Stream for ServerWebSocket<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 ServerWebSocket<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 ServerWebSocket<S>
where
S: ExtensionsRef,
{
fn extensions(&self) -> &Extensions {
self.socket.extensions()
}
}
impl<S> ServerWebSocket<S> {
#[must_use]
pub fn map_socket<T>(self, map: impl FnOnce(S) -> T) -> ServerWebSocket<T> {
ServerWebSocket {
socket: map(self.socket),
request: self.request,
}
}
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 request(&self) -> &request::Parts {
&self.request
}
pub fn into_inner(self) -> S {
self.socket
}
}
impl<S, Body> Service<Request<Body>> for WebSocketAcceptorService<S>
where
S: Clone + Service<ServerWebSocket, Output = ()>,
Body: Send + 'static,
{
type Output = Response;
type Error = S::Error;
async fn serve(&self, req: Request<Body>) -> Result<Self::Output, Self::Error> {
match self.acceptor.serve(req).await {
Ok(UpgradeResponse {
response: resp,
request: req,
extensions,
}) => {
#[cfg(not(feature = "compression"))]
if let Some(Extension::PerMessageDeflate(_)) = extensions.get_ref() {
tracing::error!(
"per-message-deflate is used but compression feature is disabled. Enable it if you wish to use this extension."
);
return Ok(StatusCode::INTERNAL_SERVER_ERROR.into_response());
}
let handler = self.service.clone();
let span = tracing::trace_root_span!(
"ws::serve",
otel.kind = "server",
url.full = %req.request_uri(),
url.path = %req.uri().path_or_root().as_ref(),
url.query = %req.uri().query_or_empty().as_ref(),
url.scheme = %req.uri().scheme_str().unwrap_or_default(),
network.protocol.name = "ws",
);
let exec = self.exec.clone().unwrap_or_default();
exec.into_spawn_task(
async move {
match upgrade::handle_upgrade(&req).await {
Ok(upgraded) => {
upgraded.extensions().extend(&extensions);
#[cfg(feature = "compression")]
let maybe_ws_config = {
let mut ws_cfg = None;
tracing::trace!("check if pmd settings have to be applied to WS cfg...");
if let Some(Extension::PerMessageDeflate(pmd_cfg)) = extensions.get_ref() {
tracing::trace!(
"apply accepted per-message-deflate cfg into WS server config: {pmd_cfg:?}"
);
ws_cfg = Some(WebSocketConfig {
per_message_deflate: Some(pmd_cfg.into()),
..Default::default()
});
}
ws_cfg
};
#[cfg(not(feature = "compression"))]
let maybe_ws_config = None;
let socket =
AsyncWebSocket::from_raw_socket(upgraded, Role::Server, maybe_ws_config)
.await;
let (parts, _) = req.into_parts();
let server_socket = ServerWebSocket {
socket,
request: parts,
};
_ = handler.serve(server_socket).await;
}
Err(e) => {
tracing::error!("ws upgrade error: {e:?}");
}
}
}
.instrument(span),
);
Ok(resp)
}
Err(resp) => Ok(resp),
}
}
}
const ECHO_SERVICE_SUB_PROTOCOL_DEFAULT_STR: &str = "echo";
pub const ECHO_SERVICE_SUB_PROTOCOL_DEFAULT: NonEmptyStr =
non_empty_str!(ECHO_SERVICE_SUB_PROTOCOL_DEFAULT_STR);
pub const ECHO_SERVICE_SUB_PROTOCOL_UPPER: NonEmptyStr = non_empty_str!("echo-upper");
pub const ECHO_SERVICE_SUB_PROTOCOL_LOWER: NonEmptyStr = non_empty_str!("echo-lower");
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct WebSocketEchoService;
impl WebSocketEchoService {
#[must_use]
pub fn new() -> Self {
Self
}
}
async fn serve_echo_web_socket<S>(mut socket: S) -> Result<(), BoxError>
where
S: ExtensionsRef
+ Stream<Item = Result<Message, ProtocolError>>
+ Sink<Message, Error = ProtocolError>
+ Unpin,
{
let protocol = socket
.extensions()
.get_ref::<headers::sec_websocket_protocol::AcceptedWebSocketProtocol>()
.map(|p| p.0.as_ref())
.unwrap_or(ECHO_SERVICE_SUB_PROTOCOL_DEFAULT_STR);
let transformer = if protocol.eq_ignore_ascii_case(&ECHO_SERVICE_SUB_PROTOCOL_LOWER) {
|msg: Message| match msg {
Message::Text(original) => Some(original.to_lowercase().into()),
msg @ Message::Binary(_) => Some(msg),
Message::Ping(_) | Message::Pong(_) | Message::Close(_) | Message::Frame(_) => None,
}
} else if protocol.eq_ignore_ascii_case(&ECHO_SERVICE_SUB_PROTOCOL_UPPER) {
|msg: Message| match msg {
Message::Text(original) => Some(original.to_uppercase().into()),
msg @ Message::Binary(_) => Some(msg),
Message::Ping(_) | Message::Pong(_) | Message::Close(_) | Message::Frame(_) => None,
}
} else {
|msg: Message| match msg {
msg @ (Message::Text(_) | Message::Binary(_)) => Some(msg),
Message::Ping(_) | Message::Pong(_) | Message::Close(_) | Message::Frame(_) => None,
}
};
loop {
let msg = socket
.next()
.await
.ok_or_else(|| BoxError::from_static_str("WebSocket closed while waiting for message"))?
.context("recv next msg")?;
if let Some(msg2) = transformer(msg) {
socket.send(msg2).await.context("echo msg back")?;
}
}
}
impl Service<AsyncWebSocket> for WebSocketEchoService {
type Output = ();
type Error = BoxError;
async fn serve(&self, socket: AsyncWebSocket) -> Result<Self::Output, Self::Error> {
serve_echo_web_socket(socket).await
}
}
impl<S> Service<ServerWebSocket<S>> for WebSocketEchoService
where
S: WebSocketIo,
{
type Output = ();
type Error = BoxError;
async fn serve(&self, socket: ServerWebSocket<S>) -> Result<Self::Output, Self::Error> {
let socket = socket.socket;
serve_echo_web_socket(socket).await
}
}
impl Service<upgrade::Upgraded> for WebSocketEchoService {
type Output = ();
type Error = BoxError;
async fn serve(&self, io: upgrade::Upgraded) -> Result<Self::Output, Self::Error> {
#[cfg(not(feature = "compression"))]
let maybe_ws_config = {
if let Some(Extension::PerMessageDeflate(_)) = io.extensions().get_ref() {
return Err(BoxError::from_static_str(
"per-message-deflate is used but compression feature is disabled. Enable it if you wish to use this extension.",
));
}
None
};
#[cfg(feature = "compression")]
let maybe_ws_config = {
let mut ws_cfg = None;
tracing::debug!("check if pmd settings have to be applied to WS cfg...");
if let Some(Extension::PerMessageDeflate(pmd_cfg)) = io.extensions().get_ref() {
tracing::debug!(
"apply accepted per-message-deflate cfg into WS server config: {pmd_cfg:?}"
);
ws_cfg = Some(WebSocketConfig {
per_message_deflate: Some(pmd_cfg.into()),
..Default::default()
});
}
ws_cfg
};
let socket = AsyncWebSocket::from_raw_socket(io, Role::Server, maybe_ws_config).await;
self.serve(socket).await
}
}
#[cfg(test)]
mod tests {
#![expect(
clippy::unreachable,
reason = "test fixtures: the `Version` matcher exhausts the variants the test setup actually produces"
)]
use headers::sec_websocket_protocol::AcceptedWebSocketProtocol;
use parking_lot::Mutex;
use rama_core::{Layer as _, ServiceInput, service::service_fn};
use rama_http::Body;
use rama_http::layer::har::{
recorder::{WebSocketCapture, WebSocketCaptureRecorder},
spec::{WebSocketMessage, WebSocketMessageType},
};
use rama_utils::str::non_empty_str;
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use super::*;
use crate::layer::har::{HARWebSocket, HARWebSocketLayer};
#[derive(Default)]
struct CaptureState {
messages: Mutex<Vec<WebSocketMessage>>,
closes: AtomicUsize,
}
struct TestCaptureRecorder(Arc<CaptureState>);
impl WebSocketCaptureRecorder for TestCaptureRecorder {
async fn record(&self, message: WebSocketMessage) -> Result<(), BoxError> {
self.0.messages.lock().push(message);
Ok(())
}
}
async fn assert_websocket_acceptor_ok(
request: Request,
acceptor: &WebSocketAcceptor,
expected_accepted_protocol: Option<AcceptedWebSocketProtocol>,
) {
let UpgradeResponse {
response: resp,
request: req,
extensions,
} = acceptor.serve(request).await.unwrap();
match req.version() {
Version::HTTP_10 | Version::HTTP_11 => {
assert_eq!(StatusCode::SWITCHING_PROTOCOLS, resp.status())
}
Version::HTTP_2 => assert_eq!(StatusCode::OK, resp.status()),
_ => unreachable!(),
}
let accepted_protocol = resp
.headers()
.typed_get::<headers::SecWebSocketProtocol>()
.map(|p| p.accept_first_protocol());
if let Some(expected_accepted_protocol) = expected_accepted_protocol {
assert_eq!(
accepted_protocol.as_ref(),
Some(&expected_accepted_protocol),
"request = {req:?}"
);
assert_eq!(
extensions.get_ref::<AcceptedWebSocketProtocol>(),
Some(&expected_accepted_protocol),
"request = {req:?}"
);
} else {
assert!(accepted_protocol.is_none());
assert!(extensions.get_ref::<AcceptedWebSocketProtocol>().is_none());
}
}
#[tokio::test]
async fn har_layer_explicitly_wraps_server_endpoint() {
type TestSocket = AsyncWebSocket<ServiceInput<tokio::io::DuplexStream>>;
let state = Arc::new(CaptureState::default());
let capture = WebSocketCapture::new(TestCaptureRecorder(state.clone()), {
let state = state.clone();
move || {
state.closes.fetch_add(1, Ordering::AcqRel);
}
});
let request = Request::new(Body::empty());
request.extensions().insert(capture);
let (request, _) = request.into_parts();
let (server_io, client_io) = tokio::io::duplex(16 * 1024);
let socket =
AsyncWebSocket::from_raw_socket(ServiceInput::new(server_io), Role::Server, None).await;
let websocket = ServerWebSocket { socket, request };
let endpoint = service_fn(
async |mut websocket: ServerWebSocket<HARWebSocket<TestSocket>>| {
let message = websocket.recv_message().await.expect("server receive");
let Message::Text(message) = message else {
panic!("expected text message");
};
websocket
.send_message(Message::text(format!("echo-{message}")))
.await
.expect("server send");
Ok::<_, std::convert::Infallible>(())
},
);
let endpoint = HARWebSocketLayer::new().into_layer(endpoint);
let endpoint_task = tokio::spawn(async move {
endpoint.serve(websocket).await.expect("serve endpoint");
});
let mut client =
AsyncWebSocket::from_raw_socket(ServiceInput::new(client_io), Role::Client, None).await;
client
.send_message(Message::text("hello"))
.await
.expect("client send");
assert_eq!(
client.recv_message().await.expect("client receive"),
Message::text("echo-hello")
);
endpoint_task.await.expect("endpoint task");
let messages = state.messages.lock();
assert_eq!(messages.len(), 2);
assert_eq!(messages[0].r#type, WebSocketMessageType::Send);
assert_eq!(messages[0].data.as_str(), "hello");
assert_eq!(messages[1].r#type, WebSocketMessageType::Receive);
assert_eq!(messages[1].data.as_str(), "echo-hello");
drop(messages);
assert_eq!(state.closes.load(Ordering::Acquire), 1);
}
async fn assert_websocket_acceptor_bad_request(request: Request, acceptor: &WebSocketAcceptor) {
let resp = acceptor.serve(request).await.unwrap_err();
assert_eq!(StatusCode::BAD_REQUEST, resp.status());
}
macro_rules! request {
(
$method:literal $version:literal $uri:literal
$(
$header_name:literal: $header_value:literal
)*
) => {
request!(
$method $version $uri
$(
$header_name: $header_value
)*
w/ []
)
};
(
$method:literal $version:literal $uri:literal
$(
$header_name:literal: $header_value:literal
)*
w/ [$($extension:expr),* $(,)?]
) => {
{
let req = Request::builder()
.uri($uri)
.version(match $version {
"HTTP/1.1" => Version::HTTP_11,
"HTTP/2" => Version::HTTP_2,
_ => unreachable!(),
})
.method(match $method {
"GET" => Method::GET,
"POST" => Method::POST,
"CONNECT" => Method::CONNECT,
_ => unreachable!(),
});
$(
let req = req.header($header_name, $header_value);
)*
$(
let req = req.extension($extension);
)*
req.body(Body::empty()).unwrap()
}
};
}
#[tokio::test]
async fn test_websocket_acceptor_default_http_2() {
let acceptor = WebSocketAcceptor::default();
assert_websocket_acceptor_bad_request(
request! {
"GET" "HTTP/2" "/"
"Connection": "upgrade"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "foobar"
},
&acceptor,
)
.await;
assert_websocket_acceptor_bad_request(
request! {
"CONNECT" "HTTP/2" "/"
w/ [
Protocol::from_static("websocket"),
]
},
&acceptor,
)
.await;
assert_websocket_acceptor_bad_request(
request! {
"GET" "HTTP/2" "/"
w/ [
Protocol::from_static("websocket"),
]
},
&acceptor,
)
.await;
assert_websocket_acceptor_ok(
request! {
"CONNECT" "HTTP/2" "/"
"Sec-WebSocket-Version": "13"
w/ [
Protocol::from_static("websocket"),
]
},
&acceptor,
None,
)
.await;
assert_websocket_acceptor_bad_request(
request! {
"CONNECT" "HTTP/2" "/"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
"Sec-WebSocket-Protocol": "client"
w/ [
Protocol::from_static("websocket"),
]
},
&acceptor,
)
.await;
}
#[tokio::test]
async fn test_websocket_acceptor_default_http_11() {
let acceptor = WebSocketAcceptor::default();
assert_websocket_acceptor_bad_request(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "foobar"
},
&acceptor,
)
.await;
assert_websocket_acceptor_bad_request(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Upgrade": "websocket"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
},
&acceptor,
)
.await;
assert_websocket_acceptor_bad_request(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "14"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
},
&acceptor,
)
.await;
assert_websocket_acceptor_bad_request(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Upgrade": "foo"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
},
&acceptor,
)
.await;
assert_websocket_acceptor_bad_request(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
},
&acceptor,
)
.await;
assert_websocket_acceptor_bad_request(
request! {
"GET" "HTTP/1.1" "/"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
},
&acceptor,
)
.await;
assert_websocket_acceptor_bad_request(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "keep-alive"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
},
&acceptor,
)
.await;
assert_websocket_acceptor_ok(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
},
&acceptor,
None,
)
.await;
}
#[tokio::test]
async fn test_websocket_accept_flex_protocols() {
let acceptor = WebSocketAcceptor::default().with_protocols_flex(true);
assert_websocket_acceptor_ok(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
},
&acceptor,
None,
)
.await;
assert_websocket_acceptor_ok(
request! {
"CONNECT" "HTTP/2" "/"
"Sec-WebSocket-Version": "13"
w/ [
Protocol::from_static("websocket"),
]
},
&acceptor,
None,
)
.await;
assert_websocket_acceptor_ok(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
"Sec-WebSocket-Protocol": "foo"
},
&acceptor,
Some(AcceptedWebSocketProtocol(non_empty_str!("foo"))),
)
.await;
assert_websocket_acceptor_ok(
request! {
"CONNECT" "HTTP/2" "/"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Protocol": "foo"
w/ [
Protocol::from_static("websocket"),
]
},
&acceptor,
Some(AcceptedWebSocketProtocol(non_empty_str!("foo"))),
)
.await;
assert_websocket_acceptor_ok(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
"Sec-WebSocket-Protocol": "foo, bar"
},
&acceptor,
Some(AcceptedWebSocketProtocol(non_empty_str!("foo"))),
)
.await;
assert_websocket_acceptor_ok(
request! {
"CONNECT" "HTTP/2" "/"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Protocol": "foo,baz, foo"
w/ [
Protocol::from_static("websocket"),
]
},
&acceptor,
Some(AcceptedWebSocketProtocol(non_empty_str!("foo"))),
)
.await;
let acceptor =
acceptor.with_protocols(headers::SecWebSocketProtocol::new(non_empty_str!("foo")));
assert_websocket_acceptor_ok(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
},
&acceptor,
None,
)
.await;
assert_websocket_acceptor_bad_request(
request! {
"CONNECT" "HTTP/2" "/"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Protocol": "baz,fo"
w/ [
Protocol::from_static("websocket"),
]
},
&acceptor,
)
.await;
}
#[tokio::test]
async fn test_websocket_accept_required_protocols() {
let acceptor = WebSocketAcceptor::default().with_protocols(headers::SecWebSocketProtocol(
non_empty_smallvec![
non_empty_str!("foo"),
non_empty_str!("a"),
non_empty_str!("b")
],
));
assert_websocket_acceptor_bad_request(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
},
&acceptor,
)
.await;
assert_websocket_acceptor_bad_request(
request! {
"CONNECT" "HTTP/2" "/"
"Sec-WebSocket-Version": "13"
w/ [
Protocol::from_static("websocket"),
]
},
&acceptor,
)
.await;
assert_websocket_acceptor_ok(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
"Sec-WebSocket-Protocol": "foo"
},
&acceptor,
Some(AcceptedWebSocketProtocol(non_empty_str!("foo"))),
)
.await;
assert_websocket_acceptor_ok(
request! {
"CONNECT" "HTTP/2" "/"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Protocol": "b"
w/ [
Protocol::from_static("websocket"),
]
},
&acceptor,
Some(AcceptedWebSocketProtocol(non_empty_str!("b"))),
)
.await;
assert_websocket_acceptor_ok(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
"Sec-WebSocket-Protocol": "test, b"
},
&acceptor,
Some(AcceptedWebSocketProtocol(non_empty_str!("b"))),
)
.await;
assert_websocket_acceptor_ok(
request! {
"CONNECT" "HTTP/2" "/"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Protocol": "a,test, c"
w/ [
Protocol::from_static("websocket"),
]
},
&acceptor,
Some(AcceptedWebSocketProtocol(non_empty_str!("a"))),
)
.await;
assert_websocket_acceptor_bad_request(
request! {
"GET" "HTTP/1.1" "/"
"Connection": "upgrade"
"Upgrade": "websocket"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ=="
"Sec-WebSocket-Protocol": "test, c"
},
&acceptor,
)
.await;
assert_websocket_acceptor_bad_request(
request! {
"CONNECT" "HTTP/2" "/"
"Sec-WebSocket-Version": "13"
"Sec-WebSocket-Protocol": "test"
w/ [
Protocol::from_static("websocket"),
]
},
&acceptor,
)
.await;
}
}