diff --git a/src/maestro/listeners/accept.rs b/src/maestro/listeners/accept.rs index 3ceae5d..75e768d 100644 --- a/src/maestro/listeners/accept.rs +++ b/src/maestro/listeners/accept.rs @@ -284,7 +284,7 @@ impl ListenerSlot { } pub(super) async fn stop(&mut self) -> Result<(), String> { - self.cancellation.cancel(); + self.request_stop(); if let Some(task) = self.task.take() { task.await.map_err(|error_value| { format!("listener {} task failed: {error_value}", self.spec.addr) @@ -305,6 +305,61 @@ impl ListenerSlot { Ok(()) } + /// Cancels admission synchronously before the shared shutdown deadline starts draining. + pub(super) fn request_stop(&self) { + self.cancellation.cancel(); + } + + /// Joins this acceptor and its WEB connections by one process shutdown deadline. + pub(super) async fn stop_until( + &mut self, + deadline: tokio::time::Instant, + ) -> Result<(), String> { + self.request_stop(); + let mut errors = Vec::new(); + if let Some(mut task) = self.task.take() { + let joined = if task.is_finished() { + Some(task.await) + } else { + match tokio::time::timeout_at(deadline, &mut task).await { + Ok(result) => Some(result), + Err(_) => { + task.abort(); + let _ = task.await; + None + } + } + }; + match joined { + Some(Ok(())) => {} + Some(Err(error_value)) => errors.push(format!( + "listener {} task failed: {error_value}", + self.spec.addr + )), + None => errors.push(format!( + "listener {} accept shutdown timed out", + self.spec.addr + )), + } + } + self.connections.close(); + if !self.connections.is_empty() + && tokio::time::timeout_at(deadline, self.connections.wait()) + .await + .is_err() + { + errors.push(format!( + "listener {} connection shutdown timed out", + self.spec.addr + )); + } + if errors.is_empty() { + Ok(()) + } else { + Err(errors.join("; ")) + } + } + pub(super) fn restart(&mut self, active_runtime: Arc>) { self.active_runtime = active_runtime.clone(); self.cancellation = CancellationToken::new(); diff --git a/src/maestro/listeners/control.rs b/src/maestro/listeners/control.rs index 8da4955..a8fc0fb 100644 --- a/src/maestro/listeners/control.rs +++ b/src/maestro/listeners/control.rs @@ -1,6 +1,7 @@ use std::collections::{BTreeMap, BTreeSet}; use std::net::SocketAddr; use std::sync::Arc; +use std::time::Duration; use arc_swap::ArcSwap; @@ -13,7 +14,7 @@ use super::bind::{BoundListeners, BoundTcpListener, PreparedTcpListener, prepare use super::plan::{ListenerBindSpec, listener_bind_plan}; #[cfg(unix)] use super::unix::UnixAcceptHandle; -use crate::web::manager::WebProcessRuntime; +use crate::web::manager::{WebProcessRuntime, WebShutdownOutcome}; use crate::web::trace::WebTraceStore; /// Process-owned listener inventory and accept-task lifecycle controller. @@ -214,24 +215,76 @@ impl ListenerManager { ); } - /// Stops and joins every accept task before sockets are released. + /// Stops every accept task and applies one deadline to the complete WEB ingress. pub(crate) async fn shutdown(&mut self) -> Result<(), String> { + if self.web_runtime.is_none() { + let mut errors = Vec::new(); + for slot in self.slots.values_mut() { + if let Err(error_value) = slot.stop().await { + errors.push(error_value); + } + } + #[cfg(unix)] + if let Some(unix) = &mut self.unix + && let Err(error_value) = unix.stop().await + { + errors.push(error_value); + } + self.slots.clear(); + #[cfg(unix)] + { + self.unix = None; + } + return if errors.is_empty() { + Ok(()) + } else { + Err(errors.join("; ")) + }; + } + let timeout_secs = self + .active_runtime + .load() + .config() + .web + .timeouts + .shutdown_secs; + let now = tokio::time::Instant::now(); + let deadline = now + .checked_add(Duration::from_secs(timeout_secs)) + .unwrap_or(now); + for slot in self.slots.values() { + slot.request_stop(); + } + #[cfg(unix)] + if let Some(unix) = &self.unix { + unix.request_stop(); + } + let Some(web_runtime) = self.web_runtime.take() else { + return Err("WEB runtime disappeared during shutdown orchestration".to_string()); + }; + let drain = web_runtime.begin_shutdown(); let mut errors = Vec::new(); - for slot in self.slots.values_mut() { - if let Err(error_value) = slot.stop().await { + let slot_waits = futures_util::future::join_all( + self.slots + .values_mut() + .map(|slot| slot.stop_until(deadline)), + ); + let (slot_results, web_outcome) = tokio::join!(slot_waits, drain.wait_until(deadline)); + for result in slot_results { + if let Err(error_value) = result { errors.push(error_value); } } #[cfg(unix)] if let Some(unix) = &mut self.unix - && let Err(error_value) = unix.stop().await + && let Err(error_value) = unix.stop_until(deadline).await { errors.push(error_value); } - self.slots.clear(); - if let Some(web_runtime) = self.web_runtime.take() { - web_runtime.shutdown().await; + if web_outcome == WebShutdownOutcome::DeadlineExceeded { + errors.push("WEB ingress shutdown deadline exceeded".to_string()); } + self.slots.clear(); #[cfg(unix)] { self.unix = None; diff --git a/src/maestro/listeners/unix.rs b/src/maestro/listeners/unix.rs index 30f08df..8159361 100644 --- a/src/maestro/listeners/unix.rs +++ b/src/maestro/listeners/unix.rs @@ -148,11 +148,42 @@ impl UnixAcceptHandle { } pub(super) async fn stop(&mut self) -> Result<(), String> { - self.cancellation.cancel(); + self.request_stop(); if let Some(task) = self.task.take() { task.await .map_err(|error_value| format!("Unix listener task failed: {error_value}"))?; } Ok(()) } + + /// Cancels Unix admission before process-owned listener waits begin. + pub(super) fn request_stop(&self) { + self.cancellation.cancel(); + } + + /// Joins the Unix acceptor by the process listener deadline. + pub(super) async fn stop_until( + &mut self, + deadline: tokio::time::Instant, + ) -> Result<(), String> { + self.request_stop(); + let Some(mut task) = self.task.take() else { + return Ok(()); + }; + if task.is_finished() { + return task + .await + .map_err(|error_value| format!("Unix listener task failed: {error_value}")); + } + match tokio::time::timeout_at(deadline, &mut task).await { + Ok(result) => { + result.map_err(|error_value| format!("Unix listener task failed: {error_value}")) + } + Err(_) => { + task.abort(); + let _ = task.await; + Err("Unix listener shutdown timed out".to_string()) + } + } + } } diff --git a/src/web/http.rs b/src/web/http.rs index 9066df5..1810074 100644 --- a/src/web/http.rs +++ b/src/web/http.rs @@ -2,7 +2,7 @@ use std::convert::Infallible; use std::error::Error; use std::net::SocketAddr; use std::sync::Arc; -use std::time::{Duration, Instant}; +use std::time::Duration; use bytes::Bytes; use http_body_util::BodyExt; @@ -13,7 +13,6 @@ use hyper::service::service_fn; use hyper::{Method, Request, Response, StatusCode}; use hyper_util::rt::{TokioIo, TokioTimer}; use ipnetwork::IpNetwork; -use parking_lot::Mutex; use tokio::net::TcpStream; use tokio_util::sync::CancellationToken; @@ -45,7 +44,7 @@ mod websocket; mod trace_tests; use crate::web::trace::{HttpTraceExchange, TraceDirection, TraceLifecycleEvent, TraceRoute}; -use activity::{ActivityBody, RequestActivity}; +use activity::{ActivityBody, ConnectionActivity, RequestActivity, RequestDeadlineHandle}; use body::{CollectBodyError, CollectedBody, RequestBody, collect_body}; use decoy::serve_decoy; use down::handle_down; @@ -80,15 +79,22 @@ pub(crate) async fn serve_connection( let max_header_bytes = config.web.limits.max_header_bytes; let header_timeout = Duration::from_secs(config.web.timeouts.header_secs); let idle_timeout = Duration::from_secs(config.web.timeouts.http_idle_secs); - let last_activity = Arc::new(Mutex::new(Instant::now())); - let service_last_activity = Arc::clone(&last_activity); + let connection_activity = ConnectionActivity::new(); + let service_activity = connection_activity.clone(); let service = service_fn(move |mut request| { let runtime = Arc::clone(&runtime); let trusted_proxy_cidrs = Arc::clone(&trusted_proxy_cidrs); - let last_activity = Arc::clone(&service_last_activity); + let connection_activity = service_activity.clone(); let client_ip_source = client_ip_source; async move { - let activity = RequestActivity::begin(last_activity); + let Some(activity) = RequestActivity::begin(connection_activity) else { + let mut response = service_unavailable(); + response + .headers_mut() + .insert(header::CONNECTION, HeaderValue::from_static("close")); + return Ok::<_, Infallible>(response); + }; + request.extensions_mut().insert(activity.deadline_handle()); let trace = runtime.trace().begin_http(&request, peer.ip()); if let Some(trace) = &trace { request.extensions_mut().insert(Arc::clone(trace)); @@ -133,9 +139,7 @@ pub(crate) async fn serve_connection( _ = cancellation.cancelled() => break, _ = &mut connection => break, _ = idle_check.tick() => { - if Instant::now().saturating_duration_since(*last_activity.lock()) - >= idle_timeout - { + if connection_activity.should_close(tokio::time::Instant::now(), idle_timeout) { break; } } @@ -404,6 +408,10 @@ fn request_trace(request: &Request) -> Option<&Arc> { request.extensions().get::>() } +fn request_deadline(request: &Request) -> Option { + request.extensions().get::().cloned() +} + fn set_trace_route(request: &Request, route: TraceRoute) { if let Some(trace) = request_trace(request) { trace.set_route(route); diff --git a/src/web/http/activity.rs b/src/web/http/activity.rs index 3b7d3fc..c84d9e9 100644 --- a/src/web/http/activity.rs +++ b/src/web/http/activity.rs @@ -1,31 +1,291 @@ use std::pin::Pin; use std::sync::Arc; use std::task::{Context, Poll}; -use std::time::Instant; +use std::time::Duration; use bytes::Bytes; use hyper::body::{Body, Frame, SizeHint}; use parking_lot::Mutex; +use tokio::time::Instant; use super::{BoxError, HttpBody}; use crate::web::trace::{HttpTraceExchange, TraceBodyState, TraceDirection}; +#[derive(Clone, Copy)] +struct DeadlineSlot { + id: u64, + deadline: Instant, +} + +struct RequestSlot { + id: u64, + deadline: Option, +} + +struct UpgradeSlot { + request_id: u64, + deadline: DeadlineSlot, +} + +struct ActivityState { + last_progress: Instant, + next_request_id: u64, + next_deadline_id: u64, + request: Option, + upgrade: Option, + failed: bool, +} + +/// Shared liveness state for one accepted HTTP connection. +#[derive(Clone)] +pub(super) struct ConnectionActivity { + state: Arc>, +} + +impl ConnectionActivity { + /// Creates activity state at the connection acceptance boundary. + pub(super) fn new() -> Self { + Self { + state: Arc::new(Mutex::new(ActivityState { + last_progress: Instant::now(), + next_request_id: 1, + next_deadline_id: 1, + request: None, + upgrade: None, + failed: false, + })), + } + } + + /// Returns whether the connection has no protected operation or recent progress. + pub(super) fn should_close(&self, now: Instant, idle: Duration) -> bool { + let state = self.state.lock(); + if state.failed { + return true; + } + let request_protected = state + .request + .as_ref() + .and_then(|request| request.deadline) + .is_some_and(|deadline| now <= deadline.deadline); + let upgrade_protected = state + .upgrade + .as_ref() + .is_some_and(|upgrade| now <= upgrade.deadline.deadline); + !request_protected + && !upgrade_protected + && now.saturating_duration_since(state.last_progress) >= idle + } + + fn fail(&self) { + self.state.lock().failed = true; + } +} + +/// Cloneable authority for protecting one request's explicitly bounded awaits. +#[derive(Clone)] +pub(super) struct RequestDeadlineHandle { + activity: ConnectionActivity, + request_id: u64, +} + +impl RequestDeadlineHandle { + /// Protects the current bounded request operation until its absolute deadline. + pub(super) fn lease_until(&self, deadline: Instant) -> Option { + let now = Instant::now(); + let mut state = self.activity.state.lock(); + if state.failed { + return None; + } + let Some(current) = state.request.as_ref() else { + state.failed = true; + return None; + }; + if current.id != self.request_id + || current + .deadline + .is_some_and(|active| now <= active.deadline) + { + state.failed = true; + return None; + } + let id = state.next_deadline_id; + let Some(next) = id.checked_add(1) else { + state.failed = true; + return None; + }; + state.next_deadline_id = next; + let Some(current) = state.request.as_mut() else { + state.failed = true; + return None; + }; + current.deadline = Some(DeadlineSlot { id, deadline }); + Some(RequestDeadlineLease { + handle: self.clone(), + deadline_id: id, + }) + } + + /// Protects the current bounded request operation for one checked duration. + pub(super) fn lease_for(&self, duration: Duration) -> Option { + let Some(deadline) = Instant::now().checked_add(duration) else { + self.activity.fail(); + return None; + }; + self.lease_until(deadline) + } + + /// Transfers idle protection to a pending Hyper upgrade operation. + pub(super) fn upgrade_until(&self, deadline: Instant) -> Option { + let now = Instant::now(); + let mut state = self.activity.state.lock(); + if state.failed + || state + .request + .as_ref() + .is_none_or(|request| request.id != self.request_id) + || state + .upgrade + .as_ref() + .is_some_and(|upgrade| now <= upgrade.deadline.deadline) + { + state.failed = true; + return None; + } + let id = state.next_deadline_id; + let Some(next) = id.checked_add(1) else { + state.failed = true; + return None; + }; + state.next_deadline_id = next; + state.upgrade = Some(UpgradeSlot { + request_id: self.request_id, + deadline: DeadlineSlot { id, deadline }, + }); + Some(UpgradeDeadlineLease { + activity: self.activity.clone(), + request_id: self.request_id, + deadline_id: id, + deadline, + }) + } +} + +/// Exact request-operation lease that cannot clear a newer deadline. +pub(super) struct RequestDeadlineLease { + handle: RequestDeadlineHandle, + deadline_id: u64, +} + +impl Drop for RequestDeadlineLease { + fn drop(&mut self) { + let mut state = self.handle.activity.state.lock(); + let matches = state.request.as_ref().is_some_and(|request| { + request.id == self.handle.request_id + && request + .deadline + .is_some_and(|deadline| deadline.id == self.deadline_id) + }); + if matches { + if let Some(request) = state.request.as_mut() { + request.deadline = None; + } + state.last_progress = Instant::now(); + } + } +} + +/// Exact pending-upgrade lease retained by the spawned upgrade future. +pub(super) struct UpgradeDeadlineLease { + activity: ConnectionActivity, + request_id: u64, + deadline_id: u64, + deadline: Instant, +} + +impl UpgradeDeadlineLease { + /// Returns the absolute deadline shared with the upgrade timeout. + pub(super) fn deadline(&self) -> Instant { + self.deadline + } +} + +impl Drop for UpgradeDeadlineLease { + fn drop(&mut self) { + let mut state = self.activity.state.lock(); + let matches = state.upgrade.as_ref().is_some_and(|upgrade| { + upgrade.request_id == self.request_id && upgrade.deadline.id == self.deadline_id + }); + if matches { + state.upgrade = None; + state.last_progress = Instant::now(); + } + } +} + /// Request lifecycle guard that refreshes HTTP connection activity on completion. pub(super) struct RequestActivity { - last_activity: Arc>, + handle: RequestDeadlineHandle, } impl RequestActivity { /// Starts activity accounting for one HTTP request. - pub(super) fn begin(last_activity: Arc>) -> Self { - *last_activity.lock() = Instant::now(); - Self { last_activity } + pub(super) fn begin(activity: ConnectionActivity) -> Option { + let mut state = activity.state.lock(); + if state.failed || state.request.is_some() { + state.failed = true; + return None; + } + let id = state.next_request_id; + let Some(next) = id.checked_add(1) else { + state.failed = true; + return None; + }; + state.next_request_id = next; + state.last_progress = Instant::now(); + state.request = Some(RequestSlot { id, deadline: None }); + drop(state); + Some(Self { + handle: RequestDeadlineHandle { + activity, + request_id: id, + }, + }) + } + + /// Returns the authority copied into request extensions for bounded awaits. + pub(super) fn deadline_handle(&self) -> RequestDeadlineHandle { + self.handle.clone() + } + + fn progress(&self) { + self.handle.activity.state.lock().last_progress = Instant::now(); + } + + fn enter_response(&mut self) { + let mut state = self.handle.activity.state.lock(); + if let Some(request) = state + .request + .as_mut() + .filter(|request| request.id == self.handle.request_id) + { + request.deadline = None; + state.last_progress = Instant::now(); + } } } impl Drop for RequestActivity { fn drop(&mut self) { - *self.last_activity.lock() = Instant::now(); + let mut state = self.handle.activity.state.lock(); + if state + .request + .as_ref() + .is_some_and(|request| request.id == self.handle.request_id) + { + state.request = None; + state.last_progress = Instant::now(); + } } } @@ -41,9 +301,10 @@ impl ActivityBody { /// Binds one response body to its request activity guard. pub(super) fn new( inner: HttpBody, - activity: RequestActivity, + mut activity: RequestActivity, trace: Option>, ) -> Self { + activity.enter_response(); Self { inner, activity, @@ -88,7 +349,7 @@ impl Body for ActivityBody { Poll::Pending => {} } if result.is_ready() { - *self.activity.last_activity.lock() = Instant::now(); + self.activity.progress(); } result } @@ -107,3 +368,60 @@ impl Drop for ActivityBody { self.finish(TraceBodyState::Aborted); } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn bounded_request_deadline_suspends_only_idle_expiry() { + let activity = ConnectionActivity::new(); + let request = RequestActivity::begin(activity.clone()).unwrap(); + let now = Instant::now(); + let lease = request + .deadline_handle() + .lease_until(now + Duration::from_secs(5)) + .unwrap(); + + assert!(!activity.should_close(now + Duration::from_secs(4), Duration::from_secs(1))); + assert!(activity.should_close(now + Duration::from_secs(6), Duration::from_secs(1))); + + drop(lease); + assert!(!activity.should_close(Instant::now(), Duration::from_secs(1))); + } + + #[test] + fn stale_request_lease_cannot_clear_a_new_request_deadline() { + let activity = ConnectionActivity::new(); + let request_a = RequestActivity::begin(activity.clone()).unwrap(); + let now = Instant::now(); + let lease_a = request_a + .deadline_handle() + .lease_until(now - Duration::from_secs(1)) + .unwrap(); + drop(request_a); + let request_b = RequestActivity::begin(activity.clone()).unwrap(); + let _lease_b = request_b + .deadline_handle() + .lease_until(now + Duration::from_secs(5)) + .unwrap(); + + drop(lease_a); + + assert!(!activity.should_close(now + Duration::from_secs(4), Duration::from_secs(1))); + } + + #[test] + fn stale_upgrade_lease_cannot_clear_its_replacement() { + let activity = ConnectionActivity::new(); + let request = RequestActivity::begin(activity.clone()).unwrap(); + let handle = request.deadline_handle(); + let now = Instant::now(); + let lease_a = handle.upgrade_until(now - Duration::from_secs(1)).unwrap(); + let _lease_b = handle.upgrade_until(now + Duration::from_secs(5)).unwrap(); + + drop(lease_a); + + assert!(!activity.should_close(now + Duration::from_secs(4), Duration::from_secs(1))); + } +} diff --git a/src/web/http/body.rs b/src/web/http/body.rs index 38b4b64..80c727c 100644 --- a/src/web/http/body.rs +++ b/src/web/http/body.rs @@ -118,6 +118,7 @@ pub(super) async fn collect_body( limit: usize, allow_empty: bool, ) -> Result { + let request_deadline = super::request_deadline(&request); let exceeds_limit = request.body().size_hint().lower() > limit as u64 || request .body() @@ -134,6 +135,14 @@ pub(super) async fn collect_body( let Some((reader_budget, body_budget)) = runtime.try_body_budget(limit) else { return Err(CollectBodyError::Limit); }; + let _deadline_lease = match request_deadline { + Some(deadline) => Some( + deadline + .lease_for(body_timeout) + .ok_or(CollectBodyError::Limit)?, + ), + None => None, + }; let body = match tokio::time::timeout(body_timeout, Limited::new(body, limit).collect()).await { Ok(Ok(body)) => body.to_bytes(), _ => { diff --git a/src/web/http/decoy.rs b/src/web/http/decoy.rs index b868e21..f7652d4 100644 --- a/src/web/http/decoy.rs +++ b/src/web/http/decoy.rs @@ -184,6 +184,7 @@ async fn proxy_to_upstream( header_timeout: Duration, runtime: &WebProcessRuntime, ) -> HttpResponse { + let request_deadline = super::request_deadline(&request); remove_hop_by_hop(request.headers_mut()); if let Ok(host) = HeaderValue::from_str(authority) { request.headers_mut().insert(header::HOST, host); @@ -197,10 +198,15 @@ async fn proxy_to_upstream( return bad_gateway(); }; *request.uri_mut() = uri; + let _deadline_lease = match lease_deadline(request_deadline.as_ref(), header_timeout) { + Ok(lease) => lease, + Err(()) => return bad_gateway(), + }; let stream = match tokio::time::timeout(header_timeout, TcpStream::connect(addr)).await { Ok(Ok(stream)) => stream, _ => return bad_gateway(), }; + drop(_deadline_lease); let max_header_bytes = runtime .active_generation() .config() @@ -209,19 +215,29 @@ async fn proxy_to_upstream( .max_header_bytes; let mut builder = hyper::client::conn::http1::Builder::new(); builder.max_buf_size(max_header_bytes); + let _deadline_lease = match lease_deadline(request_deadline.as_ref(), header_timeout) { + Ok(lease) => lease, + Err(()) => return bad_gateway(), + }; let (mut sender, connection) = match tokio::time::timeout(header_timeout, builder.handshake(TokioIo::new(stream))).await { Ok(Ok(parts)) => parts, _ => return bad_gateway(), }; + drop(_deadline_lease); runtime.spawn_auxiliary(async move { let _ = connection.await; }); + let _deadline_lease = match lease_deadline(request_deadline.as_ref(), header_timeout) { + Ok(lease) => lease, + Err(()) => return bad_gateway(), + }; let mut response = match tokio::time::timeout(header_timeout, sender.send_request(request)).await { Ok(Ok(response)) => response, _ => return bad_gateway(), }; + drop(_deadline_lease); remove_hop_by_hop(response.headers_mut()); response.map(|body| { body.map_err(|error| -> BoxError { Box::new(error) }) @@ -229,6 +245,16 @@ async fn proxy_to_upstream( }) } +fn lease_deadline( + deadline: Option<&super::activity::RequestDeadlineHandle>, + timeout: Duration, +) -> Result, ()> { + match deadline { + Some(deadline) => deadline.lease_for(timeout).map(Some).ok_or(()), + None => Ok(None), + } +} + fn sanitize_transport_request(request: &mut Request) { for name in [ header::AUTHORIZATION, diff --git a/src/web/http/down.rs b/src/web/http/down.rs index a95a36d..8156efe 100644 --- a/src/web/http/down.rs +++ b/src/web/http/down.rs @@ -76,10 +76,26 @@ pub(super) async fn handle_down( } else { None }; + let poll_timeout = match lane_id { + Some(_) => Duration::from_secs(session.timeouts().lane_open_wait_secs) + .checked_add(Duration::from_secs(session.timeouts().long_poll_secs)), + None => Some(Duration::from_secs(session.timeouts().long_poll_secs)), + }; + let Some(poll_timeout) = poll_timeout else { + return service_unavailable(); + }; + let _deadline_lease = match super::request_deadline(&request) { + Some(deadline) => match deadline.lease_for(poll_timeout) { + Some(lease) => Some(lease), + None => return service_unavailable(), + }, + None => None, + }; let result = match lane_id { Some(lane_id) => session.poll_down_lane(lane_id, cursor).await, None => session.poll_down(cursor).await, }; + drop(_deadline_lease); match result { Ok(result) if result.body.is_empty() => { let mut response = carrier_empty(StatusCode::NO_CONTENT); diff --git a/src/web/http/session_policy_tests.rs b/src/web/http/session_policy_tests.rs index 795ca8a..24176ff 100644 --- a/src/web/http/session_policy_tests.rs +++ b/src/web/http/session_policy_tests.rs @@ -1,5 +1,50 @@ use super::*; +async fn open_keepalive( + listener: &TcpListener, + runtime: &Arc, +) -> (TcpStream, CancellationToken, tokio::task::JoinHandle<()>) { + let addr = listener.local_addr().unwrap(); + let (accepted, client) = tokio::join!(listener.accept(), TcpStream::connect(addr)); + let (server, peer) = accepted.unwrap(); + let cancellation = CancellationToken::new(); + let permit = runtime.try_http_connection().unwrap(); + let task = tokio::spawn(serve_connection( + server, + peer, + WebClientIpSource::XForwardedFor, + Arc::from(["127.0.0.1/32".parse().unwrap()]), + Arc::clone(runtime), + cancellation.clone(), + permit, + )); + (client.unwrap(), cancellation, task) +} + +async fn read_http_response(client: &mut TcpStream) -> Vec { + let mut response = Vec::new(); + while !response.ends_with(b"\r\n\r\n") { + assert!(response.len() < 16 * 1024); + response.push(client.read_u8().await.unwrap()); + } + let content_length = std::str::from_utf8(&response) + .unwrap() + .lines() + .filter_map(|line| line.split_once(':')) + .find_map(|(name, value)| { + name.eq_ignore_ascii_case("content-length") + .then(|| value.trim().parse::().unwrap()) + }) + .unwrap_or(0); + let body_start = response.len(); + response.resize(body_start + content_length, 0); + client + .read_exact(&mut response[body_start..]) + .await + .unwrap(); + response +} + async fn request_with_body_delay( listener: &TcpListener, runtime: &Arc, @@ -35,6 +80,9 @@ async fn live_session_body_and_closed_token_timeouts_survive_reload() { let capability = [21u8; 32]; let mut initial_config = runtime_config(capability, WebCarrier::Https); initial_config.web.timeouts.body_secs = 3; + initial_config.web.timeouts.header_secs = 1; + initial_config.web.timeouts.http_idle_secs = 4; + initial_config.web.timeouts.long_poll_secs = 3; initial_config.web.timeouts.bootstrap_lifetime_secs = 5; let generation = test_runtime_generation(1, initial_config); let active_runtime = Arc::new(ArcSwap::from(Arc::clone(&generation))); @@ -68,6 +116,9 @@ async fn live_session_body_and_closed_token_timeouts_survive_reload() { let mut replacement_config = runtime_config(capability, WebCarrier::Https); replacement_config.web.timeouts.body_secs = 1; + replacement_config.web.timeouts.header_secs = 1; + replacement_config.web.timeouts.http_idle_secs = 2; + replacement_config.web.timeouts.long_poll_secs = 1; replacement_config.web.timeouts.bootstrap_lifetime_secs = 1; let replacement = test_runtime_generation(2, replacement_config); active_runtime.store(Arc::clone(&replacement)); @@ -84,6 +135,13 @@ async fn live_session_body_and_closed_token_timeouts_survive_reload() { assert!(retry_headers.starts_with(b"HTTP/1.1 200")); assert_eq!(response_header(retry_headers, "x-session-token"), session); + let down = format!( + "POST /api/v1/down HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {session}\r\nX-Down-Cursor: 0\r\nContent-Length: 0\r\nConnection: close\r\n\r\n" + ) + .into_bytes(); + let down_response = request(&listener, &runtime, down).await; + assert!(down_response.starts_with(b"HTTP/1.1 204")); + let close = format!( "DELETE /api/v1/session HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {session}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n" ) @@ -100,3 +158,83 @@ async fn live_session_body_and_closed_token_timeouts_survive_reload() { replacement.stop_sessions().await; replacement.stop_background_tasks().await; } + +#[tokio::test] +async fn active_body_deadline_survives_reload_on_old_keepalive_connection() { + let capability = [22u8; 32]; + let mut initial_config = runtime_config(capability, WebCarrier::Https); + initial_config.web.timeouts.header_secs = 1; + initial_config.web.timeouts.body_secs = 1; + initial_config.web.timeouts.long_poll_secs = 1; + initial_config.web.timeouts.http_idle_secs = 2; + let generation = test_runtime_generation(1, initial_config); + let active_runtime = Arc::new(ArcSwap::from(Arc::clone(&generation))); + let runtime = WebProcessRuntime::start(Arc::clone(&active_runtime)); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (mut client, cancellation, task) = open_keepalive(&listener, &runtime).await; + client + .write_all( + b"GET / HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\n\r\n", + ) + .await + .unwrap(); + assert!( + read_http_response(&mut client) + .await + .starts_with(b"HTTP/1.1 200") + ); + + let mut replacement_config = runtime_config(capability, WebCarrier::Https); + replacement_config.web.timeouts.header_secs = 1; + replacement_config.web.timeouts.body_secs = 3; + replacement_config.web.timeouts.long_poll_secs = 1; + replacement_config.web.timeouts.http_idle_secs = 4; + let replacement = test_runtime_generation(2, replacement_config); + active_runtime.store(Arc::clone(&replacement)); + let encoded = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(capability); + client + .write_all( + format!( + "GET /?bridge={encoded} HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\n\r\n" + ) + .as_bytes(), + ) + .await + .unwrap(); + let bridge = read_http_response(&mut client).await; + let (_, bridge_body) = split_response(&bridge); + let bootstrap = std::str::from_utf8(bridge_body) + .unwrap() + .split_once("bootstrap=\"") + .and_then(|(_, suffix)| suffix.split_once('"')) + .map(|(token, _)| token.to_string()) + .unwrap(); + let hello = frame::encode(FrameType::Hello, 0, &[1]); + client + .write_all( + format!( + "POST /api/v1/session HTTP/1.1\r\nHost: proxy.example.com\r\nX-Forwarded-For: 192.0.2.10\r\nAuthorization: Bearer {bootstrap}\r\nContent-Type: application/octet-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n", + hello.len() + ) + .as_bytes(), + ) + .await + .unwrap(); + tokio::time::sleep(std::time::Duration::from_millis(2200)).await; + client.write_all(&hello).await.unwrap(); + + assert!( + read_http_response(&mut client) + .await + .starts_with(b"HTTP/1.1 200") + ); + + cancellation.cancel(); + drop(client); + task.await.unwrap(); + runtime.shutdown().await; + generation.stop_sessions().await; + generation.stop_background_tasks().await; + replacement.stop_sessions().await; + replacement.stop_background_tasks().await; +} diff --git a/src/web/http/websocket.rs b/src/web/http/websocket.rs index 762b229..6709fcb 100644 --- a/src/web/http/websocket.rs +++ b/src/web/http/websocket.rs @@ -254,6 +254,7 @@ pub(super) async fn handle( runtime: Arc, vhost: Arc, ) -> HttpResponse { + let request_deadline = super::request_deadline(&request); let Some(parsed) = parse_upgrade(&request) else { return serve_decoy(request, vhost, true, &runtime).await; }; @@ -284,7 +285,16 @@ pub(super) async fn handle( }, }; let timeouts = session.timeouts().clone(); - let connection = match runtime + let admission_lease = match request_deadline.as_ref() { + Some(deadline) => { + match deadline.lease_for(Duration::from_secs(timeouts.websocket_eviction_secs)) { + Some(lease) => Some(lease), + None => return serve_decoy(request, vhost, true, &runtime).await, + } + } + None => None, + }; + let admitted = runtime .admit_websocket( session.profile_key(), session.trace_session_id(), @@ -296,11 +306,17 @@ pub(super) async fn handle( Duration::from_secs(timeouts.websocket_eviction_secs), session.carrier_cancellation(), ) - .await - { + .await; + drop(admission_lease); + let connection = match admitted { Ok(connection) => connection, Err(_) => return serve_decoy(request, vhost, true, &runtime).await, }; + if let Some(reservation) = lane_reservation.as_mut() + && reservation.bind(connection.id()).is_err() + { + return serve_decoy(request, vhost, true, &runtime).await; + } if let Some(reservation) = probe_reservation.as_mut() && reservation.bind(connection.id()).is_err() { @@ -323,6 +339,20 @@ pub(super) async fn handle( trace.bind_identity(session.trace_identity()); trace.register_redaction(parsed.protocol.as_bytes()); } + let upgrade_deadline = match request_deadline.as_ref() { + Some(deadline) => { + let Some(until) = tokio::time::Instant::now() + .checked_add(Duration::from_secs(timeouts.websocket_upgrade_secs)) + else { + return serve_decoy(request, vhost, true, &runtime).await; + }; + match deadline.upgrade_until(until) { + Some(lease) => Some(lease), + None => return serve_decoy(request, vhost, true, &runtime).await, + } + } + None => None, + }; let on_upgrade = hyper::upgrade::on(&mut request); let protocol = parsed.protocol; let accept = parsed.accept; @@ -331,6 +361,7 @@ pub(super) async fn handle( runtime.spawn_auxiliary(async move { driver::run_upgraded( on_upgrade, + upgrade_deadline, driver_runtime, driver_session, connection, diff --git a/src/web/http/websocket/driver.rs b/src/web/http/websocket/driver.rs index 7ff4932..72c3d14 100644 --- a/src/web/http/websocket/driver.rs +++ b/src/web/http/websocket/driver.rs @@ -8,6 +8,7 @@ use tokio_tungstenite::tungstenite::protocol::{Message, Role, WebSocketConfig}; use tokio_util::sync::CancellationToken; use super::ConnectionIo; +use crate::web::http::activity::UpgradeDeadlineLease; use crate::web::manager::{WebProcessRuntime, WebSocketBudgetLease, WebSocketConnection}; use crate::web::session::{WebSession, WebSocketLaneReservation, WebSocketProbeReservation}; use crate::web::trace::{TraceDirection, TraceWebSocketContext}; @@ -22,8 +23,11 @@ mod lane; use io::{flush, process_multiplex, read_message, record_message, reserve_data, send}; use lane::run_lane; +// Upgrade ownership remains explicit across cancellation and reservation boundaries. +#[allow(clippy::too_many_arguments)] pub(super) async fn run_upgraded( on_upgrade: hyper::upgrade::OnUpgrade, + upgrade_deadline: Option, runtime: Arc, session: Arc, connection: WebSocketConnection, @@ -34,13 +38,15 @@ pub(super) async fn run_upgraded( ) { let cancellation = connection.cancellation(); let timeouts = session.timeouts().clone(); + let deadline = upgrade_deadline.as_ref().map_or_else( + || tokio::time::Instant::now() + Duration::from_secs(timeouts.websocket_upgrade_secs), + UpgradeDeadlineLease::deadline, + ); let upgraded = tokio::select! { _ = cancellation.cancelled() => return, - result = tokio::time::timeout( - Duration::from_secs(timeouts.websocket_upgrade_secs), - on_upgrade, - ) => result, + result = tokio::time::timeout_at(deadline, on_upgrade) => result, }; + drop(upgrade_deadline); let Ok(Ok(upgraded)) = upgraded else { return; }; @@ -95,8 +101,7 @@ pub(super) async fn run_upgraded( _ = tokio::time::timeout(eviction, socket.close(None)) => {} } if let Some(reservation) = lane_reservation { - session.close_websocket_lane(reservation.lane_id()); - drop(reservation); + session.close_websocket_lane(reservation); } else if !acknowledge_commit || session.is_carrier_committed() { session.close(); } @@ -192,10 +197,12 @@ async fn run_multiplex( session.close(); return Err(()); } - } else if acknowledge_commit && sequence > 1 && progressed { - if !session.websocket_peer_after_commit_ack(connection.id()) { - return Err(()); - } + } else if acknowledge_commit + && sequence > 1 + && progressed + && !session.websocket_peer_after_commit_ack(connection.id()) + { + return Err(()); } if !active && progressed { if !connection.mark_active() { diff --git a/src/web/http/websocket/driver/lane.rs b/src/web/http/websocket/driver/lane.rs index 18fc899..d6542a2 100644 --- a/src/web/http/websocket/driver/lane.rs +++ b/src/web/http/websocket/driver/lane.rs @@ -12,6 +12,7 @@ use crate::web::session::{WebSession, WebSocketLaneReservation}; use crate::web::trace::{TraceDirection, TraceWebSocketContext}; #[allow(clippy::too_many_arguments)] +/// Drives one exact WebSocket lane until its isolated failure boundary closes. pub(super) async fn run_lane( socket: &mut CarrierSocket, runtime: &Arc, @@ -35,7 +36,7 @@ pub(super) async fn run_lane( let maximum_message = session.limits().carrier_batch_bytes; let mut active = false; loop { - let down = session.poll_down_lane(reservation.lane_id(), cursor); + let down = session.poll_down_websocket_lane(reservation.lane_identity(), cursor); tokio::pin!(down); let event = tokio::select! { _ = cancellation.cancelled() => return Err(()), @@ -106,10 +107,12 @@ pub(super) async fn run_lane( session.close(); return Err(()); } - } else if acknowledge_commit && sequence > 1 && progressed { - if !session.websocket_peer_after_commit_ack(connection.id()) { - return Err(()); - } + } else if acknowledge_commit + && sequence > 1 + && progressed + && !session.websocket_peer_after_commit_ack(connection.id()) + { + return Err(()); } if !active && progressed { if !connection.mark_active() { diff --git a/src/web/manager.rs b/src/web/manager.rs index 3730b6f..a8ec896 100644 --- a/src/web/manager.rs +++ b/src/web/manager.rs @@ -33,6 +33,7 @@ mod carrier_outcome; mod admission; // Shutdown and expiry work remain outside request-path coordination. mod lifecycle; +pub(crate) use lifecycle::WebShutdownOutcome; // Queue and WebSocket allocations share one process-owned data-plane budget. mod budget; // WebSocket admission, replacement, and liveness are process-scoped. @@ -294,13 +295,23 @@ impl WebProcessRuntime { where F: Future + Send + 'static, { + if self.shutdown.is_cancelled() { + drop(future); + return; + } let shutdown = self.shutdown.clone(); - self.tasks.spawn(async move { + let tracked = self.tasks.track_future(async move { tokio::select! { + biased; _ = shutdown.cancelled() => {} _ = future => {} } }); + if self.shutdown.is_cancelled() { + drop(tracked); + return; + } + drop(tokio::spawn(tracked)); } /// Reserves one body reader and its declared bounded body allocation. @@ -393,7 +404,7 @@ impl WebProcessRuntime { .try_reserve_websocket(owner, bytes, WebSocketBudgetClass::Data) } - /// Admits one WebSocket with owner-first bounded replacement. + /// Admits one WebSocket with dead-first, then owner-local bounded replacement. #[allow(clippy::too_many_arguments)] pub(crate) async fn admit_websocket( self: &Arc, diff --git a/src/web/manager/budget.rs b/src/web/manager/budget.rs index 9012ae8..3c979c6 100644 --- a/src/web/manager/budget.rs +++ b/src/web/manager/budget.rs @@ -58,6 +58,21 @@ pub(crate) struct WebDataBudgetSnapshot { pub(crate) high_water_bytes: usize, } +/// Bounded owner-usage view captured before WebSocket registry selection. +pub(super) struct WebSocketFairnessSnapshot { + /// Equal byte share at the admission watermark for captured owners. + pub(super) fair_share: usize, + /// Captured shared-budget use indexed by profile owner. + pub(super) owner_bytes: HashMap, +} + +impl WebSocketFairnessSnapshot { + /// Returns the captured byte usage for one quota owner. + pub(super) fn owner_usage(&self, owner: ProfileKey) -> usize { + self.owner_bytes.get(&owner).copied().unwrap_or(0) + } +} + impl WebDataBudget { pub(super) fn new(limits: WebLimitsConfig) -> Arc { Arc::new(Self { @@ -220,16 +235,10 @@ impl WebDataBudget { self.pressured.store(true, Ordering::Release); } - pub(super) fn owner_usage(&self, owner: ProfileKey) -> usize { - self.state - .lock() - .owner_bytes - .get(&owner) - .copied() - .unwrap_or(0) - } - - pub(super) fn fair_share(&self, additional_owner: Option) -> usize { + pub(super) fn fairness_snapshot( + &self, + additional_owner: Option, + ) -> WebSocketFairnessSnapshot { let state = self.state.lock(); let mut owners = state.owner_bytes.len(); if additional_owner.is_some_and(|owner| !state.owner_bytes.contains_key(&owner)) { @@ -239,7 +248,10 @@ impl WebDataBudget { self.limits.websocket_bytes_global, self.limits.websocket_admission_watermark_pct, ); - admission / owners.max(1) + WebSocketFairnessSnapshot { + fair_share: admission / owners.max(1), + owner_bytes: state.owner_bytes.clone(), + } } pub(super) fn snapshot(&self) -> WebDataBudgetSnapshot { diff --git a/src/web/manager/lifecycle.rs b/src/web/manager/lifecycle.rs index 3832d78..c10e7b1 100644 --- a/src/web/manager/lifecycle.rs +++ b/src/web/manager/lifecycle.rs @@ -2,13 +2,30 @@ use std::net::IpAddr; use std::sync::atomic::Ordering; use std::time::{Duration, Instant}; -use tracing::info; +use tokio::time::Instant as TokioInstant; +use tracing::{info, warn}; use super::state::{ decrement_map, remember_closed_token_locked, remove_bootstrap_locked, remove_expired_locked, }; use super::{ProfileKey, TokenHash, WebProcessRuntime}; +/// Result of draining all process-owned WEB work under one absolute deadline. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum WebShutdownOutcome { + /// Every registered session and auxiliary task completed. + Drained, + /// Cancellation was asserted, but registered work remained at the deadline. + DeadlineExceeded, +} + +/// Owned shutdown snapshot retained after sessions leave the live registry. +pub(crate) struct WebShutdownDrain { + runtime: std::sync::Arc, + sessions: Vec>, + started: TokioInstant, +} + impl WebProcessRuntime { /// Removes one closed session and retains a bounded host-bound replay marker. pub(crate) fn session_finished( @@ -49,11 +66,20 @@ impl WebProcessRuntime { self.sessions_closed.fetch_add(1, Ordering::Relaxed); } - /// Stops issuance, closes all sessions, and joins bounded child work. - pub(crate) async fn shutdown(&self) { + /// Closes every WEB authority gate before any graceful wait begins. + pub(crate) fn begin_shutdown(self: &std::sync::Arc) -> WebShutdownDrain { + let started = TokioInstant::now(); self.shutdown.cancel(); self.close_websockets(); self.data_budget.close(); + self.http_connections.close(); + self.http_handlers.close(); + self.lane_polls.close(); + self.lane_aux_polls.close(); + self.body_readers.close(); + self.body_bytes.close(); + self.stream_handshakes.close(); + self.websocket_connections.close(); let sessions = { let mut state = self.state.lock(); state.closed = true; @@ -65,6 +91,24 @@ impl WebProcessRuntime { for session in &sessions { session.close(); } + self.tasks.close(); + WebShutdownDrain { + runtime: std::sync::Arc::clone(self), + sessions, + started, + } + } + + /// Stops issuance and drains all WEB work until one absolute deadline. + pub(crate) async fn shutdown_until( + self: &std::sync::Arc, + deadline: TokioInstant, + ) -> WebShutdownOutcome { + self.begin_shutdown().wait_until(deadline).await + } + + /// Stops issuance and drains WEB work under the currently configured budget. + pub(crate) async fn shutdown(self: &std::sync::Arc) -> WebShutdownOutcome { let timeout_secs = self .active_runtime .load() @@ -72,34 +116,11 @@ impl WebProcessRuntime { .web .timeouts .shutdown_secs; - let waits = async { - for session in sessions { - session.wait().await; - } - }; - let _ = tokio::time::timeout(Duration::from_secs(timeout_secs), waits).await; - self.tasks.close(); - let _ = tokio::time::timeout(Duration::from_secs(timeout_secs), self.tasks.wait()).await; - let sessions_live = self.state.lock().sessions.len(); - let streams_live = self.stream_admission.lock().streams_live; - let budget = self.data_budget.snapshot(); - info!( - target: "telemt::web", - sessions_created = self.sessions_created.load(Ordering::Relaxed), - sessions_closed = self.sessions_closed.load(Ordering::Relaxed), - sessions_live, - streams_opened = self.streams_opened.load(Ordering::Relaxed), - streams_rejected = self.streams_rejected.load(Ordering::Relaxed), - streams_live, - pending_bytes = budget.queue_bytes, - pending_items = budget.queue_items, - websocket_bytes = budget.websocket_bytes, - data_high_water_bytes = budget.high_water_bytes, - bytes_up = self.bytes_up.load(Ordering::Relaxed), - bytes_down = self.bytes_down.load(Ordering::Relaxed), - limit_hits = self.limit_hits.load(Ordering::Relaxed), - "WEB runtime stopped" - ); + let now = TokioInstant::now(); + let deadline = now + .checked_add(Duration::from_secs(timeout_secs)) + .unwrap_or(now); + self.shutdown_until(deadline).await } /// Expires credentials and closes idle sessions without holding locks across callbacks. @@ -156,3 +177,228 @@ impl WebProcessRuntime { } } } + +impl WebShutdownDrain { + /// Waits for frozen session ownership and process auxiliary tasks concurrently. + pub(crate) async fn wait_until(self, deadline: TokioInstant) -> WebShutdownOutcome { + let sessions = &self.sessions; + let session_waits = async { + for session in sessions { + session.wait().await; + } + }; + let outcome = wait_for_drain(deadline, session_waits, self.runtime.tasks.wait()).await; + self.log_outcome(outcome, deadline); + outcome + } + + fn log_outcome(&self, outcome: WebShutdownOutcome, deadline: TokioInstant) { + let sessions_live = self.runtime.state.lock().sessions.len(); + let streams_live = self.runtime.stream_admission.lock().streams_live; + let session_tasks_live = self.sessions.iter().fold(0usize, |total, session| { + total.saturating_add(session.tasks_live()) + }); + let sessions_pending = self + .sessions + .iter() + .filter(|session| session.tasks_live() != 0) + .count(); + let auxiliary_tasks_live = self.runtime.tasks.len(); + let budget = self.runtime.data_budget.snapshot(); + let budget_ms = deadline + .saturating_duration_since(self.started) + .as_millis() + .min(u128::from(u64::MAX)) as u64; + let elapsed_ms = TokioInstant::now() + .saturating_duration_since(self.started) + .as_millis() + .min(u128::from(u64::MAX)) as u64; + match outcome { + WebShutdownOutcome::Drained => info!( + target: "telemt::web", + shutdown_drained = true, + shutdown_budget_ms = budget_ms, + shutdown_elapsed_ms = elapsed_ms, + sessions_created = self.runtime.sessions_created.load(Ordering::Relaxed), + sessions_closed = self.runtime.sessions_closed.load(Ordering::Relaxed), + sessions_live, + sessions_pending, + session_tasks_live, + auxiliary_tasks_live, + streams_opened = self.runtime.streams_opened.load(Ordering::Relaxed), + streams_rejected = self.runtime.streams_rejected.load(Ordering::Relaxed), + streams_live, + pending_bytes = budget.queue_bytes, + pending_items = budget.queue_items, + websocket_bytes = budget.websocket_bytes, + data_high_water_bytes = budget.high_water_bytes, + bytes_up = self.runtime.bytes_up.load(Ordering::Relaxed), + bytes_down = self.runtime.bytes_down.load(Ordering::Relaxed), + limit_hits = self.runtime.limit_hits.load(Ordering::Relaxed), + "WEB runtime stopped" + ), + WebShutdownOutcome::DeadlineExceeded => warn!( + target: "telemt::web", + shutdown_drained = false, + shutdown_budget_ms = budget_ms, + shutdown_elapsed_ms = elapsed_ms, + sessions_created = self.runtime.sessions_created.load(Ordering::Relaxed), + sessions_closed = self.runtime.sessions_closed.load(Ordering::Relaxed), + sessions_live, + sessions_pending, + session_tasks_live, + auxiliary_tasks_live, + streams_opened = self.runtime.streams_opened.load(Ordering::Relaxed), + streams_rejected = self.runtime.streams_rejected.load(Ordering::Relaxed), + streams_live, + pending_bytes = budget.queue_bytes, + pending_items = budget.queue_items, + websocket_bytes = budget.websocket_bytes, + data_high_water_bytes = budget.high_water_bytes, + bytes_up = self.runtime.bytes_up.load(Ordering::Relaxed), + bytes_down = self.runtime.bytes_down.load(Ordering::Relaxed), + limit_hits = self.runtime.limit_hits.load(Ordering::Relaxed), + "WEB runtime shutdown deadline exceeded" + ), + } + } +} + +async fn wait_for_drain(deadline: TokioInstant, sessions: S, tasks: T) -> WebShutdownOutcome +where + S: std::future::Future, + T: std::future::Future, +{ + let waits = async { + tokio::join!(sessions, tasks); + }; + if tokio::time::timeout_at(deadline, waits).await.is_ok() { + WebShutdownOutcome::Drained + } else { + WebShutdownOutcome::DeadlineExceeded + } +} + +#[cfg(test)] +mod tests { + use std::future::Future; + use std::pin::Pin; + use std::sync::Arc; + use std::sync::atomic::{AtomicUsize, Ordering}; + use std::task::{Context, Poll}; + + use arc_swap::ArcSwap; + use tokio::sync::Notify; + + use super::*; + use crate::config::ProxyConfig; + use crate::maestro::generation::test_runtime_generation; + + struct DropProbe { + polls: Arc, + drops: Arc, + } + + impl Future for DropProbe { + type Output = (); + + fn poll(self: Pin<&mut Self>, _context: &mut Context<'_>) -> Poll { + self.polls.fetch_add(1, Ordering::AcqRel); + Poll::Pending + } + } + + impl Drop for DropProbe { + fn drop(&mut self) { + self.drops.fetch_add(1, Ordering::AcqRel); + } + } + + fn runtime() -> ( + Arc, + Arc, + ) { + let generation = test_runtime_generation(1, ProxyConfig::default()); + let runtime = WebProcessRuntime::start(Arc::new(ArcSwap::from(Arc::clone(&generation)))); + (runtime, generation) + } + + #[tokio::test(start_paused = true)] + async fn drain_uses_one_absolute_deadline_for_both_wait_groups() { + let started = TokioInstant::now(); + let outcome = wait_for_drain( + started + Duration::from_secs(5), + tokio::time::sleep(Duration::from_secs(4)), + tokio::time::sleep(Duration::from_secs(9)), + ) + .await; + + assert_eq!(outcome, WebShutdownOutcome::DeadlineExceeded); + assert_eq!(TokioInstant::now() - started, Duration::from_secs(5)); + } + + #[tokio::test(start_paused = true)] + async fn drain_returns_when_both_wait_groups_finish() { + let started = TokioInstant::now(); + let outcome = wait_for_drain( + started + Duration::from_secs(5), + tokio::time::sleep(Duration::from_secs(3)), + tokio::time::sleep(Duration::from_secs(2)), + ) + .await; + + assert_eq!(outcome, WebShutdownOutcome::Drained); + assert_eq!(TokioInstant::now() - started, Duration::from_secs(3)); + } + + #[tokio::test] + async fn post_shutdown_auxiliary_is_dropped_without_polling() { + let (runtime, generation) = runtime(); + let polls = Arc::new(AtomicUsize::new(0)); + let drops = Arc::new(AtomicUsize::new(0)); + let drain = runtime.begin_shutdown(); + + runtime.spawn_auxiliary(DropProbe { + polls: Arc::clone(&polls), + drops: Arc::clone(&drops), + }); + + assert_eq!(polls.load(Ordering::Acquire), 0); + assert_eq!(drops.load(Ordering::Acquire), 1); + assert_eq!( + drain + .wait_until(TokioInstant::now() + Duration::from_secs(1)) + .await, + WebShutdownOutcome::Drained + ); + generation.stop_sessions().await; + generation.stop_background_tasks().await; + } + + #[tokio::test(start_paused = true)] + async fn expired_deadline_still_closes_every_runtime_gate() { + let (runtime, generation) = runtime(); + let existing = runtime.try_http_connection().unwrap(); + let release = Arc::new(Notify::new()); + let release_task = Arc::clone(&release); + runtime.tasks.spawn(async move { + release_task.notified().await; + }); + tokio::task::yield_now().await; + let drain = runtime.begin_shutdown(); + + assert!(runtime.try_http_connection().is_none()); + assert!(runtime.try_http_handler().is_none()); + assert!(runtime.try_lane_poll(false).is_none()); + assert_eq!( + drain.wait_until(TokioInstant::now()).await, + WebShutdownOutcome::DeadlineExceeded + ); + + drop(existing); + release.notify_waiters(); + runtime.tasks.wait().await; + generation.stop_sessions().await; + generation.stop_background_tasks().await; + } +} diff --git a/src/web/manager/websocket.rs b/src/web/manager/websocket.rs index cee8df2..ebec5d7 100644 --- a/src/web/manager/websocket.rs +++ b/src/web/manager/websocket.rs @@ -9,6 +9,9 @@ use tokio_util::sync::CancellationToken; use super::{ManagerError, ProfileKey, WebProcessRuntime, WebSocketBudgetLease}; +// Deterministic victim ordering remains isolated from registry mutation. +mod policy; + /// One process-owned WebSocket carrier class used for eviction priority. #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] pub(crate) enum WebSocketKind { @@ -25,6 +28,7 @@ struct WebSocketClaimKey { } #[repr(u8)] +#[derive(Clone, Copy, Debug, PartialEq, Eq)] enum WebSocketPhase { Claimed, Upgraded, @@ -239,6 +243,8 @@ enum TryAdmitError { Closed, } +// Admission inputs stay explicit so quota and cancellation ownership cannot drift. +#[allow(clippy::too_many_arguments)] fn try_admit( runtime: &Arc, owner: ProfileKey, @@ -353,8 +359,7 @@ fn select_victim( excluded_id: Option, claim: bool, ) -> Option> { - let fair_share = runtime.data_budget.fair_share(Some(owner)); - let requester_usage = runtime.data_budget.owner_usage(owner); + let fairness = runtime.data_budget.fairness_snapshot(Some(owner)); let now = runtime.websocket_tick(); let mut registry = runtime.websockets.lock(); if claim && registry.evictions_in_flight >= runtime.limits.max_websocket_evictions_in_flight { @@ -366,31 +371,8 @@ fn select_victim( .filter(|entry| Some(entry.id) != excluded_id) .filter(|entry| !entry.closing.load(Ordering::Acquire)) .filter_map(|entry| { - let owner_rank = if entry.session_id == session_id { - 0 - } else if entry.owner == owner { - 1 - } else if entry.client_ip == client_ip { - 2 - } else { - if requester_usage >= fair_share - || runtime.data_budget.owner_usage(entry.owner) <= fair_share - { - return None; - } - 3 - }; - let priority = entry_priority(entry, now); - Some(( - ( - owner_rank, - priority, - entry.last_progress_tick.load(Ordering::Acquire), - entry.created_tick, - entry.id, - ), - Arc::clone(entry), - )) + policy::admission_key(entry, now, owner, session_id, client_ip, &fairness) + .map(|key| (key, Arc::clone(entry))) }) .min_by_key(|(key, _)| *key) .map(|(_, entry)| entry)?; @@ -405,6 +387,7 @@ fn select_pressure_victim( now: u64, claim: bool, ) -> Option> { + let fairness = runtime.data_budget.fairness_snapshot(None); let mut registry = runtime.websockets.lock(); if claim && registry.evictions_in_flight >= runtime.limits.max_websocket_evictions_in_flight { return None; @@ -415,12 +398,7 @@ fn select_pressure_victim( .filter(|entry| !entry.closing.load(Ordering::Acquire)) .map(|entry| { ( - ( - entry_priority(entry, now), - entry.last_progress_tick.load(Ordering::Acquire), - entry.created_tick, - entry.id, - ), + policy::pressure_key(entry, now, &fairness), Arc::clone(entry), ) }) @@ -438,18 +416,23 @@ fn claim_stale_victims(runtime: &WebProcessRuntime, now: u64) -> Vec= dead_after(entry) - }) - .take(available) + .filter(|entry| policy::victim_class(entry, now) == policy::VictimClass::Dead) .cloned() .collect::>(); + candidates.sort_unstable_by_key(|entry| { + ( + entry.last_peer_tick.load(Ordering::Acquire), + entry.created_tick, + entry.id, + ) + }); candidates .into_iter() + .take(available) .filter(|entry| claim_entry(&mut registry, entry, runtime)) .collect() } @@ -474,18 +457,6 @@ fn claim_entry( true } -fn entry_priority(entry: &WebSocketEntry, now: u64) -> u8 { - if entry.phase.load(Ordering::Acquire) < WebSocketPhase::Active as u8 - || now.saturating_sub(entry.last_peer_tick.load(Ordering::Acquire)) >= dead_after(entry) - { - 0 - } else if matches!(entry.kind, WebSocketKind::Lane(_)) { - 1 - } else { - 2 - } -} - fn dead_after(entry: &WebSocketEntry) -> u64 { entry.liveness_interval_ms.saturating_mul(2) } diff --git a/src/web/manager/websocket/policy.rs b/src/web/manager/websocket/policy.rs new file mode 100644 index 0000000..6f360cc --- /dev/null +++ b/src/web/manager/websocket/policy.rs @@ -0,0 +1,113 @@ +use std::net::IpAddr; +use std::sync::atomic::Ordering; + +use super::{WebSocketEntry, WebSocketKind, WebSocketPhase, dead_after}; +use crate::web::manager::ProfileKey; +use crate::web::manager::budget::WebSocketFairnessSnapshot; + +/// Stable total-order key used by bounded victim selection. +pub(super) type VictimKey = (u8, u8, u8, u64, u64, u64); + +/// Lifecycle class used before locality and least-recent-progress ordering. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(super) enum VictimClass { + /// Active connection whose peer-liveness deadline elapsed. + Dead, + /// Claimed or upgraded connection still bounded by its startup deadlines. + PreActive, + /// Active per-stream lane connection. + LiveLane, + /// Active multiplexed session connection. + LiveMultiplex, +} + +impl VictimClass { + fn rank(self) -> u8 { + match self { + Self::Dead | Self::PreActive => 0, + Self::LiveLane => 1, + Self::LiveMultiplex => 2, + } + } +} + +/// Classifies one non-closing connection without conflating startup with death. +pub(super) fn victim_class(entry: &WebSocketEntry, now: u64) -> VictimClass { + let phase = entry.phase.load(Ordering::Acquire); + if phase == WebSocketPhase::Active as u8 + && now.saturating_sub(entry.last_peer_tick.load(Ordering::Acquire)) >= dead_after(entry) + { + VictimClass::Dead + } else if phase < WebSocketPhase::Active as u8 { + VictimClass::PreActive + } else if matches!(entry.kind, WebSocketKind::Lane(_)) { + VictimClass::LiveLane + } else { + VictimClass::LiveMultiplex + } +} + +/// Returns one eligible admission-victim key with global dead-first ordering. +pub(super) fn admission_key( + entry: &WebSocketEntry, + now: u64, + requester_owner: ProfileKey, + requester_session: u64, + requester_ip: IpAddr, + fairness: &WebSocketFairnessSnapshot, +) -> Option { + let class = victim_class(entry, now); + if class == VictimClass::Dead { + return Some(( + 0, + 0, + 0, + entry.last_peer_tick.load(Ordering::Acquire), + entry.created_tick, + entry.id, + )); + } + let locality = if entry.session_id == requester_session { + 0 + } else if entry.owner == requester_owner { + 1 + } else if entry.client_ip == requester_ip { + 2 + } else if fairness.owner_usage(requester_owner) < fairness.fair_share + && fairness.owner_usage(entry.owner) > fairness.fair_share + { + 3 + } else { + return None; + }; + Some(( + 1, + locality, + class.rank(), + entry.last_progress_tick.load(Ordering::Acquire), + entry.created_tick, + entry.id, + )) +} + +/// Returns one pressure-victim key preferring dead and over-share owners. +pub(super) fn pressure_key( + entry: &WebSocketEntry, + now: u64, + fairness: &WebSocketFairnessSnapshot, +) -> VictimKey { + let class = victim_class(entry, now); + let dead = class == VictimClass::Dead; + ( + u8::from(!dead), + if dead { + 0 + } else { + u8::from(fairness.owner_usage(entry.owner) <= fairness.fair_share) + }, + class.rank(), + entry.last_progress_tick.load(Ordering::Acquire), + entry.created_tick, + entry.id, + ) +} diff --git a/src/web/manager/websocket/tests.rs b/src/web/manager/websocket/tests.rs index 35877d6..7591a60 100644 --- a/src/web/manager/websocket/tests.rs +++ b/src/web/manager/websocket/tests.rs @@ -1,25 +1,39 @@ use super::*; +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; -fn entry(kind: WebSocketKind, opened: bool, peer_tick: u64) -> WebSocketEntry { - let phase = if opened { - WebSocketPhase::Active - } else { - WebSocketPhase::Claimed - }; +use super::policy::{VictimClass, admission_key, pressure_key, victim_class}; +use crate::config::ProxyConfig; +use crate::maestro::generation::test_runtime_generation; +use crate::web::manager::budget::WebSocketFairnessSnapshot; +use arc_swap::ArcSwap; + +#[allow(clippy::too_many_arguments)] +fn entry( + id: u64, + owner: ProfileKey, + session_id: u64, + client_ip: &str, + kind: WebSocketKind, + phase: WebSocketPhase, + peer_tick: u64, + progress_tick: u64, +) -> WebSocketEntry { WebSocketEntry { - id: 1, - owner: [0; 32], - session_id: 1, + id, + owner, + session_id, claim: WebSocketClaimKey { session_hash: [0; 32], kind, }, - client_ip: "192.0.2.10".parse().unwrap(), + client_ip: client_ip.parse().unwrap(), kind, liveness_interval_ms: 10, created_tick: 1, last_peer_tick: AtomicU64::new(peer_tick), - last_progress_tick: AtomicU64::new(peer_tick), + last_progress_tick: AtomicU64::new(progress_tick), phase: AtomicU8::new(phase as u8), closing: AtomicBool::new(false), cancel: CancellationToken::new(), @@ -27,25 +41,259 @@ fn entry(kind: WebSocketKind, opened: bool, peer_tick: u64) -> WebSocketEntry { } } -#[test] -fn preopen_and_dead_entries_precede_live_lane_and_multiplex_victims() { - let preopen = entry(WebSocketKind::Multiplex, false, 90); - let dead = entry(WebSocketKind::Multiplex, true, 1); - let lane = entry(WebSocketKind::Lane(7), true, 90); - let multiplex = entry(WebSocketKind::Multiplex, true, 90); - - assert_eq!(entry_priority(&preopen, 100), 0); - assert_eq!(entry_priority(&dead, 100), 0); - assert_eq!(entry_priority(&lane, 100), 1); - assert_eq!(entry_priority(&multiplex, 100), 2); +fn fairness(fair_share: usize, usages: &[(ProfileKey, usize)]) -> WebSocketFairnessSnapshot { + WebSocketFairnessSnapshot { + fair_share, + owner_bytes: usages.iter().copied().collect::>(), + } } #[test] -fn dead_classification_keeps_each_connections_creation_time_interval() { - let short_interval = entry(WebSocketKind::Multiplex, true, 80); - let mut long_interval = entry(WebSocketKind::Multiplex, true, 80); - long_interval.liveness_interval_ms = 100; +fn preactive_and_dead_are_distinct_lifecycle_classes() { + let preactive = entry( + 1, + [1; 32], + 1, + "192.0.2.10", + WebSocketKind::Multiplex, + WebSocketPhase::Claimed, + 1, + 1, + ); + let dead = entry( + 2, + [1; 32], + 1, + "192.0.2.10", + WebSocketKind::Multiplex, + WebSocketPhase::Active, + 1, + 1, + ); - assert_eq!(entry_priority(&short_interval, 100), 0); - assert_eq!(entry_priority(&long_interval, 100), 2); + assert_eq!(victim_class(&preactive, 100), VictimClass::PreActive); + assert_eq!(victim_class(&dead, 100), VictimClass::Dead); +} + +#[test] +fn dead_other_session_precedes_healthy_same_session() { + let requester_owner = [1; 32]; + let usage = fairness(100, &[(requester_owner, 100), ([2; 32], 100)]); + let dead = entry( + 2, + [2; 32], + 2, + "198.51.100.10", + WebSocketKind::Multiplex, + WebSocketPhase::Active, + 1, + 1, + ); + let healthy = entry( + 1, + requester_owner, + 1, + "192.0.2.10", + WebSocketKind::Lane(7), + WebSocketPhase::Active, + 99, + 99, + ); + + let dead_key = admission_key( + &dead, + 100, + requester_owner, + 1, + "192.0.2.10".parse().unwrap(), + &usage, + ) + .unwrap(); + let healthy_key = admission_key( + &healthy, + 100, + requester_owner, + 1, + "192.0.2.10".parse().unwrap(), + &usage, + ) + .unwrap(); + + assert!(dead_key < healthy_key); +} + +#[test] +fn unrelated_live_victim_requires_opposite_fair_share_positions() { + let requester_owner = [1; 32]; + let victim_owner = [2; 32]; + let candidate = entry( + 1, + victim_owner, + 2, + "198.51.100.10", + WebSocketKind::Lane(7), + WebSocketPhase::Active, + 99, + 99, + ); + let requester_ip = "192.0.2.10".parse().unwrap(); + + assert!( + admission_key( + &candidate, + 100, + requester_owner, + 1, + requester_ip, + &fairness(100, &[(requester_owner, 99), (victim_owner, 101)]), + ) + .is_some() + ); + assert!( + admission_key( + &candidate, + 100, + requester_owner, + 1, + requester_ip, + &fairness(100, &[(requester_owner, 100), (victim_owner, 101)]), + ) + .is_none() + ); +} + +#[test] +fn pressure_prefers_over_share_owner_then_lifecycle_and_id() { + let over_owner = [1; 32]; + let under_owner = [2; 32]; + let usage = fairness(100, &[(over_owner, 101), (under_owner, 99)]); + let over = entry( + 9, + over_owner, + 1, + "192.0.2.10", + WebSocketKind::Multiplex, + WebSocketPhase::Active, + 99, + 99, + ); + let under = entry( + 1, + under_owner, + 2, + "198.51.100.10", + WebSocketKind::Lane(7), + WebSocketPhase::Active, + 90, + 90, + ); + + assert!(pressure_key(&over, 100, &usage) < pressure_key(&under, 100, &usage)); + + let equal_usage = fairness(100, &[(over_owner, 100), (under_owner, 100)]); + let preactive = entry( + 2, + under_owner, + 2, + "198.51.100.10", + WebSocketKind::Multiplex, + WebSocketPhase::Upgraded, + 99, + 99, + ); + assert!(pressure_key(&preactive, 100, &equal_usage) < pressure_key(&under, 100, &equal_usage)); + + let lower_id = entry( + 1, + under_owner, + 2, + "198.51.100.10", + WebSocketKind::Lane(7), + WebSocketPhase::Active, + 90, + 90, + ); + let higher_id = entry( + 2, + under_owner, + 2, + "198.51.100.10", + WebSocketKind::Lane(8), + WebSocketPhase::Active, + 90, + 90, + ); + + assert!( + pressure_key(&lower_id, 100, &equal_usage) < pressure_key(&higher_id, 100, &equal_usage) + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn concurrent_victim_claims_stay_bounded_and_return_to_zero() { + let config = ProxyConfig::default(); + let generation = test_runtime_generation(1, config); + let runtime = WebProcessRuntime::start(Arc::new(ArcSwap::from(Arc::clone(&generation)))); + let limit = runtime.limits.max_websocket_evictions_in_flight; + let entries = (0..limit.saturating_mul(2)) + .map(|index| { + Arc::new(entry( + index as u64 + 1, + [index as u8; 32], + index as u64 + 1, + "192.0.2.10", + WebSocketKind::Lane(index as u32 + 1), + WebSocketPhase::Active, + 1, + 1, + )) + }) + .collect::>(); + let connections = entries + .iter() + .map(|entry| WebSocketConnection { + runtime: Arc::downgrade(&runtime), + entry: Arc::clone(entry), + slot: None, + base_budget: None, + }) + .collect::>(); + { + let mut registry = runtime.websockets.lock(); + for entry in &entries { + registry.claims.insert(entry.claim, entry.id); + registry.entries.insert(entry.id, Arc::clone(entry)); + } + } + let successes = Arc::new(AtomicUsize::new(0)); + let mut tasks = Vec::new(); + for task_id in 0..100usize { + let runtime = Arc::clone(&runtime); + let entries = entries.clone(); + let successes = Arc::clone(&successes); + tasks.push(tokio::spawn(async move { + for attempt in 0..100usize { + let entry = &entries[(task_id * 100 + attempt) % entries.len()]; + { + let mut registry = runtime.websockets.lock(); + if claim_entry(&mut registry, entry, &runtime) { + successes.fetch_add(1, Ordering::AcqRel); + } + } + tokio::task::yield_now().await; + } + })); + } + for task in tasks { + task.await.unwrap(); + } + + assert_eq!(successes.load(Ordering::Acquire), limit); + assert_eq!(runtime.websockets.lock().evictions_in_flight, limit); + + drop(connections); + assert_eq!(runtime.websockets.lock().evictions_in_flight, 0); + runtime.shutdown().await; + generation.stop_sessions().await; + generation.stop_background_tasks().await; } diff --git a/src/web/session.rs b/src/web/session.rs index 1db78b8..7790413 100644 --- a/src/web/session.rs +++ b/src/web/session.rs @@ -62,6 +62,22 @@ pub(crate) struct StreamIdentity { pub(crate) instance: u64, } +/// Exact server-local identity of one carrier-lane incarnation. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) struct CarrierLaneIdentity { + /// Numeric lane identifier carried on the wire. + pub(crate) lane_id: u32, + /// Monotonic server-local incarnation of that numeric lane. + pub(crate) instance: u64, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +struct WebSocketLaneClaim { + lane: CarrierLaneIdentity, + peer_port: u16, + connection_id: Option, +} + struct StreamState { instance: u64, inbound: VecDeque, @@ -144,7 +160,7 @@ struct SessionState { carrier_lanes: HashMap, lane_open_waits: usize, next_lane_instance: u64, - websocket_lane_reservations: HashMap, + websocket_lane_reservations: HashMap, pending_bytes: usize, pending_items: usize, pending_control_bytes: usize, @@ -506,11 +522,14 @@ fn remember_closed(state: &mut SessionState, stream_id: u32, limit: usize) -> Op evicted } -fn insert_carrier_lane(state: &mut SessionState, lane_id: u32) -> Option { +fn insert_carrier_lane(state: &mut SessionState, lane_id: u32) -> Option { + if state.carrier_lanes.contains_key(&lane_id) { + return None; + } let instance = state.next_lane_instance; state.next_lane_instance = instance.checked_add(1)?; state .carrier_lanes .insert(lane_id, CarrierLane::new(instance)); - Some(instance) + Some(CarrierLaneIdentity { lane_id, instance }) } diff --git a/src/web/session/downlink_tests.rs b/src/web/session/downlink_tests.rs index b1b229c..d85e9f7 100644 --- a/src/web/session/downlink_tests.rs +++ b/src/web/session/downlink_tests.rs @@ -29,8 +29,10 @@ fn session() -> (Arc, Arc) { max_streams: 1, max_streams_per_session: 1, }); - let mut timeouts = WebTimeoutsConfig::default(); - timeouts.long_poll_secs = 1; + let timeouts = WebTimeoutsConfig { + long_poll_secs: 1, + ..WebTimeoutsConfig::default() + }; let session = WebSession::new( Arc::downgrade(&manager), [1; 32], diff --git a/src/web/session/lanes.rs b/src/web/session/lanes.rs index 7755d66..cd0a184 100644 --- a/src/web/session/lanes.rs +++ b/src/web/session/lanes.rs @@ -6,8 +6,8 @@ use tokio::sync::OwnedSemaphorePermit; use super::lane_downlink::take_lane_down_batch; use super::{ - PendingClass, PollResult, QUEUE_ITEM_COST, QueuedFrame, SessionState, WebSession, - remember_closed, + CarrierLaneIdentity, PendingClass, PollResult, QUEUE_ITEM_COST, QueuedFrame, SessionState, + WebSession, remember_closed, }; use crate::web::frame::{self, FrameType}; use crate::web::manager::ManagerError; @@ -18,15 +18,46 @@ impl WebSession { &self, lane_id: u32, cursor: u64, + ) -> Result { + self.poll_down_lane_inner(lane_id, None, cursor).await + } + + /// Polls only the exact lane incarnation owned by one WebSocket driver. + pub(crate) async fn poll_down_websocket_lane( + &self, + lane: CarrierLaneIdentity, + cursor: u64, + ) -> Result { + self.poll_down_lane_inner(lane.lane_id, Some(lane.instance), cursor) + .await + } + + async fn poll_down_lane_inner( + &self, + lane_id: u32, + expected_instance: Option, + cursor: u64, ) -> Result { if !self.carrier().uses_lanes() || lane_id > frame::MAX_STREAM_ID { return Err(ManagerError::Protocol); } - if !self.wait_for_lane_open(lane_id, cursor).await? { + let lane_ready = if let Some(expected_instance) = expected_instance { + let state = self.state.lock(); + if state.closed { + return Err(ManagerError::Closed); + } + state + .carrier_lanes + .get(&lane_id) + .is_some_and(|lane| lane.instance == expected_instance) + } else { + self.wait_for_lane_open(lane_id, cursor).await? + }; + if !lane_ready { return Ok(PollResult { body: Bytes::new(), next_cursor: cursor, - lane_closed: false, + lane_closed: expected_instance.is_some(), }); } let (instance, epoch, notify, healthy) = { @@ -43,6 +74,13 @@ impl WebSession { lane_closed: true, }); }; + if expected_instance.is_some_and(|instance| lane.instance != instance) { + return Ok(PollResult { + body: Bytes::new(), + next_cursor: cursor, + lane_closed: true, + }); + } if let Some(unacked) = &lane.unacked { if cursor == unacked.base_cursor { return Ok(PollResult { diff --git a/src/web/session/lifecycle.rs b/src/web/session/lifecycle.rs index 29c6d8b..7ed8e99 100644 --- a/src/web/session/lifecycle.rs +++ b/src/web/session/lifecycle.rs @@ -84,6 +84,11 @@ impl WebSession { } } + /// Returns the current number of registered logical-stream tasks. + pub(crate) fn tasks_live(&self) -> usize { + self.tasks_live.load(Ordering::Acquire) + } + /// Atomically closes a session only when reconnect grace is still due. pub(crate) fn close_if_due(&self, now: Instant) -> bool { let healthy = { diff --git a/src/web/session/uplink.rs b/src/web/session/uplink.rs index ca4b824..0e8cdc9 100644 --- a/src/web/session/uplink.rs +++ b/src/web/session/uplink.rs @@ -166,6 +166,8 @@ impl WebSession { result } + // Batch application keeps every transactional accumulator explicit. + #[allow(clippy::too_many_arguments)] pub(super) fn apply_batch_locked( self: &Arc, state: &mut SessionState, diff --git a/src/web/session/websocket.rs b/src/web/session/websocket.rs index 3ef0edd..9f09adb 100644 --- a/src/web/session/websocket.rs +++ b/src/web/session/websocket.rs @@ -4,7 +4,10 @@ use std::time::Instant; use sha2::{Digest, Sha256}; use super::uplink::{AppliedProgress, inbound_reservation, validate_batch}; -use super::{PendingClass, WebSession, inbound_queue_cost, insert_carrier_lane}; +use super::{ + CarrierLaneIdentity, PendingClass, StreamIdentity, WebSession, WebSocketLaneClaim, + inbound_queue_cost, insert_carrier_lane, +}; use crate::config::WebCarrier; use crate::web::frame; use crate::web::manager::ManagerError; @@ -12,9 +15,19 @@ use crate::web::manager::ManagerError; /// Pre-OPEN stream quota and synthetic tuple ownership for one WebSocket lane. pub(crate) struct WebSocketLaneReservation { session: Arc, - lane_id: u32, - peer_port: u16, - transferred: bool, + 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. @@ -57,28 +70,99 @@ impl Drop for WebSocketProbeReservation { impl WebSocketLaneReservation { /// Returns the logical stream owned by this connection. pub(crate) fn lane_id(&self) -> u32 { - self.lane_id + self.claim.lane.lane_id } - fn transfer_to_stream(&mut self) { - let removed = self - .session - .state - .lock() - .websocket_lane_reservations - .remove(&self.lane_id); - if removed == Some(self.peer_port) { - self.transferred = true; + /// 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); } + let mut state = self.session.state.lock(); + if state.closed + || 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); + } + let mut state = self.session.state.lock(); + if 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) { - if !self.transferred { - self.session - .release_websocket_lane_reservation(self.lane_id, self.peer_port); - } + self.release(); } } @@ -162,9 +246,7 @@ impl WebSession { ); return Err(ManagerError::Limit); } - state.websocket_lane_reservations.insert(lane_id, peer_port); - if insert_carrier_lane(&mut state, lane_id).is_none() { - state.websocket_lane_reservations.remove(&lane_id); + let Some(lane) = insert_carrier_lane(&mut state, lane_id) else { state.active_peer_ports.remove(&peer_port); manager.release_stream( self.profile_key, @@ -173,14 +255,37 @@ impl WebSession { peer_port, ); return Err(ManagerError::Protocol); + }; + let claim = WebSocketLaneClaim { + lane, + peer_port, + connection_id: None, + }; + let inserted = match state.websocket_lane_reservations.entry(lane_id) { + std::collections::hash_map::Entry::Vacant(entry) => { + entry.insert(claim); + true + } + std::collections::hash_map::Entry::Occupied(_) => false, + }; + if !inserted { + self.release_lane_locked(&mut state, lane_id); + state.active_peer_ports.remove(&peer_port); + manager.release_stream( + self.profile_key, + self.client_ip, + self.profile.public_addr, + peer_port, + ); + return Err(ManagerError::Concurrent); } drop(state); self.lane_open_notify.notify_waiters(); Ok(WebSocketLaneReservation { session: Arc::clone(self), - lane_id, - peer_port, - transferred: false, + claim, + stream: None, + phase: WebSocketLaneReservationPhase::Reserved, }) } @@ -192,12 +297,16 @@ impl WebSession { body: &[u8], ) -> Result { if !Arc::ptr_eq(self, &reservation.session) - || reservation.lane_id == 0 - || reservation.lane_id > frame::MAX_STREAM_ID + || reservation.lane_id() == 0 + || reservation.lane_id() > frame::MAX_STREAM_ID + || !matches!( + reservation.phase, + WebSocketLaneReservationPhase::Bound | WebSocketLaneReservationPhase::StreamOwned + ) { return Err(ManagerError::Protocol); } - let lane_id = reservation.lane_id; + let lane_id = reservation.lane_id(); let frames = frame::parse_all(body, &self.limits).map_err(|_| ManagerError::Protocol)?; if frames .iter() @@ -216,8 +325,19 @@ impl WebSession { return Err(ManagerError::Closed); } self.ensure_carrier_active_locked(&state)?; - if !reservation.transferred - && state.websocket_lane_reservations.get(&lane_id) != Some(&reservation.peer_port) + if state + .carrier_lanes + .get(&lane_id) + .is_none_or(|lane| lane.instance != reservation.claim.lane.instance) + || (reservation.phase == WebSocketLaneReservationPhase::Bound + && state.websocket_lane_reservations.get(&lane_id) != Some(&reservation.claim)) + || (reservation.phase == WebSocketLaneReservationPhase::StreamOwned + && reservation.stream.is_none_or(|stream| { + state + .streams + .get(&stream.id) + .is_none_or(|current| current.instance != stream.instance) + })) { return Err(ManagerError::Closed); } @@ -252,8 +372,8 @@ impl WebSession { let mut unused_bytes = reserve_bytes; let mut unused_items = reserve_items; let mut progress = AppliedProgress::default(); - let mut reserved_open = - (!reservation.transferred).then_some((lane_id, reservation.peer_port)); + let mut reserved_open = (reservation.phase == WebSocketLaneReservationPhase::Bound) + .then_some((lane_id, reservation.claim.peer_port)); let applied = self.apply_batch_locked( &mut state, &frames, @@ -287,15 +407,18 @@ impl WebSession { self.finish_carrier_health(); } for completion in opened { - if completion.stream.id != lane_id || completion.peer_port != reservation.peer_port { + let stream = completion.stream; + if stream.id != lane_id || completion.peer_port != reservation.claim.peer_port { return Err(ManagerError::Protocol); } + reservation.transfer_to_stream(stream)?; + reservation.mark_stream_owned(stream)?; if !self.spawn_stream(completion, true) { + reservation.retain_after_rejected_spawn(); return Err(ManagerError::Limit); } - reservation.transfer_to_stream(); } - if !reservation.transferred { + if reservation.phase != WebSocketLaneReservationPhase::StreamOwned { return Err(ManagerError::Protocol); } if let Some(manager) = self.manager.upgrade() { @@ -304,49 +427,87 @@ impl WebSession { Ok(progressed) } - /// Ends one failed or disconnected lane without closing its parent session. - pub(crate) fn close_websocket_lane(&self, lane_id: u32) { - let reserved = { - let mut state = self.state.lock(); - let reserved = state.websocket_lane_reservations.remove(&lane_id); - if let Some(stream) = state.streams.remove(&lane_id) { - state.closing_streams.insert(lane_id, stream.instance); - let (bytes, items) = inbound_queue_cost(&stream.inbound); - self.release_locked(&mut state, bytes, items, false); - if let Some(waker) = stream.read_waker { - waker.wake(); - } - if let Some(waker) = stream.write_waker { - waker.wake(); - } - } - self.remember_closed_locked(&mut state, lane_id); - self.release_lane_locked(&mut state, lane_id); - reserved - }; - if let Some(peer_port) = reserved { - self.release_websocket_lane_reservation(lane_id, peer_port); + /// Ends one exact failed or disconnected lane without closing its parent session. + pub(crate) fn close_websocket_lane(&self, mut reservation: WebSocketLaneReservation) { + if std::ptr::eq(self, Arc::as_ptr(&reservation.session)) { + reservation.release(); } - self.lane_open_notify.notify_waiters(); } - fn release_websocket_lane_reservation(&self, lane_id: u32, peer_port: u16) { - let removed = { + fn release_websocket_lane_claim( + &self, + claim: WebSocketLaneClaim, + stream: Option, + stream_owned: bool, + ) { + let release_port = { let mut state = self.state.lock(); - if state.websocket_lane_reservations.get(&lane_id) == Some(&peer_port) { - state.websocket_lane_reservations.remove(&lane_id); + let lane_matches = state + .carrier_lanes + .get(&claim.lane.lane_id) + .is_some_and(|lane| lane.instance == claim.lane.instance); + if !lane_matches && !state.closed { + return; } - self.release_lane_locked(&mut state, lane_id); - state.active_peer_ports.remove(&peer_port) + let release_port = if let Some(stream) = stream { + if stream.id != claim.lane.lane_id { + return; + } + let current_stream = state + .streams + .get(&stream.id) + .is_some_and(|current| current.instance == stream.instance); + if current_stream { + let Some(stream_state) = state.streams.remove(&stream.id) else { + return; + }; + state + .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); + if let Some(waker) = stream_state.read_waker { + waker.wake(); + } + if let Some(waker) = stream_state.write_waker { + waker.wake(); + } + false + } else if stream_owned { + false + } else { + state.active_peer_ports.remove(&claim.peer_port) + } + } else { + if state.websocket_lane_reservations.get(&claim.lane.lane_id) != Some(&claim) { + return; + } + state + .websocket_lane_reservations + .remove(&claim.lane.lane_id); + state.active_peer_ports.remove(&claim.peer_port) + }; + if lane_matches { + self.remember_closed_locked(&mut state, 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); + } + } + release_port }; - if removed && let Some(manager) = self.manager.upgrade() { + if release_port && let Some(manager) = self.manager.upgrade() { manager.release_stream( self.profile_key, self.client_ip, self.profile.public_addr, - peer_port, + claim.peer_port, ); } + self.lane_open_notify.notify_waiters(); } } diff --git a/src/web/session/websocket/tests.rs b/src/web/session/websocket/tests.rs index 7b2ab6b..c1d8b37 100644 --- a/src/web/session/websocket/tests.rs +++ b/src/web/session/websocket/tests.rs @@ -81,10 +81,48 @@ fn runtime(admission: bool) -> TestRuntime { } } +fn detach_stale_lane(runtime: &TestRuntime, reservation: &WebSocketLaneReservation) { + let claim = reservation.claim; + { + let mut state = runtime.session.state.lock(); + assert_eq!( + state + .websocket_lane_reservations + .remove(&claim.lane.lane_id), + Some(claim) + ); + runtime + .session + .release_lane_locked(&mut state, 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, + runtime.session.profile.public_addr, + claim.peer_port, + ); +} + +fn replacement_lane( + runtime: &TestRuntime, + stale: &WebSocketLaneReservation, +) -> WebSocketLaneReservation { + detach_stale_lane(runtime, stale); + let mut replacement = runtime + .session + .reserve_websocket_lane(stale.lane_id()) + .unwrap(); + replacement.bind(2).unwrap(); + assert_ne!(replacement.lane_identity(), stale.lane_identity()); + replacement +} + #[tokio::test] async fn rejected_open_retains_stream_quota_until_lane_socket_teardown() { let runtime = runtime(false); let mut reservation = runtime.session.reserve_websocket_lane(7).unwrap(); + reservation.bind(1).unwrap(); let open = frame::encode(FrameType::Open, 7, &[]); assert_eq!( @@ -93,6 +131,19 @@ async fn rejected_open_retains_stream_quota_until_lane_socket_teardown() { .process_websocket_lane(&mut reservation, 1, &open), Err(ManagerError::Limit), ); + assert_eq!( + reservation.phase, + WebSocketLaneReservationPhase::Transferred + ); + assert!(reservation.stream.is_some()); + assert!( + !runtime + .session + .state + .lock() + .websocket_lane_reservations + .contains_key(&7) + ); assert!( runtime .manager @@ -105,8 +156,124 @@ async fn rejected_open_retains_stream_quota_until_lane_socket_teardown() { .is_none() ); - runtime.session.close_websocket_lane(7); + runtime.session.close_websocket_lane(reservation); + let peer_port = runtime + .manager + .try_acquire_stream( + runtime.session.profile_key, + runtime.session.profile.max_streams, + runtime.session.client_ip, + runtime.session.profile.public_addr, + ) + .unwrap(); + runtime.manager.release_stream( + runtime.session.profile_key, + runtime.session.client_ip, + runtime.session.profile.public_addr, + peer_port, + ); + runtime.shutdown().await; +} + +#[tokio::test] +async fn closed_session_releases_bound_lane_quota_on_reservation_drop() { + let runtime = runtime(true); + let mut reservation = runtime.session.reserve_websocket_lane(7).unwrap(); + reservation.bind(1).unwrap(); + + runtime.session.close(); drop(reservation); + + assert!(runtime.session.state.lock().active_peer_ports.is_empty()); + let peer_port = runtime + .manager + .try_acquire_stream( + runtime.session.profile_key, + runtime.session.profile.max_streams, + runtime.session.client_ip, + runtime.session.profile.public_addr, + ) + .unwrap(); + runtime.manager.release_stream( + runtime.session.profile_key, + runtime.session.client_ip, + runtime.session.profile.public_addr, + peer_port, + ); + runtime.shutdown().await; +} + +#[tokio::test] +async fn closed_session_releases_transferred_rejected_lane_quota() { + let runtime = runtime(false); + let mut reservation = runtime.session.reserve_websocket_lane(7).unwrap(); + reservation.bind(1).unwrap(); + let open = frame::encode(FrameType::Open, 7, &[]); + assert_eq!( + runtime + .session + .process_websocket_lane(&mut reservation, 1, &open), + Err(ManagerError::Limit), + ); + assert_eq!( + reservation.phase, + WebSocketLaneReservationPhase::Transferred + ); + + runtime.session.close(); + drop(reservation); + + assert!(runtime.session.state.lock().active_peer_ports.is_empty()); + let peer_port = runtime + .manager + .try_acquire_stream( + runtime.session.profile_key, + runtime.session.profile.max_streams, + runtime.session.client_ip, + runtime.session.profile.public_addr, + ) + .unwrap(); + runtime.manager.release_stream( + runtime.session.profile_key, + runtime.session.client_ip, + runtime.session.profile.public_addr, + peer_port, + ); + runtime.shutdown().await; +} + +#[tokio::test] +async fn closed_session_keeps_stream_owned_quota_until_task_completion() { + let runtime = runtime(true); + let mut reservation = runtime.session.reserve_websocket_lane(7).unwrap(); + reservation.bind(1).unwrap(); + let open = frame::encode(FrameType::Open, 7, &[]); + assert_eq!( + runtime + .session + .process_websocket_lane(&mut reservation, 1, &open), + Ok(true), + ); + assert_eq!( + reservation.phase, + WebSocketLaneReservationPhase::StreamOwned + ); + + runtime.session.close(); + drop(reservation); + + assert!( + runtime + .manager + .try_acquire_stream( + runtime.session.profile_key, + runtime.session.profile.max_streams, + runtime.session.client_ip, + runtime.session.profile.public_addr, + ) + .is_none() + ); + runtime.session.wait().await; let peer_port = runtime .manager .try_acquire_stream( @@ -129,6 +296,7 @@ async fn rejected_open_retains_stream_quota_until_lane_socket_teardown() { async fn malformed_lane_message_does_not_close_sibling_session_state() { let runtime = runtime(true); let mut reservation = runtime.session.reserve_websocket_lane(7).unwrap(); + reservation.bind(1).unwrap(); let data = frame::encode(FrameType::Data, 7, &[1]); assert_eq!( @@ -139,8 +307,122 @@ async fn malformed_lane_message_does_not_close_sibling_session_state() { ); assert!(!runtime.session.state.lock().closed); - runtime.session.close_websocket_lane(7); - drop(reservation); + runtime.session.close_websocket_lane(reservation); assert!(runtime.session.reserve_websocket_lane(8).is_ok()); runtime.shutdown().await; } + +#[tokio::test] +async fn stale_websocket_poll_does_not_close_reused_lane_instance() { + let runtime = runtime(true); + let mut stale = runtime.session.reserve_websocket_lane(7).unwrap(); + stale.bind(1).unwrap(); + let stale_identity = stale.lane_identity(); + let replacement = replacement_lane(&runtime, &stale); + + let result = runtime + .session + .poll_down_websocket_lane(stale_identity, u64::MAX) + .await + .unwrap(); + + assert!(result.lane_closed); + assert!(!runtime.session.state.lock().closed); + assert_eq!( + runtime + .session + .state + .lock() + .carrier_lanes + .get(&7) + .map(|lane| lane.instance), + Some(replacement.lane_identity().instance) + ); + drop(stale); + runtime.session.close_websocket_lane(replacement); + runtime.shutdown().await; +} + +#[tokio::test] +async fn stale_close_preserves_replacement_lane_and_tuple() { + let runtime = runtime(true); + let mut stale = runtime.session.reserve_websocket_lane(7).unwrap(); + stale.bind(1).unwrap(); + let replacement = replacement_lane(&runtime, &stale); + let replacement_claim = replacement.claim; + stale.phase = WebSocketLaneReservationPhase::Transferred; + stale.stream = Some(StreamIdentity { id: 7, instance: 1 }); + + runtime.session.close_websocket_lane(stale); + + { + let state = runtime.session.state.lock(); + assert_eq!( + state.websocket_lane_reservations.get(&7), + Some(&replacement_claim) + ); + assert!( + state + .active_peer_ports + .contains(&replacement_claim.peer_port) + ); + assert_eq!( + state.carrier_lanes.get(&7).map(|lane| lane.instance), + Some(replacement_claim.lane.instance) + ); + } + runtime.session.close_websocket_lane(replacement); + runtime.shutdown().await; +} + +#[tokio::test] +async fn stale_reservation_drop_preserves_replacement_claim() { + let runtime = runtime(true); + let mut stale = runtime.session.reserve_websocket_lane(7).unwrap(); + stale.bind(1).unwrap(); + let replacement = replacement_lane(&runtime, &stale); + let replacement_claim = replacement.claim; + + drop(stale); + + { + let state = runtime.session.state.lock(); + assert_eq!( + state.websocket_lane_reservations.get(&7), + Some(&replacement_claim) + ); + assert!( + state + .active_peer_ports + .contains(&replacement_claim.peer_port) + ); + } + runtime.session.close_websocket_lane(replacement); + runtime.shutdown().await; +} + +#[tokio::test] +async fn stale_transfer_cannot_remove_current_reservation() { + let runtime = runtime(true); + let mut stale = runtime.session.reserve_websocket_lane(7).unwrap(); + stale.bind(1).unwrap(); + let replacement = replacement_lane(&runtime, &stale); + let replacement_claim = replacement.claim; + + assert_eq!( + stale.transfer_to_stream(StreamIdentity { id: 7, instance: 1 }), + Err(ManagerError::Closed) + ); + assert_eq!( + runtime + .session + .state + .lock() + .websocket_lane_reservations + .get(&7), + Some(&replacement_claim) + ); + drop(stale); + runtime.session.close_websocket_lane(replacement); + runtime.shutdown().await; +}