WEB: websocket + websocket-lanes as Carrier

This commit is contained in:
Alexey
2026-08-26 09:16:06 +03:00
parent 8e577ec5ca
commit d2edd90479
50 changed files with 3642 additions and 350 deletions
+10 -3
View File
@@ -9,9 +9,7 @@ use hyper::Request;
use hyper::body::{Body, Frame, Incoming, SizeHint};
use crate::web::manager::WebProcessRuntime;
use crate::web::trace::{
HttpTraceExchange, TraceBodyState, TraceDirection,
};
use crate::web::trace::{HttpTraceExchange, TraceBodyState, TraceDirection};
/// Incoming request body wrapper that observes frames without changing streaming semantics.
pub(super) struct RequestBody {
@@ -30,6 +28,15 @@ impl RequestBody {
}
}
/// Completes observation for a request whose Hyper body is already empty.
pub(super) fn finish_empty(&mut self) -> bool {
if !self.inner.is_end_stream() {
return false;
}
self.finish(TraceBodyState::Complete);
true
}
fn finish(&mut self, state: TraceBodyState) {
if self.terminal {
return;
+1
View File
@@ -236,6 +236,7 @@ fn sanitize_transport_request<B>(request: &mut Request<B>) {
header::CONTENT_TYPE,
header::UPGRADE,
HeaderName::from_static("sec-websocket-key"),
HeaderName::from_static("sec-websocket-extensions"),
HeaderName::from_static("sec-websocket-protocol"),
HeaderName::from_static("sec-websocket-version"),
HeaderName::from_static("x-down-cursor"),
+3
View File
@@ -32,6 +32,9 @@ pub(super) async fn handle_down(
let Ok(session) = runtime.get_session(token_hash, &vhost.host) else {
return serve_decoy(request, vhost, true, &runtime).await;
};
if session.carrier().uses_websocket() {
return serve_decoy(request, vhost, true, &runtime).await;
}
if let Some(trace) = request_trace(&request) {
trace.set_route(TraceRoute::Downlink);
trace.bind_identity(session.trace_identity());
+2 -4
View File
@@ -9,16 +9,14 @@ use crate::config::WebCarrier;
use crate::web::frame;
/// Validates and resolves the optional carrier lane header.
pub(super) fn carrier_lane<B>(
request: &Request<B>,
carrier: WebCarrier,
) -> Option<Option<u32>> {
pub(super) fn carrier_lane<B>(request: &Request<B>, carrier: WebCarrier) -> Option<Option<u32>> {
match carrier {
WebCarrier::Https => (!request.headers().contains_key("x-lane-id")).then_some(None),
WebCarrier::HttpsLanes => canonical_u64_header(request, "x-lane-id")
.and_then(|value| u32::try_from(value).ok())
.filter(|value| *value <= frame::MAX_STREAM_ID)
.map(Some),
WebCarrier::Websocket | WebCarrier::WebsocketLanes => None,
}
}
+4 -1
View File
@@ -232,7 +232,10 @@ async fn rejected_bridge_bootstrap_falls_back_to_uncacheable_static_index() {
let fallback_response = request(&listener, &runtime, bridge_request()).await;
let (fallback_headers, fallback_body) = split_response(&fallback_response);
assert!(fallback_headers.starts_with(b"HTTP/1.1 200"));
assert_eq!(response_header(fallback_headers, "cache-control"), "no-store");
assert_eq!(
response_header(fallback_headers, "cache-control"),
"no-store"
);
assert_eq!(fallback_body, b"<!doctype html><title>decoy</title>");
runtime.shutdown().await;
+461
View File
@@ -0,0 +1,461 @@
use std::net::SocketAddr;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use base64::Engine as _;
use bytes::Bytes;
use hyper::header::{self, HeaderName, HeaderValue};
use hyper::{Method, Request, StatusCode};
use ipnetwork::IpNetwork;
use sha1::{Digest as _, Sha1};
use sha2::Sha256;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::net::TcpStream;
use tokio::sync::OwnedSemaphorePermit;
use super::body::RequestBody;
use super::decoy::serve_decoy;
use super::request::client_ip;
use super::response::{full_response, insert_header};
use super::{HttpResponse, request_trace, set_trace_route};
use crate::config::{WebCarrier, WebClientIpSource, WebRuntimeVhost};
use crate::web::manager::{TokenHash, WebProcessRuntime, WebSocketKind};
use crate::web::trace::TraceRoute;
// Codec buffers and fixed driver state are charged before HTTP 101 commits.
const BASE_BUDGET_BYTES: usize = 132 * 1024;
const WEBSOCKET_GUID: &[u8] = b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11";
/// Accepted connection IO retains the process HTTP slot after an upgrade.
pub(super) struct ConnectionIo {
stream: TcpStream,
_connection_permit: OwnedSemaphorePermit,
websocket_read: Option<WebSocketReadBoundary>,
}
impl ConnectionIo {
pub(super) fn new(stream: TcpStream, connection_permit: OwnedSemaphorePermit) -> Self {
Self {
stream,
_connection_permit: connection_permit,
websocket_read: None,
}
}
async fn readable(&self) -> std::io::Result<()> {
if self
.websocket_read
.as_ref()
.is_some_and(WebSocketReadBoundary::has_buffered)
{
return Ok(());
}
self.stream.readable().await
}
fn enable_websocket(&mut self, buffered: Bytes) {
self.websocket_read = Some(WebSocketReadBoundary::new(buffered));
}
fn websocket_fragmented_message(&self) -> bool {
self.websocket_read
.as_ref()
.is_some_and(WebSocketReadBoundary::fragmented_message)
}
}
// Frame-boundary reads keep the kernel readiness gate authoritative even when
// Hyper read-ahead or one TCP packet contains several WebSocket messages.
enum WebSocketReadState {
Header {
bytes: [u8; 14],
filled: usize,
target: usize,
},
Payload {
remaining: usize,
},
}
struct WebSocketReadBoundary {
buffered: Bytes,
state: WebSocketReadState,
fragmented_message: bool,
}
impl WebSocketReadBoundary {
fn new(buffered: Bytes) -> Self {
Self {
buffered,
state: Self::new_header(),
fragmented_message: false,
}
}
fn has_buffered(&self) -> bool {
!self.buffered.is_empty()
}
fn fragmented_message(&self) -> bool {
self.fragmented_message
}
fn maximum_read(&self, requested: usize) -> usize {
let boundary = match self.state {
WebSocketReadState::Header { filled, target, .. } => target.saturating_sub(filled),
WebSocketReadState::Payload { remaining } => remaining,
};
requested.min(boundary)
}
fn observe(&mut self, bytes: &[u8]) {
match &mut self.state {
WebSocketReadState::Header {
bytes: header,
filled,
target,
} => {
debug_assert!(bytes.len() <= target.saturating_sub(*filled));
header[*filled..*filled + bytes.len()].copy_from_slice(bytes);
*filled += bytes.len();
if *filled == 2 && *target == 2 {
let extended = match header[1] & 0x7f {
126 => 2,
127 => 8,
_ => 0,
};
let mask = usize::from(header[1] & 0x80 != 0) * 4;
*target = 2 + extended + mask;
}
if *filled == *target {
let finished = header[0] & 0x80 != 0;
match header[0] & 0x0f {
0 if finished => self.fragmented_message = false,
1 | 2 if !finished => self.fragmented_message = true,
_ => {}
}
let payload = match header[1] & 0x7f {
value @ 0..=125 => usize::from(value),
126 => usize::from(u16::from_be_bytes([header[2], header[3]])),
127 => usize::try_from(u64::from_be_bytes([
header[2], header[3], header[4], header[5], header[6], header[7],
header[8], header[9],
]))
.unwrap_or(usize::MAX),
_ => unreachable!(),
};
self.state = if payload == 0 {
Self::new_header()
} else {
WebSocketReadState::Payload { remaining: payload }
};
}
}
WebSocketReadState::Payload { remaining } => {
debug_assert!(bytes.len() <= *remaining);
*remaining -= bytes.len();
if *remaining == 0 {
self.state = Self::new_header();
}
}
}
}
fn new_header() -> WebSocketReadState {
WebSocketReadState::Header {
bytes: [0; 14],
filled: 0,
target: 2,
}
}
}
impl AsyncRead for ConnectionIo {
fn poll_read(
self: Pin<&mut Self>,
context: &mut Context<'_>,
buffer: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
let this = self.get_mut();
let Some(boundary) = this.websocket_read.as_mut() else {
return Pin::new(&mut this.stream).poll_read(context, buffer);
};
let maximum = boundary.maximum_read(buffer.remaining());
if maximum == 0 {
return Poll::Ready(Ok(()));
}
if !boundary.buffered.is_empty() {
let bytes = boundary
.buffered
.split_to(maximum.min(boundary.buffered.len()));
boundary.observe(&bytes);
buffer.put_slice(&bytes);
return Poll::Ready(Ok(()));
}
let unfilled = buffer.initialize_unfilled_to(maximum);
let mut limited = ReadBuf::new(unfilled);
match Pin::new(&mut this.stream).poll_read(context, &mut limited) {
Poll::Ready(Ok(())) => {
let filled = limited.filled().len();
boundary.observe(limited.filled());
drop(limited);
buffer.advance(filled);
Poll::Ready(Ok(()))
}
Poll::Ready(Err(error)) => Poll::Ready(Err(error)),
Poll::Pending => Poll::Pending,
}
}
}
impl AsyncWrite for ConnectionIo {
fn poll_write(
mut self: Pin<&mut Self>,
context: &mut Context<'_>,
bytes: &[u8],
) -> Poll<std::io::Result<usize>> {
Pin::new(&mut self.stream).poll_write(context, bytes)
}
fn poll_flush(
mut self: Pin<&mut Self>,
context: &mut Context<'_>,
) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.stream).poll_flush(context)
}
fn poll_shutdown(
mut self: Pin<&mut Self>,
context: &mut Context<'_>,
) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.stream).poll_shutdown(context)
}
}
enum ParsedCarrier {
Multiplex,
Lane(u32),
}
struct ParsedUpgrade {
token_hash: TokenHash,
protocol: String,
accept: String,
carrier: ParsedCarrier,
}
pub(super) async fn handle(
mut request: Request<RequestBody>,
peer: SocketAddr,
client_ip_source: WebClientIpSource,
trusted_proxy_cidrs: &[IpNetwork],
runtime: Arc<WebProcessRuntime>,
vhost: Arc<WebRuntimeVhost>,
) -> HttpResponse {
let Some(parsed) = parse_upgrade(&request) else {
return serve_decoy(request, vhost, true, &runtime).await;
};
if !request.body_mut().finish_empty() {
return serve_decoy(request, vhost, true, &runtime).await;
}
let Some(effective_ip) = client_ip(&request, peer, client_ip_source, trusted_proxy_cidrs)
else {
return serve_decoy(request, vhost, true, &runtime).await;
};
let Ok(session) = runtime.get_session(parsed.token_hash, &vhost.host) else {
return serve_decoy(request, vhost, true, &runtime).await;
};
let kind = match (parsed.carrier, session.carrier()) {
(ParsedCarrier::Multiplex, WebCarrier::Websocket) => WebSocketKind::Multiplex,
(ParsedCarrier::Lane(lane_id), WebCarrier::WebsocketLanes) => WebSocketKind::Lane(lane_id),
_ => return serve_decoy(request, vhost, true, &runtime).await,
};
let mut lane_reservation = match kind {
WebSocketKind::Multiplex => None,
WebSocketKind::Lane(lane_id) => match session.reserve_websocket_lane(lane_id) {
Ok(reservation) => Some(reservation),
Err(_) => return serve_decoy(request, vhost, true, &runtime).await,
},
};
let timeouts = runtime.active_generation().config().web.timeouts.clone();
let connection = match runtime
.admit_websocket(
session.profile_key(),
session.trace_session_id(),
effective_ip,
kind,
BASE_BUDGET_BYTES,
Duration::from_secs(timeouts.long_poll_secs),
Duration::from_secs(timeouts.websocket_eviction_secs),
)
.await
{
Ok(connection) => connection,
Err(_) => return serve_decoy(request, vhost, true, &runtime).await,
};
let trace_context = runtime.trace().websocket_context(
&request,
peer.ip(),
effective_ip,
connection.id(),
match kind {
WebSocketKind::Multiplex => None,
WebSocketKind::Lane(lane_id) => Some(lane_id),
},
|| session.trace_identity(),
);
if let Some(trace) = request_trace(&request) {
trace.set_effective_ip(effective_ip);
trace.set_route(TraceRoute::Websocket);
trace.bind_identity(session.trace_identity());
trace.register_redaction(parsed.protocol.as_bytes());
}
let on_upgrade = hyper::upgrade::on(&mut request);
let protocol = parsed.protocol;
let accept = parsed.accept;
let driver_runtime = Arc::clone(&runtime);
let driver_session = Arc::clone(&session);
runtime.spawn_auxiliary(async move {
driver::run_upgraded(
on_upgrade,
driver_runtime,
driver_session,
connection,
lane_reservation.take(),
trace_context,
)
.await;
});
set_trace_route(&request, TraceRoute::Websocket);
let mut response = full_response(StatusCode::SWITCHING_PROTOCOLS, Bytes::new());
response.headers_mut().remove(header::CONTENT_LENGTH);
response
.headers_mut()
.insert(header::CONNECTION, HeaderValue::from_static("Upgrade"));
response
.headers_mut()
.insert(header::UPGRADE, HeaderValue::from_static("websocket"));
insert_header(
&mut response,
HeaderName::from_static("sec-websocket-accept"),
&accept,
);
insert_header(
&mut response,
HeaderName::from_static("sec-websocket-protocol"),
&protocol,
);
response
}
fn parse_upgrade<B>(request: &Request<B>) -> Option<ParsedUpgrade> {
if request.method() != Method::GET
|| request.uri().query().is_some()
|| request.headers().contains_key(header::AUTHORIZATION)
|| request.headers().contains_key(header::CONTENT_LENGTH)
|| request.headers().contains_key(header::TRANSFER_ENCODING)
|| !single_header_eq(request, header::UPGRADE, "websocket")
|| !single_header_eq(request, "sec-websocket-version", "13")
|| !header_has_token(request, header::CONNECTION, "upgrade")
{
return None;
}
let key = single_header(request, "sec-websocket-key")?;
let decoded_key = base64::engine::general_purpose::STANDARD.decode(key).ok()?;
if decoded_key.len() != 16
|| base64::engine::general_purpose::STANDARD.encode(&decoded_key) != key
{
return None;
}
let protocol = single_header(request, "sec-websocket-protocol")?;
if protocol
.bytes()
.any(|value| value == b',' || value.is_ascii_whitespace())
{
return None;
}
let (token, carrier) = if let Some(token) = protocol.strip_prefix("tproxy-v1.") {
(token, ParsedCarrier::Multiplex)
} else if let Some(lane) = protocol.strip_prefix("tproxy-lane-v1.") {
let (token, lane_id) = lane.split_once('.')?;
if lane_id.is_empty()
|| lane_id.starts_with('+')
|| (lane_id.len() > 1 && lane_id.starts_with('0'))
{
return None;
}
let lane_id = lane_id
.parse::<u32>()
.ok()
.filter(|value| (1..=crate::web::frame::MAX_STREAM_ID).contains(value))?;
(token, ParsedCarrier::Lane(lane_id))
} else {
return None;
};
let raw_token = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(token)
.ok()?;
if raw_token.len() != 32
|| token.len() != 43
|| base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(&raw_token) != token
{
return None;
}
let token_hash = Sha256::digest(&raw_token).into();
let mut accept = Sha1::new();
accept.update(key.as_bytes());
accept.update(WEBSOCKET_GUID);
Some(ParsedUpgrade {
token_hash,
protocol: protocol.to_string(),
accept: base64::engine::general_purpose::STANDARD.encode(accept.finalize()),
carrier,
})
}
fn single_header<'a, B>(
request: &'a Request<B>,
name: impl hyper::header::AsHeaderName,
) -> Option<&'a str> {
let mut values = request.headers().get_all(name).iter();
let value = values.next()?.to_str().ok()?;
values.next().is_none().then_some(value)
}
fn single_header_eq<B>(
request: &Request<B>,
name: impl hyper::header::AsHeaderName,
expected: &str,
) -> bool {
single_header(request, name).is_some_and(|value| value.eq_ignore_ascii_case(expected))
}
fn header_has_token<B>(
request: &Request<B>,
name: impl hyper::header::AsHeaderName,
expected: &str,
) -> bool {
let mut found = false;
for value in request.headers().get_all(name) {
let Ok(value) = value.to_str() else {
return false;
};
for token in value.split(',').map(str::trim) {
if token.eq_ignore_ascii_case(expected) {
if found {
return false;
}
found = true;
}
}
}
found
}
// Ordered WebSocket message relay and deadlines are isolated from handshake parsing.
mod driver;
#[cfg(test)]
mod tests;
+638
View File
@@ -0,0 +1,638 @@
use std::sync::Arc;
use std::time::{Duration, Instant};
use bytes::Bytes;
use futures_util::{SinkExt, StreamExt};
use hyper_util::rt::TokioIo;
use tokio_tungstenite::WebSocketStream;
use tokio_tungstenite::tungstenite::protocol::{Message, Role, WebSocketConfig};
use tokio_util::sync::CancellationToken;
use super::ConnectionIo;
use crate::web::manager::{
ManagerError, WebProcessRuntime, WebSocketBudgetLease, WebSocketConnection,
};
use crate::web::session::{WebSession, WebSocketLaneReservation};
use crate::web::trace::{TraceDirection, TraceWebSocketContext};
const READ_BUFFER_BYTES: usize = 64 * 1024;
const WRITE_BUFFER_BYTES: usize = 64 * 1024;
pub(super) async fn run_upgraded(
on_upgrade: hyper::upgrade::OnUpgrade,
runtime: Arc<WebProcessRuntime>,
session: Arc<WebSession>,
connection: WebSocketConnection,
mut lane_reservation: Option<WebSocketLaneReservation>,
trace: Option<TraceWebSocketContext>,
) {
let Ok(upgraded) = on_upgrade.await else {
return;
};
let Ok(parts) = upgraded.downcast::<TokioIo<ConnectionIo>>() else {
return;
};
let mut io = parts.io.into_inner();
io.enable_websocket(parts.read_buf);
let limits = runtime.active_generation().config().web.limits.clone();
let config = WebSocketConfig::default()
.read_buffer_size(READ_BUFFER_BYTES)
.write_buffer_size(WRITE_BUFFER_BYTES)
.max_write_buffer_size(
WRITE_BUFFER_BYTES
.saturating_add(limits.carrier_batch_bytes)
.saturating_add(1024),
)
.max_message_size(Some(limits.carrier_batch_bytes))
.max_frame_size(Some(limits.carrier_batch_bytes));
let mut socket = WebSocketStream::from_raw_socket(io, Role::Server, Some(config)).await;
connection.mark_opened();
let cancellation = connection.cancellation();
if let Some(reservation) = lane_reservation.as_mut() {
let _ = run_lane(
&mut socket,
&runtime,
&session,
&connection,
reservation,
cancellation.clone(),
trace.as_ref(),
)
.await;
} else {
let _ = run_multiplex(
&mut socket,
&runtime,
&session,
&connection,
cancellation.clone(),
trace.as_ref(),
)
.await;
}
let eviction = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_eviction_secs,
);
let _ = tokio::time::timeout(eviction, socket.close(None)).await;
if let Some(reservation) = lane_reservation {
session.close_websocket_lane(reservation.lane_id());
drop(reservation);
} else {
session.close();
}
}
type CarrierSocket = WebSocketStream<ConnectionIo>;
async fn run_multiplex(
socket: &mut CarrierSocket,
runtime: &Arc<WebProcessRuntime>,
session: &Arc<WebSession>,
connection: &WebSocketConnection,
cancellation: CancellationToken,
trace: Option<&TraceWebSocketContext>,
) -> Result<(), ()> {
let mut sequence = 1u64;
let mut cursor = 0u64;
// The lease survives cancelled select branches and control frames interleaved
// inside one fragmented data message.
let mut read_budget = None;
let liveness_interval = connection.liveness_interval();
let mut next_ping = Instant::now() + liveness_interval;
loop {
let down = session.poll_down(cursor);
tokio::pin!(down);
let event = tokio::select! {
_ = cancellation.cancelled() => return Err(()),
_ = tokio::time::sleep_until(next_ping.into()) => DriverEvent::Liveness,
incoming = read_message(
socket,
runtime,
session.profile_key(),
&cancellation,
&mut read_budget,
) => {
DriverEvent::Incoming(incoming?)
}
down = &mut down => DriverEvent::Down(down.map_err(|_| ())?),
};
match event {
DriverEvent::Incoming((message, _budget)) => match message {
Message::Binary(body) => {
let started = Instant::now();
let result =
process_multiplex(runtime, session, sequence, &body, &cancellation).await;
record_message(
runtime,
trace,
TraceDirection::Request,
"binary",
&body,
started,
);
result?;
sequence = sequence.checked_add(1).ok_or(())?;
connection.mark_peer_activity();
next_ping = Instant::now() + liveness_interval;
}
Message::Pong(payload) => {
record_message(
runtime,
trace,
TraceDirection::Request,
"pong",
&payload,
Instant::now(),
);
connection.mark_peer_activity();
next_ping = Instant::now() + liveness_interval;
}
Message::Ping(payload) => {
let started = Instant::now();
flush(socket, runtime).await?;
record_message(
runtime,
trace,
TraceDirection::Request,
"ping",
&payload,
started,
);
record_message(
runtime,
trace,
TraceDirection::Response,
"pong",
&payload,
started,
);
connection.mark_peer_activity();
next_ping = Instant::now() + liveness_interval;
}
Message::Close(_) => {
record_message(
runtime,
trace,
TraceDirection::Request,
"close",
&[],
Instant::now(),
);
return Ok(());
}
Message::Text(text) => {
record_message(
runtime,
trace,
TraceDirection::Request,
"text",
text.as_bytes(),
Instant::now(),
);
return Err(());
}
Message::Frame(_) => return Err(()),
},
DriverEvent::Down(result) => {
if result.body.is_empty() {
let started = Instant::now();
send(socket, runtime, Message::Ping(Bytes::new())).await?;
record_message(
runtime,
trace,
TraceDirection::Response,
"ping",
&[],
started,
);
next_ping = Instant::now() + liveness_interval;
} else {
let _budget = reserve_data(
runtime,
session.profile_key(),
result.body.len(),
&cancellation,
)
.await?;
let body = result.body;
let started = Instant::now();
if trace.is_some() {
send(socket, runtime, Message::Binary(body.clone())).await?;
record_message(
runtime,
trace,
TraceDirection::Response,
"binary",
&body,
started,
);
} else {
send(socket, runtime, Message::Binary(body)).await?;
}
connection.mark_progress();
}
cursor = result.next_cursor;
}
DriverEvent::Liveness => {
let started = Instant::now();
send(socket, runtime, Message::Ping(Bytes::new())).await?;
record_message(
runtime,
trace,
TraceDirection::Response,
"ping",
&[],
started,
);
next_ping = Instant::now() + liveness_interval;
}
}
}
}
async fn run_lane(
socket: &mut CarrierSocket,
runtime: &Arc<WebProcessRuntime>,
session: &Arc<WebSession>,
connection: &WebSocketConnection,
reservation: &mut WebSocketLaneReservation,
cancellation: CancellationToken,
trace: Option<&TraceWebSocketContext>,
) -> Result<(), ()> {
let mut sequence = 1u64;
let mut cursor = 0u64;
// Lane reads use the same cancellation-safe fragmented-message ownership.
let mut read_budget = None;
let liveness_interval = connection.liveness_interval();
let mut next_ping = Instant::now() + liveness_interval;
loop {
let down = session.poll_down_lane(reservation.lane_id(), cursor);
tokio::pin!(down);
let event = tokio::select! {
_ = cancellation.cancelled() => return Err(()),
_ = tokio::time::sleep_until(next_ping.into()) => DriverEvent::Liveness,
incoming = read_message(
socket,
runtime,
session.profile_key(),
&cancellation,
&mut read_budget,
) => {
DriverEvent::Incoming(incoming?)
}
down = &mut down => DriverEvent::Down(down.map_err(|_| ())?),
};
match event {
DriverEvent::Incoming((message, _budget)) => match message {
Message::Binary(body) => {
let started = Instant::now();
let result = process_lane(
runtime,
session,
reservation,
sequence,
&body,
&cancellation,
)
.await;
record_message(
runtime,
trace,
TraceDirection::Request,
"binary",
&body,
started,
);
result?;
sequence = sequence.checked_add(1).ok_or(())?;
connection.mark_peer_activity();
next_ping = Instant::now() + liveness_interval;
}
Message::Pong(payload) => {
record_message(
runtime,
trace,
TraceDirection::Request,
"pong",
&payload,
Instant::now(),
);
connection.mark_peer_activity();
next_ping = Instant::now() + liveness_interval;
}
Message::Ping(payload) => {
let started = Instant::now();
flush(socket, runtime).await?;
record_message(
runtime,
trace,
TraceDirection::Request,
"ping",
&payload,
started,
);
record_message(
runtime,
trace,
TraceDirection::Response,
"pong",
&payload,
started,
);
connection.mark_peer_activity();
next_ping = Instant::now() + liveness_interval;
}
Message::Close(_) => {
record_message(
runtime,
trace,
TraceDirection::Request,
"close",
&[],
Instant::now(),
);
return Ok(());
}
Message::Text(text) => {
record_message(
runtime,
trace,
TraceDirection::Request,
"text",
text.as_bytes(),
Instant::now(),
);
return Err(());
}
Message::Frame(_) => return Err(()),
},
DriverEvent::Down(result) => {
if result.lane_closed {
return Ok(());
}
if result.body.is_empty() {
let started = Instant::now();
send(socket, runtime, Message::Ping(Bytes::new())).await?;
record_message(
runtime,
trace,
TraceDirection::Response,
"ping",
&[],
started,
);
next_ping = Instant::now() + liveness_interval;
} else {
let _budget = reserve_data(
runtime,
session.profile_key(),
result.body.len(),
&cancellation,
)
.await?;
let body = result.body;
let started = Instant::now();
if trace.is_some() {
send(socket, runtime, Message::Binary(body.clone())).await?;
record_message(
runtime,
trace,
TraceDirection::Response,
"binary",
&body,
started,
);
} else {
send(socket, runtime, Message::Binary(body)).await?;
}
connection.mark_progress();
}
cursor = result.next_cursor;
}
DriverEvent::Liveness => {
let started = Instant::now();
send(socket, runtime, Message::Ping(Bytes::new())).await?;
record_message(
runtime,
trace,
TraceDirection::Response,
"ping",
&[],
started,
);
next_ping = Instant::now() + liveness_interval;
}
}
}
}
enum DriverEvent {
Incoming((Message, Option<WebSocketBudgetLease>)),
Down(crate::web::session::PollResult),
Liveness,
}
async fn read_message(
socket: &mut CarrierSocket,
runtime: &Arc<WebProcessRuntime>,
owner: crate::web::manager::ProfileKey,
cancellation: &CancellationToken,
retained_budget: &mut Option<WebSocketBudgetLease>,
) -> Result<(Message, Option<WebSocketBudgetLease>), ()> {
tokio::select! {
_ = cancellation.cancelled() => return Err(()),
ready = socket.get_ref().readable() => ready.map_err(|_| ())?,
}
if retained_budget.is_none() {
let maximum = runtime
.active_generation()
.config()
.web
.limits
.carrier_batch_bytes;
*retained_budget = Some(reserve_data(runtime, owner, maximum, cancellation).await?);
}
let message = tokio::select! {
_ = cancellation.cancelled() => return Err(()),
message = socket.next() => message.ok_or(())?.map_err(|_| ())?,
};
if socket.get_ref().websocket_fragmented_message() {
return Ok((message, None));
}
let mut budget = retained_budget.take().ok_or(())?;
budget.shrink_to(message.len());
Ok((message, Some(budget)))
}
async fn reserve_data(
runtime: &Arc<WebProcessRuntime>,
owner: crate::web::manager::ProfileKey,
bytes: usize,
cancellation: &CancellationToken,
) -> Result<WebSocketBudgetLease, ()> {
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_backpressure_secs,
);
tokio::time::timeout(timeout, async {
loop {
let notify = runtime.budget_notify();
let notified = notify.notified();
if let Some(budget) = runtime.try_websocket_data_budget(owner, bytes.max(1)) {
return Ok(budget);
}
tokio::select! {
_ = cancellation.cancelled() => return Err(()),
_ = notified => {}
}
}
})
.await
.map_err(|_| ())?
}
async fn process_multiplex(
runtime: &Arc<WebProcessRuntime>,
session: &Arc<WebSession>,
sequence: u64,
body: &[u8],
cancellation: &CancellationToken,
) -> Result<(), ()> {
retry_backpressure(runtime, cancellation, || {
session.process_up(sequence, body).map(|_| ())
})
.await
}
async fn process_lane(
runtime: &Arc<WebProcessRuntime>,
session: &Arc<WebSession>,
reservation: &mut WebSocketLaneReservation,
sequence: u64,
body: &[u8],
cancellation: &CancellationToken,
) -> Result<(), ()> {
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_backpressure_secs,
);
tokio::time::timeout(timeout, async {
loop {
let notify = runtime.budget_notify();
let notified = notify.notified();
match session.process_websocket_lane(reservation, sequence, body) {
Ok(()) => return Ok(()),
Err(ManagerError::Backpressure) => {}
Err(_) => return Err(()),
}
tokio::select! {
_ = cancellation.cancelled() => return Err(()),
_ = notified => {}
}
}
})
.await
.map_err(|_| ())?
}
async fn retry_backpressure<F>(
runtime: &Arc<WebProcessRuntime>,
cancellation: &CancellationToken,
mut operation: F,
) -> Result<(), ()>
where
F: FnMut() -> Result<(), ManagerError>,
{
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_backpressure_secs,
);
tokio::time::timeout(timeout, async {
loop {
let notify = runtime.budget_notify();
let notified = notify.notified();
match operation() {
Ok(()) => return Ok(()),
Err(ManagerError::Backpressure) => {}
Err(_) => return Err(()),
}
tokio::select! {
_ = cancellation.cancelled() => return Err(()),
_ = notified => {}
}
}
})
.await
.map_err(|_| ())?
}
async fn send(
socket: &mut CarrierSocket,
runtime: &WebProcessRuntime,
message: Message,
) -> Result<(), ()> {
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_write_secs,
);
tokio::time::timeout(timeout, socket.send(message))
.await
.map_err(|_| ())?
.map_err(|_| ())
}
async fn flush(socket: &mut CarrierSocket, runtime: &WebProcessRuntime) -> Result<(), ()> {
let timeout = Duration::from_secs(
runtime
.active_generation()
.config()
.web
.timeouts
.websocket_write_secs,
);
tokio::time::timeout(timeout, socket.flush())
.await
.map_err(|_| ())?
.map_err(|_| ())
}
fn record_message(
runtime: &WebProcessRuntime,
trace: Option<&TraceWebSocketContext>,
direction: TraceDirection,
message_type: &'static str,
payload: &[u8],
started: Instant,
) {
let Some(trace) = trace else {
return;
};
runtime.trace().record_websocket_message(
trace,
direction,
message_type,
payload,
started.elapsed().as_micros().min(u128::from(u64::MAX)) as u64,
);
}
+311
View File
@@ -0,0 +1,311 @@
use super::*;
use std::sync::Arc;
use std::time::Duration;
use arc_swap::ArcSwap;
use futures_util::{SinkExt, StreamExt};
use sha2::{Digest, Sha256};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio_tungstenite::WebSocketStream;
use tokio_tungstenite::tungstenite::protocol::{Message, Role};
use tokio_util::sync::CancellationToken;
use crate::maestro::generation::{RuntimeGeneration, test_runtime_generation};
use crate::web::frame::{self, FrameType};
use crate::web::http::tests::runtime_config;
use crate::web::manager::WebProcessRuntime;
fn request(protocol: &str) -> Request<()> {
Request::builder()
.method(Method::GET)
.uri("/api/v1/ws")
.header(header::CONNECTION, "keep-alive, Upgrade")
.header(header::UPGRADE, "websocket")
.header("sec-websocket-version", "13")
.header("sec-websocket-key", "dGhlIHNhbXBsZSBub25jZQ==")
.header("sec-websocket-protocol", protocol)
.body(())
.unwrap()
}
#[test]
fn canonical_multiplex_and_lane_protocols_are_accepted() {
let token = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode([7u8; 32]);
let multiplex = parse_upgrade(&request(&format!("tproxy-v1.{token}"))).unwrap();
assert!(matches!(multiplex.carrier, ParsedCarrier::Multiplex));
assert_eq!(multiplex.accept, "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=");
let lane = parse_upgrade(&request(&format!("tproxy-lane-v1.{token}.16777215"))).unwrap();
assert!(matches!(lane.carrier, ParsedCarrier::Lane(16_777_215)));
}
#[test]
fn aliases_authorization_and_request_bodies_are_rejected_before_upgrade() {
let token = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode([9u8; 32]);
assert!(parse_upgrade(&request(&format!("tproxy-lane-v1.{token}.01"))).is_none());
assert!(parse_upgrade(&request(&format!("tproxy-lane-v1.{token}.0"))).is_none());
let with_authorization = Request::builder()
.method(Method::GET)
.uri("/api/v1/ws")
.header(header::CONNECTION, "Upgrade")
.header(header::UPGRADE, "websocket")
.header(header::AUTHORIZATION, "Bearer hidden")
.header("sec-websocket-version", "13")
.header("sec-websocket-key", "dGhlIHNhbXBsZSBub25jZQ==")
.header("sec-websocket-protocol", format!("tproxy-v1.{token}"))
.body(())
.unwrap();
assert!(parse_upgrade(&with_authorization).is_none());
let with_body = Request::builder()
.method(Method::GET)
.uri("/api/v1/ws")
.header(header::CONNECTION, "Upgrade")
.header(header::UPGRADE, "websocket")
.header(header::CONTENT_LENGTH, "1")
.header("sec-websocket-version", "13")
.header("sec-websocket-key", "dGhlIHNhbXBsZSBub25jZQ==")
.header("sec-websocket-protocol", format!("tproxy-v1.{token}"))
.body(())
.unwrap();
assert!(parse_upgrade(&with_body).is_none());
}
struct LiveRuntime {
runtime: Arc<WebProcessRuntime>,
generation: Arc<RuntimeGeneration>,
}
impl LiveRuntime {
async fn shutdown(self) {
self.runtime.shutdown().await;
self.generation.stop_sessions().await;
self.generation.stop_background_tasks().await;
}
}
fn live_runtime(carrier: WebCarrier) -> LiveRuntime {
live_runtime_with_long_poll(carrier, 1)
}
fn live_runtime_with_long_poll(carrier: WebCarrier, long_poll_secs: u64) -> LiveRuntime {
let mut config = runtime_config([31; 32], carrier);
config.web.timeouts.long_poll_secs = long_poll_secs;
config.web.timeouts.websocket_write_secs = 2;
config.web.timeouts.websocket_backpressure_secs = 2;
config.web.timeouts.websocket_eviction_secs = 1;
let generation = test_runtime_generation(1, config);
let runtime = WebProcessRuntime::start(Arc::new(ArcSwap::from(Arc::clone(&generation))));
LiveRuntime {
runtime,
generation,
}
}
fn create_session(runtime: &Arc<WebProcessRuntime>) -> (String, TokenHash) {
let profile = runtime
.active_generation()
.config()
.web
.runtime
.as_ref()
.unwrap()
.profiles[0]
.clone();
let client_ip = "192.0.2.10".parse().unwrap();
let bootstrap = runtime.issue_bootstrap(profile, client_ip).unwrap().token;
let raw = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(&bootstrap)
.unwrap();
let bootstrap_hash = Sha256::digest(raw).into();
let hello = frame::encode(FrameType::Hello, 0, &[1]);
let session = runtime
.create_session(bootstrap_hash, "proxy.example.com", client_ip, &hello)
.unwrap()
.token;
let raw = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(&session)
.unwrap();
let session_hash = Sha256::digest(raw).into();
(session, session_hash)
}
async fn upgrade(
listener: &TcpListener,
runtime: &Arc<WebProcessRuntime>,
protocol: &str,
) -> WebSocketStream<TcpStream> {
let addr = listener.local_addr().unwrap();
let (accepted, client) = tokio::join!(listener.accept(), TcpStream::connect(addr));
let (server, peer) = accepted.unwrap();
let mut client = client.unwrap();
let permit = runtime.try_http_connection().unwrap();
tokio::spawn(super::super::serve_connection(
server,
peer,
WebClientIpSource::XForwardedFor,
Arc::from(["127.0.0.1/32".parse().unwrap()]),
Arc::clone(runtime),
CancellationToken::new(),
permit,
));
let request = format!(
"GET /api/v1/ws HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nConnection: Upgrade\r\nUpgrade: websocket\r\nSec-WebSocket-Version: 13\r\nSec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\nSec-WebSocket-Protocol: {protocol}\r\nCookie: browser-state=allowed\r\n\r\n"
);
client.write_all(request.as_bytes()).await.unwrap();
let mut response = Vec::new();
while !response.ends_with(b"\r\n\r\n") {
let byte = client.read_u8().await.unwrap();
response.push(byte);
assert!(response.len() <= 16 * 1024);
}
assert!(response.starts_with(b"HTTP/1.1 101"));
assert!(
std::str::from_utf8(&response)
.unwrap()
.contains(&format!("sec-websocket-protocol: {protocol}"))
);
WebSocketStream::from_raw_socket(client, Role::Client, None).await
}
fn masked_message(opcode: u8, payload: &[u8], mask: [u8; 4]) -> Vec<u8> {
masked_frame(true, opcode, payload, mask)
}
fn masked_frame(finished: bool, opcode: u8, payload: &[u8], mask: [u8; 4]) -> Vec<u8> {
assert!(payload.len() < 126);
let mut encoded = Vec::with_capacity(payload.len() + 6);
encoded.push(u8::from(finished) << 7 | opcode);
encoded.push(0x80 | payload.len() as u8);
encoded.extend_from_slice(&mask);
encoded.extend(
payload
.iter()
.enumerate()
.map(|(index, value)| value ^ mask[index % mask.len()]),
);
encoded
}
#[tokio::test]
async fn multiplex_upgrade_relays_binary_and_transport_control_messages() {
let live = live_runtime(WebCarrier::Websocket);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let (session, _) = create_session(&live.runtime);
let protocol = format!("tproxy-v1.{session}");
let mut socket = upgrade(&listener, &live.runtime, &protocol).await;
socket
.send(Message::Binary(frame::encode(FrameType::Pong, 0, &[])))
.await
.unwrap();
socket
.send(Message::Ping(Bytes::from_static(b"live")))
.await
.unwrap();
let response = tokio::time::timeout(Duration::from_secs(2), socket.next())
.await
.unwrap()
.unwrap()
.unwrap();
assert_eq!(response, Message::Pong(Bytes::from_static(b"live")));
let _ = socket.close(None).await;
live.shutdown().await;
}
#[tokio::test]
async fn coalesced_websocket_messages_do_not_wait_for_new_tcp_readiness() {
let live = live_runtime_with_long_poll(WebCarrier::Websocket, 10);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let (session, _) = create_session(&live.runtime);
let protocol = format!("tproxy-v1.{session}");
let mut socket = upgrade(&listener, &live.runtime, &protocol).await;
let mut wire = masked_message(0x02, &frame::encode(FrameType::Pong, 0, &[]), [1, 2, 3, 4]);
wire.extend_from_slice(&masked_message(0x09, b"coalesced", [5, 6, 7, 8]));
socket.get_mut().write_all(&wire).await.unwrap();
let response = tokio::time::timeout(Duration::from_secs(2), socket.next())
.await
.unwrap()
.unwrap()
.unwrap();
assert_eq!(response, Message::Pong(Bytes::from_static(b"coalesced")));
let _ = socket.close(None).await;
live.shutdown().await;
}
#[tokio::test]
async fn fragmented_message_budget_survives_interleaved_control_frames() {
let live = live_runtime_with_long_poll(WebCarrier::Websocket, 10);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let (session, _) = create_session(&live.runtime);
let protocol = format!("tproxy-v1.{session}");
let mut socket = upgrade(&listener, &live.runtime, &protocol).await;
let body = frame::encode(FrameType::Pong, 0, &[]);
let mut wire = masked_frame(false, 0x02, &body[..4], [1, 2, 3, 4]);
wire.extend_from_slice(&masked_message(0x09, b"mid", [5, 6, 7, 8]));
wire.extend_from_slice(&masked_frame(true, 0x00, &body[4..], [9, 10, 11, 12]));
wire.extend_from_slice(&masked_message(0x09, b"after", [13, 14, 15, 16]));
socket.get_mut().write_all(&wire).await.unwrap();
for expected in [b"mid".as_slice(), b"after".as_slice()] {
let response = tokio::time::timeout(Duration::from_secs(2), socket.next())
.await
.unwrap()
.unwrap()
.unwrap();
assert_eq!(response, Message::Pong(Bytes::copy_from_slice(expected)));
}
let _ = socket.close(None).await;
live.shutdown().await;
}
#[tokio::test]
async fn malformed_websocket_lane_closes_only_that_lane() {
let live = live_runtime(WebCarrier::WebsocketLanes);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let (session, session_hash) = create_session(&live.runtime);
let first_protocol = format!("tproxy-lane-v1.{session}.7");
let mut first = upgrade(&listener, &live.runtime, &first_protocol).await;
first
.send(Message::Binary(frame::encode(FrameType::Data, 7, &[1])))
.await
.unwrap();
let closed = tokio::time::timeout(Duration::from_secs(2), first.next())
.await
.unwrap()
.unwrap()
.unwrap();
assert!(matches!(closed, Message::Close(_)));
assert!(
live.runtime
.get_session(session_hash, "proxy.example.com")
.is_ok()
);
let second_protocol = format!("tproxy-lane-v1.{session}.8");
let mut second = upgrade(&listener, &live.runtime, &second_protocol).await;
second
.send(Message::Binary(frame::encode(FrameType::Open, 8, &[])))
.await
.unwrap();
second
.send(Message::Ping(Bytes::from_static(b"lane")))
.await
.unwrap();
let pong = tokio::time::timeout(Duration::from_secs(2), second.next())
.await
.unwrap()
.unwrap()
.unwrap();
assert_eq!(pong, Message::Pong(Bytes::from_static(b"lane")));
let _ = second.close(None).await;
live.shutdown().await;
}