mirror of
https://github.com/telemt/telemt.git
synced 2026-09-08 11:34:10 +03:00
WEB: websocket + websocket-lanes as Carrier
This commit is contained in:
+10
-3
@@ -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;
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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());
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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;
|
||||
@@ -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,
|
||||
);
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
Reference in New Issue
Block a user