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
+3 -6
View File
@@ -499,12 +499,9 @@ async fn handle(
let result: Result<Response<Full<Bytes>>, ApiFailure> = async { let result: Result<Response<Full<Bytes>>, ApiFailure> = async {
match (method.as_str(), normalized_path) { match (method.as_str(), normalized_path) {
("GET", "/web-status") => Ok(web_status::render( ("GET", "/web-status") => {
query.as_deref(), Ok(web_status::render(query.as_deref(), &shared.web_trace, &cfg.web.debug).await)
&shared.web_trace, }
&cfg.web.debug,
)
.await),
("GET", "/v1/health") => { ("GET", "/v1/health") => {
let revision = current_revision(&shared.config_path).await?; let revision = current_revision(&shared.config_path).await?;
let data = HealthData { let data = HealthData {
+61 -80
View File
@@ -1,7 +1,6 @@
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::sync::Arc; use std::sync::Arc;
use base64::Engine as _;
use http_body_util::Full; use http_body_util::Full;
use hyper::body::Bytes; use hyper::body::Bytes;
use hyper::header::{self, HeaderValue}; use hyper::header::{self, HeaderValue};
@@ -9,16 +8,17 @@ use hyper::{Response, StatusCode};
use tokio::sync::OwnedSemaphorePermit; use tokio::sync::OwnedSemaphorePermit;
use crate::config::WebDebugConfig; use crate::config::WebDebugConfig;
use crate::web::trace::{ use crate::web::trace::{StoredTraceRecord, TraceRecord, TraceRecordKind, WebTraceStore};
StoredTraceRecord, TraceRecord, TraceRecordKind, WebTraceStore,
};
const MAX_PAGE_BYTES: usize = 8 * 1024 * 1024; const MAX_PAGE_BYTES: usize = 8 * 1024 * 1024;
const MAX_GROUPS: usize = 1024; const MAX_GROUPS: usize = 1024;
// Record-detail rendering remains isolated from filtering and page layout.
mod details;
// Query parsing and matching remain independent from bounded HTML rendering. // Query parsing and matching remain independent from bounded HTML rendering.
mod query; mod query;
use details::{push_body, push_frames, push_headers};
use query::{GroupBy, StatusQuery, client_ip, parse_query, record_matches}; use query::{GroupBy, StatusQuery, client_ip, parse_query, record_matches};
struct GroupSummary { struct GroupSummary {
@@ -119,11 +119,18 @@ fn push_page_start(html: &mut String) {
fn push_filter_form(html: &mut String, query: &StatusQuery) { fn push_filter_form(html: &mut String, query: &StatusQuery) {
html.push_str("<section><h2>Filters</h2><form method=\"get\" action=\"/web-status\">"); html.push_str("<section><h2>Filters</h2><form method=\"get\" action=\"/web-status\">");
input(html, "window_secs", &query.window_secs.to_string()); input(html, "window_secs", &query.window_secs.to_string());
input(html, "ip", &query.ip.map(|value| value.to_string()).unwrap_or_default()); input(
html,
"ip",
&query.ip.map(|value| value.to_string()).unwrap_or_default(),
);
input( input(
html, html,
"session", "session",
&query.session.map(|value| value.to_string()).unwrap_or_default(), &query
.session
.map(|value| value.to_string())
.unwrap_or_default(),
); );
input( input(
html, html,
@@ -133,7 +140,12 @@ fn push_filter_form(html: &mut String, query: &StatusQuery) {
input(html, "key", query.key.as_deref().unwrap_or_default()); input(html, "key", query.key.as_deref().unwrap_or_default());
input(html, "limit", &query.limit.to_string()); input(html, "limit", &query.limit.to_string());
html.push_str("<label>group_by<select name=\"group_by\" multiple size=\"4\">"); html.push_str("<label>group_by<select name=\"group_by\" multiple size=\"4\">");
for group in [GroupBy::Ip, GroupBy::Session, GroupBy::UserAgent, GroupBy::Key] { for group in [
GroupBy::Ip,
GroupBy::Session,
GroupBy::UserAgent,
GroupBy::Key,
] {
html.push_str("<option value=\""); html.push_str("<option value=\"");
html.push_str(group.as_str()); html.push_str(group.as_str());
if query.group_by.contains(&group) { if query.group_by.contains(&group) {
@@ -165,11 +177,7 @@ fn summary_row(html: &mut String, name: &str, value: &str) {
html.push_str("</td></tr>"); html.push_str("</td></tr>");
} }
fn push_groups( fn push_groups(html: &mut String, records: &[Arc<StoredTraceRecord>], groups: &[GroupBy]) {
html: &mut String,
records: &[Arc<StoredTraceRecord>],
groups: &[GroupBy],
) {
let mut summaries = BTreeMap::<Vec<String>, GroupSummary>::new(); let mut summaries = BTreeMap::<Vec<String>, GroupSummary>::new();
let mut overflow = 0usize; let mut overflow = 0usize;
for stored in records { for stored in records {
@@ -275,7 +283,15 @@ fn push_record(html: &mut String, record: &TraceRecord) {
"http", "http",
http.route.as_str(), http.route.as_str(),
http.method.as_str(), http.method.as_str(),
http.status.map(|value| value.to_string()).unwrap_or_else(|| "-".to_string()), http.status
.map(|value| value.to_string())
.unwrap_or_else(|| "-".to_string()),
),
TraceRecordKind::Websocket(message) => (
"websocket",
message.direction.as_str(),
message.message_type,
message.payload_bytes.to_string(),
), ),
TraceRecordKind::Lifecycle(event) => ( TraceRecordKind::Lifecycle(event) => (
"lifecycle", "lifecycle",
@@ -327,32 +343,40 @@ fn push_record(html: &mut String, record: &TraceRecord) {
html.push_str(&option_u64(timings.response_body_us)); html.push_str(&option_u64(timings.response_body_us));
html.push_str(" us\n(kernel flush and TCP ACK are not observed)</pre>"); html.push_str(" us\n(kernel flush and TCP ACK are not observed)</pre>");
} }
if !http.frames.is_empty() { push_frames(html, &http.frames);
html.push_str("<h3>frames</h3><table><tr><th>dir</th><th>type</th><th>stream/lane</th><th>payload</th><th>WINDOW</th><th>error</th></tr>"); }
for frame in &http.frames { TraceRecordKind::Websocket(message) => {
html.push_str("<tr>"); html.push_str("<pre>connection: ");
for value in [ html.push_str(&message.connection_id.to_string());
frame.direction.as_str().to_string(), html.push_str("\nlane: ");
frame.frame_type.unwrap_or("-").to_string(), html.push_str(
frame.stream_id.map(|v| v.to_string()).unwrap_or_else(|| "-".to_string()), &message
frame.payload_len.map(|v| v.to_string()).unwrap_or_else(|| "-".to_string()), .lane_id
frame.window_delta.map(|v| v.to_string()).unwrap_or_else(|| "-".to_string()), .map(|value| value.to_string())
frame.parse_error.unwrap_or("-").to_string(), .unwrap_or_else(|| "-".to_string()),
] { );
html.push_str("<td>"); html.push_str("\ndirection: ");
escape(html, &value); html.push_str(message.direction.as_str());
html.push_str("</td>"); html.push_str("\nmessage: ");
} html.push_str(message.message_type);
html.push_str("</tr>"); html.push_str("\npayload bytes: ");
} html.push_str(&message.payload_bytes.to_string());
html.push_str("</table>"); html.push_str("\nduration: ");
} html.push_str(&option_u64(message.duration_us));
html.push_str(" us</pre>");
push_body(html, "message body", message.body.as_ref());
push_frames(html, &message.frames);
} }
TraceRecordKind::Lifecycle(event) => { TraceRecordKind::Lifecycle(event) => {
html.push_str("<pre>event: "); html.push_str("<pre>event: ");
html.push_str(event.event.as_str()); html.push_str(event.event.as_str());
html.push_str("\nstream: "); html.push_str("\nstream: ");
html.push_str(&event.stream_id.map(|v| v.to_string()).unwrap_or_else(|| "-".to_string())); html.push_str(
&event
.stream_id
.map(|v| v.to_string())
.unwrap_or_else(|| "-".to_string()),
);
html.push_str("\nreason: "); html.push_str("\nreason: ");
html.push_str(event.reason.unwrap_or("-")); html.push_str(event.reason.unwrap_or("-"));
html.push_str("</pre>"); html.push_str("</pre>");
@@ -361,51 +385,6 @@ fn push_record(html: &mut String, record: &TraceRecord) {
html.push_str("</details></td></tr>"); html.push_str("</details></td></tr>");
} }
fn push_headers(html: &mut String, title: &str, headers: &[crate::web::trace::TraceHeader]) {
html.push_str("<h3>");
escape(html, title);
html.push_str("</h3><pre>");
for header in headers {
escape(html, &header.name);
html.push_str(": ");
escape(html, header.value.as_deref().unwrap_or("[value omitted]"));
html.push('\n');
}
html.push_str("</pre>");
}
fn push_body(
html: &mut String,
title: &str,
body: Option<&crate::web::trace::TraceBodySnapshot>,
) {
html.push_str("<h3>");
escape(html, title);
html.push_str("</h3>");
let Some(body) = body else {
html.push_str("<p class=\"muted\">capture off</p>");
return;
};
html.push_str("<p>observed=");
html.push_str(&body.observed_bytes.to_string());
html.push_str(" captured=");
html.push_str(&body.captured.len().to_string());
html.push_str(" state=");
html.push_str(body.state.as_str());
html.push_str(" truncated=");
html.push_str(yes_no(body.truncated));
html.push_str("</p><pre>");
let available = MAX_PAGE_BYTES.saturating_sub(html.len()).saturating_sub(4096);
let raw_limit = available.saturating_mul(3) / 4;
let shown = body.captured.len().min(raw_limit);
base64::engine::general_purpose::STANDARD
.encode_string(&body.captured[..shown], html);
if shown < body.captured.len() {
html.push_str("\n[page output truncated]");
}
html.push_str("</pre>");
}
fn pagination_url(query: &StatusQuery, before_seq: u64) -> String { fn pagination_url(query: &StatusQuery, before_seq: u64) -> String {
let mut serializer = url::form_urlencoded::Serializer::new(String::from("/web-status?")); let mut serializer = url::form_urlencoded::Serializer::new(String::from("/web-status?"));
serializer.append_pair("window_secs", &query.window_secs.to_string()); serializer.append_pair("window_secs", &query.window_secs.to_string());
@@ -436,7 +415,9 @@ fn format_time(epoch_millis: u64) -> String {
} }
fn option_u64(value: Option<u64>) -> String { fn option_u64(value: Option<u64>) -> String {
value.map(|value| value.to_string()).unwrap_or_else(|| "-".to_string()) value
.map(|value| value.to_string())
.unwrap_or_else(|| "-".to_string())
} }
fn body_mode(policy: &WebDebugConfig) -> &'static str { fn body_mode(policy: &WebDebugConfig) -> &'static str {
+86
View File
@@ -0,0 +1,86 @@
use base64::Engine as _;
use super::{MAX_PAGE_BYTES, escape, yes_no};
pub(super) fn push_frames(html: &mut String, frames: &[crate::web::trace::TraceFrame]) {
if frames.is_empty() {
return;
}
html.push_str("<h3>frames</h3><table><tr><th>dir</th><th>type</th><th>stream/lane</th><th>payload</th><th>WINDOW</th><th>error</th></tr>");
for frame in frames {
html.push_str("<tr>");
for value in [
frame.direction.as_str().to_string(),
frame.frame_type.unwrap_or("-").to_string(),
frame
.stream_id
.map(|value| value.to_string())
.unwrap_or_else(|| "-".to_string()),
frame
.payload_len
.map(|value| value.to_string())
.unwrap_or_else(|| "-".to_string()),
frame
.window_delta
.map(|value| value.to_string())
.unwrap_or_else(|| "-".to_string()),
frame.parse_error.unwrap_or("-").to_string(),
] {
html.push_str("<td>");
escape(html, &value);
html.push_str("</td>");
}
html.push_str("</tr>");
}
html.push_str("</table>");
}
pub(super) fn push_headers(
html: &mut String,
title: &str,
headers: &[crate::web::trace::TraceHeader],
) {
html.push_str("<h3>");
escape(html, title);
html.push_str("</h3><pre>");
for header in headers {
escape(html, &header.name);
html.push_str(": ");
escape(html, header.value.as_deref().unwrap_or("[value omitted]"));
html.push('\n');
}
html.push_str("</pre>");
}
pub(super) fn push_body(
html: &mut String,
title: &str,
body: Option<&crate::web::trace::TraceBodySnapshot>,
) {
html.push_str("<h3>");
escape(html, title);
html.push_str("</h3>");
let Some(body) = body else {
html.push_str("<p class=\"muted\">capture off</p>");
return;
};
html.push_str("<p>observed=");
html.push_str(&body.observed_bytes.to_string());
html.push_str(" captured=");
html.push_str(&body.captured.len().to_string());
html.push_str(" state=");
html.push_str(body.state.as_str());
html.push_str(" truncated=");
html.push_str(yes_no(body.truncated));
html.push_str("</p><pre>");
let available = MAX_PAGE_BYTES
.saturating_sub(html.len())
.saturating_sub(4096);
let raw_limit = available.saturating_mul(3) / 4;
let shown = body.captured.len().min(raw_limit);
base64::engine::general_purpose::STANDARD.encode_string(&body.captured[..shown], html);
if shown < body.captured.len() {
html.push_str("\n[page output truncated]");
}
html.push_str("</pre>");
}
+4 -7
View File
@@ -101,8 +101,9 @@ pub(super) fn parse_query(
query.key = Some(value.to_string()); query.key = Some(value.to_string());
} }
"group_by" => { "group_by" => {
let group = GroupBy::parse(value) let group = GroupBy::parse(value).ok_or_else(|| {
.ok_or_else(|| "group_by must be ip, session, user_agent, or key".to_string())?; "group_by must be ip, session, user_agent, or key".to_string()
})?;
if query.group_by.contains(&group) { if query.group_by.contains(&group) {
return Err("group_by values must not repeat".to_string()); return Err("group_by values must not repeat".to_string());
} }
@@ -140,11 +141,7 @@ fn parse_positive_u64(value: &str, field: &str) -> Result<u64, String> {
} }
/// Applies the complete filter predicate to one immutable record. /// Applies the complete filter predicate to one immutable record.
pub(super) fn record_matches( pub(super) fn record_matches(record: &TraceRecord, query: &StatusQuery, since_millis: u64) -> bool {
record: &TraceRecord,
query: &StatusQuery,
since_millis: u64,
) -> bool {
!(record.epoch_millis < since_millis !(record.epoch_millis < since_millis
|| query.before_seq.is_some_and(|before| record.seq >= before) || query.before_seq.is_some_and(|before| record.seq >= before)
|| query.record.is_some_and(|seq| record.seq != seq) || query.record.is_some_and(|seq| record.seq != seq)
+5 -1
View File
@@ -47,7 +47,11 @@ async fn renderer_filters_groups_and_sets_control_plane_security_headers() {
.await; .await;
assert_eq!(response.status(), StatusCode::OK); assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.headers()[header::CACHE_CONTROL], "no-store"); assert_eq!(response.headers()[header::CACHE_CONTROL], "no-store");
assert!(response.headers().contains_key(header::CONTENT_SECURITY_POLICY)); assert!(
response
.headers()
.contains_key(header::CONTENT_SECURITY_POLICY)
);
let body = response.into_body().collect().await.unwrap().to_bytes(); let body = response.into_body().collect().await.unwrap().to_bytes();
let body = std::str::from_utf8(&body).unwrap(); let body = std::str::from_utf8(&body).unwrap();
assert!(body.contains("session_created")); assert!(body.contains("session_created"));
+4 -1
View File
@@ -114,7 +114,10 @@ fn web_debug_prefix_requiring_deferred_capacity_is_not_hot_applied() {
new.web.debug.body_prefix_bytes = 3 * 1024 * 1024; new.web.debug.body_prefix_bytes = 3 * 1024 * 1024;
let applied = overlay_hot_fields(&old, &new); let applied = overlay_hot_fields(&old, &new);
assert_eq!(applied.web.limits.max_body_bytes, old.web.limits.max_body_bytes); assert_eq!(
applied.web.limits.max_body_bytes,
old.web.limits.max_body_bytes
);
assert_eq!( assert_eq!(
applied.web.debug.body_prefix_bytes, applied.web.debug.body_prefix_bytes,
old.web.debug.body_prefix_bytes old.web.debug.body_prefix_bytes
+1 -2
View File
@@ -50,8 +50,7 @@ pub(super) fn rebuild(config: &mut ProxyConfig) -> Result<()> {
client_secret(auth_entry.secret, profile.secret_mode); client_secret(auth_entry.secret, profile.secret_mode);
let capability = let capability =
derive_web_capability(&client_secret[..client_secret_len], vhost.host.as_bytes())?; derive_web_capability(&client_secret[..client_secret_len], vhost.host.as_bytes())?;
let key_fingerprint = let key_fingerprint = debug_key_fingerprint(&client_secret[..client_secret_len]);
debug_key_fingerprint(&client_secret[..client_secret_len]);
if !capabilities.insert(capability) { if !capabilities.insert(capability) {
return Err(ProxyError::Config(format!( return Err(ProxyError::Config(format!(
"WEB vhost `{}` contains profiles with the same client capability", "WEB vhost `{}` contains profiles with the same client capability",
+10 -1
View File
@@ -259,7 +259,9 @@ const LISTENER_CONFIG_KEYS: &[&str] = &[
"web_trusted_proxy_cidrs", "web_trusted_proxy_cidrs",
]; ];
const WEB_CONFIG_KEYS: &[&str] = &["enabled", "carrier", "debug", "limits", "timeouts", "vhosts"]; const WEB_CONFIG_KEYS: &[&str] = &[
"enabled", "carrier", "debug", "limits", "timeouts", "vhosts",
];
const WEB_LIMITS_CONFIG_KEYS: &[&str] = &[ const WEB_LIMITS_CONFIG_KEYS: &[&str] = &[
"max_header_bytes", "max_header_bytes",
@@ -269,6 +271,10 @@ const WEB_LIMITS_CONFIG_KEYS: &[&str] = &[
"max_frames_per_body", "max_frames_per_body",
"max_http_connections", "max_http_connections",
"max_http_handlers", "max_http_handlers",
"websocket_bytes_global",
"websocket_admission_watermark_pct",
"websocket_eviction_watermark_pct",
"websocket_http_connection_reserve",
"max_body_readers", "max_body_readers",
"max_body_bytes_global", "max_body_bytes_global",
"max_sessions_global", "max_sessions_global",
@@ -319,6 +325,9 @@ const WEB_TIMEOUTS_CONFIG_KEYS: &[&str] = &[
"body_secs", "body_secs",
"stream_handshake_secs", "stream_handshake_secs",
"long_poll_secs", "long_poll_secs",
"websocket_write_secs",
"websocket_backpressure_secs",
"websocket_eviction_secs",
"bootstrap_lifetime_secs", "bootstrap_lifetime_secs",
"reconnect_grace_secs", "reconnect_grace_secs",
"http_idle_secs", "http_idle_secs",
+9
View File
@@ -6,6 +6,8 @@ use super::*;
mod debug; mod debug;
// Memory-envelope arithmetic remains isolated from protocol validation. // Memory-envelope arithmetic remains isolated from protocol validation.
mod memory; mod memory;
// WebSocket transport policy is validated independently from HTTP body policy.
mod websocket;
const WEB_FRAME_HEADER_BYTES: usize = 8; const WEB_FRAME_HEADER_BYTES: usize = 8;
const WEB_QUEUE_ITEM_COST: usize = 256; const WEB_QUEUE_ITEM_COST: usize = 256;
@@ -69,6 +71,7 @@ pub(super) fn validate(config: &mut ProxyConfig) -> Result<()> {
return config_error("web.carrier=https-lanes requires web.limits.max_http_handlers >= 2"); return config_error("web.carrier=https-lanes requires web.limits.max_http_handlers >= 2");
} }
validate_timeouts(&config.web.timeouts)?; validate_timeouts(&config.web.timeouts)?;
websocket::validate(config.web.carrier, &config.web.limits, &config.web.timeouts)?;
validate_vhosts(config)?; validate_vhosts(config)?;
Ok(()) Ok(())
} }
@@ -327,6 +330,12 @@ fn validate_timeouts(timeouts: &WebTimeoutsConfig) -> Result<()> {
("body_secs", timeouts.body_secs), ("body_secs", timeouts.body_secs),
("stream_handshake_secs", timeouts.stream_handshake_secs), ("stream_handshake_secs", timeouts.stream_handshake_secs),
("long_poll_secs", timeouts.long_poll_secs), ("long_poll_secs", timeouts.long_poll_secs),
("websocket_write_secs", timeouts.websocket_write_secs),
(
"websocket_backpressure_secs",
timeouts.websocket_backpressure_secs,
),
("websocket_eviction_secs", timeouts.websocket_eviction_secs),
("bootstrap_lifetime_secs", timeouts.bootstrap_lifetime_secs), ("bootstrap_lifetime_secs", timeouts.bootstrap_lifetime_secs),
("reconnect_grace_secs", timeouts.reconnect_grace_secs), ("reconnect_grace_secs", timeouts.reconnect_grace_secs),
("http_idle_secs", timeouts.http_idle_secs), ("http_idle_secs", timeouts.http_idle_secs),
+3 -1
View File
@@ -18,7 +18,9 @@ pub(super) fn validate(policy: &WebDebugConfig, limits: &WebLimitsConfig) -> Res
); );
} }
if policy.body_prefix_bytes > limits.max_body_bytes { if policy.body_prefix_bytes > limits.max_body_bytes {
return config_error("web.debug.body_prefix_bytes must not exceed web.limits.max_body_bytes"); return config_error(
"web.debug.body_prefix_bytes must not exceed web.limits.max_body_bytes",
);
} }
if policy.body_prefix_bytes > limits.debug_bytes_global if policy.body_prefix_bytes > limits.debug_bytes_global
|| policy.decoy_body_prefix_bytes > limits.debug_bytes_global || policy.decoy_body_prefix_bytes > limits.debug_bytes_global
+1 -3
View File
@@ -46,9 +46,7 @@ pub(super) fn validate(limits: &WebLimitsConfig) -> Result<()> {
.checked_mul(WEB_DEBUG_GROUP_SCRATCH_BYTES) .checked_mul(WEB_DEBUG_GROUP_SCRATCH_BYTES)
.and_then(|scratch| value.checked_add(scratch)) .and_then(|scratch| value.checked_add(scratch))
}) })
.ok_or_else(|| { .ok_or_else(|| ProxyError::Config("web.debug reservations overflowed usize".to_string()))?;
ProxyError::Config("web.debug reservations overflowed usize".to_string())
})?;
let reserved = limits let reserved = limits
.pending_bytes_global .pending_bytes_global
.checked_add(limits.max_body_bytes_global) .checked_add(limits.max_body_bytes_global)
+83
View File
@@ -0,0 +1,83 @@
use super::*;
const MAX_WEBSOCKET_BATCH_BYTES: usize = 2 * 1024 * 1024;
const WEBSOCKET_IO_BUFFER_BYTES: usize = 64 * 1024;
const WEBSOCKET_DRIVER_OVERHEAD_BYTES: usize = 4 * 1024;
const WEBSOCKET_FRAME_OVERHEAD_BYTES: usize = 14;
/// Validates WebSocket admission, memory, and deadline invariants.
pub(super) fn validate(
carrier: WebCarrier,
limits: &WebLimitsConfig,
timeouts: &WebTimeoutsConfig,
) -> Result<()> {
if !(1..100).contains(&limits.websocket_admission_watermark_pct)
|| !(1..100).contains(&limits.websocket_eviction_watermark_pct)
|| limits.websocket_admission_watermark_pct >= limits.websocket_eviction_watermark_pct
{
return config_error(
"web.limits WebSocket watermarks must satisfy 1 <= admission < eviction < 100",
);
}
if limits.websocket_bytes_global == 0 {
return config_error("web.limits.websocket_bytes_global must be > 0");
}
if timeouts.websocket_eviction_secs > timeouts.websocket_write_secs {
return config_error(
"web.timeouts.websocket_eviction_secs must not exceed websocket_write_secs",
);
}
if !carrier.uses_websocket() {
return Ok(());
}
if limits.carrier_batch_bytes > MAX_WEBSOCKET_BATCH_BYTES {
return config_error(
"WebSocket carriers require web.limits.carrier_batch_bytes <= 2097152",
);
}
if limits.websocket_http_connection_reserve == 0
|| limits.websocket_http_connection_reserve >= limits.max_http_connections
{
return config_error(
"WebSocket carriers require websocket_http_connection_reserve within [1, max_http_connections)",
);
}
let socket_base = WEBSOCKET_IO_BUFFER_BYTES
.checked_mul(2)
.and_then(|value| value.checked_add(WEBSOCKET_DRIVER_OVERHEAD_BYTES))
.ok_or_else(|| ProxyError::Config("WebSocket base reservation overflowed usize".into()))?;
let minimum_websocket_progress = limits
.carrier_batch_bytes
.checked_add(WEBSOCKET_FRAME_OVERHEAD_BYTES)
.and_then(|value| value.checked_mul(2))
.and_then(|value| value.checked_add(socket_base))
.ok_or_else(|| {
ProxyError::Config("WebSocket progress reservation overflowed usize".into())
})?;
if limits.websocket_bytes_global < minimum_websocket_progress {
return config_error(
"web.limits.websocket_bytes_global must preserve one socket read and write",
);
}
let data_bytes = limits
.pending_bytes_global
.saturating_sub(limits.control_bytes_global);
let queue_progress = limits
.max_body_bytes
.checked_add(
limits
.max_frames_per_body
.checked_mul(WEB_QUEUE_ITEM_COST)
.ok_or_else(|| {
ProxyError::Config("WEB queue progress reservation overflowed usize".into())
})?,
)
.and_then(|value| value.checked_add(limits.carrier_batch_bytes))
.ok_or_else(|| {
ProxyError::Config("WEB queue progress reservation overflowed usize".into())
})?;
if limits.websocket_bytes_global > data_bytes.saturating_sub(queue_progress) {
return config_error("web.limits.websocket_bytes_global must leave bounded queue progress");
}
Ok(())
}
+56 -4
View File
@@ -64,10 +64,13 @@ fn web_debug_table_uses_debug_name_and_bounded_defaults() {
assert_eq!(config.web.debug.default_window_secs, 180); assert_eq!(config.web.debug.default_window_secs, 180);
assert_eq!(config.web.debug.max_window_secs, 900); assert_eq!(config.web.debug.max_window_secs, 900);
let old_name = format!("[general]\nconfig_strict = true\n{}", WEB_CONFIG.replace( let old_name = format!(
"[[web.vhosts]]", "[general]\nconfig_strict = true\n{}",
"[web.trace]\nenabled = true\n\n[[web.vhosts]]", WEB_CONFIG.replace(
)); "[[web.vhosts]]",
"[web.trace]\nenabled = true\n\n[[web.vhosts]]",
)
);
let error = load_config_error_from_temp_toml(&old_name); let error = load_config_error_from_temp_toml(&old_name);
assert!(error.contains("web.trace")); assert!(error.contains("web.trace"));
} }
@@ -150,3 +153,52 @@ fn web_ipv6_decoy_uses_a_valid_http_authority() {
}; };
assert_eq!(authority, "[::1]:18081"); assert_eq!(authority, "[::1]:18081");
} }
#[test]
fn websocket_carriers_build_runtime_profiles_with_bounded_defaults() {
for (name, carrier) in [
("websocket", WebCarrier::Websocket),
("websocket-lanes", WebCarrier::WebsocketLanes),
] {
let configured = WEB_CONFIG.replace("https-lanes", name);
let config = load_config_from_temp_toml(&configured);
let profile = &config.web.runtime.unwrap().profiles[0];
assert_eq!(profile.carrier, carrier);
assert_eq!(config.web.limits.websocket_bytes_global, 256 * 1024 * 1024);
assert_eq!(config.web.limits.websocket_admission_watermark_pct, 75);
assert_eq!(config.web.limits.websocket_eviction_watermark_pct, 90);
assert_eq!(config.web.limits.websocket_http_connection_reserve, 64);
assert_eq!(config.web.timeouts.websocket_write_secs, 30);
assert_eq!(config.web.timeouts.websocket_backpressure_secs, 30);
assert_eq!(config.web.timeouts.websocket_eviction_secs, 1);
}
}
#[test]
fn websocket_limits_reject_ambiguous_or_nonprogressing_policy() {
let reversed_watermarks = WEB_CONFIG.replace(
"carrier = \"https-lanes\"",
"carrier = \"websocket\"\n\n[web.limits]\nwebsocket_admission_watermark_pct = 90\nwebsocket_eviction_watermark_pct = 75",
);
assert!(
load_config_error_from_temp_toml(&reversed_watermarks).contains("WebSocket watermarks")
);
let no_http_reserve = WEB_CONFIG.replace(
"carrier = \"https-lanes\"",
"carrier = \"websocket\"\n\n[web.limits]\nwebsocket_http_connection_reserve = 0",
);
assert!(
load_config_error_from_temp_toml(&no_http_reserve)
.contains("websocket_http_connection_reserve")
);
let oversized_batch = WEB_CONFIG.replace(
"carrier = \"https-lanes\"",
"carrier = \"websocket\"\n\n[web.limits]\nmax_body_bytes = 4194304\ncarrier_batch_bytes = 4194304\nmax_body_readers = 16",
);
assert!(
load_config_error_from_temp_toml(&oversized_batch)
.contains("carrier_batch_bytes <= 2097152")
);
}
+2 -2
View File
@@ -54,12 +54,12 @@ pub use web::{
WebCarrier, WebConfig, WebDecoyConfig, WebLimitsConfig, WebProfileConfig, WebSecretMode, WebCarrier, WebConfig, WebDecoyConfig, WebLimitsConfig, WebProfileConfig, WebSecretMode,
WebTimeoutsConfig, WebVhostConfig, WebTimeoutsConfig, WebVhostConfig,
}; };
pub use web_debug::{WebDebugBodyCapture, WebDebugConfig};
pub(crate) use web_debug::web_debug_fits_limits;
pub(crate) use web::{ pub(crate) use web::{
WebRuntimeConfig, WebRuntimeDecoy, WebRuntimeProfile, WebRuntimeVhost, WebStaticAsset, WebRuntimeConfig, WebRuntimeDecoy, WebRuntimeProfile, WebRuntimeVhost, WebStaticAsset,
WebStaticSite, WebStaticSite,
}; };
pub(crate) use web_debug::web_debug_fits_limits;
pub use web_debug::{WebDebugBodyCapture, WebDebugConfig};
fn default_quota_state_path() -> PathBuf { fn default_quota_state_path() -> PathBuf {
PathBuf::from("telemt.limit.json") PathBuf::from("telemt.limit.json")
+65 -1
View File
@@ -18,7 +18,7 @@ pub enum WebSecretMode {
Dd, Dd,
} }
/// HTTP carrier selected for newly issued WEB bridge sessions. /// Carrier selected for newly issued WEB bridge sessions.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash, Serialize, Deserialize)] #[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")] #[serde(rename_all = "kebab-case")]
pub enum WebCarrier { pub enum WebCarrier {
@@ -27,6 +27,10 @@ pub enum WebCarrier {
Https, Https,
/// Give every logical stream independent HTTPS sequencing and polling state. /// Give every logical stream independent HTTPS sequencing and polling state.
HttpsLanes, HttpsLanes,
/// Multiplex all logical streams over one ordered WebSocket.
Websocket,
/// Give every logical stream an independently owned WebSocket lane.
WebsocketLanes,
} }
impl WebCarrier { impl WebCarrier {
@@ -35,8 +39,25 @@ impl WebCarrier {
match self { match self {
Self::Https => "https", Self::Https => "https",
Self::HttpsLanes => "https-lanes", Self::HttpsLanes => "https-lanes",
Self::Websocket => "websocket",
Self::WebsocketLanes => "websocket-lanes",
} }
} }
/// Returns whether one carrier owns independent state per logical stream.
pub(crate) const fn uses_lanes(self) -> bool {
matches!(self, Self::HttpsLanes | Self::WebsocketLanes)
}
/// Returns whether carrier messages use RFC 6455 instead of HTTP bodies.
pub(crate) const fn uses_websocket(self) -> bool {
matches!(self, Self::Websocket | Self::WebsocketLanes)
}
/// Returns whether all logical streams share one carrier state machine.
pub(crate) const fn is_multiplexed(self) -> bool {
matches!(self, Self::Https | Self::Websocket)
}
} }
/// One access user explicitly exposed through a WEB virtual host. /// One access user explicitly exposed through a WEB virtual host.
@@ -114,6 +135,18 @@ pub struct WebLimitsConfig {
/// Process-wide concurrently executing HTTP handler ceiling. /// Process-wide concurrently executing HTTP handler ceiling.
#[serde(default = "default_web_max_http_handlers")] #[serde(default = "default_web_max_http_handlers")]
pub max_http_handlers: usize, pub max_http_handlers: usize,
/// Process-wide transient WebSocket byte sub-budget inside pending bytes.
#[serde(default = "default_web_websocket_bytes_global")]
pub websocket_bytes_global: usize,
/// WebSocket usage percentage above which ordinary admission uses replacement.
#[serde(default = "default_web_websocket_admission_watermark_pct")]
pub websocket_admission_watermark_pct: u8,
/// WebSocket usage percentage that triggers pressure eviction.
#[serde(default = "default_web_websocket_eviction_watermark_pct")]
pub websocket_eviction_watermark_pct: u8,
/// Accepted HTTP connections that WebSocket upgrades must leave available.
#[serde(default = "default_web_websocket_http_connection_reserve")]
pub websocket_http_connection_reserve: usize,
/// Process-wide concurrently collected request body ceiling. /// Process-wide concurrently collected request body ceiling.
#[serde(default = "default_web_max_body_readers")] #[serde(default = "default_web_max_body_readers")]
pub max_body_readers: usize, pub max_body_readers: usize,
@@ -216,6 +249,10 @@ impl Default for WebLimitsConfig {
max_frames_per_body: default_web_max_frames_per_body(), max_frames_per_body: default_web_max_frames_per_body(),
max_http_connections: default_web_max_http_connections(), max_http_connections: default_web_max_http_connections(),
max_http_handlers: default_web_max_http_handlers(), max_http_handlers: default_web_max_http_handlers(),
websocket_bytes_global: default_web_websocket_bytes_global(),
websocket_admission_watermark_pct: default_web_websocket_admission_watermark_pct(),
websocket_eviction_watermark_pct: default_web_websocket_eviction_watermark_pct(),
websocket_http_connection_reserve: default_web_websocket_http_connection_reserve(),
max_body_readers: default_web_max_body_readers(), max_body_readers: default_web_max_body_readers(),
max_body_bytes_global: default_web_max_body_bytes_global(), max_body_bytes_global: default_web_max_body_bytes_global(),
max_sessions_global: default_web_max_sessions_global(), max_sessions_global: default_web_max_sessions_global(),
@@ -265,6 +302,15 @@ pub struct WebTimeoutsConfig {
/// Maximum wait for one empty downlink long poll. /// Maximum wait for one empty downlink long poll.
#[serde(default = "default_web_long_poll_timeout_secs")] #[serde(default = "default_web_long_poll_timeout_secs")]
pub long_poll_secs: u64, pub long_poll_secs: u64,
/// Maximum wait for one WebSocket write to complete.
#[serde(default = "default_web_websocket_write_secs")]
pub websocket_write_secs: u64,
/// Maximum wait for WebSocket queue or byte-budget progress.
#[serde(default = "default_web_websocket_backpressure_secs")]
pub websocket_backpressure_secs: u64,
/// Maximum graceful close wait for an evicted WebSocket.
#[serde(default = "default_web_websocket_eviction_secs")]
pub websocket_eviction_secs: u64,
/// Lifetime of an unused bootstrap credential and closed-token replay marker. /// Lifetime of an unused bootstrap credential and closed-token replay marker.
#[serde(default = "default_web_bootstrap_lifetime_secs")] #[serde(default = "default_web_bootstrap_lifetime_secs")]
pub bootstrap_lifetime_secs: u64, pub bootstrap_lifetime_secs: u64,
@@ -289,6 +335,9 @@ impl Default for WebTimeoutsConfig {
body_secs: default_web_body_timeout_secs(), body_secs: default_web_body_timeout_secs(),
stream_handshake_secs: default_web_stream_handshake_timeout_secs(), stream_handshake_secs: default_web_stream_handshake_timeout_secs(),
long_poll_secs: default_web_long_poll_timeout_secs(), long_poll_secs: default_web_long_poll_timeout_secs(),
websocket_write_secs: default_web_websocket_write_secs(),
websocket_backpressure_secs: default_web_websocket_backpressure_secs(),
websocket_eviction_secs: default_web_websocket_eviction_secs(),
bootstrap_lifetime_secs: default_web_bootstrap_lifetime_secs(), bootstrap_lifetime_secs: default_web_bootstrap_lifetime_secs(),
reconnect_grace_secs: default_web_reconnect_grace_secs(), reconnect_grace_secs: default_web_reconnect_grace_secs(),
http_idle_secs: default_web_http_idle_secs(), http_idle_secs: default_web_http_idle_secs(),
@@ -418,6 +467,14 @@ macro_rules! u32_default {
}; };
} }
macro_rules! u8_default {
($name:ident, $value:expr) => {
fn $name() -> u8 {
$value
}
};
}
macro_rules! u64_default { macro_rules! u64_default {
($name:ident, $value:expr) => { ($name:ident, $value:expr) => {
fn $name() -> u64 { fn $name() -> u64 {
@@ -433,6 +490,10 @@ usize_default!(default_web_carrier_batch_bytes, 2 * 1024 * 1024);
usize_default!(default_web_max_frames_per_body, 4096); usize_default!(default_web_max_frames_per_body, 4096);
usize_default!(default_web_max_http_connections, 1024); usize_default!(default_web_max_http_connections, 1024);
usize_default!(default_web_max_http_handlers, 512); usize_default!(default_web_max_http_handlers, 512);
usize_default!(default_web_websocket_bytes_global, 256 * 1024 * 1024);
u8_default!(default_web_websocket_admission_watermark_pct, 75);
u8_default!(default_web_websocket_eviction_watermark_pct, 90);
usize_default!(default_web_websocket_http_connection_reserve, 64);
usize_default!(default_web_max_body_readers, 32); usize_default!(default_web_max_body_readers, 32);
usize_default!(default_web_max_body_bytes_global, 64 * 1024 * 1024); usize_default!(default_web_max_body_bytes_global, 64 * 1024 * 1024);
usize_default!(default_web_max_sessions_global, 128); usize_default!(default_web_max_sessions_global, 128);
@@ -467,6 +528,9 @@ u64_default!(default_web_header_timeout_secs, 10);
u64_default!(default_web_body_timeout_secs, 30); u64_default!(default_web_body_timeout_secs, 30);
u64_default!(default_web_stream_handshake_timeout_secs, 10); u64_default!(default_web_stream_handshake_timeout_secs, 10);
u64_default!(default_web_long_poll_timeout_secs, 25); u64_default!(default_web_long_poll_timeout_secs, 25);
u64_default!(default_web_websocket_write_secs, 30);
u64_default!(default_web_websocket_backpressure_secs, 30);
u64_default!(default_web_websocket_eviction_secs, 1);
u64_default!(default_web_bootstrap_lifetime_secs, 120); u64_default!(default_web_bootstrap_lifetime_secs, 120);
u64_default!(default_web_reconnect_grace_secs, 120); u64_default!(default_web_reconnect_grace_secs, 120);
u64_default!(default_web_http_idle_secs, 75); u64_default!(default_web_http_idle_secs, 75);
+1 -4
View File
@@ -90,10 +90,7 @@ fn default_max_window_secs() -> u64 {
} }
/// Checks whether a hot debug policy fits restart-frozen process capacities. /// Checks whether a hot debug policy fits restart-frozen process capacities.
pub(crate) fn web_debug_fits_limits( pub(crate) fn web_debug_fits_limits(policy: &WebDebugConfig, limits: &WebLimitsConfig) -> bool {
policy: &WebDebugConfig,
limits: &WebLimitsConfig,
) -> bool {
policy.body_prefix_bytes <= limits.max_body_bytes policy.body_prefix_bytes <= limits.max_body_bytes
&& policy.body_prefix_bytes <= limits.debug_bytes_global && policy.body_prefix_bytes <= limits.debug_bytes_global
&& policy.decoy_body_prefix_bytes <= limits.debug_bytes_global && policy.decoy_body_prefix_bytes <= limits.debug_bytes_global
+10 -2
View File
@@ -246,11 +246,19 @@ impl RuntimeGeneration {
/// Registers a session only while admission remains open. /// Registers a session only while admission remains open.
pub(crate) fn spawn_session<F>(&self, future: F) -> bool pub(crate) fn spawn_session<F>(&self, future: F) -> bool
where
F: Future<Output = ()> + Send + 'static,
{
self.try_spawn_session(future).is_ok()
}
/// Registers one session or returns its unpolled future to the caller.
pub(crate) fn try_spawn_session<F>(&self, future: F) -> Result<(), F>
where where
F: Future<Output = ()> + Send + 'static, F: Future<Output = ()> + Send + 'static,
{ {
let Some(_registration) = self.session_admission.try_register() else { let Some(_registration) = self.session_admission.try_register() else {
return false; return Err(future);
}; };
let cancel = self.session_cancel.clone(); let cancel = self.session_cancel.clone();
self.sessions.spawn(async move { self.sessions.spawn(async move {
@@ -259,7 +267,7 @@ impl RuntimeGeneration {
_ = future => {} _ = future => {}
} }
}); });
true Ok(())
} }
/// Closes admission while preserving already registered sessions. /// Closes admission while preserving already registered sessions.
+2 -5
View File
@@ -318,11 +318,8 @@ pub(super) async fn run_telemt_core(
active_runtime_tx.send_replace(Some(active_runtime.clone())); active_runtime_tx.send_replace(Some(active_runtime.clone()));
runtime_tasks::mark_runtime_ready(&startup_tracker).await; runtime_tasks::mark_runtime_ready(&startup_tracker).await;
let listener_manager = listeners::ListenerManager::start( let listener_manager =
bound, listeners::ListenerManager::start(bound, active_runtime.clone(), web_trace.clone());
active_runtime.clone(),
web_trace.clone(),
);
let reload_supervisor = reload_supervisor::ReloadSupervisor::spawn( let reload_supervisor = reload_supervisor::ReloadSupervisor::spawn(
active_runtime.clone(), active_runtime.clone(),
reload_control, reload_control,
@@ -18,9 +18,7 @@ async fn make_pool() -> (Arc<MePool>, Arc<SecureRandom>) {
make_pool_with_decision(NetworkDecision::default()).await make_pool_with_decision(NetworkDecision::default()).await
} }
async fn make_pool_with_decision( async fn make_pool_with_decision(decision: NetworkDecision) -> (Arc<MePool>, Arc<SecureRandom>) {
decision: NetworkDecision,
) -> (Arc<MePool>, Arc<SecureRandom>) {
let general = GeneralConfig { let general = GeneralConfig {
me_route_no_writer_mode: MeRouteNoWriterMode::AsyncRecoveryFailfast, me_route_no_writer_mode: MeRouteNoWriterMode::AsyncRecoveryFailfast,
me_route_no_writer_wait_ms: 50, me_route_no_writer_wait_ms: 50,
+108 -10
View File
@@ -59,18 +59,20 @@ const batchLimit=__BATCH_LIMIT__,queueLimit=__QUEUE_LIMIT__,queueItemLimit=__QUE
const laneQueueLimit=Math.min(queueLimit,8388608),laneItemLimit=Math.min(queueItemLimit,1024),closedLaneLimit=4096; const laneQueueLimit=Math.min(queueLimit,8388608),laneItemLimit=Math.min(queueItemLimit,1024),closedLaneLimit=4096;
const fragment=location.hash,androidNonce=/^#android=([A-Za-z0-9_-]{43})$/.exec(fragment)?.[1]||''; const fragment=location.hash,androidNonce=/^#android=([A-Za-z0-9_-]{43})$/.exec(fragment)?.[1]||'';
history.replaceState(null,'',location.pathname); history.replaceState(null,'',location.pathname);
let initialized=false,closed=false,port=null,sessionToken='',createStarted=false; let initialized=false,closed=false,port=null,sessionToken='',createStarted=false,socket=null,socketReady=false;
let queuedBytes=0,queuedItems=0,upSequence=1,downCursor='0',upRunning=false,pollController=null; let queuedBytes=0,queuedItems=0,upSequence=1,downCursor='0',upRunning=false,pollController=null;
const pending=[],upPending=[],lanes=new Map(),closedLanes=new Set(),closedLaneOrder=[]; const pending=[],upPending=[],lanes=new Map(),closedLanes=new Set(),closedLaneOrder=[];
const status=state=>{if(port&&!closed)port.postMessage({t:'status',state})}; const status=state=>{if(port&&!closed)port.postMessage({t:'status',state})};
const pause=milliseconds=>new Promise(resolve=>setTimeout(resolve,milliseconds)); const pause=milliseconds=>new Promise(resolve=>setTimeout(resolve,milliseconds));
const socketURL=()=>relayOrigin.replace(/^https:/,'wss:')+'/api/v1/ws';
const options=(method,token,body,headers,signal,keepalive)=>({ const options=(method,token,body,headers,signal,keepalive)=>({
method,body,signal,keepalive:!!keepalive,mode:'same-origin',credentials:'omit',cache:'no-store',redirect:'error',referrerPolicy:'no-referrer', method,body,signal,keepalive:!!keepalive,mode:'same-origin',credentials:'omit',cache:'no-store',redirect:'error',referrerPolicy:'no-referrer',
headers:Object.assign(token?{Authorization:'Bearer '+token}:{},body?{'Content-Type':'application/octet-stream'}:{},headers||{}) headers:Object.assign(token?{Authorization:'Bearer '+token}:{},body?{'Content-Type':'application/octet-stream'}:{},headers||{})
}); });
function reserve(data,lane){ function reserve(data,lane){
if(!data.byteLength||data.byteLength>queueLimit-queuedBytes||queuedItems>=queueItemLimit)return false; let buffered=socket?socket.bufferedAmount:0;for(const value of lanes.values())if(value.socket)buffered+=value.socket.bufferedAmount;
if(lane&&(data.byteLength>laneQueueLimit-lane.bytes||lane.items>=laneItemLimit))return false; if(!data.byteLength||data.byteLength>queueLimit-queuedBytes-buffered||queuedItems>=queueItemLimit)return false;
if(lane&&(data.byteLength>laneQueueLimit-lane.bytes-(lane.socket?lane.socket.bufferedAmount:0)||lane.items>=laneItemLimit))return false;
queuedBytes+=data.byteLength;queuedItems++;if(lane){lane.bytes+=data.byteLength;lane.items++}return true; queuedBytes+=data.byteLength;queuedItems++;if(lane){lane.bytes+=data.byteLength;lane.items++}return true;
} }
function release(bytes,items,lane){queuedBytes-=bytes;queuedItems-=items;if(lane){lane.bytes-=bytes;lane.items-=items}} function release(bytes,items,lane){queuedBytes-=bytes;queuedItems-=items;if(lane){lane.bytes-=bytes;lane.items-=items}}
@@ -154,12 +156,17 @@ async function createSession(first){
const welcome=await response.arrayBuffer(); const welcome=await response.arrayBuffer();
port.postMessage(welcome,[welcome]);status('connected'); port.postMessage(welcome,[welcome]);status('connected');
if(carrier==='https-lanes')ensureLane(0); if(carrier==='https-lanes')ensureLane(0);
if(carrier==='websocket')openSocket();
for(const data of pending.splice(0)){release(data.byteLength,1,null);queueCarrier(data)} for(const data of pending.splice(0)){release(data.byteLength,1,null);queueCarrier(data)}
if(carrier==='https')poll();else pollLane(lanes.get(0)); if(carrier==='https')poll();else if(carrier==='https-lanes')pollLane(lanes.get(0));
}catch(error){fail()} }catch(error){fail()}
} }
function queueCarrier(data){ function queueCarrier(data){
try{if(carrier==='https')queueUp(data);else for(const value of splitFrames(data))queueLane(value)}catch(error){fail()} try{
if(carrier==='https')queueUp(data);
else if(carrier==='websocket')queueSocket(data);
else for(const value of splitFrames(data))queueLane(value);
}catch(error){fail()}
} }
function queueUp(data){if(!reserve(data,null)){fail();return}upPending.push(data);runUp()} function queueUp(data){if(!reserve(data,null)){fail();return}upPending.push(data);runUp()}
async function runUp(){ async function runUp(){
@@ -174,6 +181,31 @@ async function runUp(){
}catch(error){fail()} }catch(error){fail()}
finally{upRunning=false;if(!closed&&sessionToken&&upPending.length)runUp()} finally{upRunning=false;if(!closed&&sessionToken&&upPending.length)runUp()}
} }
function openSocket(){
if(socket||closed)return;socket=new WebSocket(socketURL(),'tproxy-v1.'+sessionToken);socket.binaryType='arraybuffer';
socket.onopen=()=>{if(closed)return;socketReady=true;status('connected');runSocketUp()};
socket.onmessage=event=>{
if(closed||!(event.data instanceof ArrayBuffer)){fail();return}
try{const bound=frameBound(event.data,4096,batchLimit);if(bound.bytes!==event.data.byteLength)throw new Error('invalid frame batch')}catch(error){fail();return}
port.postMessage({t:'traffic',up:0,down:event.data.byteLength});port.postMessage(event.data,[event.data]);status('connected');
};
socket.onerror=()=>{};socket.onclose=()=>{socketReady=false;if(!closed)fail()};
}
function queueSocket(data){if(!reserve(data,null)){fail();return}upPending.push(data);runSocketUp()}
async function waitSocket(next,size,limit){
while(!closed&&next.readyState===WebSocket.OPEN&&next.bufferedAmount>limit-size)await pause(10);
if(closed||next.readyState!==WebSocket.OPEN)throw new Error('websocket closed');
}
async function runSocketUp(){
if(upRunning||!socketReady)return;upRunning=true;
try{
while(!closed&&socketReady&&upPending.length){
const batch=joinPending(upPending,null);await waitSocket(socket,batch.total,queueLimit);socket.send(batch.body);
release(batch.total,batch.count,null);port.postMessage({t:'traffic',up:batch.total,down:0});
}
}catch(error){if(!closed)fail()}
finally{upRunning=false;if(!closed&&socketReady&&upPending.length)runSocketUp()}
}
async function poll(){ async function poll(){
while(!closed&&sessionToken){ while(!closed&&sessionToken){
try{ try{
@@ -190,7 +222,7 @@ async function poll(){
} }
function ensureLane(id){ function ensureLane(id){
let lane=lanes.get(id); let lane=lanes.get(id);
if(!lane){lane={id,sequence:1,cursor:'0',pending:[],bytes:0,items:0,running:false,polling:false,controller:null};lanes.set(id,lane)} if(!lane){lane={id,sequence:1,cursor:'0',pending:[],bytes:0,items:0,running:false,polling:false,controller:null,socket:null,ready:false,remoteClosed:false};lanes.set(id,lane)}
return lane; return lane;
} }
function rememberLaneClosed(id){ function rememberLaneClosed(id){
@@ -198,10 +230,13 @@ function rememberLaneClosed(id){
if(closedLaneOrder.length===closedLaneLimit)closedLanes.delete(closedLaneOrder.shift()); if(closedLaneOrder.length===closedLaneLimit)closedLanes.delete(closedLaneOrder.shift());
closedLanes.add(id);closedLaneOrder.push(id); closedLanes.add(id);closedLaneOrder.push(id);
} }
function finishLane(lane){ function closeFrame(id){const value=new Uint8Array(8);value[0]=3;value[1]=(id>>>16)&255;value[2]=(id>>>8)&255;value[3]=id&255;return value.buffer}
function finishLane(lane,notifyClient){
if(lanes.get(lane.id)!==lane)return; if(lanes.get(lane.id)!==lane)return;
if(lane.socket&&lane.socket.readyState<WebSocket.CLOSING)lane.socket.close();
if(lane.bytes||lane.items)release(lane.bytes,lane.items,lane); if(lane.bytes||lane.items)release(lane.bytes,lane.items,lane);
lane.pending.length=0;lanes.delete(lane.id);rememberLaneClosed(lane.id); lane.pending.length=0;lanes.delete(lane.id);rememberLaneClosed(lane.id);
if(notifyClient&&!lane.remoteClosed&&port){const frame=closeFrame(lane.id);port.postMessage(frame,[frame])}
} }
function queueLane(value){ function queueLane(value){
let lane=lanes.get(value.id); let lane=lanes.get(value.id);
@@ -210,7 +245,29 @@ function queueLane(value){
if(!lane&&value.type!==1)throw new Error('lane did not begin with OPEN'); if(!lane&&value.type!==1)throw new Error('lane did not begin with OPEN');
lane=lane||ensureLane(value.id); lane=lane||ensureLane(value.id);
if(!reserve(value.data,lane)){fail();return} if(!reserve(value.data,lane)){fail();return}
lane.pending.push(value.data);runLaneUp(lane); lane.pending.push(value.data);
if(carrier==='websocket-lanes'){openLaneSocket(lane);runLaneSocketUp(lane)}else runLaneUp(lane);
}
function openLaneSocket(lane){
if(lane.socket||closed)return;lane.socket=new WebSocket(socketURL(),'tproxy-lane-v1.'+sessionToken+'.'+String(lane.id));lane.socket.binaryType='arraybuffer';
lane.socket.onopen=()=>{if(closed||lanes.get(lane.id)!==lane)return;lane.ready=true;status('connected');runLaneSocketUp(lane)};
lane.socket.onmessage=event=>{
if(closed||lanes.get(lane.id)!==lane||!(event.data instanceof ArrayBuffer)){finishLane(lane,true);return}
let values;try{values=splitFrames(event.data);for(const value of values)if(value.id!==lane.id)throw new Error('cross-lane frame')}catch(error){finishLane(lane,true);return}
if(values.some(value=>value.type===3))lane.remoteClosed=true;
port.postMessage({t:'traffic',up:0,down:event.data.byteLength});port.postMessage(event.data,[event.data]);status('connected');
};
lane.socket.onerror=()=>{};lane.socket.onclose=()=>{lane.ready=false;lane.socket=null;if(!closed)finishLane(lane,true)};
}
async function runLaneSocketUp(lane){
if(lane.running||!lane.ready)return;lane.running=true;
try{
while(!closed&&lane.ready&&lanes.get(lane.id)===lane&&lane.pending.length){
const batch=joinPending(lane.pending,lane);await waitSocket(lane.socket,batch.total,laneQueueLimit);lane.socket.send(batch.body);
release(batch.total,batch.count,lane);port.postMessage({t:'traffic',up:batch.total,down:0});
}
}catch(error){if(!closed)finishLane(lane,true)}
finally{lane.running=false;if(!closed&&lane.ready&&lane.pending.length)runLaneSocketUp(lane)}
} }
async function runLaneUp(lane){ async function runLaneUp(lane){
if(lane.running)return;lane.running=true; if(lane.running)return;lane.running=true;
@@ -232,7 +289,7 @@ async function pollLane(lane){
const controller=new AbortController(),laneID=String(lane.id);lane.controller=controller; const controller=new AbortController(),laneID=String(lane.id);lane.controller=controller;
const response=await request('/api/v1/down',()=>options('POST',sessionToken,null,{'X-Down-Cursor':lane.cursor,'X-Lane-ID':laneID},controller.signal)); const response=await request('/api/v1/down',()=>options('POST',sessionToken,null,{'X-Down-Cursor':lane.cursor,'X-Lane-ID':laneID},controller.signal));
if(response.status===204){ if(response.status===204){
if(response.headers.get('X-Lane-Closed')==='1'){finishLane(lane);return} if(response.headers.get('X-Lane-Closed')==='1'){finishLane(lane,false);return}
status('connected');continue; status('connected');continue;
} }
if(response.status!==200)throw new Error('lane downlink rejected'); if(response.status!==200)throw new Error('lane downlink rejected');
@@ -250,7 +307,7 @@ function deleteSession(){
} }
function close(notifyServer){ function close(notifyServer){
if(closed)return;closed=true;if(pollController)pollController.abort(); if(closed)return;closed=true;if(pollController)pollController.abort();
for(const lane of lanes.values())if(lane.controller)lane.controller.abort(); if(socket)socket.close();for(const lane of lanes.values()){if(lane.controller)lane.controller.abort();if(lane.socket)lane.socket.close()}
if(notifyServer)deleteSession();pending.length=0;upPending.length=0; if(notifyServer)deleteSession();pending.length=0;upPending.length=0;
for(const lane of lanes.values())lane.pending.length=0;lanes.clear();queuedBytes=0;queuedItems=0;if(port)port.close(); for(const lane of lanes.values())lane.pending.length=0;lanes.clear();queuedBytes=0;queuedItems=0;if(port)port.close();
} }
@@ -357,4 +414,45 @@ mod tests {
"the iOS native carrier cannot parse a comma-declared single-quoted bootstrap" "the iOS native carrier cannot parse a comma-declared single-quoted bootstrap"
); );
} }
#[test]
fn rendered_page_advertises_exact_websocket_carriers() {
let websocket = render(
"proxy.example.com",
"CCCCCCCCCCCCCCCCCCCCCCCCCCCCCCCCCCCCCCCCCCC",
2 * 1024 * 1024,
32 * 1024 * 1024,
16 * 1024,
WebCarrier::Websocket,
&SecureRandom::new(),
);
assert!(websocket.body.contains("carrier='websocket'"));
assert!(
websocket
.body
.contains("new WebSocket(socketURL(),'tproxy-v1.'+sessionToken)")
);
assert!(
websocket
.content_security_policy
.contains("connect-src 'self' wss://proxy.example.com")
);
let lanes = render(
"proxy.example.com",
"DDDDDDDDDDDDDDDDDDDDDDDDDDDDDDDDDDDDDDDDDDD",
2 * 1024 * 1024,
32 * 1024 * 1024,
16 * 1024,
WebCarrier::WebsocketLanes,
&SecureRandom::new(),
);
assert!(lanes.body.contains("carrier='websocket-lanes'"));
assert!(
lanes
.body
.contains("'tproxy-lane-v1.'+sessionToken+'.'+String(lane.id)")
);
assert!(!lanes.body.contains("__"));
}
} }
+29 -9
View File
@@ -5,8 +5,8 @@ use std::sync::Arc;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use bytes::Bytes; use bytes::Bytes;
use http_body_util::combinators::UnsyncBoxBody;
use http_body_util::BodyExt; use http_body_util::BodyExt;
use http_body_util::combinators::UnsyncBoxBody;
use hyper::header::{self, HeaderName, HeaderValue}; use hyper::header::{self, HeaderName, HeaderValue};
use hyper::server::conn::http1; use hyper::server::conn::http1;
use hyper::service::service_fn; use hyper::service::service_fn;
@@ -34,13 +34,16 @@ mod down;
mod request; mod request;
// Carrier response construction and lane-header helpers are shared by handlers. // Carrier response construction and lane-header helpers are shared by handlers.
mod response; mod response;
// RFC 6455 upgrade validation and carrier drivers remain isolated from HTTP routing.
#[cfg(test)] #[cfg(test)]
mod tests; mod tests;
mod websocket;
// Enabled-debug integration coverage remains separate from carrier behavior tests. // Enabled-debug integration coverage remains separate from carrier behavior tests.
#[cfg(test)] #[cfg(test)]
#[path = "http/trace_tests.rs"] #[path = "http/trace_tests.rs"]
mod trace_tests; mod trace_tests;
use crate::web::trace::{HttpTraceExchange, TraceDirection, TraceLifecycleEvent, TraceRoute};
use activity::{ActivityBody, RequestActivity}; use activity::{ActivityBody, RequestActivity};
use body::{CollectBodyError, CollectedBody, RequestBody, collect_body}; use body::{CollectBodyError, CollectedBody, RequestBody, collect_body};
use decoy::serve_decoy; use decoy::serve_decoy;
@@ -53,9 +56,6 @@ use response::{
bad_gateway, carrier_empty, carrier_headers, carrier_lane, full_response, generic_not_found, bad_gateway, carrier_empty, carrier_headers, carrier_lane, full_response, generic_not_found,
insert_header, service_unavailable, insert_header, service_unavailable,
}; };
use crate::web::trace::{
HttpTraceExchange, TraceDirection, TraceLifecycleEvent, TraceRoute,
};
type BoxError = Box<dyn Error + Send + Sync>; type BoxError = Box<dyn Error + Send + Sync>;
type HttpBody = UnsyncBoxBody<Bytes, BoxError>; type HttpBody = UnsyncBoxBody<Bytes, BoxError>;
@@ -63,6 +63,7 @@ type HttpResponse = Response<HttpBody>;
const CREATE_BODY_LIMIT: usize = 64; const CREATE_BODY_LIMIT: usize = 64;
const TRANSPORT_PATHS: [&str; 3] = ["/api/v1/session", "/api/v1/up", "/api/v1/down"]; const TRANSPORT_PATHS: [&str; 3] = ["/api/v1/session", "/api/v1/up", "/api/v1/down"];
const WEBSOCKET_PATH: &str = "/api/v1/ws";
/// Serves one bounded HTTP/1.1 connection accepted from an external TLS terminator. /// Serves one bounded HTTP/1.1 connection accepted from an external TLS terminator.
pub(crate) async fn serve_connection( pub(crate) async fn serve_connection(
@@ -107,8 +108,8 @@ pub(crate) async fn serve_connection(
if let Some(trace) = &trace { if let Some(trace) = &trace {
trace.response_ready(&response); trace.response_ready(&response);
} }
let response = response let response =
.map(|body| ActivityBody::new(body, activity, trace).boxed_unsync()); response.map(|body| ActivityBody::new(body, activity, trace).boxed_unsync());
Ok::<_, Infallible>(response) Ok::<_, Infallible>(response)
} }
}); });
@@ -117,7 +118,11 @@ pub(crate) async fn serve_connection(
.header_read_timeout(header_timeout) .header_read_timeout(header_timeout)
.max_buf_size(max_header_bytes) .max_buf_size(max_header_bytes)
.keep_alive(true) .keep_alive(true)
.serve_connection(TokioIo::new(stream), service); .serve_connection(
TokioIo::new(websocket::ConnectionIo::new(stream, connection_permit)),
service,
)
.with_upgrades();
tokio::pin!(connection); tokio::pin!(connection);
let mut idle_check = tokio::time::interval((idle_timeout / 2).max(Duration::from_secs(1))); let mut idle_check = tokio::time::interval((idle_timeout / 2).max(Duration::from_secs(1)));
idle_check.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); idle_check.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
@@ -135,7 +140,6 @@ pub(crate) async fn serve_connection(
} }
} }
} }
drop(connection_permit);
} }
async fn handle_request( async fn handle_request(
@@ -163,6 +167,17 @@ async fn handle_request(
return generic_not_found(); return generic_not_found();
}; };
let path = request.uri().path(); let path = request.uri().path();
if path == WEBSOCKET_PATH {
return websocket::handle(
request,
peer,
client_ip_source,
trusted_proxy_cidrs,
runtime,
vhost,
)
.await;
}
if TRANSPORT_PATHS.contains(&path) { if TRANSPORT_PATHS.contains(&path) {
return handle_api( return handle_api(
request, request,
@@ -394,7 +409,9 @@ async fn handle_session(
); );
response response
} }
Err(error @ (ManagerError::Limit | ManagerError::Backpressure | ManagerError::Concurrent)) => { Err(
error @ (ManagerError::Limit | ManagerError::Backpressure | ManagerError::Concurrent),
) => {
runtime.trace().record_profile_lifecycle( runtime.trace().record_profile_lifecycle(
client_ip, client_ip,
Some(trace_session_id), Some(trace_session_id),
@@ -435,6 +452,9 @@ async fn handle_up(
let Ok(session) = runtime.get_session(token_hash, &vhost.host) else { let Ok(session) = runtime.get_session(token_hash, &vhost.host) else {
return serve_decoy(request, vhost, true, &runtime).await; 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) { if let Some(trace) = request_trace(&request) {
trace.set_route(TraceRoute::Uplink); trace.set_route(TraceRoute::Uplink);
trace.bind_identity(session.trace_identity()); trace.bind_identity(session.trace_identity());
+10 -3
View File
@@ -9,9 +9,7 @@ use hyper::Request;
use hyper::body::{Body, Frame, Incoming, SizeHint}; use hyper::body::{Body, Frame, Incoming, SizeHint};
use crate::web::manager::WebProcessRuntime; use crate::web::manager::WebProcessRuntime;
use crate::web::trace::{ use crate::web::trace::{HttpTraceExchange, TraceBodyState, TraceDirection};
HttpTraceExchange, TraceBodyState, TraceDirection,
};
/// Incoming request body wrapper that observes frames without changing streaming semantics. /// Incoming request body wrapper that observes frames without changing streaming semantics.
pub(super) struct RequestBody { 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) { fn finish(&mut self, state: TraceBodyState) {
if self.terminal { if self.terminal {
return; return;
+1
View File
@@ -236,6 +236,7 @@ fn sanitize_transport_request<B>(request: &mut Request<B>) {
header::CONTENT_TYPE, header::CONTENT_TYPE,
header::UPGRADE, header::UPGRADE,
HeaderName::from_static("sec-websocket-key"), HeaderName::from_static("sec-websocket-key"),
HeaderName::from_static("sec-websocket-extensions"),
HeaderName::from_static("sec-websocket-protocol"), HeaderName::from_static("sec-websocket-protocol"),
HeaderName::from_static("sec-websocket-version"), HeaderName::from_static("sec-websocket-version"),
HeaderName::from_static("x-down-cursor"), 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 { let Ok(session) = runtime.get_session(token_hash, &vhost.host) else {
return serve_decoy(request, vhost, true, &runtime).await; 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) { if let Some(trace) = request_trace(&request) {
trace.set_route(TraceRoute::Downlink); trace.set_route(TraceRoute::Downlink);
trace.bind_identity(session.trace_identity()); trace.bind_identity(session.trace_identity());
+2 -4
View File
@@ -9,16 +9,14 @@ use crate::config::WebCarrier;
use crate::web::frame; use crate::web::frame;
/// Validates and resolves the optional carrier lane header. /// Validates and resolves the optional carrier lane header.
pub(super) fn carrier_lane<B>( pub(super) fn carrier_lane<B>(request: &Request<B>, carrier: WebCarrier) -> Option<Option<u32>> {
request: &Request<B>,
carrier: WebCarrier,
) -> Option<Option<u32>> {
match carrier { match carrier {
WebCarrier::Https => (!request.headers().contains_key("x-lane-id")).then_some(None), WebCarrier::Https => (!request.headers().contains_key("x-lane-id")).then_some(None),
WebCarrier::HttpsLanes => canonical_u64_header(request, "x-lane-id") WebCarrier::HttpsLanes => canonical_u64_header(request, "x-lane-id")
.and_then(|value| u32::try_from(value).ok()) .and_then(|value| u32::try_from(value).ok())
.filter(|value| *value <= frame::MAX_STREAM_ID) .filter(|value| *value <= frame::MAX_STREAM_ID)
.map(Some), .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_response = request(&listener, &runtime, bridge_request()).await;
let (fallback_headers, fallback_body) = split_response(&fallback_response); let (fallback_headers, fallback_body) = split_response(&fallback_response);
assert!(fallback_headers.starts_with(b"HTTP/1.1 200")); 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>"); assert_eq!(fallback_body, b"<!doctype html><title>decoy</title>");
runtime.shutdown().await; 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;
}
+83 -73
View File
@@ -1,6 +1,7 @@
use std::future::Future; use std::future::Future;
use std::net::IpAddr;
use std::sync::Arc; use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration; use std::time::Duration;
use arc_swap::ArcSwap; use arc_swap::ArcSwap;
@@ -21,7 +22,14 @@ mod credentials;
mod admission; mod admission;
// Shutdown and expiry work remain outside request-path coordination. // Shutdown and expiry work remain outside request-path coordination.
mod lifecycle; mod lifecycle;
use state::{ManagerState, control_item_reserve}; // Queue and WebSocket allocations share one process-owned data-plane budget.
mod budget;
// WebSocket admission, replacement, and liveness are process-scoped.
mod websocket;
pub(crate) use budget::WebSocketBudgetLease;
use budget::{WebDataBudget, WebSocketBudgetClass};
use state::ManagerState;
pub(crate) use websocket::{WebSocketConnection, WebSocketKind};
const TOKEN_BYTES: usize = 32; const TOKEN_BYTES: usize = 32;
const CLEANUP_INTERVAL: Duration = Duration::from_secs(1); const CLEANUP_INTERVAL: Duration = Duration::from_secs(1);
@@ -76,8 +84,12 @@ pub(crate) struct WebProcessRuntime {
body_readers: Arc<Semaphore>, body_readers: Arc<Semaphore>,
body_bytes: Arc<Semaphore>, body_bytes: Arc<Semaphore>,
stream_handshakes: Arc<Semaphore>, stream_handshakes: Arc<Semaphore>,
budget_notify: Arc<Notify>, websocket_connections: Arc<Semaphore>,
budget_saturated: AtomicBool, websockets: Mutex<websocket::WebSocketRegistry>,
websocket_next_id: AtomicU64,
websocket_clock: std::time::Instant,
websocket_notify: Arc<Notify>,
data_budget: Arc<WebDataBudget>,
shutdown: CancellationToken, shutdown: CancellationToken,
tasks: TaskTracker, tasks: TaskTracker,
sessions_created: AtomicU64, sessions_created: AtomicU64,
@@ -104,6 +116,9 @@ impl WebProcessRuntime {
trace: Arc<WebTraceStore>, trace: Arc<WebTraceStore>,
) -> Arc<Self> { ) -> Arc<Self> {
let limits = active_runtime.load().config().web.limits.clone(); let limits = active_runtime.load().config().web.limits.clone();
let websocket_connections = limits
.max_http_connections
.saturating_sub(limits.websocket_http_connection_reserve);
let runtime = Arc::new(Self { let runtime = Arc::new(Self {
active_runtime, active_runtime,
trace, trace,
@@ -113,10 +128,14 @@ impl WebProcessRuntime {
body_readers: Arc::new(Semaphore::new(limits.max_body_readers)), body_readers: Arc::new(Semaphore::new(limits.max_body_readers)),
body_bytes: Arc::new(Semaphore::new(limits.max_body_bytes_global)), body_bytes: Arc::new(Semaphore::new(limits.max_body_bytes_global)),
stream_handshakes: Arc::new(Semaphore::new(limits.max_stream_handshakes)), stream_handshakes: Arc::new(Semaphore::new(limits.max_stream_handshakes)),
websocket_connections: Arc::new(Semaphore::new(websocket_connections)),
websockets: Mutex::new(websocket::WebSocketRegistry::default()),
websocket_next_id: AtomicU64::new(1),
websocket_clock: std::time::Instant::now(),
websocket_notify: Arc::new(Notify::new()),
data_budget: WebDataBudget::new(limits.clone()),
limits, limits,
state: Mutex::new(ManagerState::default()), state: Mutex::new(ManagerState::default()),
budget_notify: Arc::new(Notify::new()),
budget_saturated: AtomicBool::new(false),
shutdown: CancellationToken::new(), shutdown: CancellationToken::new(),
tasks: TaskTracker::new(), tasks: TaskTracker::new(),
sessions_created: AtomicU64::new(0), sessions_created: AtomicU64::new(0),
@@ -235,89 +254,80 @@ impl WebProcessRuntime {
/// Reserves bounded process-wide queue capacity for data or control traffic. /// Reserves bounded process-wide queue capacity for data or control traffic.
pub(crate) fn try_reserve_pending( pub(crate) fn try_reserve_pending(
&self, &self,
owner: ProfileKey,
bytes: usize, bytes: usize,
items: usize, items: usize,
control: bool, control: bool,
downlink: bool, downlink: bool,
) -> bool { ) -> bool {
let mut state = self.state.lock(); if !self
let data_byte_limit = self .data_budget
.limits .try_reserve_queue(owner, bytes, items, control, downlink)
.pending_bytes_global {
.saturating_sub(self.limits.control_bytes_global);
let control_item_reserve = control_item_reserve(&self.limits);
let data_item_limit = self
.limits
.pending_items_global
.saturating_sub(control_item_reserve);
if state.closed {
return false;
}
let fits = if control {
bytes <= self.limits.control_bytes_global
&& items <= control_item_reserve
&& state.pending_bytes <= self.limits.pending_bytes_global.saturating_sub(bytes)
&& state.pending_items <= self.limits.pending_items_global.saturating_sub(items)
&& state.pending_control_bytes
<= self.limits.control_bytes_global.saturating_sub(bytes)
&& state.pending_control_items <= control_item_reserve.saturating_sub(items)
} else {
let data_bytes = state
.pending_bytes
.saturating_sub(state.pending_control_bytes);
let data_items = state
.pending_items
.saturating_sub(state.pending_control_items);
let (byte_limit, item_limit) = if downlink {
let uplink_bytes = self.limits.max_body_bytes.saturating_add(
self.limits
.max_frames_per_body
.saturating_mul(crate::web::session::QUEUE_ITEM_COST),
);
(
data_byte_limit.saturating_sub(uplink_bytes),
data_item_limit.saturating_sub(self.limits.max_frames_per_body),
)
} else {
(data_byte_limit, data_item_limit)
};
bytes <= byte_limit
&& items <= item_limit
&& data_bytes <= byte_limit - bytes
&& data_items <= item_limit - items
};
if !fits {
self.budget_saturated.store(true, Ordering::Release);
self.record_limit_hit(); self.record_limit_hit();
return false; return false;
} }
state.pending_bytes += bytes;
state.pending_items += items;
if control {
state.pending_control_bytes += bytes;
state.pending_control_items += items;
}
true true
} }
/// Releases process-wide queue capacity and wakes blocked relay writers. /// Releases process-wide queue capacity and wakes blocked relay writers.
pub(crate) fn release_pending(&self, bytes: usize, items: usize, control: bool) { pub(crate) fn release_pending(
let mut state = self.state.lock(); &self,
state.pending_bytes = state.pending_bytes.saturating_sub(bytes); owner: ProfileKey,
state.pending_items = state.pending_items.saturating_sub(items); bytes: usize,
if control { items: usize,
state.pending_control_bytes = state.pending_control_bytes.saturating_sub(bytes); control: bool,
state.pending_control_items = state.pending_control_items.saturating_sub(items); ) {
} self.data_budget.release_queue(owner, bytes, items, control);
drop(state);
if self.budget_saturated.swap(false, Ordering::AcqRel) {
self.budget_notify.notify_waiters();
}
} }
/// Returns the shared notification source for global queue capacity changes. /// Returns the shared notification source for global queue capacity changes.
pub(crate) fn budget_notify(&self) -> Arc<Notify> { pub(crate) fn budget_notify(&self) -> Arc<Notify> {
Arc::clone(&self.budget_notify) self.data_budget.notify()
}
/// Reserves fixed WebSocket driver memory below the admission watermark.
pub(crate) fn try_websocket_base_budget(
&self,
owner: ProfileKey,
bytes: usize,
) -> Option<WebSocketBudgetLease> {
self.data_budget
.try_reserve_websocket(owner, bytes, WebSocketBudgetClass::Base)
}
/// Reserves one transient WebSocket message below the eviction watermark.
pub(crate) fn try_websocket_data_budget(
&self,
owner: ProfileKey,
bytes: usize,
) -> Option<WebSocketBudgetLease> {
self.data_budget
.try_reserve_websocket(owner, bytes, WebSocketBudgetClass::Data)
}
/// Admits one WebSocket with owner-first bounded replacement.
pub(crate) async fn admit_websocket(
self: &Arc<Self>,
owner: ProfileKey,
session_id: u64,
client_ip: IpAddr,
kind: WebSocketKind,
base_bytes: usize,
liveness_interval: Duration,
eviction_timeout: Duration,
) -> Result<WebSocketConnection, ManagerError> {
websocket::admit(
self,
owner,
session_id,
client_ip,
kind,
base_bytes,
liveness_interval,
eviction_timeout,
)
.await
} }
/// Accounts one successfully committed carrier uplink body. /// Accounts one successfully committed carrier uplink body.
+16 -6
View File
@@ -84,7 +84,7 @@ mod tests {
async fn global_downlink_budget_preserves_one_maximum_uplink_batch() { async fn global_downlink_budget_preserves_one_maximum_uplink_batch() {
let generation = test_runtime_generation(1, ProxyConfig::default()); let generation = test_runtime_generation(1, ProxyConfig::default());
let runtime = WebProcessRuntime::start(Arc::new(ArcSwap::from(generation))); let runtime = WebProcessRuntime::start(Arc::new(ArcSwap::from(generation)));
let control_items = super::super::state::control_item_reserve(&runtime.limits); let control_items = super::super::budget::control_item_reserve(&runtime.limits);
let data_bytes = runtime let data_bytes = runtime
.limits .limits
.pending_bytes_global .pending_bytes_global
@@ -97,20 +97,30 @@ mod tests {
.limits .limits
.max_body_bytes .max_body_bytes
.saturating_add(runtime.limits.max_frames_per_body * QUEUE_ITEM_COST); .saturating_add(runtime.limits.max_frames_per_body * QUEUE_ITEM_COST);
let downlink_bytes = data_bytes - uplink_bytes; let websocket_bytes = runtime.limits.carrier_batch_bytes;
let downlink_bytes = data_bytes - uplink_bytes - websocket_bytes;
let downlink_items = data_items - runtime.limits.max_frames_per_body; let downlink_items = data_items - runtime.limits.max_frames_per_body;
assert!(runtime.try_reserve_pending(downlink_bytes, downlink_items, false, true,)); assert!(runtime.try_reserve_pending([0; 32], downlink_bytes, downlink_items, false, true,));
assert!(runtime.try_reserve_pending( assert!(runtime.try_reserve_pending(
[0; 32],
uplink_bytes, uplink_bytes,
runtime.limits.max_frames_per_body, runtime.limits.max_frames_per_body,
false, false,
false, false,
)); ));
assert!(!runtime.try_reserve_pending(1, 1, false, true)); let websocket = runtime.try_websocket_data_budget([0; 32], websocket_bytes);
assert!(websocket.is_some());
assert!(!runtime.try_reserve_pending([0; 32], 1, 1, false, true));
runtime.release_pending(downlink_bytes, downlink_items, false); drop(websocket);
runtime.release_pending(uplink_bytes, runtime.limits.max_frames_per_body, false); runtime.release_pending([0; 32], downlink_bytes, downlink_items, false);
runtime.release_pending(
[0; 32],
uplink_bytes,
runtime.limits.max_frames_per_body,
false,
);
runtime.shutdown().await; runtime.shutdown().await;
} }
} }
+366
View File
@@ -0,0 +1,366 @@
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use parking_lot::Mutex;
use tokio::sync::Notify;
use super::ProfileKey;
use crate::config::WebLimitsConfig;
use crate::web::session::QUEUE_ITEM_COST;
/// WebSocket allocation class with a distinct pressure watermark.
#[derive(Clone, Copy)]
pub(super) enum WebSocketBudgetClass {
/// Long-lived codec and driver memory acquired before an upgrade commits.
Base,
/// One bounded inbound message or outbound write staging allocation.
Data,
}
#[derive(Default)]
struct BudgetState {
queue_bytes: usize,
queue_items: usize,
queue_control_bytes: usize,
queue_control_items: usize,
websocket_bytes: usize,
owner_bytes: HashMap<ProfileKey, usize>,
high_water_bytes: usize,
closed: bool,
}
/// Process-owned byte and item governor shared by queues and WebSocket I/O.
pub(super) struct WebDataBudget {
limits: WebLimitsConfig,
state: Mutex<BudgetState>,
notify: Arc<Notify>,
pressured: AtomicBool,
}
/// One exact WebSocket allocation released on every cancellation path.
pub(crate) struct WebSocketBudgetLease {
budget: Arc<WebDataBudget>,
owner: ProfileKey,
bytes: usize,
}
/// Lock-free diagnostic snapshot of one short locked budget state.
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct WebDataBudgetSnapshot {
/// Total queue bytes currently retained.
pub(crate) queue_bytes: usize,
/// Total queue items currently retained.
pub(crate) queue_items: usize,
/// Total WebSocket bytes currently retained.
pub(crate) websocket_bytes: usize,
/// Largest combined byte usage observed since process start.
pub(crate) high_water_bytes: usize,
}
impl WebDataBudget {
pub(super) fn new(limits: WebLimitsConfig) -> Arc<Self> {
Arc::new(Self {
limits,
state: Mutex::new(BudgetState::default()),
notify: Arc::new(Notify::new()),
pressured: AtomicBool::new(false),
})
}
pub(super) fn try_reserve_queue(
&self,
owner: ProfileKey,
bytes: usize,
items: usize,
control: bool,
downlink: bool,
) -> bool {
let mut state = self.state.lock();
if state.closed {
return false;
}
let control_item_reserve = control_item_reserve(&self.limits);
let data_byte_limit = self
.limits
.pending_bytes_global
.saturating_sub(self.limits.control_bytes_global);
let data_item_limit = self
.limits
.pending_items_global
.saturating_sub(control_item_reserve);
let (fits, websocket_byte_pressure) = if control {
let byte_pressure = state.websocket_bytes != 0
&& state.queue_bytes.saturating_add(state.websocket_bytes)
> self.limits.pending_bytes_global.saturating_sub(bytes);
let fits = bytes <= self.limits.control_bytes_global
&& items <= control_item_reserve
&& state.queue_bytes.saturating_add(state.websocket_bytes)
<= self.limits.pending_bytes_global.saturating_sub(bytes)
&& state.queue_items <= self.limits.pending_items_global.saturating_sub(items)
&& state.queue_control_bytes
<= self.limits.control_bytes_global.saturating_sub(bytes)
&& state.queue_control_items <= control_item_reserve.saturating_sub(items);
(fits, byte_pressure)
} else {
let queue_data_bytes = state.queue_bytes.saturating_sub(state.queue_control_bytes);
let queue_data_items = state.queue_items.saturating_sub(state.queue_control_items);
let (byte_limit, item_limit) = if downlink {
let uplink_bytes = self.limits.max_body_bytes.saturating_add(
self.limits
.max_frames_per_body
.saturating_mul(QUEUE_ITEM_COST),
);
(
data_byte_limit
.saturating_sub(uplink_bytes)
.saturating_sub(self.limits.carrier_batch_bytes),
data_item_limit.saturating_sub(self.limits.max_frames_per_body),
)
} else {
(data_byte_limit, data_item_limit)
};
let byte_pressure = state.websocket_bytes != 0
&& queue_data_bytes.saturating_add(state.websocket_bytes)
> byte_limit.saturating_sub(bytes);
let fits = bytes <= byte_limit
&& items <= item_limit
&& queue_data_bytes.saturating_add(state.websocket_bytes)
<= byte_limit.saturating_sub(bytes)
&& queue_data_items <= item_limit.saturating_sub(items);
(fits, byte_pressure)
};
if !fits {
if websocket_byte_pressure {
self.pressured.store(true, Ordering::Release);
}
return false;
}
state.queue_bytes += bytes;
state.queue_items += items;
if control {
state.queue_control_bytes += bytes;
state.queue_control_items += items;
}
add_owner(&mut state.owner_bytes, owner, bytes);
update_high_water(&mut state);
true
}
pub(super) fn release_queue(
&self,
owner: ProfileKey,
bytes: usize,
items: usize,
control: bool,
) {
let mut state = self.state.lock();
state.queue_bytes = state.queue_bytes.saturating_sub(bytes);
state.queue_items = state.queue_items.saturating_sub(items);
if control {
state.queue_control_bytes = state.queue_control_bytes.saturating_sub(bytes);
state.queue_control_items = state.queue_control_items.saturating_sub(items);
}
remove_owner(&mut state.owner_bytes, owner, bytes);
drop(state);
self.pressured.store(false, Ordering::Release);
self.notify.notify_waiters();
}
pub(super) fn try_reserve_websocket(
self: &Arc<Self>,
owner: ProfileKey,
bytes: usize,
class: WebSocketBudgetClass,
) -> Option<WebSocketBudgetLease> {
let mut state = self.state.lock();
if state.closed || bytes == 0 {
return None;
}
let websocket_limit = match class {
WebSocketBudgetClass::Base => watermark(
self.limits.websocket_bytes_global,
self.limits.websocket_admission_watermark_pct,
),
WebSocketBudgetClass::Data => watermark(
self.limits.websocket_bytes_global,
self.limits.websocket_eviction_watermark_pct,
),
};
let data_byte_limit = self
.limits
.pending_bytes_global
.saturating_sub(self.limits.control_bytes_global);
let queue_data_bytes = state.queue_bytes.saturating_sub(state.queue_control_bytes);
if state.websocket_bytes > websocket_limit.saturating_sub(bytes)
|| queue_data_bytes.saturating_add(state.websocket_bytes)
> data_byte_limit.saturating_sub(bytes)
{
self.pressured.store(true, Ordering::Release);
return None;
}
state.websocket_bytes += bytes;
add_owner(&mut state.owner_bytes, owner, bytes);
update_high_water(&mut state);
Some(WebSocketBudgetLease {
budget: Arc::clone(self),
owner,
bytes,
})
}
pub(super) fn notify(&self) -> Arc<Notify> {
Arc::clone(&self.notify)
}
pub(super) fn take_pressure(&self) -> bool {
self.pressured.swap(false, Ordering::AcqRel)
}
pub(super) fn owner_usage(&self, owner: ProfileKey) -> usize {
self.state
.lock()
.owner_bytes
.get(&owner)
.copied()
.unwrap_or(0)
}
pub(super) fn fair_share(&self, additional_owner: Option<ProfileKey>) -> usize {
let state = self.state.lock();
let mut owners = state.owner_bytes.len();
if additional_owner.is_some_and(|owner| !state.owner_bytes.contains_key(&owner)) {
owners += 1;
}
let admission = watermark(
self.limits.websocket_bytes_global,
self.limits.websocket_admission_watermark_pct,
);
admission / owners.max(1)
}
pub(super) fn snapshot(&self) -> WebDataBudgetSnapshot {
let state = self.state.lock();
WebDataBudgetSnapshot {
queue_bytes: state.queue_bytes,
queue_items: state.queue_items,
websocket_bytes: state.websocket_bytes,
high_water_bytes: state.high_water_bytes,
}
}
pub(super) fn close(&self) {
self.state.lock().closed = true;
self.notify.notify_waiters();
}
fn release_websocket(&self, owner: ProfileKey, bytes: usize) {
let mut state = self.state.lock();
state.websocket_bytes = state.websocket_bytes.saturating_sub(bytes);
remove_owner(&mut state.owner_bytes, owner, bytes);
drop(state);
self.pressured.store(false, Ordering::Release);
self.notify.notify_waiters();
}
}
impl WebSocketBudgetLease {
/// Releases unused worst-case capacity after one message is assembled.
pub(crate) fn shrink_to(&mut self, bytes: usize) {
let bytes = bytes.min(self.bytes);
let released = self.bytes - bytes;
if released == 0 {
return;
}
self.bytes = bytes;
self.budget.release_websocket(self.owner, released);
}
}
impl Drop for WebSocketBudgetLease {
fn drop(&mut self) {
self.budget.release_websocket(self.owner, self.bytes);
}
}
fn watermark(limit: usize, percentage: u8) -> usize {
limit.saturating_mul(usize::from(percentage)) / 100
}
pub(super) fn control_item_reserve(limits: &WebLimitsConfig) -> usize {
limits
.max_sessions_global
.saturating_mul(16usize.saturating_add(limits.max_streams_per_session.saturating_mul(3)))
}
fn add_owner(values: &mut HashMap<ProfileKey, usize>, owner: ProfileKey, bytes: usize) {
*values.entry(owner).or_insert(0) += bytes;
}
fn remove_owner(values: &mut HashMap<ProfileKey, usize>, owner: ProfileKey, bytes: usize) {
let remove = if let Some(value) = values.get_mut(&owner) {
*value = value.saturating_sub(bytes);
*value == 0
} else {
false
};
if remove {
values.remove(&owner);
}
}
fn update_high_water(state: &mut BudgetState) {
state.high_water_bytes = state
.high_water_bytes
.max(state.queue_bytes.saturating_add(state.websocket_bytes));
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn downlink_reservation_preserves_one_uplink_and_websocket_batch() {
let limits = WebLimitsConfig::default();
let uplink_bytes = limits
.max_body_bytes
.saturating_add(limits.max_frames_per_body.saturating_mul(QUEUE_ITEM_COST));
let downlink_bytes = limits
.pending_bytes_global
.saturating_sub(limits.control_bytes_global)
.saturating_sub(uplink_bytes)
.saturating_sub(limits.carrier_batch_bytes);
let budget = WebDataBudget::new(limits);
assert!(budget.try_reserve_queue([1; 32], downlink_bytes, 1, false, true));
assert!(!budget.try_reserve_queue([1; 32], 1, 1, false, true));
}
#[test]
fn item_limit_rejection_does_not_request_websocket_eviction() {
let limits = WebLimitsConfig::default();
let rejected_items = limits.pending_items_global.saturating_add(1);
let budget = WebDataBudget::new(limits);
let _websocket = budget
.try_reserve_websocket([1; 32], 1, WebSocketBudgetClass::Data)
.unwrap();
assert!(!budget.try_reserve_queue([2; 32], 1, rejected_items, false, false));
assert!(!budget.take_pressure());
}
#[test]
fn websocket_byte_conflict_requests_pressure_eviction() {
let limits = WebLimitsConfig::default();
let data_bytes = limits
.pending_bytes_global
.saturating_sub(limits.control_bytes_global);
let budget = WebDataBudget::new(limits);
let _websocket = budget
.try_reserve_websocket([1; 32], 1, WebSocketBudgetClass::Data)
.unwrap();
assert!(!budget.try_reserve_queue([2; 32], data_bytes, 1, false, false));
assert!(budget.take_pressure());
}
}
+10 -9
View File
@@ -69,6 +69,8 @@ impl WebProcessRuntime {
/// Stops issuance, closes all sessions, and joins bounded child work. /// Stops issuance, closes all sessions, and joins bounded child work.
pub(crate) async fn shutdown(&self) { pub(crate) async fn shutdown(&self) {
self.shutdown.cancel(); self.shutdown.cancel();
self.close_websockets();
self.data_budget.close();
let sessions = { let sessions = {
let mut state = self.state.lock(); let mut state = self.state.lock();
state.closed = true; state.closed = true;
@@ -94,15 +96,11 @@ impl WebProcessRuntime {
let _ = tokio::time::timeout(Duration::from_secs(timeout_secs), waits).await; let _ = tokio::time::timeout(Duration::from_secs(timeout_secs), waits).await;
self.tasks.close(); self.tasks.close();
let _ = tokio::time::timeout(Duration::from_secs(timeout_secs), self.tasks.wait()).await; let _ = tokio::time::timeout(Duration::from_secs(timeout_secs), self.tasks.wait()).await;
let (sessions_live, streams_live, pending_bytes, pending_items) = { let (sessions_live, streams_live) = {
let state = self.state.lock(); let state = self.state.lock();
( (state.sessions.len(), state.streams_live)
state.sessions.len(),
state.streams_live,
state.pending_bytes,
state.pending_items,
)
}; };
let budget = self.data_budget.snapshot();
info!( info!(
target: "telemt::web", target: "telemt::web",
sessions_created = self.sessions_created.load(Ordering::Relaxed), sessions_created = self.sessions_created.load(Ordering::Relaxed),
@@ -111,8 +109,10 @@ impl WebProcessRuntime {
streams_opened = self.streams_opened.load(Ordering::Relaxed), streams_opened = self.streams_opened.load(Ordering::Relaxed),
streams_rejected = self.streams_rejected.load(Ordering::Relaxed), streams_rejected = self.streams_rejected.load(Ordering::Relaxed),
streams_live, streams_live,
pending_bytes, pending_bytes = budget.queue_bytes,
pending_items, pending_items = budget.queue_items,
websocket_bytes = budget.websocket_bytes,
data_high_water_bytes = budget.high_water_bytes,
bytes_up = self.bytes_up.load(Ordering::Relaxed), bytes_up = self.bytes_up.load(Ordering::Relaxed),
bytes_down = self.bytes_down.load(Ordering::Relaxed), bytes_down = self.bytes_down.load(Ordering::Relaxed),
limit_hits = self.limit_hits.load(Ordering::Relaxed), limit_hits = self.limit_hits.load(Ordering::Relaxed),
@@ -122,6 +122,7 @@ impl WebProcessRuntime {
/// Expires credentials and closes idle sessions without holding locks across callbacks. /// Expires credentials and closes idle sessions without holding locks across callbacks.
pub(super) fn cleanup(&self) { pub(super) fn cleanup(&self) {
self.cleanup_websockets();
let now = Instant::now(); let now = Instant::now();
let sessions = { let sessions = {
let mut state = self.state.lock(); let mut state = self.state.lock();
+15 -18
View File
@@ -8,10 +8,12 @@ use sha2::{Digest, Sha256};
use zeroize::Zeroizing; use zeroize::Zeroizing;
use super::{ProfileKey, TOKEN_BYTES, TokenHash}; use super::{ProfileKey, TOKEN_BYTES, TokenHash};
use crate::config::{WebLimitsConfig, WebRuntimeConfig, WebRuntimeProfile}; use crate::config::{WebRuntimeConfig, WebRuntimeProfile};
use crate::maestro::generation::RuntimeGeneration; use crate::maestro::generation::RuntimeGeneration;
use crate::web::session::WebSession; use crate::web::session::WebSession;
const WEB_PROFILE_OWNER_CONTEXT: &[u8] = b"telemt-web-profile-owner-v1\0";
/// One issued bootstrap and optional idempotent session-creation replay state. /// One issued bootstrap and optional idempotent session-creation replay state.
pub(super) struct Bootstrap { pub(super) struct Bootstrap {
/// Credential and replay-state expiry deadline. /// Credential and replay-state expiry deadline.
@@ -74,14 +76,6 @@ pub(super) struct ManagerState {
/// Process-wide live relay-task count. /// Process-wide live relay-task count.
pub(super) streams_live: usize, pub(super) streams_live: usize,
stream_ports: HashMap<(IpAddr, SocketAddr), StreamPortState>, stream_ports: HashMap<(IpAddr, SocketAddr), StreamPortState>,
/// Total process-wide queued byte reservation.
pub(super) pending_bytes: usize,
/// Total process-wide queued item reservation.
pub(super) pending_items: usize,
/// Portion of queued bytes charged to the control reserve.
pub(super) pending_control_bytes: usize,
/// Portion of queued items charged to the control reserve.
pub(super) pending_control_items: usize,
/// Bootstrap issuance rate limiter. /// Bootstrap issuance rate limiter.
pub(super) bootstrap_rate: RateState, pub(super) bootstrap_rate: RateState,
/// Session creation rate limiter. /// Session creation rate limiter.
@@ -112,9 +106,19 @@ pub(super) fn new_unique_token(
None None
} }
/// Returns the precomputed capability as the stable process profile key. /// Derives a secret-independent quota owner stable across capability rotation.
pub(super) fn profile_key(profile: &WebRuntimeProfile) -> ProfileKey { pub(super) fn profile_key(profile: &WebRuntimeProfile) -> ProfileKey {
profile.capability let mut digest = Sha256::new();
digest.update(WEB_PROFILE_OWNER_CONTEXT);
digest.update((profile.host.len() as u64).to_be_bytes());
digest.update(profile.host.as_bytes());
digest.update((profile.user.len() as u64).to_be_bytes());
digest.update(profile.user.as_bytes());
digest.update([match profile.secret_mode {
crate::config::WebSecretMode::Plain => 0,
crate::config::WebSecretMode::Dd => 1,
}]);
digest.finalize().into()
} }
/// Re-resolves an issued profile against the active generation without weakening identity. /// Re-resolves an issued profile against the active generation without weakening identity.
@@ -211,13 +215,6 @@ where
} }
} }
/// Computes the process-wide item reserve required for session control progress.
pub(super) fn control_item_reserve(limits: &WebLimitsConfig) -> usize {
limits
.max_sessions_global
.saturating_mul(16usize.saturating_add(limits.max_streams_per_session.saturating_mul(3)))
}
/// Allocates a non-zero source port unique among live streams for one KDF route. /// Allocates a non-zero source port unique among live streams for one KDF route.
pub(super) fn allocate_stream_port( pub(super) fn allocate_stream_port(
state: &mut ManagerState, state: &mut ManagerState,
+313
View File
@@ -0,0 +1,313 @@
use std::collections::HashMap;
use std::net::IpAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::time::Duration;
use tokio::sync::OwnedSemaphorePermit;
use tokio_util::sync::CancellationToken;
use super::{ManagerError, ProfileKey, WebProcessRuntime, WebSocketBudgetLease};
/// One process-owned WebSocket carrier class used for eviction priority.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum WebSocketKind {
/// One connection multiplexes every logical stream in a session.
Multiplex,
/// One connection owns exactly one logical stream lane.
Lane(u32),
}
pub(super) struct WebSocketEntry {
id: u64,
owner: ProfileKey,
session_id: u64,
client_ip: IpAddr,
kind: WebSocketKind,
liveness_interval_ms: u64,
created_tick: u64,
last_peer_tick: AtomicU64,
last_progress_tick: AtomicU64,
opened: AtomicBool,
cancel: CancellationToken,
}
#[derive(Default)]
pub(super) struct WebSocketRegistry {
entries: HashMap<u64, Arc<WebSocketEntry>>,
closed: bool,
}
/// Exact process-owned admission retained through the upgraded socket lifetime.
pub(crate) struct WebSocketConnection {
runtime: std::sync::Weak<WebProcessRuntime>,
entry: Arc<WebSocketEntry>,
slot: Option<OwnedSemaphorePermit>,
base_budget: Option<WebSocketBudgetLease>,
}
impl WebSocketConnection {
/// Returns the cancellation signal used by shutdown and pressure eviction.
pub(crate) fn cancellation(&self) -> CancellationToken {
self.entry.cancel.clone()
}
/// Returns the process-unique connection identifier used only for debugging.
pub(crate) fn id(&self) -> u64 {
self.entry.id
}
/// Returns the creation-time transport liveness interval.
pub(crate) fn liveness_interval(&self) -> Duration {
Duration::from_millis(self.entry.liveness_interval_ms)
}
/// Marks successful ownership transfer from HTTP to the WebSocket codec.
pub(crate) fn mark_opened(&self) {
self.entry.opened.store(true, Ordering::Release);
self.mark_progress();
}
/// Refreshes the peer-liveness deadline after any received WebSocket message.
pub(crate) fn mark_peer_activity(&self) {
if let Some(runtime) = self.runtime.upgrade() {
let now = runtime.websocket_tick();
self.entry.last_peer_tick.store(now, Ordering::Release);
self.entry.last_progress_tick.store(now, Ordering::Release);
}
}
/// Refreshes least-recently-progressed ordering after a committed write.
pub(crate) fn mark_progress(&self) {
if let Some(runtime) = self.runtime.upgrade() {
self.entry
.last_progress_tick
.store(runtime.websocket_tick(), Ordering::Release);
}
}
}
impl Drop for WebSocketConnection {
fn drop(&mut self) {
if let Some(runtime) = self.runtime.upgrade() {
runtime.websockets.lock().entries.remove(&self.entry.id);
drop(self.base_budget.take());
drop(self.slot.take());
runtime.websocket_notify.notify_waiters();
}
}
}
pub(super) async fn admit(
runtime: &Arc<WebProcessRuntime>,
owner: ProfileKey,
session_id: u64,
client_ip: IpAddr,
kind: WebSocketKind,
base_bytes: usize,
liveness_interval: Duration,
eviction_timeout: Duration,
) -> Result<WebSocketConnection, ManagerError> {
let liveness_interval_ms = liveness_interval.as_millis().min(u128::from(u64::MAX)) as u64;
if let Some(connection) = try_admit(
runtime,
owner,
session_id,
client_ip,
kind,
base_bytes,
liveness_interval_ms,
) {
return Ok(connection);
}
let Some(victim) = select_victim(runtime, owner, session_id, client_ip, None) else {
runtime.record_limit_hit();
return Err(ManagerError::Limit);
};
let released = runtime.websocket_notify.notified();
victim.cancel.cancel();
let _ = tokio::time::timeout(eviction_timeout, released).await;
try_admit(
runtime,
owner,
session_id,
client_ip,
kind,
base_bytes,
liveness_interval_ms,
)
.ok_or_else(|| {
runtime.record_limit_hit();
ManagerError::Limit
})
}
fn try_admit(
runtime: &Arc<WebProcessRuntime>,
owner: ProfileKey,
session_id: u64,
client_ip: IpAddr,
kind: WebSocketKind,
base_bytes: usize,
liveness_interval_ms: u64,
) -> Option<WebSocketConnection> {
let slot = Arc::clone(&runtime.websocket_connections)
.try_acquire_owned()
.ok()?;
let base_budget = runtime.try_websocket_base_budget(owner, base_bytes)?;
let id = runtime.websocket_next_id.fetch_add(1, Ordering::Relaxed);
let now = runtime.websocket_tick();
let entry = Arc::new(WebSocketEntry {
id,
owner,
session_id,
client_ip,
kind,
liveness_interval_ms,
created_tick: now,
last_peer_tick: AtomicU64::new(now),
last_progress_tick: AtomicU64::new(now),
opened: AtomicBool::new(false),
cancel: CancellationToken::new(),
});
let mut registry = runtime.websockets.lock();
if registry.closed {
return None;
}
registry.entries.insert(id, Arc::clone(&entry));
drop(registry);
Some(WebSocketConnection {
runtime: Arc::downgrade(runtime),
entry,
slot: Some(slot),
base_budget: Some(base_budget),
})
}
impl WebProcessRuntime {
pub(super) fn websocket_tick(&self) -> u64 {
self.websocket_clock.elapsed().as_millis() as u64
}
pub(super) fn cleanup_websockets(&self) {
let now = self.websocket_tick();
let mut victims = self
.websockets
.lock()
.entries
.values()
.filter(|entry| {
now.saturating_sub(entry.last_peer_tick.load(Ordering::Acquire))
>= dead_after(entry)
})
.cloned()
.collect::<Vec<_>>();
if victims.is_empty()
&& self.data_budget.take_pressure()
&& let Some(victim) = select_pressure_victim(self, now)
{
victims.push(victim);
}
for victim in victims {
victim.cancel.cancel();
}
}
pub(super) fn close_websockets(&self) {
let victims = {
let mut registry = self.websockets.lock();
registry.closed = true;
registry.entries.values().cloned().collect::<Vec<_>>()
};
for victim in victims {
victim.cancel.cancel();
}
}
}
fn select_victim(
runtime: &WebProcessRuntime,
owner: ProfileKey,
session_id: u64,
client_ip: IpAddr,
excluded_id: Option<u64>,
) -> Option<Arc<WebSocketEntry>> {
let fair_share = runtime.data_budget.fair_share(Some(owner));
let requester_usage = runtime.data_budget.owner_usage(owner);
let now = runtime.websocket_tick();
runtime
.websockets
.lock()
.entries
.values()
.filter(|entry| Some(entry.id) != excluded_id)
.filter_map(|entry| {
let owner_rank = if entry.session_id == session_id {
0
} else if entry.owner == owner {
1
} else if entry.client_ip == client_ip {
2
} else {
if requester_usage >= fair_share
|| runtime.data_budget.owner_usage(entry.owner) <= fair_share
{
return None;
}
3
};
let priority = entry_priority(entry, now);
Some((
(
owner_rank,
priority,
entry.last_progress_tick.load(Ordering::Acquire),
entry.created_tick,
entry.id,
),
Arc::clone(entry),
))
})
.min_by_key(|(key, _)| *key)
.map(|(_, entry)| entry)
}
fn select_pressure_victim(runtime: &WebProcessRuntime, now: u64) -> Option<Arc<WebSocketEntry>> {
runtime
.websockets
.lock()
.entries
.values()
.map(|entry| {
(
(
entry_priority(entry, now),
entry.last_progress_tick.load(Ordering::Acquire),
entry.created_tick,
entry.id,
),
Arc::clone(entry),
)
})
.min_by_key(|(key, _)| *key)
.map(|(_, entry)| entry)
}
fn entry_priority(entry: &WebSocketEntry, now: u64) -> u8 {
if !entry.opened.load(Ordering::Acquire)
|| now.saturating_sub(entry.last_peer_tick.load(Ordering::Acquire)) >= dead_after(entry)
{
0
} else if matches!(entry.kind, WebSocketKind::Lane(_)) {
1
} else {
2
}
}
fn dead_after(entry: &WebSocketEntry) -> u64 {
entry.liveness_interval_ms.saturating_mul(2)
}
#[cfg(test)]
mod tests;
+40
View File
@@ -0,0 +1,40 @@
use super::*;
fn entry(kind: WebSocketKind, opened: bool, peer_tick: u64) -> WebSocketEntry {
WebSocketEntry {
id: 1,
owner: [0; 32],
session_id: 1,
client_ip: "192.0.2.10".parse().unwrap(),
kind,
liveness_interval_ms: 10,
created_tick: 1,
last_peer_tick: AtomicU64::new(peer_tick),
last_progress_tick: AtomicU64::new(peer_tick),
opened: AtomicBool::new(opened),
cancel: CancellationToken::new(),
}
}
#[test]
fn preopen_and_dead_entries_precede_live_lane_and_multiplex_victims() {
let preopen = entry(WebSocketKind::Multiplex, false, 90);
let dead = entry(WebSocketKind::Multiplex, true, 1);
let lane = entry(WebSocketKind::Lane(7), true, 90);
let multiplex = entry(WebSocketKind::Multiplex, true, 90);
assert_eq!(entry_priority(&preopen, 100), 0);
assert_eq!(entry_priority(&dead, 100), 0);
assert_eq!(entry_priority(&lane, 100), 1);
assert_eq!(entry_priority(&multiplex, 100), 2);
}
#[test]
fn dead_classification_keeps_each_connections_creation_time_interval() {
let short_interval = entry(WebSocketKind::Multiplex, true, 80);
let mut long_interval = entry(WebSocketKind::Multiplex, true, 80);
long_interval.liveness_interval_ms = 100;
assert_eq!(entry_priority(&short_interval, 100), 0);
assert_eq!(entry_priority(&long_interval, 100), 2);
}
+19 -4
View File
@@ -22,6 +22,9 @@ mod backend;
mod downlink; mod downlink;
// Lane carrier state isolates request sequencing and downlink replay per logical stream. // Lane carrier state isolates request sequencing and downlink replay per logical stream.
mod lanes; mod lanes;
// WebSocket carrier state owns pre-OPEN lane reservations and failure isolation.
mod websocket;
pub(crate) use websocket::WebSocketLaneReservation;
// Uplink batches own exactly-once sequencing and client-frame validation. // Uplink batches own exactly-once sequencing and client-frame validation.
mod uplink; mod uplink;
@@ -107,6 +110,7 @@ struct SessionState {
last_up_sequence: u64, last_up_sequence: u64,
last_up_digest: TokenHash, last_up_digest: TokenHash,
carrier_lanes: HashMap<u32, CarrierLane>, carrier_lanes: HashMap<u32, CarrierLane>,
websocket_lane_reservations: HashMap<u32, u16>,
pending_bytes: usize, pending_bytes: usize,
pending_items: usize, pending_items: usize,
pending_control_bytes: usize, pending_control_bytes: usize,
@@ -183,6 +187,7 @@ impl WebSession {
last_up_sequence: 0, last_up_sequence: 0,
last_up_digest: [0; 32], last_up_digest: [0; 32],
carrier_lanes, carrier_lanes,
websocket_lane_reservations: HashMap::new(),
pending_bytes: 0, pending_bytes: 0,
pending_items: 0, pending_items: 0,
pending_control_bytes: 0, pending_control_bytes: 0,
@@ -214,6 +219,16 @@ impl WebSession {
self.profile.carrier self.profile.carrier
} }
/// Returns the stable quota owner without exposing profile credentials.
pub(crate) fn profile_key(&self) -> ProfileKey {
self.profile_key
}
/// Returns the process-unique non-secret trace identifier.
pub(crate) fn trace_session_id(&self) -> u64 {
self.trace_session_id
}
/// Returns a cloned non-secret identity only for enabled debug capture. /// Returns a cloned non-secret identity only for enabled debug capture.
pub(crate) fn trace_identity(&self) -> crate::web::trace::TraceIdentity { pub(crate) fn trace_identity(&self) -> crate::web::trace::TraceIdentity {
crate::web::trace::TraceIdentity::from_profile(self.trace_session_id, &self.profile) crate::web::trace::TraceIdentity::from_profile(self.trace_session_id, &self.profile)
@@ -273,12 +288,12 @@ impl WebSession {
(data_bytes, data_items, control_bytes, control_items) (data_bytes, data_items, control_bytes, control_items)
}; };
self.cancel.cancel(); self.cancel.cancel();
if self.carrier() == WebCarrier::Https { if self.carrier().is_multiplexed() {
self.down_notify.notify_waiters(); self.down_notify.notify_waiters();
} }
if let Some(manager) = self.manager.upgrade() { if let Some(manager) = self.manager.upgrade() {
manager.release_pending(data_bytes, data_items, false); manager.release_pending(self.profile_key, data_bytes, data_items, false);
manager.release_pending(control_bytes, control_items, true); manager.release_pending(self.profile_key, control_bytes, control_items, true);
if !self.finished.swap(true, Ordering::AcqRel) { if !self.finished.swap(true, Ordering::AcqRel) {
self.trace_lifecycle( self.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::SessionClosed, crate::web::trace::TraceLifecycleEvent::SessionClosed,
@@ -394,7 +409,7 @@ impl WebSession {
stream.send_credit -= count as u64; stream.send_credit -= count as u64;
state.last_activity = Instant::now(); state.last_activity = Instant::now();
drop(state); drop(state);
if self.carrier() == WebCarrier::Https { if self.carrier().is_multiplexed() {
self.down_notify.notify_waiters(); self.down_notify.notify_waiters();
} }
Poll::Ready(Ok(count)) Poll::Ready(Ok(count))
+64 -21
View File
@@ -1,6 +1,6 @@
use std::io; use std::io;
use std::sync::Arc; use std::sync::Arc;
use std::sync::atomic::Ordering; use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration; use std::time::Duration;
use crate::proxy::shared_state::ConntrackClosePolicy; use crate::proxy::shared_state::ConntrackClosePolicy;
@@ -15,10 +15,15 @@ mod tests;
impl WebSession { impl WebSession {
/// Starts one owned inner handshake and relay task for an admitted stream. /// Starts one owned inner handshake and relay task for an admitted stream.
pub(super) fn spawn_stream(self: &Arc<Self>, stream_id: u32, peer_port: u16) { pub(super) fn spawn_stream(
self: &Arc<Self>,
stream_id: u32,
peer_port: u16,
retain_reservation_on_reject: bool,
) -> bool {
let Some(manager) = self.manager.upgrade() else { let Some(manager) = self.manager.upgrade() else {
self.stream_finished(stream_id, peer_port); self.stream_rejected_before_spawn(stream_id, peer_port, retain_reservation_on_reject);
return; return false;
}; };
let generation = manager.active_generation(); let generation = manager.active_generation();
if !*generation.admission_rx.borrow() { if !*generation.admission_rx.borrow() {
@@ -27,8 +32,8 @@ impl WebSession {
Some(stream_id), Some(stream_id),
Some("admission_closed"), Some("admission_closed"),
); );
self.stream_finished(stream_id, peer_port); self.stream_rejected_before_spawn(stream_id, peer_port, retain_reservation_on_reject);
return; return false;
} }
let Ok(connection_permit) = generation.max_connections.clone().try_acquire_owned() else { let Ok(connection_permit) = generation.max_connections.clone().try_acquire_owned() else {
manager.record_stream_rejected(); manager.record_stream_rejected();
@@ -37,21 +42,24 @@ impl WebSession {
Some(stream_id), Some(stream_id),
Some("connection_limit"), Some("connection_limit"),
); );
self.stream_finished(stream_id, peer_port); self.stream_rejected_before_spawn(stream_id, peer_port, retain_reservation_on_reject);
return; return false;
}; };
let deps = generation.client_runtime_deps(); let deps = generation.client_runtime_deps();
let replay_checker = Arc::clone(&generation.replay_checker); let replay_checker = Arc::clone(&generation.replay_checker);
let session = Arc::clone(self); let session = Arc::clone(self);
let cancel = self.cancel.clone(); let cancel = self.cancel.clone();
let retain_rejected = Arc::new(AtomicBool::new(false));
self.tasks_live.fetch_add(1, Ordering::AcqRel); self.tasks_live.fetch_add(1, Ordering::AcqRel);
let spawned = generation.spawn_session(async move { let completion = StreamCompletion {
session: Arc::clone(&session),
stream_id,
peer_port,
retain_rejected: Arc::clone(&retain_rejected),
};
let future = async move {
let _connection_permit = connection_permit; let _connection_permit = connection_permit;
let _completion = StreamCompletion { let _completion = completion;
session: Arc::clone(&session),
stream_id,
peer_port,
};
session.trace_lifecycle( session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::StreamAdmitted, crate::web::trace::TraceLifecycleEvent::StreamAdmitted,
Some(stream_id), Some(stream_id),
@@ -69,16 +77,41 @@ impl WebSession {
peer_port, peer_port,
) => {} ) => {}
} }
}); };
if !spawned { if let Err(future) = generation.try_spawn_session(future) {
retain_rejected.store(retain_reservation_on_reject, Ordering::Release);
self.trace_lifecycle( self.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::StreamRejected, crate::web::trace::TraceLifecycleEvent::StreamRejected,
Some(stream_id), Some(stream_id),
Some("generation_closed"), Some("generation_closed"),
); );
self.tasks_live.fetch_sub(1, Ordering::AcqRel); drop(future);
return false;
}
true
}
fn stream_rejected_before_spawn(
&self,
stream_id: u32,
peer_port: u16,
retain_reservation: bool,
) {
if !retain_reservation {
self.stream_finished(stream_id, peer_port); self.stream_finished(stream_id, peer_port);
self.tasks_done.notify_waiters(); return;
}
let queued = {
let mut state = self.state.lock();
state.streams.remove(&stream_id).map(|stream| {
let (bytes, items) = inbound_queue_cost(&stream.inbound);
self.release_locked(&mut state, bytes, items, false);
self.remember_closed_locked(&mut state, stream_id);
self.queue_control_locked(&mut state, FrameType::Close, stream_id, &[])
})
};
if queued.is_some_and(|queued| !queued) {
self.close();
} }
} }
@@ -106,7 +139,7 @@ impl WebSession {
if !queued { if !queued {
self.close(); self.close();
} }
if self.carrier() == crate::config::WebCarrier::Https { if self.carrier().is_multiplexed() {
self.down_notify.notify_waiters(); self.down_notify.notify_waiters();
} }
} }
@@ -117,6 +150,7 @@ struct StreamCompletion {
session: Arc<WebSession>, session: Arc<WebSession>,
stream_id: u32, stream_id: u32,
peer_port: u16, peer_port: u16,
retain_rejected: Arc<AtomicBool>,
} }
impl Drop for StreamCompletion { impl Drop for StreamCompletion {
@@ -126,7 +160,12 @@ impl Drop for StreamCompletion {
Some(self.stream_id), Some(self.stream_id),
None, None,
); );
self.session.stream_finished(self.stream_id, self.peer_port); if self.retain_rejected.load(Ordering::Acquire) {
self.session
.stream_rejected_before_spawn(self.stream_id, self.peer_port, true);
} else {
self.session.stream_finished(self.stream_id, self.peer_port);
}
if self.session.tasks_live.fetch_sub(1, Ordering::AcqRel) == 1 { if self.session.tasks_live.fetch_sub(1, Ordering::AcqRel) == 1 {
self.session.tasks_done.notify_waiters(); self.session.tasks_done.notify_waiters();
} }
@@ -263,6 +302,10 @@ async fn run_stream(
session.trace_lifecycle( session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::RelayEnded, crate::web::trace::TraceLifecycleEvent::RelayEnded,
Some(stream_id), Some(stream_id),
Some(if relay_result.is_ok() { "completed" } else { "error" }), Some(if relay_result.is_ok() {
"completed"
} else {
"error"
}),
); );
} }
+7 -9
View File
@@ -13,9 +13,7 @@ use crate::config::{
WebSecretMode, WebSecretMode,
}; };
use crate::crypto::{AesCtr, sha256}; use crate::crypto::{AesCtr, sha256};
use crate::maestro::generation::{ use crate::maestro::generation::{RuntimeGeneration, test_runtime_generation_with_admission};
RuntimeGeneration, test_runtime_generation_with_admission,
};
use crate::protocol::constants::{ use crate::protocol::constants::{
DC_IDX_POS, HANDSHAKE_LEN, IV_LEN, PREKEY_LEN, PROTO_TAG_POS, ProtoTag, SKIP_LEN, DC_IDX_POS, HANDSHAKE_LEN, IV_LEN, PREKEY_LEN, PROTO_TAG_POS, ProtoTag, SKIP_LEN,
}; };
@@ -39,10 +37,11 @@ impl TestRuntime {
) -> Result<u64, ManagerError> { ) -> Result<u64, ManagerError> {
let encoded = frame::encode(frame_type, stream_id, payload); let encoded = frame::encode(frame_type, stream_id, payload);
match self.session.carrier() { match self.session.carrier() {
WebCarrier::Https => self.session.process_up(sequence, &encoded), WebCarrier::Https | WebCarrier::Websocket => {
WebCarrier::HttpsLanes => { self.session.process_up(sequence, &encoded)
self.session }
.process_up_lane(stream_id, sequence, &encoded) WebCarrier::HttpsLanes | WebCarrier::WebsocketLanes => {
self.session.process_up_lane(stream_id, sequence, &encoded)
} }
} }
} }
@@ -152,8 +151,7 @@ fn valid_plain_handshake() -> [u8; HANDSHAKE_LEN] {
let mut cipher = AesCtr::new(&dec_key, u128::from_be_bytes(dec_iv)); let mut cipher = AesCtr::new(&dec_key, u128::from_be_bytes(dec_iv));
let keystream = cipher.encrypt(&[0u8; HANDSHAKE_LEN]); let keystream = cipher.encrypt(&[0u8; HANDSHAKE_LEN]);
let mut plaintext = [0u8; HANDSHAKE_LEN]; let mut plaintext = [0u8; HANDSHAKE_LEN];
plaintext[PROTO_TAG_POS..PROTO_TAG_POS + 4] plaintext[PROTO_TAG_POS..PROTO_TAG_POS + 4].copy_from_slice(&ProtoTag::Intermediate.to_bytes());
.copy_from_slice(&ProtoTag::Intermediate.to_bytes());
plaintext[DC_IDX_POS..DC_IDX_POS + 2].copy_from_slice(&2i16.to_le_bytes()); plaintext[DC_IDX_POS..DC_IDX_POS + 2].copy_from_slice(&2i16.to_le_bytes());
for index in PROTO_TAG_POS..HANDSHAKE_LEN { for index in PROTO_TAG_POS..HANDSHAKE_LEN {
handshake[index] = plaintext[index] ^ keystream[index]; handshake[index] = plaintext[index] ^ keystream[index];
+15 -8
View File
@@ -5,14 +5,13 @@ use bytes::{BufMut, Bytes, BytesMut};
use super::{ use super::{
DownBatch, PendingClass, PollResult, QUEUE_ITEM_COST, QueuedFrame, SessionState, WebSession, DownBatch, PendingClass, PollResult, QUEUE_ITEM_COST, QueuedFrame, SessionState, WebSession,
}; };
use crate::config::WebCarrier;
use crate::web::frame::{self, FrameType}; use crate::web::frame::{self, FrameType};
use crate::web::manager::ManagerError; use crate::web::manager::ManagerError;
impl WebSession { impl WebSession {
/// Polls pending downlink frames with cursor replay and newest-poll-wins semantics. /// Polls pending downlink frames with cursor replay and newest-poll-wins semantics.
pub(crate) async fn poll_down(&self, cursor: u64) -> Result<PollResult, ManagerError> { pub(crate) async fn poll_down(&self, cursor: u64) -> Result<PollResult, ManagerError> {
if self.carrier() != WebCarrier::Https { if !self.carrier().is_multiplexed() {
return Err(ManagerError::Protocol); return Err(ManagerError::Protocol);
} }
let epoch = { let epoch = {
@@ -167,7 +166,13 @@ impl WebSession {
let Some(manager) = self.manager.upgrade() else { let Some(manager) = self.manager.upgrade() else {
return false; return false;
}; };
if !manager.try_reserve_pending(bytes, items, control, class == PendingClass::Downlink) { if !manager.try_reserve_pending(
self.profile_key,
bytes,
items,
control,
class == PendingClass::Downlink,
) {
return false; return false;
} }
state.pending_bytes += bytes; state.pending_bytes += bytes;
@@ -194,7 +199,7 @@ impl WebSession {
state.pending_control_items = state.pending_control_items.saturating_sub(items); state.pending_control_items = state.pending_control_items.saturating_sub(items);
} }
if let Some(manager) = self.manager.upgrade() { if let Some(manager) = self.manager.upgrade() {
manager.release_pending(bytes, items, control); manager.release_pending(self.profile_key, bytes, items, control);
} }
} }
@@ -208,7 +213,7 @@ impl WebSession {
if amount == 0 { if amount == 0 {
return true; return true;
} }
if self.carrier() == WebCarrier::HttpsLanes { if self.carrier().uses_lanes() {
return self.queue_control_locked( return self.queue_control_locked(
state, state,
FrameType::Window, FrameType::Window,
@@ -257,7 +262,7 @@ impl WebSession {
stream_id: u32, stream_id: u32,
payload: &[u8], payload: &[u8],
) -> bool { ) -> bool {
if self.carrier() == WebCarrier::HttpsLanes { if self.carrier().uses_lanes() {
return self.queue_frame_locked(state, FrameType::Data, stream_id, payload, false); return self.queue_frame_locked(state, FrameType::Data, stream_id, payload, false);
} }
let can_coalesce = state.pending_frames.back().is_some_and(|last| { let can_coalesce = state.pending_frames.back().is_some_and(|last| {
@@ -290,7 +295,7 @@ impl WebSession {
payload: &[u8], payload: &[u8],
control: bool, control: bool,
) -> bool { ) -> bool {
if self.carrier() == WebCarrier::HttpsLanes { if self.carrier().uses_lanes() {
return self.queue_lane_frame_locked(state, frame_type, stream_id, payload, control); return self.queue_lane_frame_locked(state, frame_type, stream_id, payload, control);
} }
let cost = frame::HEADER_BYTES + payload.len() + QUEUE_ITEM_COST; let cost = frame::HEADER_BYTES + payload.len() + QUEUE_ITEM_COST;
@@ -409,7 +414,9 @@ mod tests {
use std::net::SocketAddr; use std::net::SocketAddr;
use std::sync::Arc; use std::sync::Arc;
use crate::config::{WebLimitsConfig, WebRuntimeProfile, WebSecretMode, WebTimeoutsConfig}; use crate::config::{
WebCarrier, WebLimitsConfig, WebRuntimeProfile, WebSecretMode, WebTimeoutsConfig,
};
use crate::web::manager::WebProcessRuntime; use crate::web::manager::WebProcessRuntime;
fn session() -> Arc<WebSession> { fn session() -> Arc<WebSession> {
+5 -4
View File
@@ -113,6 +113,7 @@ impl WebSession {
&mut state, &mut state,
&frames, &frames,
&mut opened, &mut opened,
&mut None,
&mut unused_bytes, &mut unused_bytes,
&mut unused_items, &mut unused_items,
); );
@@ -137,7 +138,7 @@ impl WebSession {
return result; return result;
} }
for (stream_id, peer_port) in opened { for (stream_id, peer_port) in opened {
self.spawn_stream(stream_id, peer_port); self.spawn_stream(stream_id, peer_port, false);
} }
if let Some(manager) = self.manager.upgrade() { if let Some(manager) = self.manager.upgrade() {
manager.record_up(body.len()); manager.record_up(body.len());
@@ -151,7 +152,7 @@ impl WebSession {
lane_id: u32, lane_id: u32,
cursor: u64, cursor: u64,
) -> Result<PollResult, ManagerError> { ) -> Result<PollResult, ManagerError> {
if self.carrier() != WebCarrier::HttpsLanes || lane_id > frame::MAX_STREAM_ID { if !self.carrier().uses_lanes() || lane_id > frame::MAX_STREAM_ID {
return Err(ManagerError::Protocol); return Err(ManagerError::Protocol);
} }
let (epoch, notify) = { let (epoch, notify) = {
@@ -404,7 +405,7 @@ impl WebSession {
pub(super) fn remember_closed_locked(&self, state: &mut SessionState, stream_id: u32) { pub(super) fn remember_closed_locked(&self, state: &mut SessionState, stream_id: u32) {
let evicted = remember_closed(state, stream_id, self.limits.max_tombstones_per_session); let evicted = remember_closed(state, stream_id, self.limits.max_tombstones_per_session);
if self.carrier() != WebCarrier::HttpsLanes { if !self.carrier().uses_lanes() {
return; return;
} }
if let Some(evicted) = evicted { if let Some(evicted) = evicted {
@@ -415,7 +416,7 @@ impl WebSession {
} }
} }
fn release_lane_locked(&self, state: &mut SessionState, lane_id: u32) { pub(super) fn release_lane_locked(&self, state: &mut SessionState, lane_id: u32) {
let Some(mut lane) = state.carrier_lanes.remove(&lane_id) else { let Some(mut lane) = state.carrier_lanes.remove(&lane_id) else {
return; return;
}; };
+29 -8
View File
@@ -11,7 +11,6 @@ use super::{
InboundChunk, PendingClass, QUEUE_ITEM_COST, SessionState, StreamState, WebSession, InboundChunk, PendingClass, QUEUE_ITEM_COST, SessionState, StreamState, WebSession,
inbound_queue_cost, inbound_queue_cost,
}; };
use crate::config::WebCarrier;
use crate::web::frame::{self, Frame, FrameType}; use crate::web::frame::{self, Frame, FrameType};
use crate::web::manager::{ManagerError, TokenHash}; use crate::web::manager::{ManagerError, TokenHash};
@@ -22,7 +21,7 @@ impl WebSession {
sequence: u64, sequence: u64,
body: &[u8], body: &[u8],
) -> Result<u64, ManagerError> { ) -> Result<u64, ManagerError> {
if self.carrier() != WebCarrier::Https { if !self.carrier().is_multiplexed() {
return Err(ManagerError::Protocol); return Err(ManagerError::Protocol);
} }
if self if self
@@ -90,6 +89,7 @@ impl WebSession {
&mut state, &mut state,
&frames, &frames,
&mut opened, &mut opened,
&mut None,
&mut unused_bytes, &mut unused_bytes,
&mut unused_items, &mut unused_items,
); );
@@ -113,7 +113,7 @@ impl WebSession {
return result; return result;
} }
for (stream_id, peer_port) in opened { for (stream_id, peer_port) in opened {
self.spawn_stream(stream_id, peer_port); self.spawn_stream(stream_id, peer_port, false);
} }
if let Some(manager) = self.manager.upgrade() { if let Some(manager) = self.manager.upgrade() {
manager.record_up(body.len()); manager.record_up(body.len());
@@ -126,6 +126,7 @@ impl WebSession {
state: &mut SessionState, state: &mut SessionState,
frames: &[Frame<'_>], frames: &[Frame<'_>],
opened: &mut Vec<(u32, u16)>, opened: &mut Vec<(u32, u16)>,
reserved_open: &mut Option<(u32, u16)>,
unused_bytes: &mut usize, unused_bytes: &mut usize,
unused_items: &mut usize, unused_items: &mut usize,
) -> bool { ) -> bool {
@@ -136,13 +137,31 @@ impl WebSession {
let was_closed = state.closed_streams.contains(&value.stream_id); let was_closed = state.closed_streams.contains(&value.stream_id);
match value.frame_type { match value.frame_type {
FrameType::Open => { FrameType::Open => {
let Some(peer_port) = self.reserve_stream_locked(state) else { let peer_port = match reserved_open.take() {
self.remember_closed_locked(state, value.stream_id); Some((reserved_stream_id, peer_port))
if !self.queue_control_locked(state, FrameType::Close, value.stream_id, &[]) if reserved_stream_id == value.stream_id =>
{ {
peer_port
}
Some(reserved) => {
*reserved_open = Some(reserved);
return false; return false;
} }
continue; None => {
let Some(peer_port) = self.reserve_stream_locked(state) else {
self.remember_closed_locked(state, value.stream_id);
if !self.queue_control_locked(
state,
FrameType::Close,
value.stream_id,
&[],
) {
return false;
}
continue;
};
peer_port
}
}; };
state.streams.insert( state.streams.insert(
value.stream_id, value.stream_id,
@@ -331,7 +350,9 @@ mod tests {
use super::*; use super::*;
use std::net::SocketAddr; use std::net::SocketAddr;
use crate::config::{WebLimitsConfig, WebRuntimeProfile, WebSecretMode, WebTimeoutsConfig}; use crate::config::{
WebCarrier, WebLimitsConfig, WebRuntimeProfile, WebSecretMode, WebTimeoutsConfig,
};
use crate::web::manager::WebProcessRuntime; use crate::web::manager::WebProcessRuntime;
fn session() -> Arc<WebSession> { fn session() -> Arc<WebSession> {
+250
View File
@@ -0,0 +1,250 @@
use std::sync::Arc;
use std::time::Instant;
use sha2::{Digest, Sha256};
use super::uplink::{inbound_reservation, validate_batch};
use super::{CarrierLane, PendingClass, WebSession, inbound_queue_cost};
use crate::config::WebCarrier;
use crate::web::frame;
use crate::web::manager::ManagerError;
/// Pre-OPEN stream quota and synthetic tuple ownership for one WebSocket lane.
pub(crate) struct WebSocketLaneReservation {
session: Arc<WebSession>,
lane_id: u32,
peer_port: u16,
transferred: bool,
}
impl WebSocketLaneReservation {
/// Returns the logical stream owned by this connection.
pub(crate) fn lane_id(&self) -> u32 {
self.lane_id
}
fn transfer_to_stream(&mut self) {
let removed = self
.session
.state
.lock()
.websocket_lane_reservations
.remove(&self.lane_id);
if removed == Some(self.peer_port) {
self.transferred = true;
}
}
}
impl Drop for WebSocketLaneReservation {
fn drop(&mut self) {
if !self.transferred {
self.session
.release_websocket_lane_reservation(self.lane_id, self.peer_port);
}
}
}
impl WebSession {
/// Acquires stream quota and tuple ownership before a lane returns HTTP 101.
pub(crate) fn reserve_websocket_lane(
self: &Arc<Self>,
lane_id: u32,
) -> Result<WebSocketLaneReservation, ManagerError> {
if self.carrier() != WebCarrier::WebsocketLanes
|| lane_id == 0
|| lane_id > frame::MAX_STREAM_ID
{
return Err(ManagerError::Protocol);
}
let mut state = self.state.lock();
if state.closed {
return Err(ManagerError::Closed);
}
if state.active_peer_ports.len() >= self.profile.max_streams_per_session
|| state.streams.contains_key(&lane_id)
|| state.closed_streams.contains(&lane_id)
|| state.websocket_lane_reservations.contains_key(&lane_id)
{
return Err(ManagerError::Limit);
}
let Some(manager) = self.manager.upgrade() else {
return Err(ManagerError::Closed);
};
let Some(peer_port) = manager.try_acquire_stream(
self.profile_key,
self.profile.max_streams,
self.client_ip,
self.profile.public_addr,
) else {
return Err(ManagerError::Limit);
};
if !state.active_peer_ports.insert(peer_port) {
manager.release_stream(
self.profile_key,
self.client_ip,
self.profile.public_addr,
peer_port,
);
return Err(ManagerError::Limit);
}
state.websocket_lane_reservations.insert(lane_id, peer_port);
state.carrier_lanes.insert(lane_id, CarrierLane::new());
Ok(WebSocketLaneReservation {
session: Arc::clone(self),
lane_id,
peer_port,
transferred: false,
})
}
/// Applies one ordered WebSocket lane message without closing sibling lanes.
pub(crate) fn process_websocket_lane(
self: &Arc<Self>,
reservation: &mut WebSocketLaneReservation,
sequence: u64,
body: &[u8],
) -> Result<(), ManagerError> {
if !Arc::ptr_eq(self, &reservation.session)
|| reservation.lane_id == 0
|| reservation.lane_id > frame::MAX_STREAM_ID
{
return Err(ManagerError::Protocol);
}
let lane_id = reservation.lane_id;
let frames = frame::parse_all(body, &self.limits).map_err(|_| ManagerError::Protocol)?;
if frames
.iter()
.copied()
.any(|value| value.stream_id != lane_id || frame::validate_client_shape(value).is_err())
{
return Err(ManagerError::Protocol);
}
let digest = Sha256::digest(body).into();
let mut opened = Vec::new();
let result = {
let mut state = self.state.lock();
if state.closed {
return Err(ManagerError::Closed);
}
if !reservation.transferred
&& state.websocket_lane_reservations.get(&lane_id) != Some(&reservation.peer_port)
{
return Err(ManagerError::Closed);
}
let Some(lane) = state.carrier_lanes.get_mut(&lane_id) else {
return Err(ManagerError::Closed);
};
if sequence == 0 || sequence != lane.last_up_sequence.saturating_add(1) {
return Err(ManagerError::Protocol);
}
if lane.up_active {
return Err(ManagerError::Concurrent);
}
lane.up_active = true;
if !validate_batch(&state, &frames) {
if let Some(lane) = state.carrier_lanes.get_mut(&lane_id) {
lane.up_active = false;
}
return Err(ManagerError::Protocol);
}
let (reserve_bytes, reserve_items) = inbound_reservation(&state, &frames);
if !self.reserve_locked(
&mut state,
reserve_bytes,
reserve_items,
PendingClass::Uplink,
) {
if let Some(lane) = state.carrier_lanes.get_mut(&lane_id) {
lane.up_active = false;
}
return Err(ManagerError::Backpressure);
}
let mut unused_bytes = reserve_bytes;
let mut unused_items = reserve_items;
let mut reserved_open =
(!reservation.transferred).then_some((lane_id, reservation.peer_port));
let applied = self.apply_batch_locked(
&mut state,
&frames,
&mut opened,
&mut reserved_open,
&mut unused_bytes,
&mut unused_items,
);
self.release_locked(&mut state, unused_bytes, unused_items, false);
if let Some(lane) = state.carrier_lanes.get_mut(&lane_id) {
lane.up_active = false;
if applied {
lane.last_up_sequence = sequence;
lane.last_up_digest = digest;
}
}
state.last_activity = Instant::now();
applied.then_some(()).ok_or(ManagerError::Protocol)
};
result?;
for (stream_id, peer_port) in opened {
if stream_id != lane_id || peer_port != reservation.peer_port {
return Err(ManagerError::Protocol);
}
if !self.spawn_stream(stream_id, peer_port, true) {
return Err(ManagerError::Limit);
}
reservation.transfer_to_stream();
}
if !reservation.transferred {
return Err(ManagerError::Protocol);
}
if let Some(manager) = self.manager.upgrade() {
manager.record_up(body.len());
}
Ok(())
}
/// Ends one failed or disconnected lane without closing its parent session.
pub(crate) fn close_websocket_lane(&self, lane_id: u32) {
let reserved = {
let mut state = self.state.lock();
let reserved = state.websocket_lane_reservations.remove(&lane_id);
if let Some(stream) = state.streams.remove(&lane_id) {
let (bytes, items) = inbound_queue_cost(&stream.inbound);
self.release_locked(&mut state, bytes, items, false);
if let Some(waker) = stream.read_waker {
waker.wake();
}
if let Some(waker) = stream.write_waker {
waker.wake();
}
}
self.remember_closed_locked(&mut state, lane_id);
self.release_lane_locked(&mut state, lane_id);
reserved
};
if let Some(peer_port) = reserved {
self.release_websocket_lane_reservation(lane_id, peer_port);
}
}
fn release_websocket_lane_reservation(&self, lane_id: u32, peer_port: u16) {
let removed = {
let mut state = self.state.lock();
if state.websocket_lane_reservations.get(&lane_id) == Some(&peer_port) {
state.websocket_lane_reservations.remove(&lane_id);
}
self.release_lane_locked(&mut state, lane_id);
state.active_peer_ports.remove(&peer_port)
};
if removed && let Some(manager) = self.manager.upgrade() {
manager.release_stream(
self.profile_key,
self.client_ip,
self.profile.public_addr,
peer_port,
);
}
}
}
#[cfg(test)]
mod tests;
+135
View File
@@ -0,0 +1,135 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use arc_swap::ArcSwap;
use tokio::sync::watch;
use super::*;
use crate::config::{ProxyConfig, WebRuntimeConfig, WebRuntimeProfile, WebSecretMode};
use crate::maestro::generation::{RuntimeGeneration, test_runtime_generation_with_admission};
use crate::web::frame::FrameType;
use crate::web::manager::WebProcessRuntime;
struct TestRuntime {
session: Arc<WebSession>,
manager: Arc<WebProcessRuntime>,
generation: Arc<RuntimeGeneration>,
}
impl TestRuntime {
async fn shutdown(self) {
self.session.close();
self.session.wait().await;
self.manager.shutdown().await;
self.generation.stop_sessions().await;
self.generation.stop_background_tasks().await;
}
}
fn runtime(admission: bool) -> TestRuntime {
let profile = Arc::new(WebRuntimeProfile {
host: "proxy.example.com".to_string(),
public_addr: "203.0.113.10:443".parse().unwrap(),
user: "default".to_string(),
secret_mode: WebSecretMode::Plain,
carrier: WebCarrier::WebsocketLanes,
capability: [7; 32],
key_fingerprint: "0000000000000000".to_string(),
max_sessions: 2,
max_streams: 1,
max_streams_per_session: 1,
});
let mut config = ProxyConfig::default();
config.web.enabled = true;
config.web.carrier = WebCarrier::WebsocketLanes;
config.web.timeouts.shutdown_secs = 1;
config.web.runtime = Some(Arc::new(WebRuntimeConfig {
vhosts: BTreeMap::new(),
profiles: vec![Arc::clone(&profile)],
}));
config.rebuild_runtime_user_auth().unwrap();
let limits = config.web.limits.clone();
let timeouts = config.web.timeouts.clone();
let (_admission_tx, admission_rx) = watch::channel(admission);
let generation = test_runtime_generation_with_admission(1, config, admission_rx);
let manager = WebProcessRuntime::start(Arc::new(ArcSwap::from(Arc::clone(&generation))));
let session = WebSession::new(
Arc::downgrade(&manager),
[8; 32],
"192.0.2.10".parse().unwrap(),
1,
profile,
[7; 32],
limits,
timeouts,
);
TestRuntime {
session,
manager,
generation,
}
}
#[tokio::test]
async fn rejected_open_retains_stream_quota_until_lane_socket_teardown() {
let runtime = runtime(false);
let mut reservation = runtime.session.reserve_websocket_lane(7).unwrap();
let open = frame::encode(FrameType::Open, 7, &[]);
assert_eq!(
runtime
.session
.process_websocket_lane(&mut reservation, 1, &open),
Err(ManagerError::Limit),
);
assert!(
runtime
.manager
.try_acquire_stream(
runtime.session.profile_key,
runtime.session.profile.max_streams,
runtime.session.client_ip,
runtime.session.profile.public_addr,
)
.is_none()
);
runtime.session.close_websocket_lane(7);
drop(reservation);
let peer_port = runtime
.manager
.try_acquire_stream(
runtime.session.profile_key,
runtime.session.profile.max_streams,
runtime.session.client_ip,
runtime.session.profile.public_addr,
)
.unwrap();
runtime.manager.release_stream(
runtime.session.profile_key,
runtime.session.client_ip,
runtime.session.profile.public_addr,
peer_port,
);
runtime.shutdown().await;
}
#[tokio::test]
async fn malformed_lane_message_does_not_close_sibling_session_state() {
let runtime = runtime(true);
let mut reservation = runtime.session.reserve_websocket_lane(7).unwrap();
let data = frame::encode(FrameType::Data, 7, &[1]);
assert_eq!(
runtime
.session
.process_websocket_lane(&mut reservation, 1, &data),
Err(ManagerError::Protocol),
);
assert!(!runtime.session.state.lock().closed);
runtime.session.close_websocket_lane(7);
drop(reservation);
assert!(runtime.session.reserve_websocket_lane(8).is_ok());
runtime.shutdown().await;
}
+22 -22
View File
@@ -139,7 +139,10 @@ impl HttpTraceExchange {
/// Binds non-secret profile and process session identity. /// Binds non-secret profile and process session identity.
pub(crate) fn bind_profile(&self, profile: &WebRuntimeProfile, session_id: u64) { pub(crate) fn bind_profile(&self, profile: &WebRuntimeProfile, session_id: u64) {
let dynamic = profile.user.len().saturating_add(profile.key_fingerprint.len()); let dynamic = profile
.user
.len()
.saturating_add(profile.key_fingerprint.len());
let mut state = self.state.lock(); let mut state = self.state.lock();
state.identity.session_id = Some(session_id); state.identity.session_id = Some(session_id);
if self.reserve(dynamic) { if self.reserve(dynamic) {
@@ -150,12 +153,11 @@ impl HttpTraceExchange {
/// Binds an already resolved non-secret session identity. /// Binds an already resolved non-secret session identity.
pub(crate) fn bind_identity(&self, identity: TraceIdentity) { pub(crate) fn bind_identity(&self, identity: TraceIdentity) {
let dynamic = identity.user.as_ref().map_or(0, String::len).saturating_add( let dynamic = identity
identity .user
.key_fingerprint .as_ref()
.as_ref() .map_or(0, String::len)
.map_or(0, String::len), .saturating_add(identity.key_fingerprint.as_ref().map_or(0, String::len));
);
let mut state = self.state.lock(); let mut state = self.state.lock();
state.identity.session_id = identity.session_id; state.identity.session_id = identity.session_id;
if self.reserve(dynamic) { if self.reserve(dynamic) {
@@ -223,11 +225,8 @@ impl HttpTraceExchange {
body.truncated |= !data.is_empty(); body.truncated |= !data.is_empty();
return; return;
} }
let Some(limit) = capture_limit( let Some(limit) = capture_limit(&self.policy, route, self.store.max_carrier_body_bytes())
&self.policy, else {
route,
self.store.max_carrier_body_bytes(),
) else {
return; return;
}; };
if body.captured.len() >= limit { if body.captured.len() >= limit {
@@ -267,7 +266,9 @@ impl HttpTraceExchange {
if self.policy.capture_timings { if self.policy.capture_timings {
match direction { match direction {
TraceDirection::Request => state.timings.request_body_us = Some(self.elapsed_us()), TraceDirection::Request => state.timings.request_body_us = Some(self.elapsed_us()),
TraceDirection::Response => state.timings.response_body_us = Some(self.elapsed_us()), TraceDirection::Response => {
state.timings.response_body_us = Some(self.elapsed_us())
}
} }
} }
drop(state); drop(state);
@@ -386,10 +387,7 @@ impl HttpTraceExchange {
} }
fn elapsed_us(&self) -> u64 { fn elapsed_us(&self) -> u64 {
self.started self.started.elapsed().as_micros().min(u128::from(u64::MAX)) as u64
.elapsed()
.as_micros()
.min(u128::from(u64::MAX)) as u64
} }
} }
@@ -402,10 +400,7 @@ impl Drop for HttpTraceExchange {
} }
} }
fn body_snapshot( fn body_snapshot(policy: &WebDebugConfig, body: &mut BodyCapture) -> Option<TraceBodySnapshot> {
policy: &WebDebugConfig,
body: &mut BodyCapture,
) -> Option<TraceBodySnapshot> {
(policy.body_capture != WebDebugBodyCapture::Off).then(|| TraceBodySnapshot { (policy.body_capture != WebDebugBodyCapture::Off).then(|| TraceBodySnapshot {
observed_bytes: body.observed_bytes, observed_bytes: body.observed_bytes,
captured: std::mem::take(&mut body.captured), captured: std::mem::take(&mut body.captured),
@@ -475,7 +470,12 @@ mod tests {
let request_body = http.request_body.as_ref().unwrap(); let request_body = http.request_body.as_ref().unwrap();
let response_body = http.response_body.as_ref().unwrap(); let response_body = http.response_body.as_ref().unwrap();
for secret in [request_token.as_bytes(), capability.as_bytes()] { for secret in [request_token.as_bytes(), capability.as_bytes()] {
assert!(!request_body.captured.windows(secret.len()).any(|value| value == secret)); assert!(
!request_body
.captured
.windows(secret.len())
.any(|value| value == secret)
);
} }
assert!( assert!(
!response_body !response_body
+2 -2
View File
@@ -12,6 +12,6 @@ mod types;
pub(crate) use exchange::HttpTraceExchange; pub(crate) use exchange::HttpTraceExchange;
pub(crate) use store::{StoredTraceRecord, WebTraceStore, epoch_millis as store_epoch_millis}; pub(crate) use store::{StoredTraceRecord, WebTraceStore, epoch_millis as store_epoch_millis};
pub(crate) use types::{ pub(crate) use types::{
TraceBodySnapshot, TraceBodyState, TraceDirection, TraceHeader, TraceIdentity, TraceBodySnapshot, TraceBodyState, TraceDirection, TraceFrame, TraceHeader, TraceIdentity,
TraceLifecycleEvent, TraceRecord, TraceRecordKind, TraceRoute, TraceLifecycleEvent, TraceRecord, TraceRecordKind, TraceRoute, TraceWebSocketContext,
}; };
+4 -3
View File
@@ -24,7 +24,9 @@ pub(super) fn request_dynamic_bytes<B>(
request request
.headers() .headers()
.get(header::USER_AGENT) .get(header::USER_AGENT)
.map_or(0, |value| lossy_text_reservation(value.as_bytes(), USER_AGENT_MAX_BYTES)), .map_or(0, |value| {
lossy_text_reservation(value.as_bytes(), USER_AGENT_MAX_BYTES)
}),
) )
.saturating_add( .saturating_add(
policy policy
@@ -56,8 +58,7 @@ pub(super) fn sanitized_headers(headers: &hyper::HeaderMap) -> Vec<TraceHeader>
.iter() .iter()
.map(|(name, value)| TraceHeader { .map(|(name, value)| TraceHeader {
name: name.as_str().to_string(), name: name.as_str().to_string(),
value: header_value_allowed(name) value: header_value_allowed(name).then(|| bounded_text(value.as_bytes(), 4096)),
.then(|| bounded_text(value.as_bytes(), 4096)),
}) })
.collect() .collect()
} }
+10 -13
View File
@@ -13,6 +13,9 @@ use super::types::{
}; };
use crate::config::{WebDebugConfig, WebLimitsConfig}; use crate::config::{WebDebugConfig, WebLimitsConfig};
// WebSocket message capture is isolated from HTTP exchange storage.
mod websocket;
const BASE_RECORD_RESERVATION: usize = 1024; const BASE_RECORD_RESERVATION: usize = 1024;
struct RingState { struct RingState {
@@ -66,6 +69,7 @@ pub(crate) struct WebTraceStore {
records_capacity: usize, records_capacity: usize,
bytes_capacity: usize, bytes_capacity: usize,
max_carrier_body_bytes: usize, max_carrier_body_bytes: usize,
frame_limits: WebLimitsConfig,
used_bytes: Arc<AtomicUsize>, used_bytes: Arc<AtomicUsize>,
ring: Mutex<RingState>, ring: Mutex<RingState>,
next_record_seq: AtomicU64, next_record_seq: AtomicU64,
@@ -86,7 +90,8 @@ impl WebTraceStore {
epoch: AtomicU64::new(1), epoch: AtomicU64::new(1),
records_capacity: limits.debug_records_capacity, records_capacity: limits.debug_records_capacity,
bytes_capacity: limits.debug_bytes_global, bytes_capacity: limits.debug_bytes_global,
max_carrier_body_bytes: limits.max_body_bytes, max_carrier_body_bytes: limits.max_body_bytes.max(limits.carrier_batch_bytes),
frame_limits: limits.clone(),
used_bytes: Arc::new(AtomicUsize::new(0)), used_bytes: Arc::new(AtomicUsize::new(0)),
ring: Mutex::new(RingState { ring: Mutex::new(RingState {
records: VecDeque::with_capacity(limits.debug_records_capacity), records: VecDeque::with_capacity(limits.debug_records_capacity),
@@ -171,22 +176,14 @@ impl WebTraceStore {
.user .user
.as_ref() .as_ref()
.map_or(0, String::len) .map_or(0, String::len)
.checked_add( .checked_add(identity.key_fingerprint.as_ref().map_or(0, String::len));
identity let Some(reservation) =
.key_fingerprint identity_bytes.and_then(|bytes| BASE_RECORD_RESERVATION.checked_add(bytes))
.as_ref()
.map_or(0, String::len),
);
let Some(reservation) = identity_bytes
.and_then(|bytes| BASE_RECORD_RESERVATION.checked_add(bytes))
else { else {
self.record_truncation(); self.record_truncation();
return; return;
}; };
if !policy.enabled if !policy.enabled || !policy.capture_lifecycle || !self.try_reserve_record(reservation) {
|| !policy.capture_lifecycle
|| !self.try_reserve_record(reservation)
{
return; return;
} }
let record = TraceRecord { let record = TraceRecord {
+220
View File
@@ -0,0 +1,220 @@
use std::net::IpAddr;
use std::sync::atomic::Ordering;
use crate::config::WebDebugBodyCapture;
use crate::web::frame::{self, FrameType};
use super::super::types::{
TraceBodySnapshot, TraceBodyState, TraceDirection, TraceFrame, TraceIdentity, TraceRecord,
TraceRecordKind, TraceWebSocketContext, TraceWebSocketRecord,
};
use super::{BASE_RECORD_RESERVATION, WebTraceStore, epoch_millis};
impl WebTraceStore {
/// Builds connection metadata only while WebSocket debugging is enabled.
pub(crate) fn websocket_context<B, F>(
&self,
request: &hyper::Request<B>,
peer_ip: IpAddr,
effective_ip: IpAddr,
connection_id: u64,
lane_id: Option<u32>,
identity: F,
) -> Option<TraceWebSocketContext>
where
F: FnOnce() -> TraceIdentity,
{
if !self.enabled.load(Ordering::Acquire) || !self.policy.load().enabled {
return None;
}
let user_agent = request
.headers()
.get(hyper::header::USER_AGENT)
.map(|value| super::super::sanitize::bounded_text(value.as_bytes(), 512));
Some(TraceWebSocketContext {
connection_id,
peer_ip,
effective_ip,
user_agent,
identity: identity(),
lane_id,
})
}
/// Records one policy-bounded WebSocket message without retaining credentials.
pub(crate) fn record_websocket_message(
&self,
context: &TraceWebSocketContext,
direction: TraceDirection,
message_type: &'static str,
payload: &[u8],
duration_us: u64,
) {
if !self.enabled.load(Ordering::Acquire) {
return;
}
let epoch = self.epoch.load(Ordering::Acquire);
let policy = self.policy.load_full();
if !policy.enabled {
return;
}
let capture_limit = super::super::sanitize::capture_limit(
&policy,
super::super::types::TraceRoute::Websocket,
self.max_carrier_body_bytes,
)
.unwrap_or(0);
let capture_bytes = payload.len().min(capture_limit);
let frame_reservation = policy
.capture_frames
.then(|| {
payload
.len()
.div_ceil(frame::HEADER_BYTES)
.clamp(1, self.frame_limits.max_frames_per_body)
.saturating_mul(std::mem::size_of::<TraceFrame>())
})
.unwrap_or(0);
let identity_bytes = context
.identity
.user
.as_ref()
.map_or(0, String::len)
.saturating_add(
context
.identity
.key_fingerprint
.as_ref()
.map_or(0, String::len),
);
let reservation = BASE_RECORD_RESERVATION
.saturating_add(identity_bytes)
.saturating_add(context.user_agent.as_ref().map_or(0, String::len))
.saturating_add(capture_bytes)
.saturating_add(frame_reservation);
if !self.try_reserve_record(reservation) {
return;
}
let body = (policy.body_capture != WebDebugBodyCapture::Off).then(|| TraceBodySnapshot {
observed_bytes: payload.len() as u64,
captured: payload[..capture_bytes].to_vec(),
truncated: capture_limit != 0 && capture_bytes < payload.len(),
state: TraceBodyState::Complete,
});
let frames = if policy.capture_frames && message_type == "binary" {
websocket_frames(direction, payload, &self.frame_limits)
} else {
Vec::new()
};
let record = TraceRecord {
seq: self.next_record_seq(),
epoch_millis: epoch_millis(),
peer_ip: Some(context.peer_ip),
effective_ip: Some(context.effective_ip),
user_agent: context.user_agent.clone(),
identity: context.identity.clone(),
kind: TraceRecordKind::Websocket(TraceWebSocketRecord {
direction,
message_type,
payload_bytes: payload.len(),
body,
frames,
duration_us: policy.capture_timings.then_some(duration_us),
connection_id: context.connection_id,
lane_id: context.lane_id,
}),
};
if !self.try_commit(record, reservation, epoch) {
self.release(reservation);
}
}
}
fn websocket_frames(
direction: TraceDirection,
payload: &[u8],
limits: &crate::config::WebLimitsConfig,
) -> Vec<TraceFrame> {
match frame::parse_all(payload, limits) {
Ok(frames) => frames
.into_iter()
.map(|value| TraceFrame {
direction,
frame_type: Some(super::super::sanitize::frame_type_name(value.frame_type)),
stream_id: Some(value.stream_id),
payload_len: Some(value.payload.len()),
window_delta: (value.frame_type == FrameType::Window)
.then(|| frame::window_amount(value.payload).ok())
.flatten(),
parse_error: None,
})
.collect(),
Err(error) => vec![TraceFrame {
direction,
frame_type: None,
stream_id: None,
payload_len: None,
window_delta: None,
parse_error: Some(super::super::sanitize::frame_error_name(error)),
}],
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::{WebDebugConfig, WebLimitsConfig};
use crate::web::frame::FrameType;
#[test]
fn websocket_message_capture_retains_bounded_identity_body_timing_and_frames() {
let mut policy = WebDebugConfig::default();
policy.enabled = true;
policy.capture_frames = true;
policy.capture_timings = true;
policy.body_capture = WebDebugBodyCapture::Full;
let mut limits = WebLimitsConfig::default();
limits.debug_records_capacity = 4;
limits.debug_bytes_global = 64 * 1024;
let store = WebTraceStore::new(policy, &limits);
let request = hyper::Request::builder()
.header(hyper::header::USER_AGENT, "trace-client")
.body(())
.unwrap();
let context = store
.websocket_context(
&request,
"127.0.0.1".parse().unwrap(),
"192.0.2.10".parse().unwrap(),
17,
Some(7),
|| TraceIdentity {
session_id: Some(42),
user: Some("alice".to_string()),
key_fingerprint: Some("0123456789abcdef".to_string()),
},
)
.unwrap();
let payload = crate::web::frame::encode(FrameType::Pong, 0, &[]);
store.record_websocket_message(&context, TraceDirection::Request, "binary", &payload, 123);
let records = store.snapshot_matching(|_| true);
assert_eq!(records.len(), 1);
assert_eq!(
records[0].record.user_agent.as_deref(),
Some("trace-client")
);
let TraceRecordKind::Websocket(message) = &records[0].record.kind else {
panic!("expected WebSocket trace");
};
assert_eq!(message.connection_id, 17);
assert_eq!(message.lane_id, Some(7));
assert_eq!(message.duration_us, Some(123));
assert_eq!(
message.body.as_ref().unwrap().captured.as_slice(),
payload.as_ref()
);
assert_eq!(message.frames.len(), 1);
assert_eq!(message.frames[0].frame_type, Some("PONG"));
}
}
+43
View File
@@ -15,6 +15,8 @@ pub(crate) enum TraceRoute {
Uplink, Uplink,
/// Carrier downlink exchange. /// Carrier downlink exchange.
Downlink, Downlink,
/// WebSocket upgrade handshake.
Websocket,
} }
impl TraceRoute { impl TraceRoute {
@@ -27,6 +29,7 @@ impl TraceRoute {
Self::Session => "session", Self::Session => "session",
Self::Uplink => "uplink", Self::Uplink => "uplink",
Self::Downlink => "downlink", Self::Downlink => "downlink",
Self::Websocket => "websocket",
} }
} }
} }
@@ -180,6 +183,44 @@ pub(crate) struct TraceHttpRecord {
pub(crate) timings: Option<TraceTimings>, pub(crate) timings: Option<TraceTimings>,
} }
/// Stable non-secret metadata retained across one WebSocket connection.
#[derive(Clone, Debug)]
pub(crate) struct TraceWebSocketContext {
/// Process-unique connection identifier.
pub(crate) connection_id: u64,
/// Direct listener peer address.
pub(crate) peer_ip: IpAddr,
/// Trusted effective client address.
pub(crate) effective_ip: IpAddr,
/// Bounded user-agent copied only while debugging is enabled.
pub(crate) user_agent: Option<String>,
/// Session and profile identity without credentials.
pub(crate) identity: TraceIdentity,
/// Logical lane identifier for websocket-lanes.
pub(crate) lane_id: Option<u32>,
}
/// One bounded ordered WebSocket message observation.
#[derive(Debug)]
pub(crate) struct TraceWebSocketRecord {
/// Wire direction of this message.
pub(crate) direction: TraceDirection,
/// Closed RFC 6455 message category.
pub(crate) message_type: &'static str,
/// Total payload bytes observed.
pub(crate) payload_bytes: usize,
/// Policy-bounded message body observation.
pub(crate) body: Option<TraceBodySnapshot>,
/// Parsed carrier frames for binary messages.
pub(crate) frames: Vec<TraceFrame>,
/// Message processing or write duration when timing capture is enabled.
pub(crate) duration_us: Option<u64>,
/// Process-unique owning WebSocket connection.
pub(crate) connection_id: u64,
/// Logical lane identifier for websocket-lanes.
pub(crate) lane_id: Option<u32>,
}
/// Closed WEB lifecycle event category. /// Closed WEB lifecycle event category.
#[derive(Clone, Copy, Debug, PartialEq, Eq)] #[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum TraceLifecycleEvent { pub(crate) enum TraceLifecycleEvent {
@@ -257,6 +298,8 @@ pub(crate) struct TraceLifecycleRecord {
pub(crate) enum TraceRecordKind { pub(crate) enum TraceRecordKind {
/// HTTP request-to-response exchange. /// HTTP request-to-response exchange.
Http(TraceHttpRecord), Http(TraceHttpRecord),
/// One ordered WebSocket message.
Websocket(TraceWebSocketRecord),
/// Session or stream lifecycle event. /// Session or stream lifecycle event.
Lifecycle(TraceLifecycleRecord), Lifecycle(TraceLifecycleRecord),
} }