use super::Request;
use super::body::HyperResponseBody;
use super::disconnect::DisconnectSignal;
use super::rejection::Rejected;
use super::response::HeaderPair;
use super::router::WsHandler;
use super::server_lifecycle::{
ConnectionLifecycle, ConnectionPermit, ServerControl, UpgradeRegistrar, UpgradeRegistration,
};
use super::websocket::WsConn;
use std::ops::ControlFlow;
use std::sync::Arc;
type ClientWs =
tokio_tungstenite::WebSocketStream<hyper_util::rt::TokioIo<hyper::upgrade::Upgraded>>;
type WsFrameMessage = tokio_tungstenite::tungstenite::protocol::Message;
type WsFrame = Option<Result<WsFrameMessage, tokio_tungstenite::tungstenite::Error>>;
type WsClose = tokio_tungstenite::tungstenite::protocol::CloseFrame;
pub(super) enum WsUpgrade {
Ready(hyper::upgrade::OnUpgrade, Box<str>),
Rejected(WsHandshakeError),
}
pub(super) enum WsHandshakeError {
BadRequest,
UnsupportedVersion,
}
pub(super) fn extract_ws_upgrade(req: &mut hyper::Request<hyper::body::Incoming>) -> WsUpgrade {
let accept_key = match validate_ws_handshake(req) {
Ok(key) => tokio_tungstenite::tungstenite::handshake::derive_accept_key(key.as_bytes()),
Err(error) => return WsUpgrade::Rejected(error),
};
WsUpgrade::Ready(hyper::upgrade::on(req), accept_key.into())
}
fn validate_ws_handshake(
request: &hyper::Request<hyper::body::Incoming>,
) -> Result<&hyper::header::HeaderValue, WsHandshakeError> {
if request.method() != hyper::Method::GET || request.version() != hyper::Version::HTTP_11 {
return Err(WsHandshakeError::BadRequest);
}
let headers = request.headers();
let asks_to_upgrade =
is_ws_upgrade_head(headers) && header_contains_token(headers, "connection", "upgrade");
match asks_to_upgrade {
true => {}
false => return Err(WsHandshakeError::BadRequest),
}
validate_ws_version(headers)?;
validate_ws_subprotocols(headers)?;
let key = match single_header(headers, "sec-websocket-key") {
Some(key) if valid_ws_key(key.as_bytes()) => key,
_ => return Err(WsHandshakeError::BadRequest),
};
Ok(key)
}
fn validate_ws_subprotocols(headers: &hyper::HeaderMap) -> Result<(), WsHandshakeError> {
for value in headers.get_all("sec-websocket-protocol") {
let value = value.to_str().map_err(|_| WsHandshakeError::BadRequest)?;
if !value.split(',').map(str::trim).all(is_http_token) {
return Err(WsHandshakeError::BadRequest);
}
}
Ok(())
}
fn validate_ws_version(headers: &hyper::HeaderMap) -> Result<(), WsHandshakeError> {
let mut versions = headers.get_all("sec-websocket-version").iter();
let version = match versions.next() {
Some(version) => version,
None => return Err(WsHandshakeError::BadRequest),
};
match (version == "13", versions.next()) {
(true, None) => Ok(()),
_ => Err(WsHandshakeError::UnsupportedVersion),
}
}
pub(super) fn is_ws_upgrade_head(headers: &hyper::HeaderMap) -> bool {
single_header_equals(headers, "upgrade", "websocket")
}
pub(super) fn is_ws_upgrade_request(req: &Request) -> bool {
single_value_equals(named_request_headers(req, "upgrade"), "websocket")
}
fn single_header<'a>(
headers: &'a hyper::HeaderMap,
name: &'static str,
) -> Option<&'a hyper::header::HeaderValue> {
let mut values = headers.get_all(name).iter();
match (values.next(), values.next()) {
(Some(value), None) => Some(value),
_ => None,
}
}
fn single_header_equals(headers: &hyper::HeaderMap, name: &'static str, expected: &str) -> bool {
single_value_equals(
headers
.get_all(name)
.iter()
.map(|value| value.to_str().unwrap_or("")),
expected,
)
}
fn single_value_equals<'a>(mut values: impl Iterator<Item = &'a str>, expected: &str) -> bool {
match (values.next(), values.next()) {
(Some(value), None) => value.eq_ignore_ascii_case(expected),
_ => false,
}
}
fn header_contains_token(headers: &hyper::HeaderMap, name: &'static str, expected: &str) -> bool {
headers
.get_all(name)
.iter()
.try_fold(false, |found, value| {
value
.to_str()
.ok()?
.split(',')
.try_fold(found, |seen, token| {
let token = token.trim_matches([' ', '\t']);
is_http_token(token).then_some(seen || token.eq_ignore_ascii_case(expected))
})
})
.is_some_and(|found| found)
}
fn valid_ws_key(key: &[u8]) -> bool {
match key {
[symbols @ .., b'=', b'='] if symbols.len() == 22 => {
symbols.iter().copied().all(is_base64_symbol)
&& symbols
.last()
.copied()
.and_then(base64_value)
.is_some_and(|value| value & 0x0f == 0)
}
_ => false,
}
}
const fn is_base64_symbol(byte: u8) -> bool {
base64_value(byte).is_some()
}
const fn base64_value(byte: u8) -> Option<u8> {
match byte {
b'A'..=b'Z' => Some(byte - b'A'),
b'a'..=b'z' => Some(byte - b'a' + 26),
b'0'..=b'9' => Some(byte - b'0' + 52),
b'+' => Some(62),
b'/' => Some(63),
_ => None,
}
}
pub(super) fn check_ws_origin(req: &Request) -> Option<Rejected> {
let origin = match unique_request_header(req, "origin") {
HeaderPresence::Absent => return None,
HeaderPresence::Unique(origin) => origin,
HeaderPresence::Repeated => {
return rejected_origin("handshake carries more than one Origin header");
}
};
let host = match unique_request_header(req, "host") {
HeaderPresence::Unique(host) => host,
HeaderPresence::Absent | HeaderPresence::Repeated => {
return rejected_origin("handshake states no single Host to match the Origin against");
}
};
match origin_matches_host(origin, host) {
true => None,
false => rejected_origin("handshake Origin does not match the requested Host"),
}
}
enum HeaderPresence<'a> {
Absent,
Unique(&'a str),
Repeated,
}
fn unique_request_header<'a>(req: &'a Request, name: &'static str) -> HeaderPresence<'a> {
let mut values = named_request_headers(req, name);
match (values.next(), values.next()) {
(Some(value), None) => HeaderPresence::Unique(value),
(None, _) => HeaderPresence::Absent,
(Some(_), Some(_)) => HeaderPresence::Repeated,
}
}
fn named_request_headers<'a>(
req: &'a Request,
name: &'static str,
) -> impl Iterator<Item = &'a str> {
req.headers()
.filter_map(move |(candidate, value)| candidate.eq_ignore_ascii_case(name).then_some(value))
}
fn rejected_origin(detail: &'static str) -> Option<Rejected> {
Some(Rejected::ws_origin_rejected(detail))
}
fn origin_matches_host(origin: &str, host: &str) -> bool {
let (scheme, origin_authority) = match origin.split_once("://") {
Some(parts) => parts,
None => return false,
};
let default_port = match scheme {
value if value.eq_ignore_ascii_case("http") => 80,
value if value.eq_ignore_ascii_case("https") => 443,
_ => return false,
};
let origin = match parse_authority(origin_authority) {
Some(authority) => authority,
None => return false,
};
let host = match parse_authority(host) {
Some(authority) => authority,
None => return false,
};
let host_matches = origin
.authority
.host()
.eq_ignore_ascii_case(host.authority.host());
let port_matches = match host.port {
Some(host_port) => host_port == origin.port.unwrap_or(default_port),
None => origin
.port
.is_none_or(|origin_port| origin_port == default_port),
};
host_matches && port_matches
}
struct ParsedAuthority {
authority: hyper::http::uri::Authority,
port: Option<u16>,
}
fn parse_authority(value: &str) -> Option<ParsedAuthority> {
match value
.bytes()
.any(|byte| matches!(byte, b'@' | b'/' | b'?' | b'#' | b',' | b' ' | b'\t'))
{
true => return None,
false => {}
}
let authority: hyper::http::uri::Authority = value.parse().ok()?;
match authority.host().is_empty() {
true => return None,
false => {}
}
let port = explicit_authority_port(&authority)?;
Some(ParsedAuthority { authority, port })
}
fn explicit_authority_port(authority: &hyper::http::uri::Authority) -> Option<Option<u16>> {
let value = authority.as_str();
let has_separator = match value.starts_with('[') {
true => value.contains("]:"),
false => value.contains(':'),
};
match (has_separator, authority.port_u16()) {
(false, None) => Some(None),
(true, Some(port)) => Some(Some(port)),
_ => None,
}
}
fn ws_upgrade_pair(
ws_upgrade: WsUpgrade,
) -> Result<(hyper::upgrade::OnUpgrade, Box<str>), WsHandshakeError> {
match ws_upgrade {
WsUpgrade::Ready(on_upgrade, accept_key) => Ok((on_upgrade, accept_key)),
WsUpgrade::Rejected(error) => Err(error),
}
}
struct WsHandoff<'a> {
on_upgrade: hyper::upgrade::OnUpgrade,
subprotocol: Option<&'a str>,
response: hyper::Response<HyperResponseBody>,
permit: Arc<ConnectionPermit>,
handoff: DisconnectSignal,
}
pub(super) struct WsRefusal {
pub(super) rejected: Rejected,
pub(super) subprotocol: Option<Box<str>>,
}
impl WsRefusal {
fn unnegotiated(rejected: Rejected) -> Self {
Self {
rejected,
subprotocol: None,
}
}
fn negotiated(rejected: Rejected, subprotocol: Option<&str>) -> Self {
Self {
rejected,
subprotocol: subprotocol.map(Box::from),
}
}
}
enum WsHandoffOutcome<'a> {
Ready(WsHandoff<'a>),
Refused(WsRefusal),
}
fn prepare_ws_handoff<'a>(
ws_upgrade: WsUpgrade,
req: &'a Request,
lifecycle: &ConnectionLifecycle,
) -> WsHandoffOutcome<'a> {
let (on_upgrade, accept_key) = match ws_upgrade_pair(ws_upgrade) {
Ok(pair) => pair,
Err(error) => {
return WsHandoffOutcome::Refused(WsRefusal::unnegotiated(ws_handshake_rejection(
error,
)));
}
};
let subprotocol = extract_ws_subprotocol(req);
let response = match ws_switching_protocols(accept_key.as_ref(), subprotocol) {
Ok(response) => response,
Err(error) => {
return WsHandoffOutcome::Refused(WsRefusal::negotiated(
Rejected::ws_upgrade_unbuildable(error),
subprotocol,
));
}
};
WsHandoffOutcome::Ready(WsHandoff {
on_upgrade,
subprotocol,
response,
permit: lifecycle.permit(),
handoff: req.on_disconnect(),
})
}
pub(super) async fn handle_ws_upgrade(
ws_upgrade: WsUpgrade,
handler: WsHandler,
req: Request,
buffer_size: usize,
lifecycle: &ConnectionLifecycle,
) -> Result<hyper::Response<HyperResponseBody>, WsRefusal> {
let prepared = match prepare_ws_handoff(ws_upgrade, &req, lifecycle) {
WsHandoffOutcome::Ready(prepared) => prepared,
WsHandoffOutcome::Refused(refusal) => return Err(refusal),
};
let selected: Option<Box<str>> = prepared.subprotocol.map(Box::from);
let WsHandoff {
on_upgrade,
response,
permit,
handoff,
..
} = prepared;
let script = lifecycle.script();
own_upgrade_bridge(lifecycle, response, &handoff, move |attachment| {
bridge_ws_handler(
on_upgrade,
handler,
req,
buffer_size,
attachment,
script,
permit,
)
})
.await
.map_err(|rejected| WsRefusal {
rejected,
subprotocol: selected,
})
}
struct BridgeAttachment {
control: tokio::sync::watch::Receiver<ServerControl>,
dispatch: super::server_lifecycle::UpgradeDispatchGate,
}
impl BridgeAttachment {
fn split(
attachment: Option<Self>,
) -> (
Option<tokio::sync::watch::Receiver<ServerControl>>,
Option<super::server_lifecycle::UpgradeDispatchGate>,
) {
match attachment {
Some(Self { control, dispatch }) => (Some(control), Some(dispatch)),
None => (None, None),
}
}
}
async fn own_upgrade_bridge<F, Fut>(
lifecycle: &ConnectionLifecycle,
response: hyper::Response<HyperResponseBody>,
handoff: &DisconnectSignal,
build_bridge: F,
) -> Result<hyper::Response<HyperResponseBody>, Rejected>
where
F: FnOnce(Option<BridgeAttachment>) -> Fut,
Fut: std::future::Future<Output = ()> + Send + 'static,
{
let registrar = match lifecycle.upgrade_registrar() {
Some(registrar) => registrar,
None => {
detach_bridge(build_bridge(None));
return Ok(commit_upgrade(response, handoff));
}
};
let attachment = BridgeAttachment {
control: registrar.control(),
dispatch: registrar.dispatch_gate(),
};
let (gate, start) = tokio::sync::oneshot::channel();
let handle = spawn_gated_bridge(start, build_bridge(Some(attachment)));
complete_upgrade_registration(registrar, handle, gate, response, handoff).await
}
fn detach_bridge<F>(bridge: F)
where
F: std::future::Future<Output = ()> + Send + 'static,
{
drop(tokio::spawn(bridge));
}
fn commit_upgrade(
response: hyper::Response<HyperResponseBody>,
handoff: &DisconnectSignal,
) -> hyper::Response<HyperResponseBody> {
handoff.complete();
response
}
fn spawn_gated_bridge<F>(
start: tokio::sync::oneshot::Receiver<()>,
bridge: F,
) -> tokio::task::JoinHandle<()>
where
F: std::future::Future<Output = ()> + Send + 'static,
{
tokio::spawn(async move {
match start.await {
Ok(()) => bridge.await,
Err(_) => {}
}
})
}
async fn await_upgrade(
on_upgrade: hyper::upgrade::OnUpgrade,
context: &str,
) -> Option<hyper::upgrade::Upgraded> {
match on_upgrade.await {
Ok(u) => Some(u),
Err(e) => {
tracing::warn!(error = %e, "{context}");
None
}
}
}
async fn upgrade_client_ws(
on_upgrade: hyper::upgrade::OnUpgrade,
context: &str,
) -> Option<ClientWs> {
let upgraded = await_upgrade(on_upgrade, context).await?;
Some(
tokio_tungstenite::WebSocketStream::from_raw_socket(
hyper_util::rt::TokioIo::new(upgraded),
tokio_tungstenite::tungstenite::protocol::Role::Server,
None,
)
.await,
)
}
async fn commit_dispatch(
gate: Option<super::server_lifecycle::UpgradeDispatchGate>,
stream: &mut ClientWs,
) -> ControlFlow<()> {
let committed = match gate {
Some(gate) => gate.committed().await,
None => true,
};
match committed {
true => ControlFlow::Continue(()),
false => {
shutdown_client_transport(stream).await;
ControlFlow::Break(())
}
}
}
type OpenBridge = (
Option<tokio::sync::watch::Receiver<ServerControl>>,
ClientWs,
);
async fn open_bridge(
on_upgrade: hyper::upgrade::OnUpgrade,
attachment: Option<BridgeAttachment>,
context: &str,
) -> Option<OpenBridge> {
let (control, dispatch) = BridgeAttachment::split(attachment);
let mut stream = upgrade_client_ws(on_upgrade, context).await?;
match commit_dispatch(dispatch, &mut stream).await {
ControlFlow::Break(()) => None,
ControlFlow::Continue(()) => Some((control, stream)),
}
}
async fn bridge_ws_handler(
on_upgrade: hyper::upgrade::OnUpgrade,
handler: WsHandler,
req: Request,
buffer_size: usize,
attachment: Option<BridgeAttachment>,
script: Option<Arc<super::mock::LifecycleScript>>,
permit: Arc<ConnectionPermit>,
) {
let opened = open_bridge(on_upgrade, attachment, "WebSocket client upgrade failed").await;
let (mut control, mut ws_stream) = match opened {
Some(opened) => opened,
None => return,
};
super::mock::LifecycleScript::pause_at(
script.as_deref(),
super::mock::LifecycleCheckpoint::WebSocketOutgoingBufferConfigured(buffer_size),
)
.await;
let (outgoing_tx, mut outgoing_rx) = tokio::sync::mpsc::channel::<WsFrameMessage>(buffer_size);
super::mock::LifecycleScript::pause_at(
script.as_deref(),
super::mock::LifecycleCheckpoint::WebSocketIncomingBufferConfigured(buffer_size),
)
.await;
let (incoming_tx, incoming_rx) = tokio::sync::mpsc::channel::<WsFrameMessage>(buffer_size);
use futures_util::StreamExt;
drop(tokio::task::spawn_blocking(move || {
let conn = WsConn::new(outgoing_tx, incoming_rx);
match crate::task::catch_panic(move || handler(&req, conn)) {
Ok(Ok(())) => {}
Ok(Err(e)) => tracing::warn!(error = %e, "WebSocket handler returned error"),
Err(error) => tracing::error!(%error, "WebSocket handler panicked"),
}
}));
loop {
let flow = tokio::select! {
biased;
mode = next_control(&mut control) => stop_direct_bridge(mode, &mut ws_stream).await,
outgoing = outgoing_rx.recv() => forward_outgoing(outgoing, &mut ws_stream).await,
incoming = ws_stream.next() => {
forward_incoming(incoming, &mut ws_stream, &incoming_tx).await
}
};
match flow {
ControlFlow::Break(()) => break,
ControlFlow::Continue(()) => {}
}
}
shutdown_client_transport(&mut ws_stream).await;
drop(permit);
}
async fn stop_direct_bridge(mode: ServerControl, stream: &mut ClientWs) -> ControlFlow<()> {
match mode {
ServerControl::Graceful => graceful_close_direct(stream).await,
ServerControl::Abort | ServerControl::Running => {}
}
ControlFlow::Break(())
}
async fn graceful_close_direct(stream: &mut ClientWs) {
send_close(stream, None).await;
drain_until_close(stream).await;
}
async fn forward_outgoing(
outgoing: Option<WsFrameMessage>,
stream: &mut ClientWs,
) -> ControlFlow<()> {
use futures_util::SinkExt;
let message = match outgoing {
Some(message) => message,
None => {
close_transport(stream).await;
return ControlFlow::Break(());
}
};
match stream.send(message).await {
Ok(()) => ControlFlow::Continue(()),
Err(error) => {
tracing::debug!(%error, "WebSocket client send failed");
ControlFlow::Break(())
}
}
}
async fn forward_incoming(
incoming: WsFrame,
stream: &mut ClientWs,
handler_tx: &tokio::sync::mpsc::Sender<WsFrameMessage>,
) -> ControlFlow<()> {
let message = next_frame(incoming, "WebSocket client bridge closed")?;
match message.is_close() {
true => {
flush_transport(stream).await;
ControlFlow::Break(())
}
false => queue_for_handler(handler_tx, message).await,
}
}
async fn queue_for_handler(
handler_tx: &tokio::sync::mpsc::Sender<WsFrameMessage>,
message: WsFrameMessage,
) -> ControlFlow<()> {
match handler_tx.send(message).await {
Ok(()) => ControlFlow::Continue(()),
Err(_) => ControlFlow::Break(()),
}
}
fn next_frame(frame: WsFrame, context: &str) -> ControlFlow<(), WsFrameMessage> {
match frame {
Some(Ok(message)) => ControlFlow::Continue(message),
Some(Err(error)) => {
tracing::debug!(%error, "{context}");
ControlFlow::Break(())
}
None => ControlFlow::Break(()),
}
}
pub(super) async fn handle_proxy_ws(
ws_upgrade: WsUpgrade,
req: Request,
backend: Arc<str>,
prefix: Arc<str>,
lifecycle: &ConnectionLifecycle,
) -> Result<hyper::Response<HyperResponseBody>, WsRefusal> {
let prepared = match prepare_ws_handoff(ws_upgrade, &req, lifecycle) {
WsHandoffOutcome::Ready(prepared) => prepared,
WsHandoffOutcome::Refused(refusal) => return Err(refusal),
};
let subprotocol = prepared.subprotocol;
let backend_ws_url = match build_backend_ws_url(req.raw_path_and_query(), &prefix, &backend) {
Ok(url) => url,
Err(rejected) => return Err(WsRefusal::negotiated(rejected, subprotocol)),
};
let forwarded_headers = collect_forwardable_ws_headers(&req, subprotocol);
let WsHandoff {
on_upgrade,
response,
permit,
handoff,
..
} = prepared;
own_upgrade_bridge(lifecycle, response, &handoff, move |attachment| {
bridge_ws_proxy(
on_upgrade,
backend_ws_url,
forwarded_headers,
attachment,
permit,
)
})
.await
.map_err(|rejected| WsRefusal::negotiated(rejected, subprotocol))
}
async fn complete_upgrade_registration(
registrar: UpgradeRegistrar,
handle: tokio::task::JoinHandle<()>,
gate: tokio::sync::oneshot::Sender<()>,
response: hyper::Response<HyperResponseBody>,
handoff: &DisconnectSignal,
) -> Result<hyper::Response<HyperResponseBody>, Rejected> {
match registrar.submit(handle).await {
UpgradeRegistration::Admitted => release_admitted_bridge(gate, response, handoff),
UpgradeRegistration::Rejected => Err(Rejected::upgrade_registration_refused()),
UpgradeRegistration::Unavailable => Err(Rejected::upgrade_registration_unavailable()),
}
}
fn release_admitted_bridge(
gate: tokio::sync::oneshot::Sender<()>,
response: hyper::Response<HyperResponseBody>,
handoff: &DisconnectSignal,
) -> Result<hyper::Response<HyperResponseBody>, Rejected> {
match gate.send(()) {
Ok(()) => Ok(commit_upgrade(response, handoff)),
Err(()) => Err(Rejected::upgrade_registration_unavailable()),
}
}
fn extract_ws_subprotocol(req: &Request) -> Option<&str> {
req.headers()
.filter(|(name, _)| name.eq_ignore_ascii_case("sec-websocket-protocol"))
.flat_map(|(_, value)| value.split(','))
.map(str::trim)
.find(|protocol| is_http_token(protocol))
}
fn is_http_token(value: &str) -> bool {
!value.is_empty()
&& value.bytes().all(|byte| {
byte.is_ascii_alphanumeric()
|| matches!(
byte,
b'!' | b'#'
| b'$'
| b'%'
| b'&'
| b'\''
| b'*'
| b'+'
| b'-'
| b'.'
| b'^'
| b'_'
| b'`'
| b'|'
| b'~'
)
})
}
fn collect_forwardable_ws_headers(req: &Request, subprotocol: Option<&str>) -> Box<[HeaderPair]> {
let headers = req
.headers()
.filter(|(name, _)| is_forwardable_ws_header(name))
.map(|(name, value)| {
(
std::borrow::Cow::Owned(name.to_owned()),
std::borrow::Cow::Owned(value.to_owned()),
)
});
let selected = subprotocol.into_iter().map(|protocol| {
(
std::borrow::Cow::Borrowed("Sec-WebSocket-Protocol"),
std::borrow::Cow::Owned(protocol.to_owned()),
)
});
headers.chain(selected).collect()
}
fn is_forwardable_ws_header(name: &str) -> bool {
match name {
n if n.eq_ignore_ascii_case("authorization") => true,
n if n.eq_ignore_ascii_case("cookie") => true,
n if n.eq_ignore_ascii_case("sec-websocket-protocol") => false,
n if n
.get(..2)
.is_some_and(|prefix| prefix.eq_ignore_ascii_case("x-"))
&& !super::async_proxy::is_forwarded_metadata(n) =>
{
true
}
_ => false,
}
}
fn build_backend_ws_url(path: &str, prefix: &str, backend: &str) -> Result<Box<str>, Rejected> {
let remainder = match super::async_proxy::strip_prefix(path, prefix) {
Some(remainder) => remainder,
None => {
return Err(unbuildable_ws_target(super::async_proxy::TRAVERSAL_SEGMENT));
}
};
match backend {
s if s.starts_with("http://") => {
Ok(format!("ws://{}{remainder}", &s["http://".len()..]).into_boxed_str())
}
s if s.starts_with("https://") => {
Ok(format!("wss://{}{remainder}", &s["https://".len()..]).into_boxed_str())
}
_ => Err(unbuildable_ws_target(
"the configured backend names no scheme this proxy can upgrade over",
)),
}
}
fn unbuildable_ws_target(detail: &'static str) -> Rejected {
Rejected::from_proxy_failure(super::async_proxy::ProxyFailure::UnbuildableTarget(detail))
}
async fn bridge_ws_proxy(
on_upgrade: hyper::upgrade::OnUpgrade,
backend_ws_url: Box<str>,
forwarded_headers: Box<[HeaderPair]>,
attachment: Option<BridgeAttachment>,
permit: Arc<ConnectionPermit>,
) {
let opened = open_bridge(
on_upgrade,
attachment,
"WebSocket proxy client upgrade failed",
)
.await;
let (mut control, mut client_ws) = match opened {
Some(opened) => opened,
None => return,
};
let backend_request = match build_ws_backend_request(&backend_ws_url, &forwarded_headers) {
Some(req) => req,
None => {
end_client_transport(&mut client_ws, Some(backend_fault_close())).await;
return;
}
};
let (mut backend_ws, _) = match tokio_tungstenite::connect_async(backend_request).await {
Ok(pair) => pair,
Err(e) => {
tracing::warn!(url = %backend_ws_url, error = %e, "WebSocket proxy backend connection failed");
end_client_transport(&mut client_ws, Some(backend_fault_close())).await;
return;
}
};
use futures_util::StreamExt;
let exit = loop {
let flow = tokio::select! {
biased;
mode = next_control(&mut control) => {
stop_proxy_bridge(mode, &mut client_ws, &mut backend_ws).await
}
message = client_ws.next() => {
owes_close(forward_client_frame(message, &mut client_ws, &mut backend_ws).await)
}
message = backend_ws.next() => {
owes_close(forward_backend_frame(message, &mut client_ws).await)
}
};
match flow {
ControlFlow::Break(exit) => break exit,
ControlFlow::Continue(()) => {}
}
};
match exit {
ProxyExit::Settled => shutdown_client_transport(&mut client_ws).await,
ProxyExit::Owed => {
close_transport(&mut backend_ws).await;
end_client_transport(&mut client_ws, None).await;
}
}
drop(permit);
}
enum ProxyExit {
Settled,
Owed,
}
fn owes_close(flow: ControlFlow<()>) -> ControlFlow<ProxyExit> {
match flow {
ControlFlow::Break(()) => ControlFlow::Break(ProxyExit::Owed),
ControlFlow::Continue(()) => ControlFlow::Continue(()),
}
}
async fn end_client_transport(stream: &mut ClientWs, reason: Option<WsClose>) {
send_close(stream, reason).await;
shutdown_client_transport(stream).await;
}
fn backend_fault_close() -> WsClose {
WsClose {
code: tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode::Error,
reason: tokio_tungstenite::tungstenite::Utf8Bytes::from_static(
"WebSocket proxy backend unavailable",
),
}
}
async fn stop_proxy_bridge<C, B>(
mode: ServerControl,
client: &mut tokio_tungstenite::WebSocketStream<C>,
backend: &mut tokio_tungstenite::WebSocketStream<B>,
) -> ControlFlow<ProxyExit>
where
C: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
B: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
match mode {
ServerControl::Graceful => {
graceful_close_proxy(client, backend).await;
ControlFlow::Break(ProxyExit::Settled)
}
ServerControl::Abort | ServerControl::Running => ControlFlow::Break(ProxyExit::Owed),
}
}
async fn graceful_close_proxy<C, B>(
client: &mut tokio_tungstenite::WebSocketStream<C>,
backend: &mut tokio_tungstenite::WebSocketStream<B>,
) where
C: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
B: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
send_close(client, None).await;
send_close(backend, None).await;
drain_proxy_close(client, backend).await;
}
async fn forward_client_frame<C, B>(
frame: WsFrame,
client: &mut tokio_tungstenite::WebSocketStream<C>,
backend: &mut tokio_tungstenite::WebSocketStream<B>,
) -> ControlFlow<()>
where
C: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
B: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
use futures_util::SinkExt;
let message = next_frame(frame, "WebSocket proxy client closed")?;
let closes = message.is_close();
match (backend.send(message).await, closes) {
(Ok(()), false) => ControlFlow::Continue(()),
(Ok(()), true) => {
forward_backend_close(client, backend).await;
ControlFlow::Break(())
}
(Err(error), _) => {
tracing::debug!(%error, "WebSocket proxy backend send failed");
ControlFlow::Break(())
}
}
}
async fn forward_backend_frame<C>(
frame: WsFrame,
client: &mut tokio_tungstenite::WebSocketStream<C>,
) -> ControlFlow<()>
where
C: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
use futures_util::SinkExt;
let message = next_frame(frame, "WebSocket proxy backend closed")?;
let closes = message.is_close();
match (client.send(message).await, closes) {
(Ok(()), false) => ControlFlow::Continue(()),
(Ok(()), true) => ControlFlow::Break(()),
(Err(error), _) => {
tracing::debug!(%error, "WebSocket proxy client send failed");
ControlFlow::Break(())
}
}
}
async fn forward_backend_close<C, B>(
client: &mut tokio_tungstenite::WebSocketStream<C>,
backend: &mut tokio_tungstenite::WebSocketStream<B>,
) where
C: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
B: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
drain_until_close(backend).await;
flush_transport(client).await;
}
async fn send_close<S>(stream: &mut tokio_tungstenite::WebSocketStream<S>, reason: Option<WsClose>)
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
use futures_util::SinkExt;
match stream.send(WsFrameMessage::Close(reason)).await {
Ok(()) => {}
Err(error) => tracing::debug!(%error, "WebSocket close frame send failed"),
}
}
async fn close_transport<S>(stream: &mut tokio_tungstenite::WebSocketStream<S>)
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
match stream.close(None).await {
Ok(()) => {}
Err(error) => tracing::debug!(%error, "WebSocket close failed"),
}
}
async fn flush_transport<S>(stream: &mut tokio_tungstenite::WebSocketStream<S>)
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
use futures_util::SinkExt;
match stream.flush().await {
Ok(()) => {}
Err(error) => tracing::debug!(%error, "WebSocket flush failed"),
}
}
async fn shutdown_client_transport<S>(stream: &mut tokio_tungstenite::WebSocketStream<S>)
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
use tokio::io::AsyncWriteExt;
match stream.get_mut().shutdown().await {
Ok(()) => {}
Err(error) => tracing::debug!(%error, "WebSocket transport shutdown failed"),
}
tokio::task::yield_now().await;
}
async fn next_control(
control: &mut Option<tokio::sync::watch::Receiver<ServerControl>>,
) -> ServerControl {
let receiver = match control {
Some(receiver) => receiver,
None => return std::future::pending().await,
};
loop {
let current = *receiver.borrow_and_update();
if current != ServerControl::Running {
return current;
}
match receiver.changed().await {
Ok(()) => {}
Err(_) => return current,
}
}
}
async fn drain_until_close<S>(stream: &mut tokio_tungstenite::WebSocketStream<S>)
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
use futures_util::StreamExt;
while let Some(result) = stream.next().await {
match result {
Ok(message) if message.is_close() => return,
Ok(_) => {}
Err(error) => {
tracing::debug!(%error, "WebSocket close drain failed");
return;
}
}
}
}
async fn drain_proxy_close<C, B>(
client: &mut tokio_tungstenite::WebSocketStream<C>,
backend: &mut tokio_tungstenite::WebSocketStream<B>,
) where
C: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
B: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
let ((), ()) = tokio::join!(drain_until_close(client), drain_until_close(backend));
}
fn build_ws_backend_request(url: &str, headers: &[HeaderPair]) -> Option<hyper::Request<()>> {
let uri: hyper::Uri = match url.parse() {
Ok(u) => u,
Err(e) => {
tracing::warn!(url = %url, error = %e, "WebSocket backend URI parse failed");
return None;
}
};
let host = match uri.authority() {
Some(authority) => authority.as_str(),
None => {
tracing::warn!(url = %url, "WebSocket backend URL names no authority");
return None;
}
};
let mut builder = hyper::Request::builder()
.uri(url)
.header("Host", host)
.header("Connection", "Upgrade")
.header("Upgrade", "websocket")
.header("Sec-WebSocket-Version", "13")
.header(
"Sec-WebSocket-Key",
tokio_tungstenite::tungstenite::handshake::client::generate_key(),
);
for (name, value) in headers {
builder = builder.header(name.as_ref(), value.as_ref());
}
match builder.body(()) {
Ok(req) => Some(req),
Err(e) => {
tracing::warn!(url = %url, error = %e, "WebSocket backend request build failed");
None
}
}
}
fn ws_handshake_rejection(error: WsHandshakeError) -> Rejected {
match error {
WsHandshakeError::BadRequest => Rejected::ws_bad_handshake(),
WsHandshakeError::UnsupportedVersion => Rejected::ws_unsupported_version(),
}
}
fn ws_switching_protocols(
accept_key: &str,
subprotocol: Option<&str>,
) -> Result<hyper::Response<HyperResponseBody>, hyper::http::Error> {
let mut builder = hyper::Response::builder()
.status(hyper::StatusCode::SWITCHING_PROTOCOLS)
.header("Upgrade", "websocket")
.header("Connection", "Upgrade")
.header("Sec-WebSocket-Accept", accept_key);
if let Some(proto) = subprotocol {
builder = builder.header("Sec-WebSocket-Protocol", proto);
}
builder.body(HyperResponseBody::Full(http_body_util::Full::new(
bytes::Bytes::new(),
)))
}