From 08109d53e88c162867e5092f48d364e00c7cb25f Mon Sep 17 00:00:00 2001 From: Alexey <247128645+axkurcom@users.noreply.github.com> Date: Thu, 24 Sep 2026 01:02:07 +0300 Subject: [PATCH] WEB Data Budget: Session effects until after state unlock deferred --- src/web/manager.rs | 102 +----- src/web/manager/admission.rs | 125 ++++--- src/web/manager/budget.rs | 183 +++++++--- src/web/manager/budget/tests.rs | 78 +++++ src/web/manager/operator_lifecycle.rs | 15 +- .../manager/operator_lifecycle/admission.rs | 75 +++- src/web/manager/operator_lifecycle/tests.rs | 43 +++ src/web/session.rs | 126 +------ src/web/session/backend.rs | 42 ++- src/web/session/downlink.rs | 168 ++++----- src/web/session/downlink/batch.rs | 115 ++++++ src/web/session/downlink_tests.rs | 149 +++++++- src/web/session/effects.rs | 221 ++++++++++++ src/web/session/lane_downlink.rs | 6 +- src/web/session/lane_uplink.rs | 34 +- src/web/session/lanes.rs | 66 ++-- src/web/session/lanes/tests.rs | 94 ++++- src/web/session/lifecycle.rs | 93 +++-- src/web/session/stream_io.rs | 238 +++++++++++++ src/web/session/stream_io/tests.rs | 331 ++++++++++++++++++ src/web/session/uplink.rs | 69 ++-- src/web/session/uplink_tests.rs | 186 ++++++++++ src/web/session/websocket.rs | 249 ++++--------- src/web/session/websocket/reservation.rs | 191 ++++++++++ src/web/session/websocket/tests.rs | 78 ++++- 25 files changed, 2325 insertions(+), 752 deletions(-) create mode 100644 src/web/manager/budget/tests.rs create mode 100644 src/web/session/downlink/batch.rs create mode 100644 src/web/session/effects.rs create mode 100644 src/web/session/stream_io.rs create mode 100644 src/web/session/stream_io/tests.rs create mode 100644 src/web/session/websocket/reservation.rs diff --git a/src/web/manager.rs b/src/web/manager.rs index d7d21dd..38b796a 100644 --- a/src/web/manager.rs +++ b/src/web/manager.rs @@ -57,7 +57,7 @@ pub(crate) use observability::{WebCapacityResourceStatus, WebCapacitySnapshot}; // Asynchronous bounded close operations isolate mutation lifecycle from HTTP requests. mod control; pub(crate) use budget::WebSocketBudgetLease; -use budget::{WebDataBudget, WebSocketBudgetClass}; +use budget::WebDataBudget; pub(crate) use control::{CloseOperationSelector, ControlError}; pub(crate) use negotiation::{ CarrierCapabilities, CarrierClientClass, CarrierFailure, CarrierLearningContext, CarrierRequest, @@ -405,106 +405,6 @@ impl WebProcessRuntime { drop(tokio::spawn(tracked)); } - /// Reserves one body reader and its declared bounded body allocation. - pub(crate) fn try_body_budget( - &self, - bytes: usize, - ) -> Option<(OwnedSemaphorePermit, OwnedSemaphorePermit)> { - let Some(bytes) = u32::try_from(bytes).ok() else { - self.record_limit_hit(); - self.telemetry - .record_rejection(WebRejectionReason::BodyBytesCapacity); - return None; - }; - let Some(reader) = Arc::clone(&self.body_readers).try_acquire_owned().ok() else { - self.record_limit_hit(); - self.telemetry - .record_rejection(WebRejectionReason::BodyReaderCapacity); - return None; - }; - let Some(body) = Arc::clone(&self.body_bytes) - .try_acquire_many_owned(bytes) - .ok() - else { - self.record_limit_hit(); - self.telemetry - .record_rejection(WebRejectionReason::BodyBytesCapacity); - return None; - }; - Some((reader, body)) - } - - /// Reserves transient bytes while one downlink batch replaces queued frames. - pub(crate) fn try_downlink_staging_budget(&self, bytes: usize) -> Option { - let bytes = u32::try_from(bytes).ok()?; - let permit = Arc::clone(&self.body_bytes) - .try_acquire_many_owned(bytes) - .ok(); - if permit.is_none() { - self.record_limit_hit(); - self.telemetry - .record_rejection(WebRejectionReason::BodyBytesCapacity); - } - permit - } - - /// Reserves bounded process-wide queue capacity for data or control traffic. - pub(crate) fn try_reserve_pending( - &self, - owner: ProfileKey, - bytes: usize, - items: usize, - control: bool, - downlink: bool, - ) -> bool { - if !self - .data_budget - .try_reserve_queue(owner, bytes, items, control, downlink) - { - self.record_limit_hit(); - self.telemetry - .record_rejection(WebRejectionReason::QueueGlobalCapacity); - return false; - } - true - } - - /// Releases process-wide queue capacity and wakes blocked relay writers. - pub(crate) fn release_pending( - &self, - owner: ProfileKey, - bytes: usize, - items: usize, - control: bool, - ) { - self.data_budget.release_queue(owner, bytes, items, control); - } - - /// Returns the shared notification source for global queue capacity changes. - pub(crate) fn budget_notify(&self) -> Arc { - 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 { - 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 { - self.data_budget - .try_reserve_websocket(owner, bytes, WebSocketBudgetClass::Data) - } - /// Admits one WebSocket with dead-first, then owner-local bounded replacement. #[allow(clippy::too_many_arguments)] pub(crate) async fn admit_websocket( diff --git a/src/web/manager/admission.rs b/src/web/manager/admission.rs index a8a7190..6ab8557 100644 --- a/src/web/manager/admission.rs +++ b/src/web/manager/admission.rs @@ -1,12 +1,16 @@ use std::net::{IpAddr, SocketAddr}; +use std::sync::Arc; use std::time::Instant; +use tokio::sync::Notify; + use super::state::{allocate_stream_port, allow_rate, decrement_map, release_stream_port}; use super::{ProfileKey, WebProcessRuntime}; use crate::web::telemetry::WebRejectionReason; impl WebProcessRuntime { - /// Reserves one process-wide and per-profile live logical-stream slot. + /// Reserves one live stream slot and immediately dispatches any operator-fence wake. + #[allow(dead_code)] pub(crate) fn try_acquire_stream( &self, profile_key: ProfileKey, @@ -14,49 +18,72 @@ impl WebProcessRuntime { client_ip: IpAddr, public_addr: SocketAddr, ) -> Result { - let _operator_admission = match self.try_operator_admission() { + let (result, notify) = self.try_acquire_stream_quiet( + profile_key, + max_streams, + client_ip, + public_addr, + ); + if let Some(notify) = notify { + notify.notify_waiters(); + } + result + } + + /// Reserves one stream while returning any operator-fence wake for deferred dispatch. + pub(crate) fn try_acquire_stream_quiet( + &self, + profile_key: ProfileKey, + max_streams: usize, + client_ip: IpAddr, + public_addr: SocketAddr, + ) -> (Result, Option>) { + let operator_admission = match self.try_operator_admission() { Ok(admission) => admission, Err(error) => { self.telemetry.record_stream_rejected(); - return Err(error); + return (Err(error), None); } }; - let now = Instant::now(); - let mut state = self.stream_admission.lock(); - if state.closed { - self.telemetry.record_stream_rejected(); - self.telemetry - .record_rejection(WebRejectionReason::RuntimeClosed); - return Err(super::ManagerError::Closed); - } - if state.streams_live >= self.limits.max_streams_global - || state - .streams_per_profile - .get(&profile_key) - .copied() - .unwrap_or(0) - >= max_streams - { - self.record_stream_rejected_reason(WebRejectionReason::StreamCapacity); - return Err(super::ManagerError::Limit); - } - if !allow_rate( - &mut state.stream_rate, - now, - self.limits.new_streams_per_minute, - self.limits.new_streams_burst, - ) { - self.record_stream_rejected_reason(WebRejectionReason::StreamRate); - return Err(super::ManagerError::Limit); - } - let Some(peer_port) = allocate_stream_port(&mut state, client_ip, public_addr) else { - self.record_stream_rejected_reason(WebRejectionReason::StreamTupleExhausted); - return Err(super::ManagerError::Limit); + let result = { + let now = Instant::now(); + let mut state = self.stream_admission.lock(); + if state.closed { + self.telemetry.record_stream_rejected(); + self.telemetry + .record_rejection(WebRejectionReason::RuntimeClosed); + Err(super::ManagerError::Closed) + } else if state.streams_live >= self.limits.max_streams_global + || state + .streams_per_profile + .get(&profile_key) + .copied() + .unwrap_or(0) + >= max_streams + { + self.record_stream_rejected_reason(WebRejectionReason::StreamCapacity); + Err(super::ManagerError::Limit) + } else if !allow_rate( + &mut state.stream_rate, + now, + self.limits.new_streams_per_minute, + self.limits.new_streams_burst, + ) { + self.record_stream_rejected_reason(WebRejectionReason::StreamRate); + Err(super::ManagerError::Limit) + } else if let Some(peer_port) = allocate_stream_port(&mut state, client_ip, public_addr) + { + state.streams_live += 1; + *state.streams_per_profile.entry(profile_key).or_insert(0) += 1; + self.telemetry.record_stream_opened(); + Ok(peer_port) + } else { + self.record_stream_rejected_reason(WebRejectionReason::StreamTupleExhausted); + Err(super::ManagerError::Limit) + } }; - state.streams_live += 1; - *state.streams_per_profile.entry(profile_key).or_insert(0) += 1; - self.telemetry.record_stream_opened(); - Ok(peer_port) + let notify = operator_admission.release_deferred(); + (result, notify) } /// Releases one live logical-stream slot after its relay task exits. @@ -67,14 +94,32 @@ impl WebProcessRuntime { public_addr: SocketAddr, peer_port: u16, ) { + if let Some(notify) = self.release_stream_quiet( + profile_key, + client_ip, + public_addr, + peer_port, + ) { + notify.notify_waiters(); + } + } + + /// Releases one stream while returning any drain wake for deferred dispatch. + pub(crate) fn release_stream_quiet( + &self, + profile_key: ProfileKey, + client_ip: IpAddr, + public_addr: SocketAddr, + peer_port: u16, + ) -> Option> { let mut state = self.stream_admission.lock(); if !release_stream_port(&mut state, client_ip, public_addr, peer_port) { - return; + return None; } state.streams_live = state.streams_live.saturating_sub(1); decrement_map(&mut state.streams_per_profile, &profile_key); drop(state); - self.notify_operator_work_changed(); + self.operator_lifecycle.work_changed_notification() } /// Records a logical stream rejected outside manager quota acquisition. diff --git a/src/web/manager/budget.rs b/src/web/manager/budget.rs index 149dd05..efae15e 100644 --- a/src/web/manager/budget.rs +++ b/src/web/manager/budget.rs @@ -3,11 +3,12 @@ use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; use parking_lot::Mutex; -use tokio::sync::Notify; +use tokio::sync::{Notify, OwnedSemaphorePermit}; -use super::ProfileKey; +use super::{ProfileKey, WebProcessRuntime}; use crate::config::WebLimitsConfig; use crate::web::session::QUEUE_ITEM_COST; +use crate::web::telemetry::WebRejectionReason; /// WebSocket allocation class with a distinct pressure watermark. #[derive(Clone, Copy)] @@ -177,6 +178,18 @@ impl WebDataBudget { items: usize, control: bool, ) { + let notify = self.release_queue_quiet(owner, bytes, items, control); + notify.notify_waiters(); + } + + /// Releases queue accounting and returns the notification capability uninvoked. + pub(super) fn release_queue_quiet( + &self, + owner: ProfileKey, + bytes: usize, + items: usize, + control: bool, + ) -> Arc { let mut state = self.state.lock(); state.queue_bytes = state.queue_bytes.saturating_sub(bytes); state.queue_items = state.queue_items.saturating_sub(items); @@ -186,7 +199,7 @@ impl WebDataBudget { } remove_owner(&mut state.owner_bytes, owner, bytes); drop(state); - self.notify.notify_waiters(); + Arc::clone(&self.notify) } pub(super) fn try_reserve_websocket( @@ -323,6 +336,120 @@ impl Drop for WebSocketBudgetLease { } } +impl WebProcessRuntime { + /// Reserves one body reader and its declared bounded body allocation. + pub(crate) fn try_body_budget( + &self, + bytes: usize, + ) -> Option<(OwnedSemaphorePermit, OwnedSemaphorePermit)> { + let Some(bytes) = u32::try_from(bytes).ok() else { + self.record_limit_hit(); + self.telemetry + .record_rejection(WebRejectionReason::BodyBytesCapacity); + return None; + }; + let Some(reader) = Arc::clone(&self.body_readers).try_acquire_owned().ok() else { + self.record_limit_hit(); + self.telemetry + .record_rejection(WebRejectionReason::BodyReaderCapacity); + return None; + }; + let Some(body) = Arc::clone(&self.body_bytes) + .try_acquire_many_owned(bytes) + .ok() + else { + self.record_limit_hit(); + self.telemetry + .record_rejection(WebRejectionReason::BodyBytesCapacity); + return None; + }; + Some((reader, body)) + } + + /// Reserves transient bytes while one downlink batch replaces queued frames. + pub(crate) fn try_downlink_staging_budget(&self, bytes: usize) -> Option { + let bytes = u32::try_from(bytes).ok()?; + let permit = Arc::clone(&self.body_bytes) + .try_acquire_many_owned(bytes) + .ok(); + if permit.is_none() { + self.record_limit_hit(); + self.telemetry + .record_rejection(WebRejectionReason::BodyBytesCapacity); + } + permit + } + + /// Reserves bounded process-wide queue capacity for data or control traffic. + pub(crate) fn try_reserve_pending( + &self, + owner: ProfileKey, + bytes: usize, + items: usize, + control: bool, + downlink: bool, + ) -> bool { + if !self + .data_budget + .try_reserve_queue(owner, bytes, items, control, downlink) + { + self.record_limit_hit(); + self.telemetry + .record_rejection(WebRejectionReason::QueueGlobalCapacity); + return false; + } + true + } + + /// Releases process-wide queue capacity and wakes blocked relay writers. + pub(crate) fn release_pending( + &self, + owner: ProfileKey, + bytes: usize, + items: usize, + control: bool, + ) { + self.data_budget.release_queue(owner, bytes, items, control); + } + + /// Releases queue accounting without invoking wake callbacks in the caller's lock scope. + pub(crate) fn release_pending_quiet( + &self, + owner: ProfileKey, + bytes: usize, + items: usize, + control: bool, + ) -> Arc { + self.data_budget + .release_queue_quiet(owner, bytes, items, control) + } + + /// Returns the shared notification source for global queue capacity changes. + pub(crate) fn budget_notify(&self) -> Arc { + 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 { + 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 { + self.data_budget + .try_reserve_websocket(owner, bytes, WebSocketBudgetClass::Data) + } +} + fn watermark(limit: usize, percentage: u8) -> usize { limit.saturating_mul(usize::from(percentage)) / 100 } @@ -356,51 +483,5 @@ fn update_high_water(state: &mut BudgetState) { } #[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()); - } -} +#[path = "budget/tests.rs"] +mod tests; diff --git a/src/web/manager/budget/tests.rs b/src/web/manager/budget/tests.rs new file mode 100644 index 0000000..d5cd9a9 --- /dev/null +++ b/src/web/manager/budget/tests.rs @@ -0,0 +1,78 @@ +use std::future::Future; +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::task::{Context, Poll, Wake, Waker}; + +use super::*; + +struct WakeCounter(AtomicUsize); + +impl Wake for WakeCounter { + fn wake(self: Arc) { + self.0.fetch_add(1, Ordering::AcqRel); + } +} + +#[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()); +} + +#[test] +fn quiet_queue_release_updates_accounting_before_notification_dispatch() { + let budget = WebDataBudget::new(WebLimitsConfig::default()); + assert!(budget.try_reserve_queue([1; 32], 64, 1, false, false)); + let counter = Arc::new(WakeCounter(AtomicUsize::new(0))); + let waker = Waker::from(Arc::clone(&counter)); + let mut context = Context::from_waker(&waker); + let mut notified = Box::pin(budget.notify.notified()); + assert!(matches!(notified.as_mut().poll(&mut context), Poll::Pending)); + + let notify = budget.release_queue_quiet([1; 32], 64, 1, false); + + assert_eq!(budget.snapshot().queue_bytes, 0); + assert_eq!(budget.snapshot().queue_items, 0); + assert_eq!(counter.0.load(Ordering::Acquire), 0); + notify.notify_waiters(); + assert_eq!(counter.0.load(Ordering::Acquire), 1); +} diff --git a/src/web/manager/operator_lifecycle.rs b/src/web/manager/operator_lifecycle.rs index 499d5f6..95e926e 100644 --- a/src/web/manager/operator_lifecycle.rs +++ b/src/web/manager/operator_lifecycle.rs @@ -43,7 +43,7 @@ pub(super) struct OperatorLifecycle { commands: AsyncMutex<()>, inner: Mutex, published: ArcSwap, - work_changed: Notify, + work_changed: Arc, next_operation_id: AtomicU64, } @@ -71,7 +71,7 @@ impl OperatorLifecycle { drain: None, }), published: ArcSwap::from_pointee(snapshot), - work_changed: Notify::new(), + work_changed: Arc::new(Notify::new()), next_operation_id: AtomicU64::new(1), } } @@ -85,11 +85,18 @@ impl OperatorLifecycle { /// Wakes an active drain after tracked work ownership changes. pub(super) fn notify_work_changed(&self) { - if self.admission.is_closed() { - self.work_changed.notify_waiters(); + if let Some(notify) = self.work_changed_notification() { + notify.notify_waiters(); } } + /// Returns the drain notification capability without invoking callbacks. + pub(super) fn work_changed_notification(&self) -> Option> { + self.admission + .is_closed() + .then(|| Arc::clone(&self.work_changed)) + } + /// Returns the lock-free lifecycle snapshot with effective config admission. pub(super) fn status(&self, config_enabled: bool) -> OperatorLifecycleStatus { let snapshot = self.published.load(); diff --git a/src/web/manager/operator_lifecycle/admission.rs b/src/web/manager/operator_lifecycle/admission.rs index 741163f..aa6f924 100644 --- a/src/web/manager/operator_lifecycle/admission.rs +++ b/src/web/manager/operator_lifecycle/admission.rs @@ -1,3 +1,4 @@ +use std::sync::Arc; use std::sync::atomic::{AtomicUsize, Ordering}; use tokio::sync::Notify; @@ -63,12 +64,13 @@ impl OperatorAdmissionRejection { /// Lock-free admission fence with bounded pre-cutover registration tracking. pub(super) struct OperatorAdmission { state: AtomicUsize, - registrations_drained: Notify, + registrations_drained: Arc, } /// RAII ownership of one synchronous pre-cutover admission section. pub(in crate::web::manager) struct OperatorRegistration<'a> { admission: &'a OperatorAdmission, + released: bool, } impl OperatorAdmission { @@ -76,7 +78,7 @@ impl OperatorAdmission { pub(super) fn new() -> Self { Self { state: AtomicUsize::new(0), - registrations_drained: Notify::new(), + registrations_drained: Arc::new(Notify::new()), } } @@ -100,7 +102,12 @@ impl OperatorAdmission { Ordering::AcqRel, Ordering::Acquire, ) { - Ok(_) => return Ok(OperatorRegistration { admission: self }), + Ok(_) => { + return Ok(OperatorRegistration { + admission: self, + released: false, + }); + } Err(observed) => state = observed, } } @@ -147,13 +154,69 @@ impl OperatorAdmission { notified.await; } } + + fn release_registration(&self) -> Option> { + let previous = self.state.fetch_sub(1, Ordering::AcqRel); + (previous & OPERATOR_REGISTRATION_COUNT == 1) + .then(|| Arc::clone(&self.registrations_drained)) + } +} + +impl OperatorRegistration<'_> { + /// Releases ownership and returns the final-registration notification uninvoked. + pub(in crate::web::manager) fn release_deferred(mut self) -> Option> { + self.released = true; + self.admission.release_registration() + } } impl Drop for OperatorRegistration<'_> { fn drop(&mut self) { - let previous = self.admission.state.fetch_sub(1, Ordering::AcqRel); - if previous & OPERATOR_REGISTRATION_COUNT == 1 { - self.admission.registrations_drained.notify_waiters(); + if self.released { + return; + } + if let Some(notify) = self.admission.release_registration() { + notify.notify_waiters(); } } } + +#[cfg(test)] +mod tests { + use std::future::Future; + use std::sync::Arc; + use std::sync::atomic::{AtomicUsize, Ordering}; + use std::task::{Context, Poll, Wake, Waker}; + + use super::*; + + struct WakeCounter(AtomicUsize); + + impl Wake for WakeCounter { + fn wake(self: Arc) { + self.0.fetch_add(1, Ordering::AcqRel); + } + } + + #[test] + fn deferred_registration_release_does_not_invoke_waiter_before_dispatch() { + let admission = OperatorAdmission::new(); + let registration = admission.try_register().unwrap(); + admission.close(OperatorAdmissionRejection::Paused); + let counter = Arc::new(WakeCounter(AtomicUsize::new(0))); + let waker = Waker::from(Arc::clone(&counter)); + let mut context = Context::from_waker(&waker); + let mut notified = Box::pin(admission.registrations_drained.notified()); + assert!(matches!(notified.as_mut().poll(&mut context), Poll::Pending)); + + let notify = registration.release_deferred().unwrap(); + + assert_eq!( + admission.state.load(Ordering::Acquire) & OPERATOR_REGISTRATION_COUNT, + 0 + ); + assert_eq!(counter.0.load(Ordering::Acquire), 0); + notify.notify_waiters(); + assert_eq!(counter.0.load(Ordering::Acquire), 1); + } +} diff --git a/src/web/manager/operator_lifecycle/tests.rs b/src/web/manager/operator_lifecycle/tests.rs index a72a705..ddce9a3 100644 --- a/src/web/manager/operator_lifecycle/tests.rs +++ b/src/web/manager/operator_lifecycle/tests.rs @@ -1,5 +1,7 @@ +use std::future::Future; use std::sync::Arc; use std::sync::atomic::{AtomicUsize, Ordering}; +use std::task::{Context, Poll, Wake, Waker}; use arc_swap::ArcSwap; use tokio::sync::Barrier; @@ -8,6 +10,14 @@ use super::*; use crate::config::ProxyConfig; use crate::maestro::generation::test_runtime_generation; +struct WakeCounter(AtomicUsize); + +impl Wake for WakeCounter { + fn wake(self: Arc) { + self.0.fetch_add(1, Ordering::AcqRel); + } +} + fn test_runtime() -> ( Arc, Arc, @@ -68,6 +78,39 @@ async fn pause_waits_for_pre_cutover_admission_and_rejects_late_registration() { stop_runtime(runtime, generation).await; } +#[tokio::test] +async fn quiet_stream_release_returns_exact_post_accounting_drain_notification() { + let (runtime, generation) = test_runtime(); + let profile_key = [1; 32]; + let client_ip = "192.0.2.10".parse().unwrap(); + let public_addr = "203.0.113.10:443".parse().unwrap(); + let peer_port = runtime + .try_acquire_stream(profile_key, 1, client_ip, public_addr) + .unwrap(); + runtime.pause_operator().await.unwrap(); + let counter = Arc::new(WakeCounter(AtomicUsize::new(0))); + let waker = Waker::from(Arc::clone(&counter)); + let mut context = Context::from_waker(&waker); + let mut notified = Box::pin(runtime.operator_lifecycle.work_changed.notified()); + assert!(matches!(notified.as_mut().poll(&mut context), Poll::Pending)); + + let notify = runtime + .release_stream_quiet(profile_key, client_ip, public_addr, peer_port) + .unwrap(); + + assert_eq!(runtime.stream_admission.lock().streams_live, 0); + assert_eq!(counter.0.load(Ordering::Acquire), 0); + assert!( + runtime + .release_stream_quiet(profile_key, client_ip, public_addr, peer_port) + .is_none() + ); + notify.notify_waiters(); + assert_eq!(counter.0.load(Ordering::Acquire), 1); + drop(notified); + stop_runtime(runtime, generation).await; +} + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] async fn pause_fence_leaves_no_late_admission_commits_under_scheduler_pressure() { const ATTEMPTS: usize = 10_000; diff --git a/src/web/session.rs b/src/web/session.rs index e200e61..8436bb5 100644 --- a/src/web/session.rs +++ b/src/web/session.rs @@ -1,19 +1,17 @@ use std::collections::{HashMap, HashSet, VecDeque}; -use std::io; use std::net::IpAddr; use std::sync::Arc; use std::sync::atomic::{AtomicBool, AtomicU8, AtomicUsize}; -use std::task::{Context, Poll, Waker}; +use std::task::Waker; use std::time::Instant; use bytes::{Bytes, BytesMut}; use parking_lot::Mutex; -use tokio::io::ReadBuf; use tokio::sync::Notify; use tokio_util::sync::CancellationToken; use crate::config::{WebCarrier, WebLimitsConfig, WebRuntimeProfile, WebTimeoutsConfig}; -use crate::web::frame::{self, FrameType}; +use crate::web::frame::FrameType; use crate::web::manager::{ CarrierClientClass, CarrierLearningContext, ProfileKey, TokenHash, WebProcessRuntime, }; @@ -21,6 +19,9 @@ use crate::proxy::user_admission::UserSessionRegistration; // Backend tasks own generation admission and authenticated MTProxy relay lifetimes. mod backend; +// Deferred callbacks preserve the session-state lock as a callback-free boundary. +mod effects; +use effects::DeferredSessionEffects; // Activity clocks separate authenticated peer leases from diagnostic progress. mod activity; use activity::SessionActivity; @@ -53,6 +54,8 @@ use lifecycle::SessionNegotiationPhase; pub(crate) use lifecycle::{SessionCloseOutcome, SessionCloseReason}; // Uplink batches own exactly-once sequencing and client-frame validation. mod uplink; +// Logical stream polling owns cancellation-safe waker registration. +mod stream_io; /// Conservative allocator and container overhead charged to every queued item. pub(crate) const QUEUE_ITEM_COST: usize = 256; @@ -414,119 +417,4 @@ impl WebSession { &self.timeouts } - /// Polls client-to-server bytes and returns consumed flow-control credit. - pub(super) fn poll_read( - &self, - stream: StreamIdentity, - cx: &mut Context<'_>, - output: &mut ReadBuf<'_>, - ) -> Poll> { - let mut state = self.state.lock(); - let (count, finished) = { - 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_state.inbound.front_mut() else { - stream_state.read_waker = Some(cx.waker().clone()); - return Poll::Pending; - }; - let available = &chunk.bytes[chunk.offset..]; - let count = available.len().min(output.remaining()); - output.put_slice(&available[..count]); - chunk.offset += count; - let finished = chunk.offset == chunk.bytes.len(); - if finished { - stream_state.inbound.pop_front(); - } - 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) { - drop(state); - self.close(SessionCloseReason::Backpressure); - return Poll::Ready(Err(io::Error::other( - "WEB session control budget exhausted", - ))); - } - Poll::Ready(Ok(())) - } - - /// Polls server-to-client writes against stream credit and bounded queues. - pub(super) fn poll_write( - &self, - stream: StreamIdentity, - cx: &mut Context<'_>, - input: &[u8], - ) -> Poll> { - if input.is_empty() { - return Poll::Ready(Ok(0)); - } - let mut state = self.state.lock(); - 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", - ))); - }; - let count = input - .len() - .min(frame::DATA_CHUNK_BYTES) - .min(self.limits.max_frame_payload_bytes) - .min(if self.carrier().uses_lanes() { - self.limits - .pending_bytes_per_lane - .saturating_sub(frame::HEADER_BYTES + QUEUE_ITEM_COST) - } else { - usize::MAX - }) - .min(stream_state.send_credit as usize); - if count == 0 { - 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) = 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) = 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_state.send_credit -= count as u64; - state.activity.touch_progress(Instant::now()); - drop(state); - if self.carrier().is_multiplexed() { - self.down_notify.notify_waiters(); - } - Poll::Ready(Ok(count)) - } - - /// Returns the process queue-capacity notification source while the manager lives. - pub(super) fn budget_notify(&self) -> Option> { - self.manager - .upgrade() - .map(|manager| manager.budget_notify()) - } } diff --git a/src/web/session/backend.rs b/src/web/session/backend.rs index 004800c..9042f43 100644 --- a/src/web/session/backend.rs +++ b/src/web/session/backend.rs @@ -106,11 +106,10 @@ impl WebSession { self.stream_finished(stream, peer_port); return; } - let queued = { - let mut state = self.state.lock(); + let queued = self.with_state_effects(|state, effects| { if state.closing_streams.get(&stream.id) == Some(&stream.instance) { state.closing_streams.remove(&stream.id); - self.remember_closed_locked(&mut state, stream.id); + self.remember_closed_locked(state, effects, stream.id); } state .streams @@ -119,21 +118,26 @@ impl WebSession { .is_some() .then(|| state.streams.remove(&stream.id)) .flatten() - .map(|stream_state| { + .map(|mut stream_state| { + if let Some(waker) = stream_state.read_waker.take() { + effects.drop_waker(waker); + } + if let Some(waker) = stream_state.write_waker.take() { + effects.drop_waker(waker); + } 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.release_locked(state, effects, bytes, items, false); + self.remember_closed_locked(state, effects, stream.id); + self.queue_control_locked(state, effects, FrameType::Close, stream.id, &[]) }) - }; + }); if queued.is_some_and(|queued| !queued) { self.close(SessionCloseReason::Backpressure); } } fn stream_finished(&self, stream: StreamIdentity, peer_port: u16) { - let (queued, reserved) = { - let mut state = self.state.lock(); + let (queued, reserved) = self.with_state_effects(|state, effects| { let reserved = state.active_peer_ports.remove(&peer_port); let current = state .streams @@ -142,18 +146,24 @@ impl WebSession { let queued = current .then(|| state.streams.remove(&stream.id)) .flatten() - .map(|stream_state| { + .map(|mut stream_state| { + if let Some(waker) = stream_state.read_waker.take() { + effects.drop_waker(waker); + } + if let Some(waker) = stream_state.write_waker.take() { + effects.drop_waker(waker); + } 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.release_locked(state, effects, bytes, items, false); + self.remember_closed_locked(state, effects, stream.id); + self.queue_control_locked(state, effects, 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); + self.remember_closed_locked(state, effects, stream.id); } (queued, reserved) - }; + }); if reserved && let Some(manager) = self.manager.upgrade() { manager.release_stream( self.profile_key, diff --git a/src/web/session/downlink.rs b/src/web/session/downlink.rs index b59b9ab..f6d64d1 100644 --- a/src/web/session/downlink.rs +++ b/src/web/session/downlink.rs @@ -3,15 +3,17 @@ use std::time::{Duration, Instant}; use bytes::{BufMut, Bytes, BytesMut}; -use super::resident::{OwnedBatchBody, PendingCounts, PendingResponseLease}; use super::{ - DownBatch, PendingClass, PollResult, QUEUE_ITEM_COST, QueuedFrame, SessionCloseReason, - SessionState, WebSession, + DeferredSessionEffects, PendingClass, PollResult, QUEUE_ITEM_COST, QueuedFrame, + SessionCloseReason, SessionState, WebSession, }; use crate::web::frame::{self, FrameType}; use crate::web::manager::ManagerError; use crate::web::telemetry::WebSessionLifecycleObservation; +// Batch staging owns transient permits and detached response leases. +mod batch; + impl WebSession { /// Polls pending downlink frames with cursor replay and newest-poll-wins semantics. pub(crate) async fn poll_down(&self, cursor: u64) -> Result { @@ -29,6 +31,7 @@ impl WebSession { if self.close_if_cancelled() { return Err(ManagerError::Closed); } + let mut effects = DeferredSessionEffects::new(); let (epoch, healthy) = { let mut state = self.state.lock(); if state.closed || self.cancel.is_cancelled() { @@ -58,7 +61,7 @@ impl WebSession { return Err(ManagerError::Protocol); } let carrier_health_eligible = unacked.carrier_health_eligible; - self.release_unacked_locked(&mut state); + self.release_unacked_locked(&mut state, &mut effects); state.carrier_health_downlink |= carrier_health_eligible; if carrier_health_eligible { state.carrier_health_activity_at = Some(Instant::now()); @@ -70,6 +73,7 @@ impl WebSession { } let Some(epoch) = state.down_epoch.checked_add(1) else { drop(state); + effects.finish(); self.close(SessionCloseReason::Protocol); return Err(ManagerError::Protocol); }; @@ -84,6 +88,7 @@ impl WebSession { let healthy = self.carrier_health_ready_locked(&mut state, Instant::now()); (state.down_epoch, healthy) }; + effects.finish(); if let Some(claim) = healthy { self.finish_carrier_health(claim); } @@ -95,6 +100,7 @@ impl WebSession { let notified = self.down_notify.notified(); tokio::pin!(notified); notified.as_mut().enable(); + let mut effects = DeferredSessionEffects::new(); { let mut state = self.state.lock(); if self.cancel.is_cancelled() { @@ -110,7 +116,11 @@ impl WebSession { }); } if !state.pending_frames.is_empty() { - let batch = match self.take_down_batch_locked(&mut state, cursor) { + let batch = match self.take_down_batch_locked( + &mut state, + &mut effects, + cursor, + ) { Ok(batch) => batch, Err(ManagerError::Backpressure) => { return Err(ManagerError::Backpressure); @@ -129,7 +139,11 @@ impl WebSession { if let Some(manager) = self.manager.upgrade() { manager.record_down(result.body.len()); } - state.unacked = Some(batch); + if let Some(previous) = state.unacked.replace(batch) { + effects.retain_batch(previous); + } + drop(state); + effects.finish(); return Ok(result); } if state.closed { @@ -274,6 +288,7 @@ impl WebSession { pub(super) fn release_locked( &self, state: &mut SessionState, + effects: &mut DeferredSessionEffects, bytes: usize, items: usize, control: bool, @@ -285,7 +300,12 @@ impl WebSession { state.pending_control_items = state.pending_control_items.saturating_sub(items); } if let Some(manager) = self.manager.upgrade() { - manager.release_pending(self.profile_key, bytes, items, control); + effects.notify(manager.release_pending_quiet( + self.profile_key, + bytes, + items, + control, + )); } } @@ -308,6 +328,7 @@ impl WebSession { pub(super) fn queue_window_locked( &self, state: &mut SessionState, + effects: &mut DeferredSessionEffects, stream_id: u32, amount: u32, ) -> bool { @@ -317,6 +338,7 @@ impl WebSession { if self.carrier().uses_lanes() { return self.queue_control_locked( state, + effects, FrameType::Window, stream_id, &frame::window_payload(amount), @@ -333,12 +355,13 @@ impl WebSession { if let Some(total) = previous.checked_add(amount) { queued.encoded[frame::HEADER_BYTES..frame::HEADER_BYTES + 4] .copy_from_slice(&total.to_be_bytes()); - self.down_notify.notify_waiters(); + effects.notify(Arc::clone(&self.down_notify)); return true; } } self.queue_control_locked( state, + effects, FrameType::Window, stream_id, &frame::window_payload(amount), @@ -349,22 +372,31 @@ impl WebSession { pub(super) fn queue_control_locked( &self, state: &mut SessionState, + effects: &mut DeferredSessionEffects, frame_type: FrameType, stream_id: u32, payload: &[u8], ) -> bool { - self.queue_frame_locked(state, frame_type, stream_id, payload, true) + self.queue_frame_locked(state, effects, frame_type, stream_id, payload, true) } /// Appends one server-to-client DATA frame under downlink data budgets. pub(super) fn queue_data_locked( &self, state: &mut SessionState, + effects: &mut DeferredSessionEffects, stream_id: u32, payload: &[u8], ) -> bool { if self.carrier().uses_lanes() { - return self.queue_frame_locked(state, FrameType::Data, stream_id, payload, false); + return self.queue_frame_locked( + state, + effects, + FrameType::Data, + stream_id, + payload, + false, + ); } let can_coalesce = state.pending_frames.back().is_some_and(|last| { last.frame_type == FrameType::Data @@ -385,19 +417,34 @@ impl WebSession { last.encoded[4..8].copy_from_slice(&payload_len.to_be_bytes()); return true; } - self.queue_frame_locked(state, FrameType::Data, stream_id, payload, false) + self.queue_frame_locked( + state, + effects, + FrameType::Data, + stream_id, + payload, + false, + ) } fn queue_frame_locked( &self, state: &mut SessionState, + effects: &mut DeferredSessionEffects, frame_type: FrameType, stream_id: u32, payload: &[u8], control: bool, ) -> bool { 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, + effects, + frame_type, + stream_id, + payload, + control, + ); } let cost = frame::HEADER_BYTES + payload.len() + QUEUE_ITEM_COST; let class = if control { @@ -426,105 +473,10 @@ impl WebSession { if frame_type == FrameType::Window { state.pending_windows.insert(stream_id, index); } - self.down_notify.notify_waiters(); + effects.notify(Arc::clone(&self.down_notify)); true } - fn take_down_batch_locked( - &self, - state: &mut SessionState, - cursor: u64, - ) -> Result { - let next_cursor = state - .down_cursor - .checked_add(1) - .ok_or(ManagerError::Protocol)?; - let mut count = 0usize; - let mut body_len = 0usize; - for queued in &state.pending_frames { - if count >= self.limits.max_frames_per_body - || (count != 0 - && body_len.saturating_add(queued.encoded.len()) - > self.limits.carrier_batch_bytes) - { - break; - } - body_len += queued.encoded.len(); - count += 1; - } - let Some(manager) = self.manager.upgrade() else { - return Err(ManagerError::Closed); - }; - let Some(_staging) = manager.try_downlink_staging_budget(body_len) else { - return Err(ManagerError::Backpressure); - }; - let mut body = BytesMut::with_capacity(body_len); - let mut data_bytes = 0usize; - let mut data_items = 0usize; - let mut control_bytes = 0usize; - let mut control_items = 0usize; - for index in 0..count { - let Some(queued) = state.pending_frames.get(index) else { - break; - }; - if queued.frame_type == FrameType::Window - && state.pending_windows.get(&queued.stream_id) == Some(&index) - { - state.pending_windows.remove(&queued.stream_id); - } - } - for _ in 0..count { - let Some(queued) = state.pending_frames.pop_front() else { - break; - }; - body.extend_from_slice(&queued.encoded); - if queued.control { - control_bytes += queued.cost; - control_items += 1; - } else { - data_bytes += queued.cost; - data_items += 1; - } - } - for index in state.pending_windows.values_mut() { - *index = index.saturating_sub(count); - } - state.down_cursor = next_cursor; - let counts = PendingCounts { - data_bytes, - data_items, - control_bytes, - control_items, - }; - let lease = PendingResponseLease::new(self, counts, None); - let body = Bytes::from_owner(OwnedBatchBody::new(body.freeze(), Arc::clone(&lease))); - Ok(DownBatch { - body, - lease, - base_cursor: cursor, - next_cursor, - data_bytes, - data_items, - control_bytes, - control_items, - carrier_health_eligible: state.negotiation_phase - == super::SessionNegotiationPhase::Committed, - }) - } - - fn release_unacked_locked(&self, state: &mut SessionState) { - let Some(batch) = state.unacked.take() else { - return; - }; - batch.lease.detach(); - self.release_local_locked(state, batch.data_bytes, batch.data_items, false); - self.release_local_locked(state, batch.control_bytes, batch.control_items, true); - for stream in state.streams.values_mut() { - if let Some(waker) = stream.write_waker.take() { - waker.wake(); - } - } - } } #[cfg(test)] diff --git a/src/web/session/downlink/batch.rs b/src/web/session/downlink/batch.rs new file mode 100644 index 0000000..bf85cf6 --- /dev/null +++ b/src/web/session/downlink/batch.rs @@ -0,0 +1,115 @@ +use std::sync::Arc; + +use bytes::{Bytes, BytesMut}; + +use super::super::resident::{OwnedBatchBody, PendingCounts, PendingResponseLease}; +use super::super::{DeferredSessionEffects, DownBatch, SessionState, WebSession}; +use crate::web::frame::FrameType; +use crate::web::manager::ManagerError; + +impl WebSession { + /// Stages one multiplexed batch while retaining its transient permit. + pub(super) fn take_down_batch_locked( + &self, + state: &mut SessionState, + effects: &mut DeferredSessionEffects, + cursor: u64, + ) -> Result { + let next_cursor = state + .down_cursor + .checked_add(1) + .ok_or(ManagerError::Protocol)?; + let mut count = 0usize; + let mut body_len = 0usize; + for queued in &state.pending_frames { + if count >= self.limits.max_frames_per_body + || (count != 0 + && body_len.saturating_add(queued.encoded.len()) + > self.limits.carrier_batch_bytes) + { + break; + } + body_len += queued.encoded.len(); + count += 1; + } + let Some(manager) = self.manager.upgrade() else { + return Err(ManagerError::Closed); + }; + let Some(staging) = manager.try_downlink_staging_budget(body_len) else { + return Err(ManagerError::Backpressure); + }; + effects.retain_staging_permit(staging); + let mut body = BytesMut::with_capacity(body_len); + let mut data_bytes = 0usize; + let mut data_items = 0usize; + let mut control_bytes = 0usize; + let mut control_items = 0usize; + for index in 0..count { + let Some(queued) = state.pending_frames.get(index) else { + break; + }; + if queued.frame_type == FrameType::Window + && state.pending_windows.get(&queued.stream_id) == Some(&index) + { + state.pending_windows.remove(&queued.stream_id); + } + } + for _ in 0..count { + let Some(queued) = state.pending_frames.pop_front() else { + break; + }; + body.extend_from_slice(&queued.encoded); + if queued.control { + control_bytes += queued.cost; + control_items += 1; + } else { + data_bytes += queued.cost; + data_items += 1; + } + } + for index in state.pending_windows.values_mut() { + *index = index.saturating_sub(count); + } + state.down_cursor = next_cursor; + let counts = PendingCounts { + data_bytes, + data_items, + control_bytes, + control_items, + }; + let lease = PendingResponseLease::new(self, counts, None); + let body = Bytes::from_owner(OwnedBatchBody::new(body.freeze(), Arc::clone(&lease))); + Ok(DownBatch { + body, + lease, + base_cursor: cursor, + next_cursor, + data_bytes, + data_items, + control_bytes, + control_items, + carrier_health_eligible: state.negotiation_phase + == super::super::SessionNegotiationPhase::Committed, + }) + } + + /// Detaches one acknowledged response and defers writer wakes and lease drop. + pub(super) fn release_unacked_locked( + &self, + state: &mut SessionState, + effects: &mut DeferredSessionEffects, + ) { + let Some(batch) = state.unacked.take() else { + return; + }; + batch.lease.detach(); + self.release_local_locked(state, batch.data_bytes, batch.data_items, false); + self.release_local_locked(state, batch.control_bytes, batch.control_items, true); + for stream in state.streams.values_mut() { + if let Some(waker) = stream.write_waker.take() { + effects.wake(waker); + } + } + effects.retain_batch(batch); + } +} diff --git a/src/web/session/downlink_tests.rs b/src/web/session/downlink_tests.rs index c444b0f..11201af 100644 --- a/src/web/session/downlink_tests.rs +++ b/src/web/session/downlink_tests.rs @@ -1,6 +1,9 @@ use super::*; +use std::future::Future; use std::net::SocketAddr; use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::task::{Context, Poll, Wake, Waker}; use arc_swap::ArcSwap; @@ -10,6 +13,20 @@ use crate::config::{ use crate::maestro::generation::test_runtime_generation; use crate::web::manager::WebProcessRuntime; +struct SessionLockProbe { + session: std::sync::Weak, + lock_was_free: Arc, +} + +impl Wake for SessionLockProbe { + fn wake(self: Arc) { + if let Some(session) = self.session.upgrade() { + self.lock_was_free + .store(session.state.try_lock().is_some(), Ordering::Release); + } + } +} + fn session() -> (Arc, Arc) { let generation = test_runtime_generation(1, ProxyConfig::default()); let manager = WebProcessRuntime::start(Arc::new(ArcSwap::from(generation))); @@ -57,8 +74,96 @@ fn session() -> (Arc, Arc) { } fn queue_close(session: &WebSession) { + session.with_state_effects(|state, effects| { + assert!(session.queue_control_locked(state, effects, FrameType::Close, 1, &[])); + }); +} + +#[tokio::test] +async fn queued_frame_notifies_only_after_releasing_session_lock() { + let (session, manager) = session(); + let lock_was_free = Arc::new(AtomicBool::new(false)); + let waker = Waker::from(Arc::new(SessionLockProbe { + session: Arc::downgrade(&session), + lock_was_free: Arc::clone(&lock_was_free), + })); + let mut context = Context::from_waker(&waker); + let mut notified = Box::pin(session.down_notify.notified()); + assert!(matches!(notified.as_mut().poll(&mut context), Poll::Pending)); + + queue_close(&session); + + assert!(lock_was_free.load(Ordering::Acquire)); + session.close(super::SessionCloseReason::ApiClose); + manager.shutdown().await; +} + +#[tokio::test] +async fn budget_release_notifies_only_after_session_accounting_and_unlock() { + let (session, manager) = session(); + queue_close(&session); + let lock_was_free = Arc::new(AtomicBool::new(false)); + let waker = Waker::from(Arc::new(SessionLockProbe { + session: Arc::downgrade(&session), + lock_was_free: Arc::clone(&lock_was_free), + })); + let mut context = Context::from_waker(&waker); + let mut notified = Box::pin(manager.budget_notify().notified_owned()); + assert!(matches!(notified.as_mut().poll(&mut context), Poll::Pending)); + + session.with_state_effects(|state, effects| { + let bytes = state.pending_control_bytes; + let items = state.pending_control_items; + state.pending_frames.clear(); + state.pending_windows.clear(); + session.release_locked(state, effects, bytes, items, true); + assert_eq!(state.pending_bytes, 0); + assert_eq!(state.pending_items, 0); + }); + + assert!(lock_was_free.load(Ordering::Acquire)); + session.close(super::SessionCloseReason::ApiClose); + manager.shutdown().await; +} + +#[tokio::test] +async fn staging_permit_wakes_waiter_after_batch_publication_and_unlock() { + let (session, manager) = session(); + queue_close(&session); + let body_len = session + .state + .lock() + .pending_frames + .front() + .map(|frame| frame.encoded.len()) + .unwrap(); + let held = manager + .try_downlink_staging_budget(session.limits.max_body_bytes_global - body_len) + .unwrap(); + let semaphore = Arc::clone(held.semaphore()); + let lock_was_free = Arc::new(AtomicBool::new(false)); + let waker = Waker::from(Arc::new(SessionLockProbe { + session: Arc::downgrade(&session), + lock_was_free: Arc::clone(&lock_was_free), + })); + let mut context = Context::from_waker(&waker); + let mut effects = DeferredSessionEffects::new(); let mut state = session.state.lock(); - assert!(session.queue_control_locked(&mut state, FrameType::Close, 1, &[])); + let batch = session + .take_down_batch_locked(&mut state, &mut effects, 0) + .unwrap(); + let mut waiter = Box::pin(semaphore.acquire_owned()); + assert!(matches!(waiter.as_mut().poll(&mut context), Poll::Pending)); + assert!(state.unacked.replace(batch).is_none()); + drop(state); + + effects.finish(); + + assert!(lock_was_free.load(Ordering::Acquire)); + drop(waiter); + drop(held); + session.close(super::SessionCloseReason::ApiClose); + manager.shutdown().await; } #[tokio::test] @@ -82,11 +187,10 @@ async fn acknowledged_response_stays_resident_until_the_last_body_clone_drops() queue_close(&session); let response = session.poll_down(0).await.unwrap(); let retained = response.body.clone(); - { - let mut state = session.state.lock(); - session.release_unacked_locked(&mut state); + session.with_state_effects(|state, effects| { + session.release_unacked_locked(state, effects); assert_eq!(state.pending_bytes, 0); - } + }); assert!(session.resident.snapshot().bytes() > 0); drop(response); assert!(session.resident.snapshot().bytes() > 0); @@ -146,6 +250,41 @@ async fn newer_poll_supersedes_older_poll_without_closing_session() { manager.shutdown().await; } +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn poll_queue_scheduler_pressure_preserves_single_batch_ownership() { + const POLLS: usize = 256; + + let (session, manager) = session(); + let mut polls = tokio::task::JoinSet::new(); + for _ in 0..POLLS { + let polling = Arc::clone(&session); + polls.spawn(async move { polling.poll_down(0).await }); + } + tokio::time::timeout(Duration::from_secs(1), async { + while session.state.lock().down_epoch < POLLS as u64 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + queue_close(&session); + + let mut nonempty = 0usize; + tokio::time::timeout(Duration::from_secs(1), async { + while let Some(result) = polls.join_next().await { + let result = result.unwrap().unwrap(); + nonempty += usize::from(!result.body.is_empty()); + } + }) + .await + .unwrap(); + + assert_eq!(nonempty, 1); + assert!(!session.state.lock().closed); + session.close(super::SessionCloseReason::ApiClose); + manager.shutdown().await; +} + #[tokio::test] async fn websocket_downlink_poll_does_not_extend_the_peer_lease() { let (session, manager) = session(); diff --git a/src/web/session/effects.rs b/src/web/session/effects.rs new file mode 100644 index 0000000..e1727f2 --- /dev/null +++ b/src/web/session/effects.rs @@ -0,0 +1,221 @@ +use std::sync::Arc; +use std::task::Waker; + +use tokio::sync::{Notify, OwnedSemaphorePermit}; + +use super::{DownBatch, SessionState, StreamState, WebSession}; + +enum DeferredSessionEffect { + Wake(Waker), + DropWaker(Waker), + Notify(Arc), +} + +enum RetainedSessionResource { + Batch(DownBatch), + StagingPermit(OwnedSemaphorePermit), + Stream(StreamState), +} + +struct DeferredItems { + first: Option, + second: Option, + spill: Vec, +} + +impl DeferredItems { + fn new() -> Self { + Self { + first: None, + second: None, + spill: Vec::new(), + } + } + + fn push(&mut self, value: T) { + if self.first.is_none() { + self.first = Some(value); + } else if self.second.is_none() { + self.second = Some(value); + } else { + self.spill.push(value); + } + } + + fn into_iter(self) -> impl Iterator { + [self.first, self.second] + .into_iter() + .flatten() + .chain(self.spill) + } + + #[cfg(test)] + fn spilled(&self) -> bool { + !self.spill.is_empty() + } +} + +/// Wake-capable work detached from one session-state transaction. +#[must_use = "deferred session effects must be finished after releasing SessionState"] +pub(super) struct DeferredSessionEffects { + retained: DeferredItems, + callbacks: DeferredItems, +} + +impl DeferredSessionEffects { + /// Creates an empty callback and retained-resource accumulator. + pub(super) fn new() -> Self { + Self { + retained: DeferredItems::new(), + callbacks: DeferredItems::new(), + } + } + + /// Defers one consuming wake until the session-state guard is gone. + pub(super) fn wake(&mut self, waker: Waker) { + self.callbacks.push(DeferredSessionEffect::Wake(waker)); + } + + /// Defers one RawWaker drop without delivering a readiness signal. + pub(super) fn drop_waker(&mut self, waker: Waker) { + self.callbacks + .push(DeferredSessionEffect::DropWaker(waker)); + } + + /// Defers one exact notification without coalescing sibling effects. + pub(super) fn notify(&mut self, notify: Arc) { + self.callbacks.push(DeferredSessionEffect::Notify(notify)); + } + + /// Retains a detached response batch until its lease can drop safely. + pub(super) fn retain_batch(&mut self, batch: DownBatch) { + self.retained + .push(RetainedSessionResource::Batch(batch)); + } + + /// Retains transient semaphore capacity until state publication completes. + pub(super) fn retain_staging_permit(&mut self, permit: OwnedSemaphorePermit) { + self.retained + .push(RetainedSessionResource::StagingPermit(permit)); + } + + /// Retains replaced stream-owned wakers for a post-unlock drop. + pub(super) fn retain_stream(&mut self, stream: StreamState) { + self.retained.push(RetainedSessionResource::Stream(stream)); + } + + /// Drops retained ownership first, then dispatches callbacks in FIFO order. + pub(super) fn finish(self) { + for retained in self.retained.into_iter() { + match retained { + RetainedSessionResource::Batch(batch) => drop(batch), + RetainedSessionResource::StagingPermit(permit) => drop(permit), + RetainedSessionResource::Stream(stream) => drop(stream), + } + } + for callback in self.callbacks.into_iter() { + match callback { + DeferredSessionEffect::Wake(waker) => waker.wake(), + DeferredSessionEffect::DropWaker(waker) => drop(waker), + DeferredSessionEffect::Notify(notify) => notify.notify_waiters(), + } + } + } + + #[cfg(test)] + /// Reports whether callback storage crossed the allocation-free inline bound. + pub(super) fn callbacks_spilled(&self) -> bool { + self.callbacks.spilled() + } +} + +impl WebSession { + /// Executes one session-state transaction and dispatches callbacks after unlock. + pub(super) fn with_state_effects( + &self, + apply: impl FnOnce(&mut SessionState, &mut DeferredSessionEffects) -> R, + ) -> R { + let mut effects = DeferredSessionEffects::new(); + let result = { + let mut state = self.state.lock(); + apply(&mut state, &mut effects) + }; + effects.finish(); + result + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + use std::sync::atomic::{AtomicUsize, Ordering}; + use std::task::{Wake, Waker}; + + use tokio::sync::Semaphore; + + use super::*; + + struct OrderedWake { + id: usize, + next: Arc, + } + + impl Wake for OrderedWake { + fn wake(self: Arc) { + assert_eq!(self.next.fetch_add(1, Ordering::AcqRel), self.id); + } + } + + struct PermitOrderWake { + semaphore: Arc, + observed_release: Arc, + } + + impl Wake for PermitOrderWake { + fn wake(self: Arc) { + self.observed_release.store( + self.semaphore.available_permits(), + Ordering::Release, + ); + } + } + + #[test] + fn two_callbacks_stay_inline_and_the_third_spills_in_fifo_order() { + let next = Arc::new(AtomicUsize::new(0)); + let mut effects = DeferredSessionEffects::new(); + for id in 0..2 { + effects.wake(Waker::from(Arc::new(OrderedWake { + id, + next: Arc::clone(&next), + }))); + } + assert!(!effects.callbacks_spilled()); + effects.wake(Waker::from(Arc::new(OrderedWake { + id: 2, + next: Arc::clone(&next), + }))); + assert!(effects.callbacks_spilled()); + + effects.finish(); + + assert_eq!(next.load(Ordering::Acquire), 3); + } + + #[test] + fn retained_resources_drop_before_callbacks_run() { + let semaphore = Arc::new(Semaphore::new(1)); + let permit = Arc::clone(&semaphore).try_acquire_owned().unwrap(); + let observed_release = Arc::new(AtomicUsize::new(0)); + let mut effects = DeferredSessionEffects::new(); + effects.retain_staging_permit(permit); + effects.wake(Waker::from(Arc::new(PermitOrderWake { + semaphore, + observed_release: Arc::clone(&observed_release), + }))); + + effects.finish(); + + assert_eq!(observed_release.load(Ordering::Acquire), 1); + } +} diff --git a/src/web/session/lane_downlink.rs b/src/web/session/lane_downlink.rs index 33d2fd8..3bfb424 100644 --- a/src/web/session/lane_downlink.rs +++ b/src/web/session/lane_downlink.rs @@ -3,7 +3,7 @@ use std::sync::Arc; use bytes::{Bytes, BytesMut}; use super::resident::{OwnedBatchBody, PendingCounts, PendingResponseLease}; -use super::{CarrierLane, DownBatch, WebSession}; +use super::{CarrierLane, DeferredSessionEffects, DownBatch, WebSession}; use crate::config::WebLimitsConfig; use crate::web::frame::FrameType; use crate::web::manager::ManagerError; @@ -13,6 +13,7 @@ pub(super) fn take_lane_down_batch( session: &WebSession, limits: &WebLimitsConfig, lane: &mut CarrierLane, + effects: &mut DeferredSessionEffects, cursor: u64, carrier_health_eligible: bool, ) -> Result { @@ -35,9 +36,10 @@ pub(super) fn take_lane_down_batch( let Some(manager) = session.manager.upgrade() else { return Err(ManagerError::Closed); }; - let Some(_staging) = manager.try_downlink_staging_budget(body_len) else { + let Some(staging) = manager.try_downlink_staging_budget(body_len) else { return Err(ManagerError::Backpressure); }; + effects.retain_staging_permit(staging); let mut body = BytesMut::with_capacity(body_len); let mut data_bytes = 0usize; let mut data_items = 0usize; diff --git a/src/web/session/lane_uplink.rs b/src/web/session/lane_uplink.rs index e167411..e635efc 100644 --- a/src/web/session/lane_uplink.rs +++ b/src/web/session/lane_uplink.rs @@ -5,7 +5,9 @@ use sha2::{Digest, Sha256}; use subtle::ConstantTimeEq; use super::uplink::{AppliedProgress, inbound_reservation, validate_batch}; -use super::{PendingClass, SessionCloseReason, WebSession, insert_carrier_lane}; +use super::{ + DeferredSessionEffects, PendingClass, SessionCloseReason, WebSession, insert_carrier_lane, +}; use crate::config::WebCarrier; use crate::web::frame::{self, Frame, FrameType}; use crate::web::manager::{ManagerError, TokenHash}; @@ -44,6 +46,7 @@ impl WebSession { let mut opened = Vec::new(); let mut committed = false; let mut healthy = None; + let mut effects = DeferredSessionEffects::new(); let result = { let mut state = self.state.lock(); if state.closed { @@ -140,13 +143,28 @@ impl WebSession { return Err(ManagerError::Backpressure); } if new_lane && insert_carrier_lane(&mut state, lane_id).is_none() { - self.release_locked(&mut state, reserve_bytes, reserve_items, false); + self.release_locked( + &mut state, + &mut effects, + reserve_bytes, + reserve_items, + false, + ); drop(state); + effects.finish(); self.close(SessionCloseReason::Protocol); return Err(ManagerError::Protocol); } let Some(lane) = state.carrier_lanes.get_mut(&lane_id) else { - self.release_locked(&mut state, reserve_bytes, reserve_items, false); + self.release_locked( + &mut state, + &mut effects, + reserve_bytes, + reserve_items, + false, + ); + drop(state); + effects.finish(); return Err(ManagerError::Closed); }; lane.up_active = true; @@ -156,13 +174,20 @@ impl WebSession { let applied = self.apply_batch_locked( &mut state, &frames, + &mut effects, &mut opened, &mut None, &mut unused_bytes, &mut unused_items, &mut progress, ); - self.release_locked(&mut state, unused_bytes, unused_items, false); + self.release_locked( + &mut state, + &mut effects, + unused_bytes, + unused_items, + false, + ); if let Some(lane) = state.carrier_lanes.get_mut(&lane_id) { lane.up_active = false; if applied { @@ -175,6 +200,7 @@ impl WebSession { } applied.then_some(sequence).ok_or(ManagerError::Closed) }; + effects.finish(); if matches!(result, Err(ManagerError::Backpressure)) { return result; } diff --git a/src/web/session/lanes.rs b/src/web/session/lanes.rs index ab3fc43..abdf41e 100644 --- a/src/web/session/lanes.rs +++ b/src/web/session/lanes.rs @@ -5,8 +5,8 @@ use bytes::{BufMut, Bytes, BytesMut}; use super::lane_downlink::take_lane_down_batch; use super::{ - CarrierLaneIdentity, PendingClass, PollResult, QUEUE_ITEM_COST, QueuedFrame, - SessionCloseReason, SessionState, WebSession, remember_closed, + CarrierLaneIdentity, DeferredSessionEffects, PendingClass, PollResult, QUEUE_ITEM_COST, + QueuedFrame, SessionCloseReason, SessionState, WebSession, remember_closed, }; use crate::web::frame::{self, FrameType}; use crate::web::manager::ManagerError; @@ -68,6 +68,7 @@ impl WebSession { lane_closed: expected_instance.is_some(), }); } + let mut effects = DeferredSessionEffects::new(); let (instance, epoch, notify, healthy) = { let mut state = self.state.lock(); if state.closed || self.cancel.is_cancelled() { @@ -146,15 +147,18 @@ impl WebSession { if let Some(stream) = state.streams.get_mut(&lane_id) && let Some(waker) = stream.write_waker.take() { - waker.wake(); + effects.wake(waker); } + effects.retain_batch(batch); } - let lane = state - .carrier_lanes - .get_mut(&lane_id) - .ok_or(ManagerError::Protocol)?; + let Some(lane) = state.carrier_lanes.get_mut(&lane_id) else { + drop(state); + effects.finish(); + return Err(ManagerError::Protocol); + }; let Some(epoch) = lane.down_epoch.checked_add(1) else { drop(state); + effects.finish(); self.close(SessionCloseReason::Protocol); return Err(ManagerError::Protocol); }; @@ -171,6 +175,7 @@ impl WebSession { let healthy = self.carrier_health_ready_locked(&mut state, Instant::now()); (instance, epoch, notify, healthy) }; + effects.finish(); if let Some(claim) = healthy { self.finish_carrier_health(claim); } @@ -182,6 +187,7 @@ impl WebSession { let notified = notify.notified(); tokio::pin!(notified); notified.as_mut().enable(); + let mut effects = DeferredSessionEffects::new(); { let mut state = self.state.lock(); if state.closed || self.cancel.is_cancelled() { @@ -217,6 +223,7 @@ impl WebSession { self, &self.limits, lane, + &mut effects, cursor, carrier_health_eligible, ) { @@ -235,8 +242,11 @@ impl WebSession { next_cursor: batch.next_cursor, lane_closed: false, }; - lane.unacked = Some(batch); + if let Some(previous) = lane.unacked.replace(batch) { + effects.retain_batch(previous); + } drop(state); + effects.finish(); if let Some(manager) = self.manager.upgrade() { manager.record_down(result.body.len()); } @@ -316,6 +326,7 @@ impl WebSession { pub(super) fn queue_lane_frame_locked( &self, state: &mut SessionState, + effects: &mut DeferredSessionEffects, frame_type: FrameType, stream_id: u32, payload: &[u8], @@ -343,7 +354,7 @@ impl WebSession { { queued.encoded[frame::HEADER_BYTES..frame::HEADER_BYTES + 4] .copy_from_slice(&total.to_be_bytes()); - lane.notify.notify_waiters(); + effects.notify(Arc::clone(&lane.notify)); return true; } } @@ -374,11 +385,11 @@ impl WebSession { return false; } let Some(lane) = state.carrier_lanes.get_mut(&stream_id) else { - self.release_locked(state, payload.len(), 0, false); + self.release_locked(state, effects, payload.len(), 0, false); return false; }; let Some(last) = lane.pending_frames.back_mut() else { - self.release_locked(state, payload.len(), 0, false); + self.release_locked(state, effects, payload.len(), 0, false); return false; }; last.encoded.extend_from_slice(payload); @@ -386,7 +397,7 @@ impl WebSession { let payload_len = (last.encoded.len() - frame::HEADER_BYTES) as u32; last.encoded[4..8].copy_from_slice(&payload_len.to_be_bytes()); lane.pending_bytes += payload.len(); - lane.notify.notify_waiters(); + effects.notify(Arc::clone(&lane.notify)); return true; } let cost = frame::HEADER_BYTES + payload.len() + QUEUE_ITEM_COST; @@ -418,7 +429,7 @@ impl WebSession { encoded.put_u32(payload.len() as u32); encoded.extend_from_slice(payload); let Some(lane) = state.carrier_lanes.get_mut(&stream_id) else { - self.release_locked(state, cost, 1, control); + self.release_locked(state, effects, cost, 1, control); return false; }; let index = lane.pending_frames.len(); @@ -436,28 +447,38 @@ impl WebSession { if frame_type == FrameType::Window { lane.pending_windows.insert(stream_id, index); } - lane.notify.notify_waiters(); + effects.notify(Arc::clone(&lane.notify)); true } - pub(super) fn remember_closed_locked(&self, state: &mut SessionState, stream_id: u32) { + pub(super) fn remember_closed_locked( + &self, + state: &mut SessionState, + effects: &mut DeferredSessionEffects, + stream_id: u32, + ) { let evicted = remember_closed(state, stream_id, self.limits.max_tombstones_per_session); if !self.carrier().uses_lanes() { return; } if let Some(evicted) = evicted { - self.release_lane_locked(state, evicted); + self.release_lane_locked(state, effects, evicted); } if let Some(lane) = state.carrier_lanes.get(&stream_id) { - lane.notify.notify_waiters(); + effects.notify(Arc::clone(&lane.notify)); } } - pub(super) fn release_lane_locked(&self, state: &mut SessionState, lane_id: u32) { + pub(super) fn release_lane_locked( + &self, + state: &mut SessionState, + effects: &mut DeferredSessionEffects, + lane_id: u32, + ) { let Some(mut lane) = state.carrier_lanes.remove(&lane_id) else { return; }; - lane.notify.notify_waiters(); + effects.notify(Arc::clone(&lane.notify)); let mut data_bytes = 0usize; let mut data_items = 0usize; let mut control_bytes = 0usize; @@ -475,10 +496,11 @@ impl WebSession { batch.lease.detach(); self.release_local_locked(state, batch.data_bytes, batch.data_items, false); self.release_local_locked(state, batch.control_bytes, batch.control_items, true); + effects.retain_batch(batch); } - self.release_locked(state, data_bytes, data_items, false); - self.release_locked(state, control_bytes, control_items, true); - self.lane_open_notify.notify_waiters(); + self.release_locked(state, effects, data_bytes, data_items, false); + self.release_locked(state, effects, control_bytes, control_items, true); + effects.notify(Arc::clone(&self.lane_open_notify)); } } diff --git a/src/web/session/lanes/tests.rs b/src/web/session/lanes/tests.rs index 020ea5f..06d568b 100644 --- a/src/web/session/lanes/tests.rs +++ b/src/web/session/lanes/tests.rs @@ -1,5 +1,8 @@ +use std::collections::VecDeque; use std::net::SocketAddr; use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::task::{Wake, Waker}; use arc_swap::ArcSwap; use bytes::BytesMut; @@ -10,7 +13,21 @@ use crate::config::{ }; use crate::maestro::generation::test_runtime_generation; use crate::web::manager::WebProcessRuntime; -use crate::web::session::{CarrierLane, insert_carrier_lane}; +use crate::web::session::{CarrierLane, StreamState, insert_carrier_lane}; + +struct SessionLockProbe { + session: std::sync::Weak, + lock_was_free: Arc, +} + +impl Wake for SessionLockProbe { + fn wake(self: Arc) { + if let Some(session) = self.session.upgrade() { + self.lock_was_free + .store(session.state.try_lock().is_some(), Ordering::Release); + } + } +} fn session_with_limits(limits: WebLimitsConfig) -> Arc { new_session(limits, std::sync::Weak::new()) @@ -96,12 +113,11 @@ async fn early_down_waits_without_creating_a_provisional_lane() { tokio::task::yield_now().await; } assert!(!session.state.lock().carrier_lanes.contains_key(&7)); - { - let mut state = session.state.lock(); - assert!(insert_carrier_lane(&mut state, 7).is_some()); + session.with_state_effects(|state, effects| { + assert!(insert_carrier_lane(state, 7).is_some()); state.closed_streams.insert(7); - assert!(session.queue_control_locked(&mut state, FrameType::Close, 7, &[])); - } + assert!(session.queue_control_locked(state, effects, FrameType::Close, 7, &[])); + }); session.lane_open_notify.notify_waiters(); let result = tokio::time::timeout(Duration::from_secs(1), poll) .await @@ -217,12 +233,11 @@ fn cross_lane_frame_is_fatal_to_https_lane_session() { #[tokio::test] async fn drained_closed_lane_replays_then_signals_completion() { let (session, manager) = session_with_manager(); - { - let mut state = session.state.lock(); + session.with_state_effects(|state, effects| { state.carrier_lanes.insert(7, CarrierLane::new(7)); state.closed_streams.insert(7); - assert!(session.queue_control_locked(&mut state, FrameType::Close, 7, &[])); - } + assert!(session.queue_control_locked(state, effects, FrameType::Close, 7, &[])); + }); let first = session.poll_down_lane(7, 0).await.unwrap(); let replay = session.poll_down_lane(7, 0).await.unwrap(); assert_eq!(first.body, replay.body); @@ -236,6 +251,56 @@ async fn drained_closed_lane_replays_then_signals_completion() { manager.shutdown().await; } +#[tokio::test] +async fn lane_ack_wakes_writer_only_after_releasing_session_lock() { + let (session, manager) = session_with_manager(); + session.with_state_effects(|state, effects| { + state.carrier_lanes.insert(7, CarrierLane::new(7)); + state.streams.insert( + 7, + StreamState { + instance: 1, + inbound: VecDeque::new(), + receive_window: frame::INITIAL_STREAM_WINDOW, + send_credit: 0, + read_waker: None, + write_waker: None, + }, + ); + assert!(session.queue_control_locked( + state, + effects, + FrameType::Window, + 7, + &frame::window_payload(1), + )); + }); + let first = session.poll_down_lane(7, 0).await.unwrap(); + let lock_was_free = Arc::new(AtomicBool::new(false)); + let waker = Waker::from(Arc::new(SessionLockProbe { + session: Arc::downgrade(&session), + lock_was_free: Arc::clone(&lock_was_free), + })); + session.with_state_effects(|state, effects| { + state.streams.get_mut(&7).unwrap().write_waker = Some(waker); + assert!(session.queue_control_locked( + state, + effects, + FrameType::Window, + 7, + &frame::window_payload(1), + )); + }); + + let second = session.poll_down_lane(7, first.next_cursor).await.unwrap(); + + assert!(lock_was_free.load(Ordering::Acquire)); + drop(first); + drop(second); + session.close(super::super::SessionCloseReason::ApiClose); + manager.shutdown().await; +} + #[test] fn tombstone_eviction_releases_lane_budget_and_accepts_late_frames() { let limits = WebLimitsConfig { @@ -243,8 +308,7 @@ fn tombstone_eviction_releases_lane_budget_and_accepts_late_frames() { ..WebLimitsConfig::default() }; let session = session_with_limits(limits); - { - let mut state = session.state.lock(); + session.with_state_effects(|state, effects| { state.carrier_lanes.insert(7, CarrierLane::new(7)); let encoded = frame::encode(FrameType::Close, 7, &[]); let cost = encoded.len() + QUEUE_ITEM_COST; @@ -264,13 +328,13 @@ fn tombstone_eviction_releases_lane_budget_and_accepts_late_frames() { state.pending_items = 1; state.pending_control_bytes = cost; state.pending_control_items = 1; - session.remember_closed_locked(&mut state, 7); + session.remember_closed_locked(state, effects, 7); state.carrier_lanes.insert(8, CarrierLane::new(8)); - session.remember_closed_locked(&mut state, 8); + session.remember_closed_locked(state, effects, 8); assert!(!state.carrier_lanes.contains_key(&7)); assert_eq!(state.pending_bytes, 0); assert_eq!(state.pending_items, 0); - } + }); let late = frame::encode(FrameType::Data, 7, b"late"); assert_eq!(session.process_up_lane(7, 7, &late), Ok(7)); assert!(!session.state.lock().closed); diff --git a/src/web/session/lifecycle.rs b/src/web/session/lifecycle.rs index 22cedf2..9307750 100644 --- a/src/web/session/lifecycle.rs +++ b/src/web/session/lifecycle.rs @@ -1,11 +1,7 @@ -use std::sync::Arc; use std::sync::atomic::Ordering; -use std::task::Waker; use std::time::{Duration, Instant}; -use tokio::sync::Notify; - -use super::WebSession; +use super::{DeferredSessionEffects, WebSession}; /// Stable terminal cause assigned by the first session-close winner. #[derive(Clone, Copy, Debug, PartialEq, Eq)] @@ -108,8 +104,7 @@ struct ReleasedQueues { recovery_closed_before_commit: bool, reason: SessionCloseReason, peer_gap: Duration, - stream_wakers: Vec, - lane_notifies: Vec>, + effects: DeferredSessionEffects, } /// Deferred queue release after manager publication linearizes a supersede. @@ -138,6 +133,7 @@ impl WebSession { /// Closes carrier state while relay tasks retain their admission until exit. pub(crate) fn close(&self, reason: SessionCloseReason) -> SessionCloseOutcome { + let effects = DeferredSessionEffects::new(); let mut state = self.state.lock(); if state.closed || state.close_requested.is_some() { return SessionCloseOutcome::AlreadyClosing; @@ -146,7 +142,7 @@ impl WebSession { state.close_requested = Some(reason); return SessionCloseOutcome::Deferred; } - let released = self.release_on_close_locked(&mut state, reason); + let released = self.release_on_close_locked(&mut state, reason, effects); drop(state); self.finish_close(released); SessionCloseOutcome::Closed @@ -171,23 +167,22 @@ impl WebSession { /// Restores an uncommitted attempt after successor admission failed. pub(crate) fn cancel_carrier_supersede(&self) { - let released = { - let mut state = self.state.lock(); - if !state.closed && state.negotiation_phase == SessionNegotiationPhase::Replacing { - state.negotiation_phase = SessionNegotiationPhase::Uncommitted; - } - state - .close_requested - .filter(|_| !state.closed) - .map(|reason| self.release_on_close_locked(&mut state, reason)) - }; - if let Some(released) = released { - self.finish_close(released); + let effects = DeferredSessionEffects::new(); + let mut state = self.state.lock(); + if !state.closed && state.negotiation_phase == SessionNegotiationPhase::Replacing { + state.negotiation_phase = SessionNegotiationPhase::Uncommitted; } + let Some(reason) = state.close_requested.filter(|_| !state.closed) else { + return; + }; + let released = self.release_on_close_locked(&mut state, reason, effects); + drop(state); + self.finish_close(released); } /// Linearizes manager publication against close requests on the old token. pub(crate) fn prepare_carrier_supersede(&self) -> Option> { + let effects = DeferredSessionEffects::new(); let mut state = self.state.lock(); if state.closed || state.negotiation_phase != SessionNegotiationPhase::Replacing @@ -196,7 +191,11 @@ impl WebSession { return None; } let released = - self.release_on_close_locked(&mut state, SessionCloseReason::CarrierSuperseded); + self.release_on_close_locked( + &mut state, + SessionCloseReason::CarrierSuperseded, + effects, + ); Some(CarrierSupersedeCompletion { session: self, released, @@ -254,6 +253,7 @@ impl WebSession { } fn begin_idle_close(&self, now: Instant) -> Option { + let effects = DeferredSessionEffects::new(); let mut state = self.state.lock(); if state.closed || state.close_requested.is_some() { return None; @@ -264,13 +264,18 @@ impl WebSession { { return None; } - Some(self.release_on_close_locked(&mut state, SessionCloseReason::PeerIdle)) + Some(self.release_on_close_locked( + &mut state, + SessionCloseReason::PeerIdle, + effects, + )) } fn release_on_close_locked( &self, state: &mut super::SessionState, reason: SessionCloseReason, + mut effects: DeferredSessionEffects, ) -> ReleasedQueues { let peer_gap = state.activity.peer_idle(Instant::now()); let closed_before_health = self.automatic_carrier @@ -282,13 +287,12 @@ impl WebSession { if reason == SessionCloseReason::CarrierSuperseded { state.negotiation_phase = SessionNegotiationPhase::Superseded; } - let mut stream_wakers = Vec::with_capacity(state.streams.len().saturating_mul(2)); for stream in state.streams.values_mut() { if let Some(waker) = stream.read_waker.take() { - stream_wakers.push(waker); + effects.wake(waker); } if let Some(waker) = stream.write_waker.take() { - stream_wakers.push(waker); + effects.wake(waker); } } state.streams.clear(); @@ -298,20 +302,21 @@ impl WebSession { batch.lease.detach(); self.release_local_locked(state, batch.data_bytes, batch.data_items, false); self.release_local_locked(state, batch.control_bytes, batch.control_items, true); + effects.retain_batch(batch); } let mut lane_data_bytes = 0usize; let mut lane_data_items = 0usize; let mut lane_control_bytes = 0usize; let mut lane_control_items = 0usize; - let mut lane_notifies = Vec::with_capacity(state.carrier_lanes.len()); for lane in state.carrier_lanes.values_mut() { - lane_notifies.push(Arc::clone(&lane.notify)); + effects.notify(std::sync::Arc::clone(&lane.notify)); if let Some(batch) = lane.unacked.take() { batch.lease.detach(); lane_data_bytes = lane_data_bytes.saturating_add(batch.data_bytes); lane_data_items = lane_data_items.saturating_add(batch.data_items); lane_control_bytes = lane_control_bytes.saturating_add(batch.control_bytes); lane_control_items = lane_control_items.saturating_add(batch.control_items); + effects.retain_batch(batch); } } self.release_local_locked(state, lane_data_bytes, lane_data_items, false); @@ -334,25 +339,11 @@ impl WebSession { recovery_closed_before_commit, reason, peer_gap, - stream_wakers, - lane_notifies, + effects, } } - fn finish_close(&self, released: ReleasedQueues) { - for waker in released.stream_wakers { - waker.wake(); - } - for notify in released.lane_notifies { - notify.notify_waiters(); - } - self.cancel.cancel(); - if self.carrier().is_multiplexed() { - self.down_notify.notify_waiters(); - } - if self.carrier().uses_lanes() { - self.lane_open_notify.notify_waiters(); - } + fn finish_close(&self, mut released: ReleasedQueues) { let manager = self.manager.upgrade(); if let Some(manager) = &manager { if released.closed_before_health { @@ -366,18 +357,26 @@ impl WebSession { crate::web::telemetry::WebBridgeRecoveryEvent::ClosedBeforeCommit, ); } - manager.release_pending( + released.effects.notify(manager.release_pending_quiet( self.profile_key, released.data_bytes, released.data_items, false, - ); - manager.release_pending( + )); + released.effects.notify(manager.release_pending_quiet( self.profile_key, released.control_bytes, released.control_items, true, - ); + )); + } + released.effects.finish(); + self.cancel.cancel(); + if self.carrier().is_multiplexed() { + self.down_notify.notify_waiters(); + } + if self.carrier().uses_lanes() { + self.lane_open_notify.notify_waiters(); } if !self.finished.swap(true, Ordering::AcqRel) { if let Some(manager) = &manager { diff --git a/src/web/session/stream_io.rs b/src/web/session/stream_io.rs new file mode 100644 index 0000000..2d71201 --- /dev/null +++ b/src/web/session/stream_io.rs @@ -0,0 +1,238 @@ +use std::io; +use std::sync::Arc; +use std::task::{Context, Poll, Waker}; +use std::time::Instant; + +use tokio::io::ReadBuf; +use tokio::sync::Notify; + +use super::{DeferredSessionEffects, QUEUE_ITEM_COST, SessionCloseReason, SessionState}; +use super::{StreamIdentity, WebSession}; +use crate::web::frame; + +enum ReadAttempt { + Ready, + Pending, + Backpressure, +} + +enum WriteAttempt { + Ready(usize), + Pending, + Closed, +} + +impl WebSession { + /// Polls client-to-server bytes and returns consumed flow-control credit. + pub(in crate::web) fn poll_read( + &self, + stream: StreamIdentity, + cx: &mut Context<'_>, + output: &mut ReadBuf<'_>, + ) -> Poll> { + let first = self.with_state_effects(|state, effects| { + self.try_read_locked(state, effects, stream, output, None) + }); + if !matches!(first, ReadAttempt::Pending) { + return self.finish_read_attempt(first); + } + + // Cloning an arbitrary RawWaker may invoke caller code, so it is never + // performed while the session state is locked. + let waker = cx.waker().clone(); + let second = self.with_state_effects(|state, effects| { + self.try_read_locked(state, effects, stream, output, Some(waker)) + }); + self.finish_read_attempt(second) + } + + fn try_read_locked( + &self, + state: &mut SessionState, + effects: &mut DeferredSessionEffects, + stream: StreamIdentity, + output: &mut ReadBuf<'_>, + prepared_waker: Option, + ) -> ReadAttempt { + let (count, finished) = { + let Some(stream_state) = state + .streams + .get_mut(&stream.id) + .filter(|state| state.instance == stream.instance) + else { + if let Some(waker) = prepared_waker { + effects.drop_waker(waker); + } + return ReadAttempt::Ready; + }; + let Some(chunk) = stream_state.inbound.front_mut() else { + if let Some(waker) = prepared_waker { + if let Some(previous) = stream_state.read_waker.replace(waker) { + effects.drop_waker(previous); + } + } + return ReadAttempt::Pending; + }; + if let Some(waker) = prepared_waker { + effects.drop_waker(waker); + } + let available = &chunk.bytes[chunk.offset..]; + let count = available.len().min(output.remaining()); + output.put_slice(&available[..count]); + chunk.offset += count; + let finished = chunk.offset == chunk.bytes.len(); + if finished { + stream_state.inbound.pop_front(); + } + 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( + state, + effects, + count + overhead, + usize::from(finished), + false, + ); + if !self.queue_window_locked(state, effects, stream.id, count as u32) { + return ReadAttempt::Backpressure; + } + ReadAttempt::Ready + } + + fn finish_read_attempt(&self, attempt: ReadAttempt) -> Poll> { + match attempt { + ReadAttempt::Ready => Poll::Ready(Ok(())), + ReadAttempt::Pending => Poll::Pending, + ReadAttempt::Backpressure => { + self.close(SessionCloseReason::Backpressure); + Poll::Ready(Err(io::Error::other( + "WEB session control budget exhausted", + ))) + } + } + } + + /// Polls server-to-client writes against stream credit and bounded queues. + pub(in crate::web) fn poll_write( + &self, + stream: StreamIdentity, + cx: &mut Context<'_>, + input: &[u8], + ) -> Poll> { + if input.is_empty() { + return Poll::Ready(Ok(0)); + } + let first = self.with_state_effects(|state, effects| { + self.try_write_locked(state, effects, stream, input, None) + }); + if !matches!(first, WriteAttempt::Pending) { + return finish_write_attempt(first); + } + + // The second locked check closes the producer-versus-registration race. + let waker = cx.waker().clone(); + let second = self.with_state_effects(|state, effects| { + self.try_write_locked(state, effects, stream, input, Some(waker)) + }); + finish_write_attempt(second) + } + + fn try_write_locked( + &self, + state: &mut SessionState, + effects: &mut DeferredSessionEffects, + stream: StreamIdentity, + input: &[u8], + prepared_waker: Option, + ) -> WriteAttempt { + let Some(stream_state) = state + .streams + .get_mut(&stream.id) + .filter(|state| state.instance == stream.instance) + else { + if let Some(waker) = prepared_waker { + effects.drop_waker(waker); + } + return WriteAttempt::Closed; + }; + let count = input + .len() + .min(frame::DATA_CHUNK_BYTES) + .min(self.limits.max_frame_payload_bytes) + .min(if self.carrier().uses_lanes() { + self.limits + .pending_bytes_per_lane + .saturating_sub(frame::HEADER_BYTES + QUEUE_ITEM_COST) + } else { + usize::MAX + }) + .min(stream_state.send_credit as usize); + if count == 0 { + install_waker(&mut stream_state.write_waker, prepared_waker, effects); + return WriteAttempt::Pending; + } + if !self.queue_data_locked(state, effects, stream.id, &input[..count]) { + if let Some(stream_state) = state + .streams + .get_mut(&stream.id) + .filter(|state| state.instance == stream.instance) + { + install_waker(&mut stream_state.write_waker, prepared_waker, effects); + } else if let Some(waker) = prepared_waker { + effects.drop_waker(waker); + } + return WriteAttempt::Pending; + } + if let Some(waker) = prepared_waker { + effects.drop_waker(waker); + } + let Some(stream_state) = state + .streams + .get_mut(&stream.id) + .filter(|state| state.instance == stream.instance) + else { + return WriteAttempt::Closed; + }; + stream_state.send_credit -= count as u64; + state.activity.touch_progress(Instant::now()); + if self.carrier().is_multiplexed() { + effects.notify(Arc::clone(&self.down_notify)); + } + WriteAttempt::Ready(count) + } + + /// Returns the process queue-capacity notification source while the manager lives. + pub(in crate::web) fn budget_notify(&self) -> Option> { + self.manager + .upgrade() + .map(|manager| manager.budget_notify()) + } +} + +fn install_waker( + slot: &mut Option, + prepared: Option, + effects: &mut DeferredSessionEffects, +) { + if let Some(prepared) = prepared + && let Some(previous) = slot.replace(prepared) + { + effects.drop_waker(previous); + } +} + +fn finish_write_attempt(attempt: WriteAttempt) -> Poll> { + match attempt { + WriteAttempt::Ready(count) => Poll::Ready(Ok(count)), + WriteAttempt::Pending => Poll::Pending, + WriteAttempt::Closed => Poll::Ready(Err(io::Error::new( + io::ErrorKind::BrokenPipe, + "WEB logical stream is closed", + ))), + } +} + +#[cfg(test)] +mod tests; diff --git a/src/web/session/stream_io/tests.rs b/src/web/session/stream_io/tests.rs new file mode 100644 index 0000000..b6da12e --- /dev/null +++ b/src/web/session/stream_io/tests.rs @@ -0,0 +1,331 @@ +use std::collections::VecDeque; +use std::mem::ManuallyDrop; +use std::net::SocketAddr; +use std::sync::{Arc, Barrier, Weak}; +use std::sync::atomic::{AtomicBool, AtomicU8, AtomicUsize, Ordering}; +use std::task::{Context, Poll, RawWaker, RawWakerVTable, Waker}; + +use bytes::Bytes; +use tokio::io::ReadBuf; + +use super::super::{InboundChunk, StreamIdentity, StreamState, WebSession}; +use crate::config::{ + WebCarrier, WebLimitsConfig, WebRuntimeProfile, WebSecretMode, WebTimeoutsConfig, +}; +use crate::web::frame; +use crate::web::frame::FrameType; +use crate::web::manager::WebProcessRuntime; + +const CLONE_ACTION_NONE: u8 = 0; +const CLONE_ACTION_INSERT_DATA: u8 = 1; + +struct CallbackProbe { + session: Weak, + stream: StreamIdentity, + clone_action: AtomicU8, + clones: AtomicUsize, + drops: AtomicUsize, + wakes: AtomicUsize, + clone_while_locked: AtomicBool, + drop_while_locked: AtomicBool, + wake_while_locked: AtomicBool, +} + +struct WakeCounter(AtomicUsize); + +impl std::task::Wake for WakeCounter { + fn wake(self: Arc) { + self.0.fetch_add(1, Ordering::AcqRel); + } +} + +impl CallbackProbe { + fn new(session: &Arc, stream: StreamIdentity, clone_action: u8) -> Arc { + Arc::new(Self { + session: Arc::downgrade(session), + stream, + clone_action: AtomicU8::new(clone_action), + clones: AtomicUsize::new(0), + drops: AtomicUsize::new(0), + wakes: AtomicUsize::new(0), + clone_while_locked: AtomicBool::new(false), + drop_while_locked: AtomicBool::new(false), + wake_while_locked: AtomicBool::new(false), + }) + } + + fn on_clone(&self) { + self.clones.fetch_add(1, Ordering::AcqRel); + let Some(session) = self.session.upgrade() else { + return; + }; + let Some(mut state) = session.state.try_lock() else { + self.clone_while_locked.store(true, Ordering::Release); + return; + }; + if self.clone_action.swap(CLONE_ACTION_NONE, Ordering::AcqRel) + == CLONE_ACTION_INSERT_DATA + && let Some(stream) = state + .streams + .get_mut(&self.stream.id) + .filter(|stream| stream.instance == self.stream.instance) + { + stream.inbound.push_back(InboundChunk { + bytes: Bytes::from_static(b"x"), + offset: 0, + }); + } + } + + fn on_drop(&self) { + self.drops.fetch_add(1, Ordering::AcqRel); + if let Some(session) = self.session.upgrade() + && session.state.try_lock().is_none() + { + self.drop_while_locked.store(true, Ordering::Release); + } + } + + fn on_wake(&self) { + self.wakes.fetch_add(1, Ordering::AcqRel); + if let Some(session) = self.session.upgrade() + && session.state.try_lock().is_none() + { + self.wake_while_locked.store(true, Ordering::Release); + } + } +} + +unsafe fn clone_probe(data: *const ()) -> RawWaker { + // SAFETY: every probe RawWaker originates from Arc::into_raw with this exact type. + let probe = unsafe { Arc::::from_raw(data.cast()) }; + probe.on_clone(); + let clone = Arc::clone(&probe); + let _ = Arc::into_raw(probe); + RawWaker::new(Arc::into_raw(clone).cast(), &PROBE_VTABLE) +} + +unsafe fn wake_probe(data: *const ()) { + // SAFETY: consuming wake reconstructs and consumes the RawWaker-owned Arc exactly once. + let probe = unsafe { Arc::::from_raw(data.cast()) }; + probe.on_wake(); +} + +unsafe fn wake_probe_by_ref(data: *const ()) { + // SAFETY: by-reference wake reconstructs the Arc without consuming its raw ownership. + let probe = ManuallyDrop::new(unsafe { Arc::::from_raw(data.cast()) }); + probe.on_wake(); +} + +unsafe fn drop_probe(data: *const ()) { + // SAFETY: dropping reconstructs and consumes the RawWaker-owned Arc exactly once. + let probe = unsafe { Arc::::from_raw(data.cast()) }; + probe.on_drop(); +} + +static PROBE_VTABLE: RawWakerVTable = + RawWakerVTable::new(clone_probe, wake_probe, wake_probe_by_ref, drop_probe); + +fn probe_waker(probe: Arc) -> Waker { + let raw = RawWaker::new(Arc::into_raw(probe).cast(), &PROBE_VTABLE); + // SAFETY: PROBE_VTABLE preserves the Arc ownership contract for every operation. + unsafe { Waker::from_raw(raw) } +} + +fn session() -> Arc { + let profile = Arc::new(WebRuntimeProfile { + host: "proxy.example.com".to_string(), + public_addr: SocketAddr::from(([203, 0, 113, 10], 443)), + user: "alice".to_string(), + secret_mode: WebSecretMode::Plain, + carrier: WebCarrier::Https, + carrier_negotiation_enabled: false, + carrier_learning: false, + carriers: Arc::from([WebCarrier::Https]), + carrier_negotiation_deadlines_secs: [3, 5, 8, 12], + capability: [0; 32], + credential_id: [0; 16], + key_fingerprint: "0000000000000000".to_string(), + max_sessions: 1, + max_streams: 1, + max_streams_per_session: 1, + }); + WebSession::new( + Weak::::new(), + [1; 32], + "192.0.2.10".parse().unwrap(), + 1, + profile, + [2; 32], + WebCarrier::Https, + 1, + [3; 32], + None, + crate::web::manager::CarrierClientClass::Legacy, + None, + false, + false, + WebLimitsConfig::default(), + WebTimeoutsConfig::default(), + None, + ) +} + +fn insert_stream(session: &WebSession, stream: StreamIdentity, send_credit: u64) { + session.state.lock().streams.insert( + stream.id, + StreamState { + instance: stream.instance, + inbound: VecDeque::new(), + receive_window: frame::INITIAL_STREAM_WINDOW, + send_credit, + read_waker: None, + write_waker: None, + }, + ); +} + +#[test] +fn read_rechecks_after_unlocked_waker_clone() { + let session = session(); + let stream = StreamIdentity { id: 1, instance: 1 }; + insert_stream(&session, stream, u64::from(frame::INITIAL_STREAM_WINDOW)); + let probe = CallbackProbe::new(&session, stream, CLONE_ACTION_INSERT_DATA); + let waker = probe_waker(Arc::clone(&probe)); + let mut context = Context::from_waker(&waker); + let mut output = []; + let mut read = ReadBuf::new(&mut output); + + assert!(matches!( + session.poll_read(stream, &mut context, &mut read), + Poll::Ready(Ok(())) + )); + assert_eq!(probe.clones.load(Ordering::Acquire), 1); + assert!(!probe.clone_while_locked.load(Ordering::Acquire)); + assert!(probe.drops.load(Ordering::Acquire) >= 1); + assert!(!probe.drop_while_locked.load(Ordering::Acquire)); +} + +#[test] +fn replacing_registered_waker_drops_the_previous_handle_after_unlock() { + let session = session(); + let stream = StreamIdentity { id: 1, instance: 1 }; + insert_stream(&session, stream, u64::from(frame::INITIAL_STREAM_WINDOW)); + let first_probe = CallbackProbe::new(&session, stream, CLONE_ACTION_NONE); + let first_waker = probe_waker(Arc::clone(&first_probe)); + let mut first_context = Context::from_waker(&first_waker); + let mut first_output = [0u8; 1]; + let mut first_read = ReadBuf::new(&mut first_output); + assert!(matches!( + session.poll_read(stream, &mut first_context, &mut first_read), + Poll::Pending + )); + let drops_before_replace = first_probe.drops.load(Ordering::Acquire); + + let second_probe = CallbackProbe::new(&session, stream, CLONE_ACTION_NONE); + let second_waker = probe_waker(Arc::clone(&second_probe)); + let mut second_context = Context::from_waker(&second_waker); + let mut second_output = [0u8; 1]; + let mut second_read = ReadBuf::new(&mut second_output); + assert!(matches!( + session.poll_read(stream, &mut second_context, &mut second_read), + Poll::Pending + )); + + assert!(first_probe.drops.load(Ordering::Acquire) > drops_before_replace); + assert!(!first_probe.drop_while_locked.load(Ordering::Acquire)); + assert!(!second_probe.clone_while_locked.load(Ordering::Acquire)); +} + +#[test] +fn write_clones_waker_only_after_releasing_session_lock() { + let session = session(); + let stream = StreamIdentity { id: 1, instance: 1 }; + insert_stream(&session, stream, 0); + let probe = CallbackProbe::new(&session, stream, CLONE_ACTION_NONE); + let waker = probe_waker(Arc::clone(&probe)); + let mut context = Context::from_waker(&waker); + + assert!(matches!( + session.poll_write(stream, &mut context, b"x"), + Poll::Pending + )); + assert_eq!(probe.clones.load(Ordering::Acquire), 1); + assert!(!probe.clone_while_locked.load(Ordering::Acquire)); + assert!(!probe.drop_while_locked.load(Ordering::Acquire)); +} + +#[test] +fn concurrent_read_registration_and_data_publication_never_lose_readiness() { + const ATTEMPTS: usize = 10_000; + + let session = session(); + let stream = StreamIdentity { id: 1, instance: 1 }; + insert_stream(&session, stream, u64::from(frame::INITIAL_STREAM_WINDOW)); + let barrier = Arc::new(Barrier::new(3)); + let outcome = Arc::new(AtomicU8::new(0)); + let wakes = Arc::new(WakeCounter(AtomicUsize::new(0))); + + std::thread::scope(|scope| { + let poll_session = Arc::clone(&session); + let poll_barrier = Arc::clone(&barrier); + let poll_outcome = Arc::clone(&outcome); + let poll_wakes = Arc::clone(&wakes); + scope.spawn(move || { + for _ in 0..ATTEMPTS { + poll_barrier.wait(); + let waker = Waker::from(Arc::clone(&poll_wakes)); + let mut context = Context::from_waker(&waker); + let mut output = []; + let mut read = ReadBuf::new(&mut output); + let value = match poll_session.poll_read(stream, &mut context, &mut read) { + Poll::Pending => 1, + Poll::Ready(Ok(())) => 2, + Poll::Ready(Err(error)) => panic!("unexpected read failure: {error}"), + }; + poll_outcome.store(value, Ordering::Release); + poll_barrier.wait(); + } + }); + + let data_session = Arc::clone(&session); + let data_barrier = Arc::clone(&barrier); + scope.spawn(move || { + let body = frame::encode(FrameType::Data, stream.id, b"x"); + for _ in 0..ATTEMPTS { + data_barrier.wait(); + let frames = frame::parse_all(&body, &data_session.limits).unwrap(); + let mut opened = Vec::new(); + let mut unused_bytes = body.len().saturating_add(super::super::QUEUE_ITEM_COST); + let mut unused_items = 1; + let mut progress = super::super::uplink::AppliedProgress::default(); + data_session.with_state_effects(|state, effects| { + assert!(data_session.apply_batch_locked( + state, + &frames, + effects, + &mut opened, + &mut None, + &mut unused_bytes, + &mut unused_items, + &mut progress, + )); + }); + data_barrier.wait(); + } + }); + + for _ in 0..ATTEMPTS { + outcome.store(0, Ordering::Release); + wakes.0.store(0, Ordering::Release); + barrier.wait(); + barrier.wait(); + let observed = outcome.load(Ordering::Acquire); + assert!(observed == 2 || wakes.0.load(Ordering::Acquire) != 0); + let mut state = session.state.lock(); + let stream_state = state.streams.get_mut(&stream.id).unwrap(); + stream_state.inbound.clear(); + assert!(stream_state.read_waker.is_none()); + } + }); +} diff --git a/src/web/session/uplink.rs b/src/web/session/uplink.rs index 9b058d5..ff77de1 100644 --- a/src/web/session/uplink.rs +++ b/src/web/session/uplink.rs @@ -9,8 +9,8 @@ use subtle::ConstantTimeEq; use super::backend::StreamCompletion; use super::{ - InboundChunk, PendingClass, QUEUE_ITEM_COST, SessionCloseReason, SessionState, StreamIdentity, - StreamState, WebSession, inbound_queue_cost, + DeferredSessionEffects, InboundChunk, PendingClass, QUEUE_ITEM_COST, SessionCloseReason, + SessionState, StreamIdentity, StreamState, WebSession, inbound_queue_cost, }; use crate::web::frame::{self, Frame, FrameType}; use crate::web::manager::{ManagerError, TokenHash}; @@ -99,6 +99,7 @@ impl WebSession { let mut opened = Vec::new(); let mut committed = false; let mut healthy = None; + let mut effects = DeferredSessionEffects::new(); let result = { let mut state = self.state.lock(); if state.closed { @@ -141,13 +142,20 @@ impl WebSession { let applied = self.apply_batch_locked( &mut state, &frames, + &mut effects, &mut opened, &mut None, &mut unused_bytes, &mut unused_items, &mut progress, ); - self.release_locked(&mut state, unused_bytes, unused_items, false); + self.release_locked( + &mut state, + &mut effects, + unused_bytes, + unused_items, + false, + ); if !applied { Err(ManagerError::Closed) } else { @@ -157,6 +165,7 @@ impl WebSession { Ok((sequence, progress.any())) } }; + effects.finish(); if matches!(result, Err(ManagerError::Backpressure)) { return result; } @@ -190,6 +199,7 @@ impl WebSession { self: &Arc, state: &mut SessionState, frames: &[Frame<'_>], + effects: &mut DeferredSessionEffects, opened: &mut Vec, reserved_open: &mut Option<(u32, u16)>, unused_bytes: &mut usize, @@ -218,10 +228,11 @@ impl WebSession { return false; } None => { - let Some(peer_port) = self.reserve_stream_locked(state) else { - self.remember_closed_locked(state, value.stream_id); + let Some(peer_port) = self.reserve_stream_locked(state, effects) else { + self.remember_closed_locked(state, effects, value.stream_id); if !self.queue_control_locked( state, + effects, FrameType::Close, value.stream_id, &[], @@ -233,7 +244,7 @@ impl WebSession { peer_port } }; - state.streams.insert( + if let Some(previous) = state.streams.insert( value.stream_id, StreamState { instance: stream.instance, @@ -243,7 +254,9 @@ impl WebSession { read_waker: None, write_waker: None, }, - ); + ) { + effects.retain_stream(previous); + } progress.accepted_open = true; opened.push(self.own_stream_task(stream, peer_port)); } @@ -261,7 +274,7 @@ impl WebSession { unused_bytes.saturating_sub(value.payload.len() + QUEUE_ITEM_COST); *unused_items = unused_items.saturating_sub(1); if let Some(waker) = stream.read_waker.take() { - waker.wake(); + effects.wake(waker); } } FrameType::Window if !was_closed => { @@ -274,7 +287,7 @@ impl WebSession { .saturating_add(u64::from(amount)) .min(u64::from(u32::MAX)); if let Some(waker) = stream.write_waker.take() { - waker.wake(); + effects.wake(waker); } } FrameType::Close if !was_closed => { @@ -285,13 +298,13 @@ impl WebSession { .closing_streams .insert(value.stream_id, stream.instance); let (bytes, items) = inbound_queue_cost(&stream.inbound); - self.release_locked(state, bytes, items, false); - self.remember_closed_locked(state, value.stream_id); + self.release_locked(state, effects, bytes, items, false); + self.remember_closed_locked(state, effects, value.stream_id); if let Some(waker) = stream.read_waker { - waker.wake(); + effects.wake(waker); } if let Some(waker) = stream.write_waker { - waker.wake(); + effects.wake(waker); } } FrameType::Data | FrameType::Window | FrameType::Close => {} @@ -301,7 +314,11 @@ impl WebSession { true } - fn reserve_stream_locked(&self, state: &mut SessionState) -> Option { + fn reserve_stream_locked( + &self, + state: &mut SessionState, + effects: &mut DeferredSessionEffects, + ) -> Option { let manager = self.manager.upgrade()?; if state.active_peer_ports.len() >= self.profile.max_streams_per_session { manager.record_stream_rejected_reason( @@ -309,23 +326,27 @@ impl WebSession { ); return None; } - let peer_port = manager - .try_acquire_stream( - self.profile_key, - self.profile.max_streams, - self.client_ip, - self.profile.public_addr, - ) - .ok()?; + let (peer_port, notify) = manager.try_acquire_stream_quiet( + self.profile_key, + self.profile.max_streams, + self.client_ip, + self.profile.public_addr, + ); + if let Some(notify) = notify { + effects.notify(notify); + } + let peer_port = peer_port.ok()?; if state.active_peer_ports.insert(peer_port) { return Some(peer_port); } - manager.release_stream( + if let Some(notify) = manager.release_stream_quiet( self.profile_key, self.client_ip, self.profile.public_addr, peer_port, - ); + ) { + effects.notify(notify); + } None } } diff --git a/src/web/session/uplink_tests.rs b/src/web/session/uplink_tests.rs index 97d0a3b..4f29894 100644 --- a/src/web/session/uplink_tests.rs +++ b/src/web/session/uplink_tests.rs @@ -130,6 +130,192 @@ fn supersede_completion_defers_stream_wake_until_finish() { assert!(lock_was_free.load(AtomicOrdering::Acquire)); } +#[test] +fn data_wakes_reader_only_after_releasing_session_lock() { + let session = session(); + let lock_was_free = Arc::new(AtomicBool::new(false)); + let waker = Waker::from(Arc::new(SessionLockProbe { + session: Arc::downgrade(&session), + lock_was_free: Arc::clone(&lock_was_free), + })); + { + let mut state = session.state.lock(); + state.streams.insert( + 1, + StreamState { + instance: 1, + inbound: VecDeque::new(), + receive_window: frame::INITIAL_STREAM_WINDOW, + send_credit: u64::from(frame::INITIAL_STREAM_WINDOW), + read_waker: Some(waker), + write_waker: None, + }, + ); + } + + let body = frame::encode(FrameType::Data, 1, &[1]); + let frames = frame::parse_all(&body, &session.limits).unwrap(); + let mut opened = Vec::new(); + let mut unused_bytes = body.len().saturating_add(QUEUE_ITEM_COST); + let mut unused_items = 1; + let mut progress = AppliedProgress::default(); + session.with_state_effects(|state, effects| { + assert!(session.apply_batch_locked( + state, + &frames, + effects, + &mut opened, + &mut None, + &mut unused_bytes, + &mut unused_items, + &mut progress, + )); + }); + assert!(lock_was_free.load(AtomicOrdering::Acquire)); +} + +#[test] +fn rejected_batch_dispatches_prior_effects_after_releasing_session_lock() { + let session = session(); + let lock_was_free = Arc::new(AtomicBool::new(false)); + let waker = Waker::from(Arc::new(SessionLockProbe { + session: Arc::downgrade(&session), + lock_was_free: Arc::clone(&lock_was_free), + })); + { + let mut state = session.state.lock(); + state.streams.insert( + 1, + StreamState { + instance: 1, + inbound: VecDeque::new(), + receive_window: frame::INITIAL_STREAM_WINDOW, + send_credit: u64::from(frame::INITIAL_STREAM_WINDOW), + read_waker: Some(waker), + write_waker: None, + }, + ); + } + + let mut body = frame::encode(FrameType::Data, 1, &[1]).to_vec(); + body.extend_from_slice(&frame::encode(FrameType::Ping, 1, &[])); + let frames = frame::parse_all(&body, &session.limits).unwrap(); + let mut opened = Vec::new(); + let mut unused_bytes = body.len().saturating_add(QUEUE_ITEM_COST); + let mut unused_items = 1; + let mut progress = AppliedProgress::default(); + let applied = session.with_state_effects(|state, effects| { + session.apply_batch_locked( + state, + &frames, + effects, + &mut opened, + &mut None, + &mut unused_bytes, + &mut unused_items, + &mut progress, + ) + }); + + assert!(!applied); + assert!(lock_was_free.load(AtomicOrdering::Acquire)); +} + +#[test] +fn window_wakes_writer_only_after_releasing_session_lock() { + let session = session(); + let lock_was_free = Arc::new(AtomicBool::new(false)); + let waker = Waker::from(Arc::new(SessionLockProbe { + session: Arc::downgrade(&session), + lock_was_free: Arc::clone(&lock_was_free), + })); + { + let mut state = session.state.lock(); + state.streams.insert( + 1, + StreamState { + instance: 1, + inbound: VecDeque::new(), + receive_window: frame::INITIAL_STREAM_WINDOW, + send_credit: 0, + read_waker: None, + write_waker: Some(waker), + }, + ); + } + + let body = frame::encode(FrameType::Window, 1, &frame::window_payload(1)); + let frames = frame::parse_all(&body, &session.limits).unwrap(); + let mut opened = Vec::new(); + let mut unused_bytes = 0; + let mut unused_items = 0; + let mut progress = AppliedProgress::default(); + session.with_state_effects(|state, effects| { + assert!(session.apply_batch_locked( + state, + &frames, + effects, + &mut opened, + &mut None, + &mut unused_bytes, + &mut unused_items, + &mut progress, + )); + }); + assert!(lock_was_free.load(AtomicOrdering::Acquire)); +} + +#[test] +fn close_frame_wakes_both_stream_halves_after_releasing_session_lock() { + let session = session(); + let read_lock_was_free = Arc::new(AtomicBool::new(false)); + let write_lock_was_free = Arc::new(AtomicBool::new(false)); + let read_waker = Waker::from(Arc::new(SessionLockProbe { + session: Arc::downgrade(&session), + lock_was_free: Arc::clone(&read_lock_was_free), + })); + let write_waker = Waker::from(Arc::new(SessionLockProbe { + session: Arc::downgrade(&session), + lock_was_free: Arc::clone(&write_lock_was_free), + })); + { + let mut state = session.state.lock(); + state.streams.insert( + 1, + StreamState { + instance: 1, + inbound: VecDeque::new(), + receive_window: frame::INITIAL_STREAM_WINDOW, + send_credit: u64::from(frame::INITIAL_STREAM_WINDOW), + read_waker: Some(read_waker), + write_waker: Some(write_waker), + }, + ); + } + let body = frame::encode(FrameType::Close, 1, &[]); + let frames = frame::parse_all(&body, &session.limits).unwrap(); + let mut opened = Vec::new(); + let mut unused_bytes = 0; + let mut unused_items = 0; + let mut progress = AppliedProgress::default(); + + session.with_state_effects(|state, effects| { + assert!(session.apply_batch_locked( + state, + &frames, + effects, + &mut opened, + &mut None, + &mut unused_bytes, + &mut unused_items, + &mut progress, + )); + }); + + assert!(read_lock_was_free.load(AtomicOrdering::Acquire)); + assert!(write_lock_was_free.load(AtomicOrdering::Acquire)); +} + #[test] fn uplink_retry_commits_only_one_exact_body() { let session = session(); diff --git a/src/web/session/websocket.rs b/src/web/session/websocket.rs index b26193d..d27b4be 100644 --- a/src/web/session/websocket.rs +++ b/src/web/session/websocket.rs @@ -5,180 +5,17 @@ use sha2::{Digest, Sha256}; use super::uplink::{AppliedProgress, inbound_reservation, validate_batch}; use super::{ - CarrierLaneIdentity, PendingClass, StreamIdentity, WebSession, WebSocketLaneClaim, + DeferredSessionEffects, PendingClass, StreamIdentity, WebSession, WebSocketLaneClaim, inbound_queue_cost, insert_carrier_lane, }; 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, - claim: WebSocketLaneClaim, - stream: Option, - phase: WebSocketLaneReservationPhase, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -enum WebSocketLaneReservationPhase { - Reserved, - Bound, - Transferred, - StreamOwned, - Closing, - Released, -} - -/// Session-wide ownership of the only automatic WebSocket carrier probe. -pub(crate) struct WebSocketProbeReservation { - session: Arc, - owner: Option, -} - -impl WebSocketProbeReservation { - /// Binds the admitted process connection to the future commit acknowledgement. - pub(crate) fn bind(&mut self, owner: u64) -> Result<(), ManagerError> { - if self.session.close_if_cancelled() { - return Err(ManagerError::Closed); - } - let mut state = self.session.state.lock(); - if state.closed - || self.session.cancel.is_cancelled() - || !state.websocket_probe_claimed - || state.websocket_commit_ack_owner.is_some() - { - return Err(ManagerError::Closed); - } - state.websocket_commit_ack_owner = Some(owner); - self.owner = Some(owner); - Ok(()) - } -} - -impl Drop for WebSocketProbeReservation { - fn drop(&mut self) { - let mut state = self.session.state.lock(); - state.websocket_probe_claimed = false; - if state.websocket_commit_ack_owner == self.owner { - state.websocket_commit_ack_owner = None; - if self.session.carrier_health_publication_state() - != super::CarrierHealthPublicationState::Published - { - state.websocket_commit_ack_written = false; - state.carrier_health_uplink = false; - state.carrier_health_activity_at = None; - } - } - } -} - -impl WebSocketLaneReservation { - /// Returns the logical stream owned by this connection. - pub(crate) fn lane_id(&self) -> u32 { - self.claim.lane.lane_id - } - - /// Returns the exact lane incarnation owned by this connection. - pub(crate) fn lane_identity(&self) -> CarrierLaneIdentity { - self.claim.lane - } - - /// Binds this pre-upgrade reservation to one admitted process connection. - pub(crate) fn bind(&mut self, connection_id: u64) -> Result<(), ManagerError> { - if self.phase != WebSocketLaneReservationPhase::Reserved { - return Err(ManagerError::Concurrent); - } - if self.session.close_if_cancelled() { - return Err(ManagerError::Closed); - } - let mut state = self.session.state.lock(); - if state.closed - || self.session.cancel.is_cancelled() - || state - .carrier_lanes - .get(&self.claim.lane.lane_id) - .is_none_or(|lane| lane.instance != self.claim.lane.instance) - { - return Err(ManagerError::Closed); - } - let Some(current) = state - .websocket_lane_reservations - .get_mut(&self.claim.lane.lane_id) - .filter(|current| **current == self.claim && current.connection_id.is_none()) - else { - return Err(ManagerError::Closed); - }; - current.connection_id = Some(connection_id); - self.claim.connection_id = Some(connection_id); - self.phase = WebSocketLaneReservationPhase::Bound; - Ok(()) - } - - fn transfer_to_stream(&mut self, stream: StreamIdentity) -> Result<(), ManagerError> { - if self.phase != WebSocketLaneReservationPhase::Bound - || stream.id != self.claim.lane.lane_id - { - return Err(ManagerError::Protocol); - } - if self.session.close_if_cancelled() { - return Err(ManagerError::Closed); - } - let mut state = self.session.state.lock(); - if self.session.cancel.is_cancelled() - || state - .carrier_lanes - .get(&self.claim.lane.lane_id) - .is_none_or(|lane| lane.instance != self.claim.lane.instance) - || state - .streams - .get(&stream.id) - .is_none_or(|current| current.instance != stream.instance) - || state - .websocket_lane_reservations - .get(&self.claim.lane.lane_id) - != Some(&self.claim) - { - return Err(ManagerError::Closed); - } - state - .websocket_lane_reservations - .remove(&self.claim.lane.lane_id); - self.stream = Some(stream); - self.phase = WebSocketLaneReservationPhase::Transferred; - Ok(()) - } - - fn mark_stream_owned(&mut self, stream: StreamIdentity) -> Result<(), ManagerError> { - if self.phase != WebSocketLaneReservationPhase::Transferred || self.stream != Some(stream) { - return Err(ManagerError::Protocol); - } - self.phase = WebSocketLaneReservationPhase::StreamOwned; - Ok(()) - } - - fn retain_after_rejected_spawn(&mut self) { - debug_assert_eq!(self.phase, WebSocketLaneReservationPhase::StreamOwned); - self.phase = WebSocketLaneReservationPhase::Transferred; - } - - fn release(&mut self) { - if self.phase == WebSocketLaneReservationPhase::Released { - return; - } - let stream_owned = self.phase == WebSocketLaneReservationPhase::StreamOwned; - self.phase = WebSocketLaneReservationPhase::Closing; - self.session - .release_websocket_lane_claim(self.claim, self.stream, stream_owned); - self.phase = WebSocketLaneReservationPhase::Released; - } -} - -impl Drop for WebSocketLaneReservation { - fn drop(&mut self) { - self.release(); - } -} +// Reservation ownership keeps pre-OPEN quota and exact lane identity transactional. +mod reservation; +pub(crate) use reservation::{WebSocketLaneReservation, WebSocketProbeReservation}; +use reservation::WebSocketLaneReservationPhase; impl WebSession { /// Reserves the only automatic WebSocket probe before any HTTP 101 response. @@ -234,6 +71,7 @@ impl WebSession { if self.close_if_cancelled() { return Err(ManagerError::Closed); } + let mut effects = DeferredSessionEffects::new(); let mut state = self.state.lock(); if state.closed || self.cancel.is_cancelled() { return Err(ManagerError::Closed); @@ -249,29 +87,48 @@ impl WebSession { let Some(manager) = self.manager.upgrade() else { return Err(ManagerError::Closed); }; - let peer_port = manager.try_acquire_stream( + let (peer_port, notify) = manager.try_acquire_stream_quiet( self.profile_key, self.profile.max_streams, self.client_ip, self.profile.public_addr, - )?; + ); + if let Some(notify) = notify { + effects.notify(notify); + } + let peer_port = match peer_port { + Ok(peer_port) => peer_port, + Err(error) => { + drop(state); + effects.finish(); + return Err(error); + } + }; if !state.active_peer_ports.insert(peer_port) { - manager.release_stream( + if let Some(notify) = manager.release_stream_quiet( self.profile_key, self.client_ip, self.profile.public_addr, peer_port, - ); + ) { + effects.notify(notify); + } + drop(state); + effects.finish(); return Err(ManagerError::Limit); } let Some(lane) = insert_carrier_lane(&mut state, lane_id) else { state.active_peer_ports.remove(&peer_port); - manager.release_stream( + if let Some(notify) = manager.release_stream_quiet( self.profile_key, self.client_ip, self.profile.public_addr, peer_port, - ); + ) { + effects.notify(notify); + } + drop(state); + effects.finish(); return Err(ManagerError::Protocol); }; let claim = WebSocketLaneClaim { @@ -287,18 +144,23 @@ impl WebSession { std::collections::hash_map::Entry::Occupied(_) => false, }; if !inserted { - self.release_lane_locked(&mut state, lane_id); + self.release_lane_locked(&mut state, &mut effects, lane_id); state.active_peer_ports.remove(&peer_port); - manager.release_stream( + if let Some(notify) = manager.release_stream_quiet( self.profile_key, self.client_ip, self.profile.public_addr, peer_port, - ); + ) { + effects.notify(notify); + } + drop(state); + effects.finish(); return Err(ManagerError::Concurrent); } + effects.notify(Arc::clone(&self.lane_open_notify)); drop(state); - self.lane_open_notify.notify_waiters(); + effects.finish(); Ok(WebSocketLaneReservation { session: Arc::clone(self), claim, @@ -340,6 +202,7 @@ impl WebSession { let mut opened = Vec::new(); let mut committed = false; let mut healthy = None; + let mut effects = DeferredSessionEffects::new(); let result = { let mut state = self.state.lock(); if state.closed { @@ -398,13 +261,20 @@ impl WebSession { let applied = self.apply_batch_locked( &mut state, &frames, + &mut effects, &mut opened, &mut reserved_open, &mut unused_bytes, &mut unused_items, &mut progress, ); - self.release_locked(&mut state, unused_bytes, unused_items, false); + self.release_locked( + &mut state, + &mut effects, + unused_bytes, + unused_items, + false, + ); if let Some(lane) = state.carrier_lanes.get_mut(&lane_id) { lane.up_active = false; if applied { @@ -424,6 +294,7 @@ impl WebSession { .then_some(progress.any()) .ok_or(ManagerError::Protocol) }; + effects.finish(); let progressed = result?; if committed { self.finish_carrier_commit(); @@ -465,6 +336,7 @@ impl WebSession { stream: Option, stream_owned: bool, ) { + let mut effects = DeferredSessionEffects::new(); let release_port = { let mut state = self.state.lock(); let lane_matches = state @@ -490,12 +362,12 @@ impl WebSession { .closing_streams .insert(claim.lane.lane_id, stream.instance); let (bytes, items) = inbound_queue_cost(&stream_state.inbound); - self.release_locked(&mut state, bytes, items, false); + self.release_locked(&mut state, &mut effects, bytes, items, false); if let Some(waker) = stream_state.read_waker { - waker.wake(); + effects.wake(waker); } if let Some(waker) = stream_state.write_waker { - waker.wake(); + effects.wake(waker); } false } else if stream_owned { @@ -513,17 +385,26 @@ impl WebSession { state.active_peer_ports.remove(&claim.peer_port) }; if lane_matches { - self.remember_closed_locked(&mut state, claim.lane.lane_id); + self.remember_closed_locked( + &mut state, + &mut effects, + claim.lane.lane_id, + ); if state .carrier_lanes .get(&claim.lane.lane_id) .is_some_and(|lane| lane.instance == claim.lane.instance) { - self.release_lane_locked(&mut state, claim.lane.lane_id); + self.release_lane_locked( + &mut state, + &mut effects, + claim.lane.lane_id, + ); } } release_port }; + effects.finish(); if release_port && let Some(manager) = self.manager.upgrade() { manager.release_stream( self.profile_key, diff --git a/src/web/session/websocket/reservation.rs b/src/web/session/websocket/reservation.rs new file mode 100644 index 0000000..d396667 --- /dev/null +++ b/src/web/session/websocket/reservation.rs @@ -0,0 +1,191 @@ +use std::sync::Arc; + +use super::super::{CarrierLaneIdentity, StreamIdentity, WebSession, WebSocketLaneClaim}; +use crate::web::manager::ManagerError; + +/// Pre-OPEN stream quota and synthetic tuple ownership for one WebSocket lane. +pub(crate) struct WebSocketLaneReservation { + /// Session whose exact lane incarnation owns the reservation. + pub(super) session: Arc, + /// Stable lane, tuple, and connection claim validated during teardown. + pub(super) claim: WebSocketLaneClaim, + /// Logical stream identity after a successful OPEN transfer. + pub(super) stream: Option, + /// Current single-owner lifecycle phase. + pub(super) phase: WebSocketLaneReservationPhase, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +/// Internal ownership phase for one exact WebSocket lane reservation. +pub(super) enum WebSocketLaneReservationPhase { + /// Stream quota and a synthetic tuple are reserved before upgrade. + Reserved, + /// The reservation is bound to an admitted WebSocket connection. + Bound, + /// OPEN transferred the tuple into session stream state. + Transferred, + /// A backend task owns stream completion and quota release. + StreamOwned, + /// Teardown is synchronously releasing exact incarnation ownership. + Closing, + /// All reservation-owned cleanup has completed. + Released, +} + +/// Session-wide ownership of the only automatic WebSocket carrier probe. +pub(crate) struct WebSocketProbeReservation { + /// Session owning the single automatic-carrier probe slot. + pub(super) session: Arc, + /// Bound process connection allowed to acknowledge commit. + pub(super) owner: Option, +} + +impl WebSocketProbeReservation { + /// Binds the admitted process connection to the future commit acknowledgement. + pub(crate) fn bind(&mut self, owner: u64) -> Result<(), ManagerError> { + if self.session.close_if_cancelled() { + return Err(ManagerError::Closed); + } + let mut state = self.session.state.lock(); + if state.closed + || self.session.cancel.is_cancelled() + || !state.websocket_probe_claimed + || state.websocket_commit_ack_owner.is_some() + { + return Err(ManagerError::Closed); + } + state.websocket_commit_ack_owner = Some(owner); + self.owner = Some(owner); + Ok(()) + } +} + +impl Drop for WebSocketProbeReservation { + fn drop(&mut self) { + let mut state = self.session.state.lock(); + state.websocket_probe_claimed = false; + if state.websocket_commit_ack_owner == self.owner { + state.websocket_commit_ack_owner = None; + if self.session.carrier_health_publication_state() + != super::super::CarrierHealthPublicationState::Published + { + state.websocket_commit_ack_written = false; + state.carrier_health_uplink = false; + state.carrier_health_activity_at = None; + } + } + } +} + +impl WebSocketLaneReservation { + /// Returns the logical stream owned by this connection. + pub(crate) fn lane_id(&self) -> u32 { + self.claim.lane.lane_id + } + + /// Returns the exact lane incarnation owned by this connection. + pub(crate) fn lane_identity(&self) -> CarrierLaneIdentity { + self.claim.lane + } + + /// Binds this pre-upgrade reservation to one admitted process connection. + pub(crate) fn bind(&mut self, connection_id: u64) -> Result<(), ManagerError> { + if self.phase != WebSocketLaneReservationPhase::Reserved { + return Err(ManagerError::Concurrent); + } + if self.session.close_if_cancelled() { + return Err(ManagerError::Closed); + } + let mut state = self.session.state.lock(); + if state.closed + || self.session.cancel.is_cancelled() + || state + .carrier_lanes + .get(&self.claim.lane.lane_id) + .is_none_or(|lane| lane.instance != self.claim.lane.instance) + { + return Err(ManagerError::Closed); + } + let Some(current) = state + .websocket_lane_reservations + .get_mut(&self.claim.lane.lane_id) + .filter(|current| **current == self.claim && current.connection_id.is_none()) + else { + return Err(ManagerError::Closed); + }; + current.connection_id = Some(connection_id); + self.claim.connection_id = Some(connection_id); + self.phase = WebSocketLaneReservationPhase::Bound; + Ok(()) + } + + pub(super) fn transfer_to_stream( + &mut self, + stream: StreamIdentity, + ) -> Result<(), ManagerError> { + if self.phase != WebSocketLaneReservationPhase::Bound + || stream.id != self.claim.lane.lane_id + { + return Err(ManagerError::Protocol); + } + if self.session.close_if_cancelled() { + return Err(ManagerError::Closed); + } + let mut state = self.session.state.lock(); + if self.session.cancel.is_cancelled() + || state + .carrier_lanes + .get(&self.claim.lane.lane_id) + .is_none_or(|lane| lane.instance != self.claim.lane.instance) + || state + .streams + .get(&stream.id) + .is_none_or(|current| current.instance != stream.instance) + || state + .websocket_lane_reservations + .get(&self.claim.lane.lane_id) + != Some(&self.claim) + { + return Err(ManagerError::Closed); + } + state + .websocket_lane_reservations + .remove(&self.claim.lane.lane_id); + self.stream = Some(stream); + self.phase = WebSocketLaneReservationPhase::Transferred; + Ok(()) + } + + pub(super) fn mark_stream_owned( + &mut self, + stream: StreamIdentity, + ) -> Result<(), ManagerError> { + if self.phase != WebSocketLaneReservationPhase::Transferred || self.stream != Some(stream) { + return Err(ManagerError::Protocol); + } + self.phase = WebSocketLaneReservationPhase::StreamOwned; + Ok(()) + } + + pub(super) fn retain_after_rejected_spawn(&mut self) { + debug_assert_eq!(self.phase, WebSocketLaneReservationPhase::StreamOwned); + self.phase = WebSocketLaneReservationPhase::Transferred; + } + + pub(super) fn release(&mut self) { + if self.phase == WebSocketLaneReservationPhase::Released { + return; + } + let stream_owned = self.phase == WebSocketLaneReservationPhase::StreamOwned; + self.phase = WebSocketLaneReservationPhase::Closing; + self.session + .release_websocket_lane_claim(self.claim, self.stream, stream_owned); + self.phase = WebSocketLaneReservationPhase::Released; + } +} + +impl Drop for WebSocketLaneReservation { + fn drop(&mut self) { + self.release(); + } +} diff --git a/src/web/session/websocket/tests.rs b/src/web/session/websocket/tests.rs index 7a4900c..ec245aa 100644 --- a/src/web/session/websocket/tests.rs +++ b/src/web/session/websocket/tests.rs @@ -1,5 +1,8 @@ use std::collections::BTreeMap; +use std::collections::VecDeque; use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::task::{Wake, Waker}; use arc_swap::ArcSwap; use tokio::sync::watch; @@ -9,6 +12,21 @@ use crate::config::{ProxyConfig, WebRuntimeConfig, WebRuntimeProfile, WebSecretM use crate::maestro::generation::{RuntimeGeneration, test_runtime_generation_with_admission}; use crate::web::frame::FrameType; use crate::web::manager::WebProcessRuntime; +use crate::web::session::StreamState; + +struct SessionLockProbe { + session: std::sync::Weak, + lock_was_free: Arc, +} + +impl Wake for SessionLockProbe { + fn wake(self: Arc) { + if let Some(session) = self.session.upgrade() { + self.lock_was_free + .store(session.state.try_lock().is_some(), Ordering::Release); + } + } +} struct TestRuntime { session: Arc, @@ -87,8 +105,7 @@ fn runtime(admission: bool) -> TestRuntime { fn detach_stale_lane(runtime: &TestRuntime, reservation: &WebSocketLaneReservation) { let claim = reservation.claim; - { - let mut state = runtime.session.state.lock(); + runtime.session.with_state_effects(|state, effects| { assert_eq!( state .websocket_lane_reservations @@ -97,9 +114,9 @@ fn detach_stale_lane(runtime: &TestRuntime, reservation: &WebSocketLaneReservati ); runtime .session - .release_lane_locked(&mut state, claim.lane.lane_id); + .release_lane_locked(state, effects, claim.lane.lane_id); assert!(state.active_peer_ports.remove(&claim.peer_port)); - } + }); runtime.manager.release_stream( runtime.session.profile_key, runtime.session.client_ip, @@ -411,6 +428,59 @@ async fn stale_reservation_drop_preserves_replacement_claim() { runtime.shutdown().await; } +#[tokio::test] +async fn exact_lane_teardown_wakes_stream_after_releasing_session_lock() { + let runtime = runtime(true); + let mut reservation = runtime.session.reserve_websocket_lane(7).unwrap(); + reservation.bind(1).unwrap(); + let claim = reservation.claim; + let stream = StreamIdentity { id: 7, instance: 1 }; + let read_lock_was_free = Arc::new(AtomicBool::new(false)); + let write_lock_was_free = Arc::new(AtomicBool::new(false)); + let read_waker = Waker::from(Arc::new(SessionLockProbe { + session: Arc::downgrade(&runtime.session), + lock_was_free: Arc::clone(&read_lock_was_free), + })); + let write_waker = Waker::from(Arc::new(SessionLockProbe { + session: Arc::downgrade(&runtime.session), + lock_was_free: Arc::clone(&write_lock_was_free), + })); + runtime.session.state.lock().streams.insert( + stream.id, + StreamState { + instance: stream.instance, + inbound: VecDeque::new(), + receive_window: frame::INITIAL_STREAM_WINDOW, + send_credit: u64::from(frame::INITIAL_STREAM_WINDOW), + read_waker: Some(read_waker), + write_waker: Some(write_waker), + }, + ); + + runtime + .session + .release_websocket_lane_claim(claim, Some(stream), false); + + assert!(read_lock_was_free.load(Ordering::Acquire)); + assert!(write_lock_was_free.load(Ordering::Acquire)); + assert!( + runtime + .session + .state + .lock() + .active_peer_ports + .remove(&claim.peer_port) + ); + runtime.manager.release_stream( + runtime.session.profile_key, + runtime.session.client_ip, + runtime.session.profile.public_addr, + claim.peer_port, + ); + drop(reservation); + runtime.shutdown().await; +} + #[tokio::test] async fn stale_transfer_cannot_remove_current_reservation() { let runtime = runtime(true);