From f1800579739778f457cbb0fa34161a4dc8b97e9c Mon Sep 17 00:00:00 2001 From: Alexey <247128645+axkurcom@users.noreply.github.com> Date: Wed, 26 Aug 2026 18:33:56 +0300 Subject: [PATCH] Carriers Auto-negotiation Aggressiveness Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com> --- src/config/load/strict_keys.rs | 10 ++ src/config/load/validate_web.rs | 16 ++- src/config/load/validate_web/negotiation.rs | 8 ++ src/config/load/validate_web/timeouts.rs | 5 + src/config/types.rs | 4 +- src/config/types/web.rs | 55 +++++++ src/config/types/web/defaults.rs | 9 ++ src/web/session.rs | 69 ++++++--- src/web/session/backend.rs | 150 +++++++++++++------- src/web/stream.rs | 14 +- 10 files changed, 260 insertions(+), 80 deletions(-) diff --git a/src/config/load/strict_keys.rs b/src/config/load/strict_keys.rs index b1a31b9..6229bb9 100644 --- a/src/config/load/strict_keys.rs +++ b/src/config/load/strict_keys.rs @@ -264,6 +264,7 @@ const WEB_CONFIG_KEYS: &[&str] = &[ "carrier", "carriers", "carrier_learning", + "carrier_negotiation_aggressiveness", "debug", "limits", "timeouts", @@ -278,10 +279,14 @@ const WEB_LIMITS_CONFIG_KEYS: &[&str] = &[ "max_frames_per_body", "max_http_connections", "max_http_handlers", + "max_lane_open_waits_per_session", + "pending_bytes_per_lane", + "pending_items_per_lane", "websocket_bytes_global", "websocket_admission_watermark_pct", "websocket_eviction_watermark_pct", "websocket_http_connection_reserve", + "max_websocket_evictions_in_flight", "max_carrier_learning_entries", "max_body_readers", "max_body_bytes_global", @@ -332,7 +337,12 @@ const WEB_TIMEOUTS_CONFIG_KEYS: &[&str] = &[ "header_secs", "body_secs", "stream_handshake_secs", + "stream_first_byte_secs", "long_poll_secs", + "lane_open_wait_secs", + "carrier_health_secs", + "websocket_upgrade_secs", + "websocket_open_secs", "websocket_write_secs", "websocket_backpressure_secs", "websocket_eviction_secs", diff --git a/src/config/load/validate_web.rs b/src/config/load/validate_web.rs index 95ed829..b1f429f 100644 --- a/src/config/load/validate_web.rs +++ b/src/config/load/validate_web.rs @@ -72,9 +72,9 @@ pub(super) fn validate(config: &mut ProxyConfig) -> Result<()> { validate_limits(&config.web.limits)?; debug::validate(&config.web.debug, &config.web.limits)?; 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( - "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)?; @@ -162,6 +162,16 @@ fn validate_limits(limits: &WebLimitsConfig) -> Result<()> { let positive = [ ("max_http_connections", limits.max_http_connections), ("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", limits.max_carrier_learning_entries, @@ -240,6 +250,8 @@ fn validate_limits(limits: &WebLimitsConfig) -> Result<()> { || limits.max_body_readers > limits.max_http_handlers || limits.pending_bytes_per_session > limits.pending_bytes_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.pending_bytes_per_session || limits.control_bytes_global > limits.pending_bytes_global diff --git a/src/config/load/validate_web/negotiation.rs b/src/config/load/validate_web/negotiation.rs index ec2c894..3ed0798 100644 --- a/src/config/load/validate_web/negotiation.rs +++ b/src/config/load/validate_web/negotiation.rs @@ -14,6 +14,14 @@ pub(super) fn validate(config: &WebConfig) -> Result> { } } 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() { return config_error( "web.carriers and the web.carrier fallback must contain at most four carriers", diff --git a/src/config/load/validate_web/timeouts.rs b/src/config/load/validate_web/timeouts.rs index 2482b89..736962a 100644 --- a/src/config/load/validate_web/timeouts.rs +++ b/src/config/load/validate_web/timeouts.rs @@ -6,7 +6,12 @@ pub(super) fn validate(timeouts: &WebTimeoutsConfig) -> Result<()> { ("header_secs", timeouts.header_secs), ("body_secs", timeouts.body_secs), ("stream_handshake_secs", timeouts.stream_handshake_secs), + ("stream_first_byte_secs", timeouts.stream_first_byte_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_backpressure_secs", diff --git a/src/config/types.rs b/src/config/types.rs index 3575b04..0b9f907 100644 --- a/src/config/types.rs +++ b/src/config/types.rs @@ -53,8 +53,8 @@ pub use server::{ }; #[allow(unused_imports)] pub use web::{ - WebConfig, WebDecoyConfig, WebLimitsConfig, WebProfileConfig, WebSecretMode, WebTimeoutsConfig, - WebVhostConfig, + WebCarrierNegotiationAggressiveness, WebConfig, WebDecoyConfig, WebLimitsConfig, + WebProfileConfig, WebSecretMode, WebTimeoutsConfig, WebVhostConfig, }; #[allow(unused_imports)] pub use web_carrier::{WebCarrier, WebCarriers}; diff --git a/src/config/types/web.rs b/src/config/types/web.rs index 7d54f85..85f1b72 100644 --- a/src/config/types/web.rs +++ b/src/config/types/web.rs @@ -98,6 +98,15 @@ pub struct WebLimitsConfig { /// Process-wide concurrently executing HTTP handler ceiling. #[serde(default = "default_web_max_http_handlers")] 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. #[serde(default = "default_web_websocket_bytes_global")] pub websocket_bytes_global: usize, @@ -110,6 +119,9 @@ pub struct WebLimitsConfig { /// Accepted HTTP connections that WebSocket upgrades must leave available. #[serde(default = "default_web_websocket_http_connection_reserve")] 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. #[serde(default = "default_web_max_carrier_learning_entries")] 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_http_connections: default_web_max_http_connections(), 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_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_websocket_evictions_in_flight: + default_web_max_websocket_evictions_in_flight(), max_carrier_learning_entries: default_web_max_carrier_learning_entries(), max_body_readers: default_web_max_body_readers(), 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. #[serde(default = "default_web_stream_handshake_timeout_secs")] 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. #[serde(default = "default_web_long_poll_timeout_secs")] 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. #[serde(default = "default_web_websocket_write_secs")] pub websocket_write_secs: u64, @@ -307,7 +339,12 @@ impl Default for WebTimeoutsConfig { header_secs: default_web_header_timeout_secs(), body_secs: default_web_body_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(), + 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_backpressure_secs: default_web_websocket_backpressure_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. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct WebConfig { @@ -338,6 +388,9 @@ pub struct WebConfig { /// Enables bounded process-local carrier learning for automatic sessions. #[serde(default = "default_web_carrier_learning")] 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. #[serde(default)] pub limits: WebLimitsConfig, @@ -381,6 +434,8 @@ impl Default for WebConfig { carrier: WebCarrier::default(), carriers: WebCarriers::default(), carrier_learning: default_web_carrier_learning(), + carrier_negotiation_aggressiveness: + WebCarrierNegotiationAggressiveness::default(), limits: WebLimitsConfig::default(), debug: WebDebugConfig::default(), timeouts: WebTimeoutsConfig::default(), diff --git a/src/config/types/web/defaults.rs b/src/config/types/web/defaults.rs index 70c83ec..ff424fe 100644 --- a/src/config/types/web/defaults.rs +++ b/src/config/types/web/defaults.rs @@ -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_http_connections, 1024); 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); 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_websocket_evictions_in_flight, 8); usize_default!(default_web_max_carrier_learning_entries, 4096); usize_default!(default_web_max_body_readers, 32); 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_body_timeout_secs, 30); 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_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_backpressure_secs, 30); u64_default!(default_web_websocket_eviction_secs, 1); diff --git a/src/web/session.rs b/src/web/session.rs index 615dd3e..3be00f8 100644 --- a/src/web/session.rs +++ b/src/web/session.rs @@ -46,7 +46,14 @@ struct InboundChunk { offset: usize, } +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) struct StreamIdentity { + pub(crate) id: u32, + pub(crate) instance: u64, +} + struct StreamState { + instance: u64, inbound: VecDeque, receive_window: u32, send_credit: u64, @@ -73,6 +80,7 @@ struct DownBatch { } struct CarrierLane { + instance: u64, pending_frames: VecDeque, pending_windows: HashMap, unacked: Option, @@ -85,8 +93,9 @@ struct CarrierLane { } impl CarrierLane { - fn new() -> Self { + fn new(instance: u64) -> Self { Self { + instance, pending_frames: VecDeque::new(), pending_windows: HashMap::new(), unacked: None, @@ -102,6 +111,8 @@ impl CarrierLane { struct SessionState { streams: HashMap, + closing_streams: HashMap, + next_stream_instance: u64, active_peer_ports: HashSet, closed_streams: HashSet, closed_order: VecDeque, @@ -113,6 +124,7 @@ struct SessionState { last_up_sequence: u64, last_up_digest: TokenHash, carrier_lanes: HashMap, + next_lane_instance: u64, websocket_lane_reservations: HashMap, pending_bytes: usize, pending_items: usize, @@ -183,8 +195,10 @@ impl WebSession { timeouts: WebTimeoutsConfig, ) -> Arc { let mut carrier_lanes = HashMap::new(); + let mut next_lane_instance = 1; 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 { manager, @@ -201,6 +215,8 @@ impl WebSession { timeouts, state: Mutex::new(SessionState { streams: HashMap::new(), + closing_streams: HashMap::new(), + next_stream_instance: 1, active_peer_ports: HashSet::new(), closed_streams: HashSet::new(), closed_order: VecDeque::new(), @@ -212,6 +228,7 @@ impl WebSession { last_up_sequence: 0, last_up_digest: [0; 32], carrier_lanes, + next_lane_instance, websocket_lane_reservations: HashMap::new(), pending_bytes: 0, pending_items: 0, @@ -329,17 +346,21 @@ impl WebSession { /// Polls client-to-server bytes and returns consumed flow-control credit. pub(super) fn poll_read( &self, - stream_id: u32, + stream: StreamIdentity, cx: &mut Context<'_>, output: &mut ReadBuf<'_>, ) -> Poll> { let mut state = self.state.lock(); 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(())); }; - let Some(chunk) = stream.inbound.front_mut() else { - stream.read_waker = Some(cx.waker().clone()); + let Some(chunk) = stream_state.inbound.front_mut() else { + stream_state.read_waker = Some(cx.waker().clone()); return Poll::Pending; }; let available = &chunk.bytes[chunk.offset..]; @@ -348,14 +369,14 @@ impl WebSession { chunk.offset += count; let finished = chunk.offset == chunk.bytes.len(); 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) }; let overhead = if finished { QUEUE_ITEM_COST } else { 0 }; 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); self.close(); return Poll::Ready(Err(io::Error::other( @@ -368,7 +389,7 @@ impl WebSession { /// Polls server-to-client writes against stream credit and bounded queues. pub(super) fn poll_write( &self, - stream_id: u32, + stream: StreamIdentity, cx: &mut Context<'_>, input: &[u8], ) -> Poll> { @@ -376,7 +397,11 @@ impl WebSession { return Poll::Ready(Ok(0)); } 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( io::ErrorKind::BrokenPipe, "WEB logical stream is closed", @@ -386,24 +411,32 @@ impl WebSession { .len() .min(frame::DATA_CHUNK_BYTES) .min(self.limits.max_frame_payload_bytes) - .min(stream.send_credit as usize); + .min(stream_state.send_credit as usize); if count == 0 { - stream.write_waker = Some(cx.waker().clone()); + stream_state.write_waker = Some(cx.waker().clone()); return Poll::Pending; } - if !self.queue_data_locked(&mut state, stream_id, &input[..count]) { - if let Some(stream) = state.streams.get_mut(&stream_id) { - stream.write_waker = Some(cx.waker().clone()); + if !self.queue_data_locked(&mut state, stream.id, &input[..count]) { + if let Some(stream_state) = state + .streams + .get_mut(&stream.id) + .filter(|state| state.instance == stream.instance) + { + stream_state.write_waker = Some(cx.waker().clone()); } 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( io::ErrorKind::BrokenPipe, "WEB logical stream is closed", ))); }; - stream.send_credit -= count as u64; + stream_state.send_credit -= count as u64; state.last_activity = Instant::now(); drop(state); if self.carrier().is_multiplexed() { diff --git a/src/web/session/backend.rs b/src/web/session/backend.rs index 2152833..3289617 100644 --- a/src/web/session/backend.rs +++ b/src/web/session/backend.rs @@ -7,7 +7,7 @@ use crate::proxy::shared_state::ConntrackClosePolicy; use crate::web::frame::FrameType; use crate::web::stream::WebLogicalStream; -use super::{WebSession, inbound_queue_cost}; +use super::{StreamIdentity, WebSession, inbound_queue_cost}; #[cfg(test)] #[path = "backend_tests.rs"] @@ -17,61 +17,64 @@ impl WebSession { /// Starts one owned inner handshake and relay task for an admitted stream. pub(super) fn spawn_stream( self: &Arc, - stream_id: u32, - peer_port: u16, + completion: StreamCompletion, retain_reservation_on_reject: bool, ) -> bool { + let stream = completion.stream; + let peer_port = completion.peer_port; 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; }; let generation = manager.active_generation(); if !*generation.admission_rx.borrow() { self.trace_lifecycle( crate::web::trace::TraceLifecycleEvent::StreamRejected, - Some(stream_id), + Some(stream.id), 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; } let Ok(connection_permit) = generation.max_connections.clone().try_acquire_owned() else { manager.record_stream_rejected(); self.trace_lifecycle( crate::web::trace::TraceLifecycleEvent::StreamRejected, - Some(stream_id), + Some(stream.id), 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; }; let deps = generation.client_runtime_deps(); let replay_checker = Arc::clone(&generation.replay_checker); let session = Arc::clone(self); let cancel = self.cancel.clone(); - let retain_rejected = Arc::new(AtomicBool::new(false)); - 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 retain_rejected = Arc::clone(&completion.retain_rejected); let future = async move { let _connection_permit = connection_permit; let _completion = completion; session.trace_lifecycle( crate::web::trace::TraceLifecycleEvent::StreamAdmitted, - Some(stream_id), + Some(stream.id), None, ); - let stream = WebLogicalStream::new(Arc::clone(&session), stream_id); + let logical_stream = WebLogicalStream::new(Arc::clone(&session), stream); tokio::select! { _ = cancel.cancelled() => {} _ = run_stream( Arc::clone(&session), - stream_id, stream, + logical_stream, deps, replay_checker, peer_port, @@ -82,7 +85,7 @@ impl WebSession { retain_rejected.store(retain_reservation_on_reject, Ordering::Release); self.trace_lifecycle( crate::web::trace::TraceLifecycleEvent::StreamRejected, - Some(stream_id), + Some(stream.id), Some("generation_closed"), ); drop(future); @@ -93,21 +96,28 @@ impl WebSession { fn stream_rejected_before_spawn( &self, - stream_id: u32, + stream: StreamIdentity, peer_port: u16, retain_reservation: bool, ) { if !retain_reservation { - self.stream_finished(stream_id, peer_port); + self.stream_finished(stream, peer_port); return; } let queued = { let mut state = self.state.lock(); - state.streams.remove(&stream_id).map(|stream| { - let (bytes, items) = inbound_queue_cost(&stream.inbound); + state + .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.remember_closed_locked(&mut state, stream_id); - self.queue_control_locked(&mut state, FrameType::Close, stream_id, &[]) + 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) { @@ -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 mut state = self.state.lock(); let reserved = state.active_peer_ports.remove(&peer_port); - let queued = state.streams.remove(&stream_id).map(|stream| { - let (bytes, items) = inbound_queue_cost(&stream.inbound); + let current = state + .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.remember_closed_locked(&mut state, stream_id); - self.queue_control_locked(&mut state, FrameType::Close, stream_id, &[]) + self.remember_closed_locked(&mut state, 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) }; if reserved && let Some(manager) = self.manager.upgrade() { @@ -146,25 +164,41 @@ impl WebSession { } } -struct StreamCompletion { +pub(super) struct StreamCompletion { session: Arc, - stream_id: u32, - peer_port: u16, + pub(super) stream: StreamIdentity, + pub(super) peer_port: u16, retain_rejected: Arc, } +impl WebSession { + pub(super) fn own_stream_task( + self: &Arc, + 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 { fn drop(&mut self) { self.session.trace_lifecycle( crate::web::trace::TraceLifecycleEvent::StreamClosed, - Some(self.stream_id), + Some(self.stream.id), None, ); if self.retain_rejected.load(Ordering::Acquire) { self.session - .stream_rejected_before_spawn(self.stream_id, self.peer_port, true); + .stream_rejected_before_spawn(self.stream, self.peer_port, true); } 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 { self.session.tasks_done.notify_waiters(); @@ -174,7 +208,7 @@ impl Drop for StreamCompletion { async fn run_stream( session: Arc, - stream_id: u32, + stream_identity: StreamIdentity, stream: WebLogicalStream, deps: crate::proxy::authenticated::ClientRuntimeDeps, replay_checker: Arc, @@ -191,13 +225,27 @@ async fn run_stream( let peer = std::net::SocketAddr::new(session.client_ip, peer_port); deps.stats.increment_connects_all(); - // A carrier may publish OPEN before the local MTProto socket writes its - // first byte. Session and stream quotas bound this idle phase without - // consuming the process-wide active-handshake budget. - if reader.read_exact(&mut handshake[..1]).await.is_err() { + // Silent OPEN ownership has a separate absolute deadline so it cannot + // consume stream and tuple quotas indefinitely before handshake admission. + let first_byte = tokio::time::timeout( + 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( crate::web::trace::TraceLifecycleEvent::HandshakeIo, - Some(stream_id), + Some(stream_identity.id), Some("first_byte_io"), ); deps.stats @@ -206,7 +254,7 @@ async fn run_stream( } session.trace_lifecycle( crate::web::trace::TraceLifecycleEvent::StreamFirstByte, - Some(stream_id), + Some(stream_identity.id), None, ); let Some(manager) = session.manager.upgrade() else { @@ -215,7 +263,7 @@ async fn run_stream( let Some(handshake_permit) = manager.try_stream_handshake() else { session.trace_lifecycle( crate::web::trace::TraceLifecycleEvent::StreamRejected, - Some(stream_id), + Some(stream_identity.id), Some("handshake_limit"), ); return; @@ -246,7 +294,7 @@ async fn run_stream( Err(_) => { session.trace_lifecycle( crate::web::trace::TraceLifecycleEvent::HandshakeTimeout, - Some(stream_id), + Some(stream_identity.id), Some("timeout"), ); deps.stats @@ -258,7 +306,7 @@ async fn run_stream( Ok(Err(_)) => { session.trace_lifecycle( crate::web::trace::TraceLifecycleEvent::HandshakeIo, - Some(stream_id), + Some(stream_identity.id), Some("io"), ); deps.stats @@ -268,7 +316,7 @@ async fn run_stream( Ok(Ok(crate::error::HandshakeResult::Success((reader, writer, success)))) => { session.trace_lifecycle( crate::web::trace::TraceLifecycleEvent::HandshakeSucceeded, - Some(stream_id), + Some(stream_identity.id), None, ); (reader, writer, success) @@ -276,7 +324,7 @@ async fn run_stream( Ok(Ok(_)) => { session.trace_lifecycle( crate::web::trace::TraceLifecycleEvent::HandshakeRejected, - Some(stream_id), + Some(stream_identity.id), Some("bad_client"), ); deps.stats @@ -286,7 +334,7 @@ async fn run_stream( }; session.trace_lifecycle( crate::web::trace::TraceLifecycleEvent::RelayStarted, - Some(stream_id), + Some(stream_identity.id), None, ); let relay_result = run_authenticated( @@ -301,7 +349,7 @@ async fn run_stream( .await; session.trace_lifecycle( crate::web::trace::TraceLifecycleEvent::RelayEnded, - Some(stream_id), + Some(stream_identity.id), Some(if relay_result.is_ok() { "completed" } else { diff --git a/src/web/stream.rs b/src/web/stream.rs index 2ad9389..69b2ce1 100644 --- a/src/web/stream.rs +++ b/src/web/stream.rs @@ -7,21 +7,21 @@ use std::task::{Context, Poll}; use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; 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. pub(crate) struct WebLogicalStream { session: Arc, - stream_id: u32, + stream: StreamIdentity, budget_wait: Option>>, } impl WebLogicalStream { /// Binds a virtual byte stream to one live carrier stream identifier. - pub(crate) fn new(session: Arc, stream_id: u32) -> Self { + pub(crate) fn new(session: Arc, stream: StreamIdentity) -> Self { Self { session, - stream_id, + stream, budget_wait: None, } } @@ -33,7 +33,7 @@ impl AsyncRead for WebLogicalStream { cx: &mut Context<'_>, output: &mut ReadBuf<'_>, ) -> Poll> { - 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<'_>, input: &[u8], ) -> Poll> { - 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() { self.budget_wait = None; return result; @@ -64,7 +64,7 @@ impl AsyncWrite for WebLogicalStream { } 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) => { self.budget_wait = None; Poll::Ready(result)