mirror of
https://github.com/telemt/telemt.git
synced 2026-09-10 20:44:08 +03:00
WEB: websocket + websocket-lanes as Carrier
This commit is contained in:
@@ -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