Carriers Auto-negotiation Aggressiveness

Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
This commit is contained in:
Alexey
2026-08-26 18:33:56 +03:00
parent 923c79796a
commit f180057973
10 changed files with 260 additions and 80 deletions
+10
View File
@@ -264,6 +264,7 @@ const WEB_CONFIG_KEYS: &[&str] = &[
"carrier", "carrier",
"carriers", "carriers",
"carrier_learning", "carrier_learning",
"carrier_negotiation_aggressiveness",
"debug", "debug",
"limits", "limits",
"timeouts", "timeouts",
@@ -278,10 +279,14 @@ 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",
"max_lane_open_waits_per_session",
"pending_bytes_per_lane",
"pending_items_per_lane",
"websocket_bytes_global", "websocket_bytes_global",
"websocket_admission_watermark_pct", "websocket_admission_watermark_pct",
"websocket_eviction_watermark_pct", "websocket_eviction_watermark_pct",
"websocket_http_connection_reserve", "websocket_http_connection_reserve",
"max_websocket_evictions_in_flight",
"max_carrier_learning_entries", "max_carrier_learning_entries",
"max_body_readers", "max_body_readers",
"max_body_bytes_global", "max_body_bytes_global",
@@ -332,7 +337,12 @@ const WEB_TIMEOUTS_CONFIG_KEYS: &[&str] = &[
"header_secs", "header_secs",
"body_secs", "body_secs",
"stream_handshake_secs", "stream_handshake_secs",
"stream_first_byte_secs",
"long_poll_secs", "long_poll_secs",
"lane_open_wait_secs",
"carrier_health_secs",
"websocket_upgrade_secs",
"websocket_open_secs",
"websocket_write_secs", "websocket_write_secs",
"websocket_backpressure_secs", "websocket_backpressure_secs",
"websocket_eviction_secs", "websocket_eviction_secs",
+14 -2
View File
@@ -72,9 +72,9 @@ pub(super) fn validate(config: &mut ProxyConfig) -> Result<()> {
validate_limits(&config.web.limits)?; validate_limits(&config.web.limits)?;
debug::validate(&config.web.debug, &config.web.limits)?; debug::validate(&config.web.debug, &config.web.limits)?;
let carriers = negotiation::validate(&config.web)?; let carriers = negotiation::validate(&config.web)?;
if carriers.contains(&WebCarrier::HttpsLanes) && config.web.limits.max_http_handlers < 2 { if carriers.contains(&WebCarrier::HttpsLanes) && config.web.limits.max_http_handlers < 4 {
return config_error( return config_error(
"WEB https-lanes candidates require web.limits.max_http_handlers >= 2", "WEB https-lanes candidates require web.limits.max_http_handlers >= 4",
); );
} }
timeouts::validate(&config.web.timeouts)?; timeouts::validate(&config.web.timeouts)?;
@@ -162,6 +162,16 @@ fn validate_limits(limits: &WebLimitsConfig) -> Result<()> {
let positive = [ let positive = [
("max_http_connections", limits.max_http_connections), ("max_http_connections", limits.max_http_connections),
("max_http_handlers", limits.max_http_handlers), ("max_http_handlers", limits.max_http_handlers),
(
"max_lane_open_waits_per_session",
limits.max_lane_open_waits_per_session,
),
("pending_bytes_per_lane", limits.pending_bytes_per_lane),
("pending_items_per_lane", limits.pending_items_per_lane),
(
"max_websocket_evictions_in_flight",
limits.max_websocket_evictions_in_flight,
),
( (
"max_carrier_learning_entries", "max_carrier_learning_entries",
limits.max_carrier_learning_entries, limits.max_carrier_learning_entries,
@@ -240,6 +250,8 @@ fn validate_limits(limits: &WebLimitsConfig) -> Result<()> {
|| limits.max_body_readers > limits.max_http_handlers || limits.max_body_readers > limits.max_http_handlers
|| limits.pending_bytes_per_session > limits.pending_bytes_global || limits.pending_bytes_per_session > limits.pending_bytes_global
|| limits.pending_items_per_session > limits.pending_items_global || limits.pending_items_per_session > limits.pending_items_global
|| limits.pending_bytes_per_lane > limits.pending_bytes_per_session
|| limits.pending_items_per_lane > limits.pending_items_per_session
|| limits.control_bytes_per_session > limits.control_bytes_global || limits.control_bytes_per_session > limits.control_bytes_global
|| limits.control_bytes_per_session > limits.pending_bytes_per_session || limits.control_bytes_per_session > limits.pending_bytes_per_session
|| limits.control_bytes_global > limits.pending_bytes_global || limits.control_bytes_global > limits.pending_bytes_global
@@ -14,6 +14,14 @@ pub(super) fn validate(config: &WebConfig) -> Result<Vec<WebCarrier>> {
} }
} }
let candidates = config.carrier_candidates(); let candidates = config.carrier_candidates();
if config.carrier_negotiation_enabled()
&& config.carrier_learning
&& config.limits.max_carrier_learning_entries < 3
{
return config_error(
"web.limits.max_carrier_learning_entries must be >= 3 when carrier learning is enabled",
);
}
if candidates.len() > WebCarrier::ALL.len() { if candidates.len() > WebCarrier::ALL.len() {
return config_error( return config_error(
"web.carriers and the web.carrier fallback must contain at most four carriers", "web.carriers and the web.carrier fallback must contain at most four carriers",
+5
View File
@@ -6,7 +6,12 @@ pub(super) fn validate(timeouts: &WebTimeoutsConfig) -> Result<()> {
("header_secs", timeouts.header_secs), ("header_secs", timeouts.header_secs),
("body_secs", timeouts.body_secs), ("body_secs", timeouts.body_secs),
("stream_handshake_secs", timeouts.stream_handshake_secs), ("stream_handshake_secs", timeouts.stream_handshake_secs),
("stream_first_byte_secs", timeouts.stream_first_byte_secs),
("long_poll_secs", timeouts.long_poll_secs), ("long_poll_secs", timeouts.long_poll_secs),
("lane_open_wait_secs", timeouts.lane_open_wait_secs),
("carrier_health_secs", timeouts.carrier_health_secs),
("websocket_upgrade_secs", timeouts.websocket_upgrade_secs),
("websocket_open_secs", timeouts.websocket_open_secs),
("websocket_write_secs", timeouts.websocket_write_secs), ("websocket_write_secs", timeouts.websocket_write_secs),
( (
"websocket_backpressure_secs", "websocket_backpressure_secs",
+2 -2
View File
@@ -53,8 +53,8 @@ pub use server::{
}; };
#[allow(unused_imports)] #[allow(unused_imports)]
pub use web::{ pub use web::{
WebConfig, WebDecoyConfig, WebLimitsConfig, WebProfileConfig, WebSecretMode, WebTimeoutsConfig, WebCarrierNegotiationAggressiveness, WebConfig, WebDecoyConfig, WebLimitsConfig,
WebVhostConfig, WebProfileConfig, WebSecretMode, WebTimeoutsConfig, WebVhostConfig,
}; };
#[allow(unused_imports)] #[allow(unused_imports)]
pub use web_carrier::{WebCarrier, WebCarriers}; pub use web_carrier::{WebCarrier, WebCarriers};
+55
View File
@@ -98,6 +98,15 @@ 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,
/// Per-session ceiling for downlink polls waiting for a lane OPEN.
#[serde(default = "default_web_max_lane_open_waits_per_session")]
pub max_lane_open_waits_per_session: usize,
/// Queued and in-flight downlink bytes allowed for one independent lane.
#[serde(default = "default_web_pending_bytes_per_lane")]
pub pending_bytes_per_lane: usize,
/// Queued and in-flight downlink items allowed for one independent lane.
#[serde(default = "default_web_pending_items_per_lane")]
pub pending_items_per_lane: usize,
/// Process-wide transient WebSocket byte sub-budget inside pending bytes. /// Process-wide transient WebSocket byte sub-budget inside pending bytes.
#[serde(default = "default_web_websocket_bytes_global")] #[serde(default = "default_web_websocket_bytes_global")]
pub websocket_bytes_global: usize, pub websocket_bytes_global: usize,
@@ -110,6 +119,9 @@ pub struct WebLimitsConfig {
/// Accepted HTTP connections that WebSocket upgrades must leave available. /// Accepted HTTP connections that WebSocket upgrades must leave available.
#[serde(default = "default_web_websocket_http_connection_reserve")] #[serde(default = "default_web_websocket_http_connection_reserve")]
pub websocket_http_connection_reserve: usize, pub websocket_http_connection_reserve: usize,
/// Concurrent pressure-eviction claims allowed process-wide.
#[serde(default = "default_web_max_websocket_evictions_in_flight")]
pub max_websocket_evictions_in_flight: usize,
/// Process-wide bounded carrier-learning evidence entry ceiling. /// Process-wide bounded carrier-learning evidence entry ceiling.
#[serde(default = "default_web_max_carrier_learning_entries")] #[serde(default = "default_web_max_carrier_learning_entries")]
pub max_carrier_learning_entries: usize, pub max_carrier_learning_entries: usize,
@@ -215,10 +227,15 @@ 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(),
max_lane_open_waits_per_session: default_web_max_lane_open_waits_per_session(),
pending_bytes_per_lane: default_web_pending_bytes_per_lane(),
pending_items_per_lane: default_web_pending_items_per_lane(),
websocket_bytes_global: default_web_websocket_bytes_global(), websocket_bytes_global: default_web_websocket_bytes_global(),
websocket_admission_watermark_pct: default_web_websocket_admission_watermark_pct(), websocket_admission_watermark_pct: default_web_websocket_admission_watermark_pct(),
websocket_eviction_watermark_pct: default_web_websocket_eviction_watermark_pct(), websocket_eviction_watermark_pct: default_web_websocket_eviction_watermark_pct(),
websocket_http_connection_reserve: default_web_websocket_http_connection_reserve(), websocket_http_connection_reserve: default_web_websocket_http_connection_reserve(),
max_websocket_evictions_in_flight:
default_web_max_websocket_evictions_in_flight(),
max_carrier_learning_entries: default_web_max_carrier_learning_entries(), max_carrier_learning_entries: default_web_max_carrier_learning_entries(),
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(),
@@ -266,9 +283,24 @@ pub struct WebTimeoutsConfig {
/// Deadline from the first inner byte through MTProxy authentication. /// Deadline from the first inner byte through MTProxy authentication.
#[serde(default = "default_web_stream_handshake_timeout_secs")] #[serde(default = "default_web_stream_handshake_timeout_secs")]
pub stream_handshake_secs: u64, pub stream_handshake_secs: u64,
/// Absolute deadline for receiving the first inner MTProxy byte.
#[serde(default = "default_web_stream_first_byte_secs")]
pub stream_first_byte_secs: u64,
/// 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,
/// Grace for a canonical downlink poll that races its lane OPEN.
#[serde(default = "default_web_lane_open_wait_secs")]
pub lane_open_wait_secs: u64,
/// Post-commit observation interval required before learning succeeds.
#[serde(default = "default_web_carrier_health_secs")]
pub carrier_health_secs: u64,
/// Maximum wait for Hyper to transfer an accepted WebSocket upgrade.
#[serde(default = "default_web_websocket_upgrade_secs")]
pub websocket_upgrade_secs: u64,
/// Absolute deadline for the first carrier binary message after upgrade.
#[serde(default = "default_web_websocket_open_secs")]
pub websocket_open_secs: u64,
/// Maximum wait for one WebSocket write to complete. /// Maximum wait for one WebSocket write to complete.
#[serde(default = "default_web_websocket_write_secs")] #[serde(default = "default_web_websocket_write_secs")]
pub websocket_write_secs: u64, pub websocket_write_secs: u64,
@@ -307,7 +339,12 @@ impl Default for WebTimeoutsConfig {
header_secs: default_web_header_timeout_secs(), header_secs: default_web_header_timeout_secs(),
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(),
stream_first_byte_secs: default_web_stream_first_byte_secs(),
long_poll_secs: default_web_long_poll_timeout_secs(), long_poll_secs: default_web_long_poll_timeout_secs(),
lane_open_wait_secs: default_web_lane_open_wait_secs(),
carrier_health_secs: default_web_carrier_health_secs(),
websocket_upgrade_secs: default_web_websocket_upgrade_secs(),
websocket_open_secs: default_web_websocket_open_secs(),
websocket_write_secs: default_web_websocket_write_secs(), websocket_write_secs: default_web_websocket_write_secs(),
websocket_backpressure_secs: default_web_websocket_backpressure_secs(), websocket_backpressure_secs: default_web_websocket_backpressure_secs(),
websocket_eviction_secs: default_web_websocket_eviction_secs(), websocket_eviction_secs: default_web_websocket_eviction_secs(),
@@ -323,6 +360,19 @@ impl Default for WebTimeoutsConfig {
} }
} }
/// Sensitivity of process-local carrier-learning evidence.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum WebCarrierNegotiationAggressiveness {
/// Require broad evidence and never rank by client IP.
#[default]
Conservative,
/// Use moderate User-Agent, client-IP, and profile thresholds.
Balanced,
/// React to the first bounded evidence sample.
Aggressive,
}
/// WEB ingress, carrier, fallback, and lifecycle configuration. /// WEB ingress, carrier, fallback, and lifecycle configuration.
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WebConfig { pub struct WebConfig {
@@ -338,6 +388,9 @@ pub struct WebConfig {
/// Enables bounded process-local carrier learning for automatic sessions. /// Enables bounded process-local carrier learning for automatic sessions.
#[serde(default = "default_web_carrier_learning")] #[serde(default = "default_web_carrier_learning")]
pub carrier_learning: bool, pub carrier_learning: bool,
/// Controls the evidence thresholds used by automatic carrier ranking.
#[serde(default)]
pub carrier_negotiation_aggressiveness: WebCarrierNegotiationAggressiveness,
/// Hard process and protocol limits. /// Hard process and protocol limits.
#[serde(default)] #[serde(default)]
pub limits: WebLimitsConfig, pub limits: WebLimitsConfig,
@@ -381,6 +434,8 @@ impl Default for WebConfig {
carrier: WebCarrier::default(), carrier: WebCarrier::default(),
carriers: WebCarriers::default(), carriers: WebCarriers::default(),
carrier_learning: default_web_carrier_learning(), carrier_learning: default_web_carrier_learning(),
carrier_negotiation_aggressiveness:
WebCarrierNegotiationAggressiveness::default(),
limits: WebLimitsConfig::default(), limits: WebLimitsConfig::default(),
debug: WebDebugConfig::default(), debug: WebDebugConfig::default(),
timeouts: WebTimeoutsConfig::default(), timeouts: WebTimeoutsConfig::default(),
+9
View File
@@ -41,10 +41,14 @@ 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_max_lane_open_waits_per_session, 16);
usize_default!(default_web_pending_bytes_per_lane, 8 * 1024 * 1024);
usize_default!(default_web_pending_items_per_lane, 1024);
usize_default!(default_web_websocket_bytes_global, 256 * 1024 * 1024); usize_default!(default_web_websocket_bytes_global, 256 * 1024 * 1024);
u8_default!(default_web_websocket_admission_watermark_pct, 75); u8_default!(default_web_websocket_admission_watermark_pct, 75);
u8_default!(default_web_websocket_eviction_watermark_pct, 90); u8_default!(default_web_websocket_eviction_watermark_pct, 90);
usize_default!(default_web_websocket_http_connection_reserve, 64); usize_default!(default_web_websocket_http_connection_reserve, 64);
usize_default!(default_web_max_websocket_evictions_in_flight, 8);
usize_default!(default_web_max_carrier_learning_entries, 4096); usize_default!(default_web_max_carrier_learning_entries, 4096);
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);
@@ -79,7 +83,12 @@ u32_default!(default_web_new_streams_burst, 512);
u64_default!(default_web_header_timeout_secs, 10); 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_stream_first_byte_secs, 30);
u64_default!(default_web_long_poll_timeout_secs, 25); u64_default!(default_web_long_poll_timeout_secs, 25);
u64_default!(default_web_lane_open_wait_secs, 2);
u64_default!(default_web_carrier_health_secs, 30);
u64_default!(default_web_websocket_upgrade_secs, 5);
u64_default!(default_web_websocket_open_secs, 15);
u64_default!(default_web_websocket_write_secs, 30); u64_default!(default_web_websocket_write_secs, 30);
u64_default!(default_web_websocket_backpressure_secs, 30); u64_default!(default_web_websocket_backpressure_secs, 30);
u64_default!(default_web_websocket_eviction_secs, 1); u64_default!(default_web_websocket_eviction_secs, 1);
+51 -18
View File
@@ -46,7 +46,14 @@ struct InboundChunk {
offset: usize, offset: usize,
} }
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) struct StreamIdentity {
pub(crate) id: u32,
pub(crate) instance: u64,
}
struct StreamState { struct StreamState {
instance: u64,
inbound: VecDeque<InboundChunk>, inbound: VecDeque<InboundChunk>,
receive_window: u32, receive_window: u32,
send_credit: u64, send_credit: u64,
@@ -73,6 +80,7 @@ struct DownBatch {
} }
struct CarrierLane { struct CarrierLane {
instance: u64,
pending_frames: VecDeque<QueuedFrame>, pending_frames: VecDeque<QueuedFrame>,
pending_windows: HashMap<u32, usize>, pending_windows: HashMap<u32, usize>,
unacked: Option<DownBatch>, unacked: Option<DownBatch>,
@@ -85,8 +93,9 @@ struct CarrierLane {
} }
impl CarrierLane { impl CarrierLane {
fn new() -> Self { fn new(instance: u64) -> Self {
Self { Self {
instance,
pending_frames: VecDeque::new(), pending_frames: VecDeque::new(),
pending_windows: HashMap::new(), pending_windows: HashMap::new(),
unacked: None, unacked: None,
@@ -102,6 +111,8 @@ impl CarrierLane {
struct SessionState { struct SessionState {
streams: HashMap<u32, StreamState>, streams: HashMap<u32, StreamState>,
closing_streams: HashMap<u32, u64>,
next_stream_instance: u64,
active_peer_ports: HashSet<u16>, active_peer_ports: HashSet<u16>,
closed_streams: HashSet<u32>, closed_streams: HashSet<u32>,
closed_order: VecDeque<u32>, closed_order: VecDeque<u32>,
@@ -113,6 +124,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>,
next_lane_instance: u64,
websocket_lane_reservations: HashMap<u32, u16>, websocket_lane_reservations: HashMap<u32, u16>,
pending_bytes: usize, pending_bytes: usize,
pending_items: usize, pending_items: usize,
@@ -183,8 +195,10 @@ impl WebSession {
timeouts: WebTimeoutsConfig, timeouts: WebTimeoutsConfig,
) -> Arc<Self> { ) -> Arc<Self> {
let mut carrier_lanes = HashMap::new(); let mut carrier_lanes = HashMap::new();
let mut next_lane_instance = 1;
if selected_carrier == WebCarrier::HttpsLanes { if selected_carrier == WebCarrier::HttpsLanes {
carrier_lanes.insert(0, CarrierLane::new()); carrier_lanes.insert(0, CarrierLane::new(next_lane_instance));
next_lane_instance += 1;
} }
Arc::new(Self { Arc::new(Self {
manager, manager,
@@ -201,6 +215,8 @@ impl WebSession {
timeouts, timeouts,
state: Mutex::new(SessionState { state: Mutex::new(SessionState {
streams: HashMap::new(), streams: HashMap::new(),
closing_streams: HashMap::new(),
next_stream_instance: 1,
active_peer_ports: HashSet::new(), active_peer_ports: HashSet::new(),
closed_streams: HashSet::new(), closed_streams: HashSet::new(),
closed_order: VecDeque::new(), closed_order: VecDeque::new(),
@@ -212,6 +228,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,
next_lane_instance,
websocket_lane_reservations: HashMap::new(), websocket_lane_reservations: HashMap::new(),
pending_bytes: 0, pending_bytes: 0,
pending_items: 0, pending_items: 0,
@@ -329,17 +346,21 @@ impl WebSession {
/// Polls client-to-server bytes and returns consumed flow-control credit. /// Polls client-to-server bytes and returns consumed flow-control credit.
pub(super) fn poll_read( pub(super) fn poll_read(
&self, &self,
stream_id: u32, stream: StreamIdentity,
cx: &mut Context<'_>, cx: &mut Context<'_>,
output: &mut ReadBuf<'_>, output: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> { ) -> Poll<io::Result<()>> {
let mut state = self.state.lock(); let mut state = self.state.lock();
let (count, finished) = { let (count, finished) = {
let Some(stream) = state.streams.get_mut(&stream_id) else { let Some(stream_state) = state
.streams
.get_mut(&stream.id)
.filter(|state| state.instance == stream.instance)
else {
return Poll::Ready(Ok(())); return Poll::Ready(Ok(()));
}; };
let Some(chunk) = stream.inbound.front_mut() else { let Some(chunk) = stream_state.inbound.front_mut() else {
stream.read_waker = Some(cx.waker().clone()); stream_state.read_waker = Some(cx.waker().clone());
return Poll::Pending; return Poll::Pending;
}; };
let available = &chunk.bytes[chunk.offset..]; let available = &chunk.bytes[chunk.offset..];
@@ -348,14 +369,14 @@ impl WebSession {
chunk.offset += count; chunk.offset += count;
let finished = chunk.offset == chunk.bytes.len(); let finished = chunk.offset == chunk.bytes.len();
if finished { if finished {
stream.inbound.pop_front(); stream_state.inbound.pop_front();
} }
stream.receive_window = stream.receive_window.saturating_add(count as u32); stream_state.receive_window = stream_state.receive_window.saturating_add(count as u32);
(count, finished) (count, finished)
}; };
let overhead = if finished { QUEUE_ITEM_COST } else { 0 }; let overhead = if finished { QUEUE_ITEM_COST } else { 0 };
self.release_locked(&mut state, count + overhead, usize::from(finished), false); self.release_locked(&mut state, count + overhead, usize::from(finished), false);
if !self.queue_window_locked(&mut state, stream_id, count as u32) { if !self.queue_window_locked(&mut state, stream.id, count as u32) {
drop(state); drop(state);
self.close(); self.close();
return Poll::Ready(Err(io::Error::other( return Poll::Ready(Err(io::Error::other(
@@ -368,7 +389,7 @@ impl WebSession {
/// Polls server-to-client writes against stream credit and bounded queues. /// Polls server-to-client writes against stream credit and bounded queues.
pub(super) fn poll_write( pub(super) fn poll_write(
&self, &self,
stream_id: u32, stream: StreamIdentity,
cx: &mut Context<'_>, cx: &mut Context<'_>,
input: &[u8], input: &[u8],
) -> Poll<io::Result<usize>> { ) -> Poll<io::Result<usize>> {
@@ -376,7 +397,11 @@ impl WebSession {
return Poll::Ready(Ok(0)); return Poll::Ready(Ok(0));
} }
let mut state = self.state.lock(); let mut state = self.state.lock();
let Some(stream) = state.streams.get_mut(&stream_id) else { let Some(stream_state) = state
.streams
.get_mut(&stream.id)
.filter(|state| state.instance == stream.instance)
else {
return Poll::Ready(Err(io::Error::new( return Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe, io::ErrorKind::BrokenPipe,
"WEB logical stream is closed", "WEB logical stream is closed",
@@ -386,24 +411,32 @@ impl WebSession {
.len() .len()
.min(frame::DATA_CHUNK_BYTES) .min(frame::DATA_CHUNK_BYTES)
.min(self.limits.max_frame_payload_bytes) .min(self.limits.max_frame_payload_bytes)
.min(stream.send_credit as usize); .min(stream_state.send_credit as usize);
if count == 0 { if count == 0 {
stream.write_waker = Some(cx.waker().clone()); stream_state.write_waker = Some(cx.waker().clone());
return Poll::Pending; return Poll::Pending;
} }
if !self.queue_data_locked(&mut state, stream_id, &input[..count]) { if !self.queue_data_locked(&mut state, stream.id, &input[..count]) {
if let Some(stream) = state.streams.get_mut(&stream_id) { if let Some(stream_state) = state
stream.write_waker = Some(cx.waker().clone()); .streams
.get_mut(&stream.id)
.filter(|state| state.instance == stream.instance)
{
stream_state.write_waker = Some(cx.waker().clone());
} }
return Poll::Pending; return Poll::Pending;
} }
let Some(stream) = state.streams.get_mut(&stream_id) else { let Some(stream_state) = state
.streams
.get_mut(&stream.id)
.filter(|state| state.instance == stream.instance)
else {
return Poll::Ready(Err(io::Error::new( return Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe, io::ErrorKind::BrokenPipe,
"WEB logical stream is closed", "WEB logical stream is closed",
))); )));
}; };
stream.send_credit -= count as u64; stream_state.send_credit -= count as u64;
state.last_activity = Instant::now(); state.last_activity = Instant::now();
drop(state); drop(state);
if self.carrier().is_multiplexed() { if self.carrier().is_multiplexed() {
+99 -51
View File
@@ -7,7 +7,7 @@ use crate::proxy::shared_state::ConntrackClosePolicy;
use crate::web::frame::FrameType; use crate::web::frame::FrameType;
use crate::web::stream::WebLogicalStream; use crate::web::stream::WebLogicalStream;
use super::{WebSession, inbound_queue_cost}; use super::{StreamIdentity, WebSession, inbound_queue_cost};
#[cfg(test)] #[cfg(test)]
#[path = "backend_tests.rs"] #[path = "backend_tests.rs"]
@@ -17,61 +17,64 @@ 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( pub(super) fn spawn_stream(
self: &Arc<Self>, self: &Arc<Self>,
stream_id: u32, completion: StreamCompletion,
peer_port: u16,
retain_reservation_on_reject: bool, retain_reservation_on_reject: bool,
) -> bool { ) -> bool {
let stream = completion.stream;
let peer_port = completion.peer_port;
let Some(manager) = self.manager.upgrade() else { let Some(manager) = self.manager.upgrade() else {
self.stream_rejected_before_spawn(stream_id, peer_port, retain_reservation_on_reject); completion
.retain_rejected
.store(retain_reservation_on_reject, Ordering::Release);
drop(completion);
return false; return false;
}; };
let generation = manager.active_generation(); let generation = manager.active_generation();
if !*generation.admission_rx.borrow() { if !*generation.admission_rx.borrow() {
self.trace_lifecycle( self.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::StreamRejected, crate::web::trace::TraceLifecycleEvent::StreamRejected,
Some(stream_id), Some(stream.id),
Some("admission_closed"), Some("admission_closed"),
); );
self.stream_rejected_before_spawn(stream_id, peer_port, retain_reservation_on_reject); completion
.retain_rejected
.store(retain_reservation_on_reject, Ordering::Release);
drop(completion);
return false; 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();
self.trace_lifecycle( self.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::StreamRejected, crate::web::trace::TraceLifecycleEvent::StreamRejected,
Some(stream_id), Some(stream.id),
Some("connection_limit"), Some("connection_limit"),
); );
self.stream_rejected_before_spawn(stream_id, peer_port, retain_reservation_on_reject); completion
.retain_rejected
.store(retain_reservation_on_reject, Ordering::Release);
drop(completion);
return false; 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)); let retain_rejected = Arc::clone(&completion.retain_rejected);
self.tasks_live.fetch_add(1, Ordering::AcqRel);
let completion = StreamCompletion {
session: Arc::clone(&session),
stream_id,
peer_port,
retain_rejected: Arc::clone(&retain_rejected),
};
let future = async move { let future = async move {
let _connection_permit = connection_permit; let _connection_permit = connection_permit;
let _completion = completion; let _completion = completion;
session.trace_lifecycle( session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::StreamAdmitted, crate::web::trace::TraceLifecycleEvent::StreamAdmitted,
Some(stream_id), Some(stream.id),
None, None,
); );
let stream = WebLogicalStream::new(Arc::clone(&session), stream_id); let logical_stream = WebLogicalStream::new(Arc::clone(&session), stream);
tokio::select! { tokio::select! {
_ = cancel.cancelled() => {} _ = cancel.cancelled() => {}
_ = run_stream( _ = run_stream(
Arc::clone(&session), Arc::clone(&session),
stream_id,
stream, stream,
logical_stream,
deps, deps,
replay_checker, replay_checker,
peer_port, peer_port,
@@ -82,7 +85,7 @@ impl WebSession {
retain_rejected.store(retain_reservation_on_reject, Ordering::Release); 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"),
); );
drop(future); drop(future);
@@ -93,21 +96,28 @@ impl WebSession {
fn stream_rejected_before_spawn( fn stream_rejected_before_spawn(
&self, &self,
stream_id: u32, stream: StreamIdentity,
peer_port: u16, peer_port: u16,
retain_reservation: bool, retain_reservation: bool,
) { ) {
if !retain_reservation { if !retain_reservation {
self.stream_finished(stream_id, peer_port); self.stream_finished(stream, peer_port);
return; return;
} }
let queued = { let queued = {
let mut state = self.state.lock(); let mut state = self.state.lock();
state.streams.remove(&stream_id).map(|stream| { state
let (bytes, items) = inbound_queue_cost(&stream.inbound); .streams
.get(&stream.id)
.filter(|state| state.instance == stream.instance)
.is_some()
.then(|| state.streams.remove(&stream.id))
.flatten()
.map(|stream_state| {
let (bytes, items) = inbound_queue_cost(&stream_state.inbound);
self.release_locked(&mut state, bytes, items, false); self.release_locked(&mut state, bytes, items, false);
self.remember_closed_locked(&mut state, stream_id); self.remember_closed_locked(&mut state, stream.id);
self.queue_control_locked(&mut state, FrameType::Close, stream_id, &[]) self.queue_control_locked(&mut state, FrameType::Close, stream.id, &[])
}) })
}; };
if queued.is_some_and(|queued| !queued) { if queued.is_some_and(|queued| !queued) {
@@ -115,16 +125,24 @@ impl WebSession {
} }
} }
fn stream_finished(&self, stream_id: u32, peer_port: u16) { fn stream_finished(&self, stream: StreamIdentity, peer_port: u16) {
let (queued, reserved) = { let (queued, reserved) = {
let mut state = self.state.lock(); let mut state = self.state.lock();
let reserved = state.active_peer_ports.remove(&peer_port); let reserved = state.active_peer_ports.remove(&peer_port);
let queued = state.streams.remove(&stream_id).map(|stream| { let current = state
let (bytes, items) = inbound_queue_cost(&stream.inbound); .streams
.get(&stream.id)
.is_some_and(|state| state.instance == stream.instance);
let queued = current.then(|| state.streams.remove(&stream.id)).flatten().map(|stream_state| {
let (bytes, items) = inbound_queue_cost(&stream_state.inbound);
self.release_locked(&mut state, bytes, items, false); self.release_locked(&mut state, bytes, items, false);
self.remember_closed_locked(&mut state, stream_id); self.remember_closed_locked(&mut state, stream.id);
self.queue_control_locked(&mut state, FrameType::Close, stream_id, &[]) self.queue_control_locked(&mut state, FrameType::Close, stream.id, &[])
}); });
if state.closing_streams.get(&stream.id) == Some(&stream.instance) {
state.closing_streams.remove(&stream.id);
self.remember_closed_locked(&mut state, stream.id);
}
(queued, reserved) (queued, reserved)
}; };
if reserved && let Some(manager) = self.manager.upgrade() { if reserved && let Some(manager) = self.manager.upgrade() {
@@ -146,25 +164,41 @@ impl WebSession {
} }
} }
struct StreamCompletion { pub(super) struct StreamCompletion {
session: Arc<WebSession>, session: Arc<WebSession>,
stream_id: u32, pub(super) stream: StreamIdentity,
peer_port: u16, pub(super) peer_port: u16,
retain_rejected: Arc<AtomicBool>, retain_rejected: Arc<AtomicBool>,
} }
impl WebSession {
pub(super) fn own_stream_task(
self: &Arc<Self>,
stream: StreamIdentity,
peer_port: u16,
) -> StreamCompletion {
self.tasks_live.fetch_add(1, Ordering::AcqRel);
StreamCompletion {
session: Arc::clone(self),
stream,
peer_port,
retain_rejected: Arc::new(AtomicBool::new(false)),
}
}
}
impl Drop for StreamCompletion { impl Drop for StreamCompletion {
fn drop(&mut self) { fn drop(&mut self) {
self.session.trace_lifecycle( self.session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::StreamClosed, crate::web::trace::TraceLifecycleEvent::StreamClosed,
Some(self.stream_id), Some(self.stream.id),
None, None,
); );
if self.retain_rejected.load(Ordering::Acquire) { if self.retain_rejected.load(Ordering::Acquire) {
self.session self.session
.stream_rejected_before_spawn(self.stream_id, self.peer_port, true); .stream_rejected_before_spawn(self.stream, self.peer_port, true);
} else { } else {
self.session.stream_finished(self.stream_id, self.peer_port); self.session.stream_finished(self.stream, 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();
@@ -174,7 +208,7 @@ impl Drop for StreamCompletion {
async fn run_stream( async fn run_stream(
session: Arc<WebSession>, session: Arc<WebSession>,
stream_id: u32, stream_identity: StreamIdentity,
stream: WebLogicalStream, stream: WebLogicalStream,
deps: crate::proxy::authenticated::ClientRuntimeDeps, deps: crate::proxy::authenticated::ClientRuntimeDeps,
replay_checker: Arc<crate::stats::ReplayChecker>, replay_checker: Arc<crate::stats::ReplayChecker>,
@@ -191,13 +225,27 @@ async fn run_stream(
let peer = std::net::SocketAddr::new(session.client_ip, peer_port); let peer = std::net::SocketAddr::new(session.client_ip, peer_port);
deps.stats.increment_connects_all(); deps.stats.increment_connects_all();
// A carrier may publish OPEN before the local MTProto socket writes its // Silent OPEN ownership has a separate absolute deadline so it cannot
// first byte. Session and stream quotas bound this idle phase without // consume stream and tuple quotas indefinitely before handshake admission.
// consuming the process-wide active-handshake budget. let first_byte = tokio::time::timeout(
if reader.read_exact(&mut handshake[..1]).await.is_err() { Duration::from_secs(session.timeouts.stream_first_byte_secs),
reader.read_exact(&mut handshake[..1]),
)
.await;
if first_byte.is_err() {
session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::HandshakeTimeout,
Some(stream_identity.id),
Some("first_byte_timeout"),
);
deps.stats
.increment_connects_bad_with_class("web_mtproto_first_byte_timeout");
return;
}
if first_byte.is_ok_and(|result| result.is_err()) {
session.trace_lifecycle( session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::HandshakeIo, crate::web::trace::TraceLifecycleEvent::HandshakeIo,
Some(stream_id), Some(stream_identity.id),
Some("first_byte_io"), Some("first_byte_io"),
); );
deps.stats deps.stats
@@ -206,7 +254,7 @@ async fn run_stream(
} }
session.trace_lifecycle( session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::StreamFirstByte, crate::web::trace::TraceLifecycleEvent::StreamFirstByte,
Some(stream_id), Some(stream_identity.id),
None, None,
); );
let Some(manager) = session.manager.upgrade() else { let Some(manager) = session.manager.upgrade() else {
@@ -215,7 +263,7 @@ async fn run_stream(
let Some(handshake_permit) = manager.try_stream_handshake() else { let Some(handshake_permit) = manager.try_stream_handshake() else {
session.trace_lifecycle( session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::StreamRejected, crate::web::trace::TraceLifecycleEvent::StreamRejected,
Some(stream_id), Some(stream_identity.id),
Some("handshake_limit"), Some("handshake_limit"),
); );
return; return;
@@ -246,7 +294,7 @@ async fn run_stream(
Err(_) => { Err(_) => {
session.trace_lifecycle( session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::HandshakeTimeout, crate::web::trace::TraceLifecycleEvent::HandshakeTimeout,
Some(stream_id), Some(stream_identity.id),
Some("timeout"), Some("timeout"),
); );
deps.stats deps.stats
@@ -258,7 +306,7 @@ async fn run_stream(
Ok(Err(_)) => { Ok(Err(_)) => {
session.trace_lifecycle( session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::HandshakeIo, crate::web::trace::TraceLifecycleEvent::HandshakeIo,
Some(stream_id), Some(stream_identity.id),
Some("io"), Some("io"),
); );
deps.stats deps.stats
@@ -268,7 +316,7 @@ async fn run_stream(
Ok(Ok(crate::error::HandshakeResult::Success((reader, writer, success)))) => { Ok(Ok(crate::error::HandshakeResult::Success((reader, writer, success)))) => {
session.trace_lifecycle( session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::HandshakeSucceeded, crate::web::trace::TraceLifecycleEvent::HandshakeSucceeded,
Some(stream_id), Some(stream_identity.id),
None, None,
); );
(reader, writer, success) (reader, writer, success)
@@ -276,7 +324,7 @@ async fn run_stream(
Ok(Ok(_)) => { Ok(Ok(_)) => {
session.trace_lifecycle( session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::HandshakeRejected, crate::web::trace::TraceLifecycleEvent::HandshakeRejected,
Some(stream_id), Some(stream_identity.id),
Some("bad_client"), Some("bad_client"),
); );
deps.stats deps.stats
@@ -286,7 +334,7 @@ async fn run_stream(
}; };
session.trace_lifecycle( session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::RelayStarted, crate::web::trace::TraceLifecycleEvent::RelayStarted,
Some(stream_id), Some(stream_identity.id),
None, None,
); );
let relay_result = run_authenticated( let relay_result = run_authenticated(
@@ -301,7 +349,7 @@ async fn run_stream(
.await; .await;
session.trace_lifecycle( session.trace_lifecycle(
crate::web::trace::TraceLifecycleEvent::RelayEnded, crate::web::trace::TraceLifecycleEvent::RelayEnded,
Some(stream_id), Some(stream_identity.id),
Some(if relay_result.is_ok() { Some(if relay_result.is_ok() {
"completed" "completed"
} else { } else {
+7 -7
View File
@@ -7,21 +7,21 @@ use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::sync::futures::OwnedNotified; use tokio::sync::futures::OwnedNotified;
use crate::web::session::WebSession; use crate::web::session::{StreamIdentity, WebSession};
/// Async byte stream that maps one WEB stream identifier onto carrier frames. /// Async byte stream that maps one WEB stream identifier onto carrier frames.
pub(crate) struct WebLogicalStream { pub(crate) struct WebLogicalStream {
session: Arc<WebSession>, session: Arc<WebSession>,
stream_id: u32, stream: StreamIdentity,
budget_wait: Option<Pin<Box<OwnedNotified>>>, budget_wait: Option<Pin<Box<OwnedNotified>>>,
} }
impl WebLogicalStream { impl WebLogicalStream {
/// Binds a virtual byte stream to one live carrier stream identifier. /// Binds a virtual byte stream to one live carrier stream identifier.
pub(crate) fn new(session: Arc<WebSession>, stream_id: u32) -> Self { pub(crate) fn new(session: Arc<WebSession>, stream: StreamIdentity) -> Self {
Self { Self {
session, session,
stream_id, stream,
budget_wait: None, budget_wait: None,
} }
} }
@@ -33,7 +33,7 @@ impl AsyncRead for WebLogicalStream {
cx: &mut Context<'_>, cx: &mut Context<'_>,
output: &mut ReadBuf<'_>, output: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> { ) -> Poll<io::Result<()>> {
self.session.poll_read(self.stream_id, cx, output) self.session.poll_read(self.stream, cx, output)
} }
} }
@@ -43,7 +43,7 @@ impl AsyncWrite for WebLogicalStream {
cx: &mut Context<'_>, cx: &mut Context<'_>,
input: &[u8], input: &[u8],
) -> Poll<io::Result<usize>> { ) -> Poll<io::Result<usize>> {
let result = self.session.poll_write(self.stream_id, cx, input); let result = self.session.poll_write(self.stream, cx, input);
if !result.is_pending() { if !result.is_pending() {
self.budget_wait = None; self.budget_wait = None;
return result; return result;
@@ -64,7 +64,7 @@ impl AsyncWrite for WebLogicalStream {
} }
self.budget_wait = None; self.budget_wait = None;
} }
match self.session.poll_write(self.stream_id, cx, input) { match self.session.poll_write(self.stream, cx, input) {
Poll::Ready(result) => { Poll::Ready(result) => {
self.budget_wait = None; self.budget_wait = None;
Poll::Ready(result) Poll::Ready(result)