WEB: Lifecycle + Lane ownership + Diag fixes

Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
This commit is contained in:
Alexey
2026-08-27 09:02:19 +03:00
parent c75cf5cc9d
commit f73f52a033
25 changed files with 2050 additions and 245 deletions
+56 -1
View File
@@ -284,7 +284,7 @@ impl ListenerSlot {
} }
pub(super) async fn stop(&mut self) -> Result<(), String> { pub(super) async fn stop(&mut self) -> Result<(), String> {
self.cancellation.cancel(); self.request_stop();
if let Some(task) = self.task.take() { if let Some(task) = self.task.take() {
task.await.map_err(|error_value| { task.await.map_err(|error_value| {
format!("listener {} task failed: {error_value}", self.spec.addr) format!("listener {} task failed: {error_value}", self.spec.addr)
@@ -305,6 +305,61 @@ impl ListenerSlot {
Ok(()) 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<ArcSwap<RuntimeGeneration>>) { pub(super) fn restart(&mut self, active_runtime: Arc<ArcSwap<RuntimeGeneration>>) {
self.active_runtime = active_runtime.clone(); self.active_runtime = active_runtime.clone();
self.cancellation = CancellationToken::new(); self.cancellation = CancellationToken::new();
+61 -8
View File
@@ -1,6 +1,7 @@
use std::collections::{BTreeMap, BTreeSet}; use std::collections::{BTreeMap, BTreeSet};
use std::net::SocketAddr; use std::net::SocketAddr;
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration;
use arc_swap::ArcSwap; use arc_swap::ArcSwap;
@@ -13,7 +14,7 @@ use super::bind::{BoundListeners, BoundTcpListener, PreparedTcpListener, prepare
use super::plan::{ListenerBindSpec, listener_bind_plan}; use super::plan::{ListenerBindSpec, listener_bind_plan};
#[cfg(unix)] #[cfg(unix)]
use super::unix::UnixAcceptHandle; use super::unix::UnixAcceptHandle;
use crate::web::manager::WebProcessRuntime; use crate::web::manager::{WebProcessRuntime, WebShutdownOutcome};
use crate::web::trace::WebTraceStore; use crate::web::trace::WebTraceStore;
/// Process-owned listener inventory and accept-task lifecycle controller. /// 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> { 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(); let mut errors = Vec::new();
for slot in self.slots.values_mut() { let slot_waits = futures_util::future::join_all(
if let Err(error_value) = slot.stop().await { 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); errors.push(error_value);
} }
} }
#[cfg(unix)] #[cfg(unix)]
if let Some(unix) = &mut self.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); errors.push(error_value);
} }
self.slots.clear(); if web_outcome == WebShutdownOutcome::DeadlineExceeded {
if let Some(web_runtime) = self.web_runtime.take() { errors.push("WEB ingress shutdown deadline exceeded".to_string());
web_runtime.shutdown().await;
} }
self.slots.clear();
#[cfg(unix)] #[cfg(unix)]
{ {
self.unix = None; self.unix = None;
+32 -1
View File
@@ -148,11 +148,42 @@ impl UnixAcceptHandle {
} }
pub(super) async fn stop(&mut self) -> Result<(), String> { pub(super) async fn stop(&mut self) -> Result<(), String> {
self.cancellation.cancel(); self.request_stop();
if let Some(task) = self.task.take() { if let Some(task) = self.task.take() {
task.await task.await
.map_err(|error_value| format!("Unix listener task failed: {error_value}"))?; .map_err(|error_value| format!("Unix listener task failed: {error_value}"))?;
} }
Ok(()) 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())
}
}
}
} }
+18 -10
View File
@@ -2,7 +2,7 @@ use std::convert::Infallible;
use std::error::Error; use std::error::Error;
use std::net::SocketAddr; use std::net::SocketAddr;
use std::sync::Arc; use std::sync::Arc;
use std::time::{Duration, Instant}; use std::time::Duration;
use bytes::Bytes; use bytes::Bytes;
use http_body_util::BodyExt; use http_body_util::BodyExt;
@@ -13,7 +13,6 @@ use hyper::service::service_fn;
use hyper::{Method, Request, Response, StatusCode}; use hyper::{Method, Request, Response, StatusCode};
use hyper_util::rt::{TokioIo, TokioTimer}; use hyper_util::rt::{TokioIo, TokioTimer};
use ipnetwork::IpNetwork; use ipnetwork::IpNetwork;
use parking_lot::Mutex;
use tokio::net::TcpStream; use tokio::net::TcpStream;
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
@@ -45,7 +44,7 @@ mod websocket;
mod trace_tests; mod trace_tests;
use crate::web::trace::{HttpTraceExchange, TraceDirection, TraceLifecycleEvent, TraceRoute}; 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 body::{CollectBodyError, CollectedBody, RequestBody, collect_body};
use decoy::serve_decoy; use decoy::serve_decoy;
use down::handle_down; 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 max_header_bytes = config.web.limits.max_header_bytes;
let header_timeout = Duration::from_secs(config.web.timeouts.header_secs); let header_timeout = Duration::from_secs(config.web.timeouts.header_secs);
let idle_timeout = Duration::from_secs(config.web.timeouts.http_idle_secs); let idle_timeout = Duration::from_secs(config.web.timeouts.http_idle_secs);
let last_activity = Arc::new(Mutex::new(Instant::now())); let connection_activity = ConnectionActivity::new();
let service_last_activity = Arc::clone(&last_activity); let service_activity = connection_activity.clone();
let service = service_fn(move |mut request| { let service = service_fn(move |mut request| {
let runtime = Arc::clone(&runtime); let runtime = Arc::clone(&runtime);
let trusted_proxy_cidrs = Arc::clone(&trusted_proxy_cidrs); 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; let client_ip_source = client_ip_source;
async move { 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()); let trace = runtime.trace().begin_http(&request, peer.ip());
if let Some(trace) = &trace { if let Some(trace) = &trace {
request.extensions_mut().insert(Arc::clone(trace)); request.extensions_mut().insert(Arc::clone(trace));
@@ -133,9 +139,7 @@ pub(crate) async fn serve_connection(
_ = cancellation.cancelled() => break, _ = cancellation.cancelled() => break,
_ = &mut connection => break, _ = &mut connection => break,
_ = idle_check.tick() => { _ = idle_check.tick() => {
if Instant::now().saturating_duration_since(*last_activity.lock()) if connection_activity.should_close(tokio::time::Instant::now(), idle_timeout) {
>= idle_timeout
{
break; break;
} }
} }
@@ -404,6 +408,10 @@ fn request_trace<B>(request: &Request<B>) -> Option<&Arc<HttpTraceExchange>> {
request.extensions().get::<Arc<HttpTraceExchange>>() request.extensions().get::<Arc<HttpTraceExchange>>()
} }
fn request_deadline<B>(request: &Request<B>) -> Option<RequestDeadlineHandle> {
request.extensions().get::<RequestDeadlineHandle>().cloned()
}
fn set_trace_route<B>(request: &Request<B>, route: TraceRoute) { fn set_trace_route<B>(request: &Request<B>, route: TraceRoute) {
if let Some(trace) = request_trace(request) { if let Some(trace) = request_trace(request) {
trace.set_route(route); trace.set_route(route);
+326 -8
View File
@@ -1,31 +1,291 @@
use std::pin::Pin; use std::pin::Pin;
use std::sync::Arc; use std::sync::Arc;
use std::task::{Context, Poll}; use std::task::{Context, Poll};
use std::time::Instant; use std::time::Duration;
use bytes::Bytes; use bytes::Bytes;
use hyper::body::{Body, Frame, SizeHint}; use hyper::body::{Body, Frame, SizeHint};
use parking_lot::Mutex; use parking_lot::Mutex;
use tokio::time::Instant;
use super::{BoxError, HttpBody}; use super::{BoxError, HttpBody};
use crate::web::trace::{HttpTraceExchange, TraceBodyState, TraceDirection}; use crate::web::trace::{HttpTraceExchange, TraceBodyState, TraceDirection};
#[derive(Clone, Copy)]
struct DeadlineSlot {
id: u64,
deadline: Instant,
}
struct RequestSlot {
id: u64,
deadline: Option<DeadlineSlot>,
}
struct UpgradeSlot {
request_id: u64,
deadline: DeadlineSlot,
}
struct ActivityState {
last_progress: Instant,
next_request_id: u64,
next_deadline_id: u64,
request: Option<RequestSlot>,
upgrade: Option<UpgradeSlot>,
failed: bool,
}
/// Shared liveness state for one accepted HTTP connection.
#[derive(Clone)]
pub(super) struct ConnectionActivity {
state: Arc<Mutex<ActivityState>>,
}
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<RequestDeadlineLease> {
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<RequestDeadlineLease> {
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<UpgradeDeadlineLease> {
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. /// Request lifecycle guard that refreshes HTTP connection activity on completion.
pub(super) struct RequestActivity { pub(super) struct RequestActivity {
last_activity: Arc<Mutex<Instant>>, handle: RequestDeadlineHandle,
} }
impl RequestActivity { impl RequestActivity {
/// Starts activity accounting for one HTTP request. /// Starts activity accounting for one HTTP request.
pub(super) fn begin(last_activity: Arc<Mutex<Instant>>) -> Self { pub(super) fn begin(activity: ConnectionActivity) -> Option<Self> {
*last_activity.lock() = Instant::now(); let mut state = activity.state.lock();
Self { last_activity } 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 { impl Drop for RequestActivity {
fn drop(&mut self) { 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. /// Binds one response body to its request activity guard.
pub(super) fn new( pub(super) fn new(
inner: HttpBody, inner: HttpBody,
activity: RequestActivity, mut activity: RequestActivity,
trace: Option<Arc<HttpTraceExchange>>, trace: Option<Arc<HttpTraceExchange>>,
) -> Self { ) -> Self {
activity.enter_response();
Self { Self {
inner, inner,
activity, activity,
@@ -88,7 +349,7 @@ impl Body for ActivityBody {
Poll::Pending => {} Poll::Pending => {}
} }
if result.is_ready() { if result.is_ready() {
*self.activity.last_activity.lock() = Instant::now(); self.activity.progress();
} }
result result
} }
@@ -107,3 +368,60 @@ impl Drop for ActivityBody {
self.finish(TraceBodyState::Aborted); 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)));
}
}
+9
View File
@@ -118,6 +118,7 @@ pub(super) async fn collect_body(
limit: usize, limit: usize,
allow_empty: bool, allow_empty: bool,
) -> Result<CollectedBody, CollectBodyError> { ) -> Result<CollectedBody, CollectBodyError> {
let request_deadline = super::request_deadline(&request);
let exceeds_limit = request.body().size_hint().lower() > limit as u64 let exceeds_limit = request.body().size_hint().lower() > limit as u64
|| request || request
.body() .body()
@@ -134,6 +135,14 @@ pub(super) async fn collect_body(
let Some((reader_budget, body_budget)) = runtime.try_body_budget(limit) else { let Some((reader_budget, body_budget)) = runtime.try_body_budget(limit) else {
return Err(CollectBodyError::Limit); 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 { let body = match tokio::time::timeout(body_timeout, Limited::new(body, limit).collect()).await {
Ok(Ok(body)) => body.to_bytes(), Ok(Ok(body)) => body.to_bytes(),
_ => { _ => {
+26
View File
@@ -184,6 +184,7 @@ async fn proxy_to_upstream(
header_timeout: Duration, header_timeout: Duration,
runtime: &WebProcessRuntime, runtime: &WebProcessRuntime,
) -> HttpResponse { ) -> HttpResponse {
let request_deadline = super::request_deadline(&request);
remove_hop_by_hop(request.headers_mut()); remove_hop_by_hop(request.headers_mut());
if let Ok(host) = HeaderValue::from_str(authority) { if let Ok(host) = HeaderValue::from_str(authority) {
request.headers_mut().insert(header::HOST, host); request.headers_mut().insert(header::HOST, host);
@@ -197,10 +198,15 @@ async fn proxy_to_upstream(
return bad_gateway(); return bad_gateway();
}; };
*request.uri_mut() = uri; *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 { let stream = match tokio::time::timeout(header_timeout, TcpStream::connect(addr)).await {
Ok(Ok(stream)) => stream, Ok(Ok(stream)) => stream,
_ => return bad_gateway(), _ => return bad_gateway(),
}; };
drop(_deadline_lease);
let max_header_bytes = runtime let max_header_bytes = runtime
.active_generation() .active_generation()
.config() .config()
@@ -209,19 +215,29 @@ async fn proxy_to_upstream(
.max_header_bytes; .max_header_bytes;
let mut builder = hyper::client::conn::http1::Builder::new(); let mut builder = hyper::client::conn::http1::Builder::new();
builder.max_buf_size(max_header_bytes); 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) = let (mut sender, connection) =
match tokio::time::timeout(header_timeout, builder.handshake(TokioIo::new(stream))).await { match tokio::time::timeout(header_timeout, builder.handshake(TokioIo::new(stream))).await {
Ok(Ok(parts)) => parts, Ok(Ok(parts)) => parts,
_ => return bad_gateway(), _ => return bad_gateway(),
}; };
drop(_deadline_lease);
runtime.spawn_auxiliary(async move { runtime.spawn_auxiliary(async move {
let _ = connection.await; 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 = let mut response =
match tokio::time::timeout(header_timeout, sender.send_request(request)).await { match tokio::time::timeout(header_timeout, sender.send_request(request)).await {
Ok(Ok(response)) => response, Ok(Ok(response)) => response,
_ => return bad_gateway(), _ => return bad_gateway(),
}; };
drop(_deadline_lease);
remove_hop_by_hop(response.headers_mut()); remove_hop_by_hop(response.headers_mut());
response.map(|body| { response.map(|body| {
body.map_err(|error| -> BoxError { Box::new(error) }) 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<Option<super::activity::RequestDeadlineLease>, ()> {
match deadline {
Some(deadline) => deadline.lease_for(timeout).map(Some).ok_or(()),
None => Ok(None),
}
}
fn sanitize_transport_request<B>(request: &mut Request<B>) { fn sanitize_transport_request<B>(request: &mut Request<B>) {
for name in [ for name in [
header::AUTHORIZATION, header::AUTHORIZATION,
+16
View File
@@ -76,10 +76,26 @@ pub(super) async fn handle_down(
} else { } else {
None 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 { let result = match lane_id {
Some(lane_id) => session.poll_down_lane(lane_id, cursor).await, Some(lane_id) => session.poll_down_lane(lane_id, cursor).await,
None => session.poll_down(cursor).await, None => session.poll_down(cursor).await,
}; };
drop(_deadline_lease);
match result { match result {
Ok(result) if result.body.is_empty() => { Ok(result) if result.body.is_empty() => {
let mut response = carrier_empty(StatusCode::NO_CONTENT); let mut response = carrier_empty(StatusCode::NO_CONTENT);
+138
View File
@@ -1,5 +1,50 @@
use super::*; use super::*;
async fn open_keepalive(
listener: &TcpListener,
runtime: &Arc<WebProcessRuntime>,
) -> (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<u8> {
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::<usize>().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( async fn request_with_body_delay(
listener: &TcpListener, listener: &TcpListener,
runtime: &Arc<WebProcessRuntime>, runtime: &Arc<WebProcessRuntime>,
@@ -35,6 +80,9 @@ async fn live_session_body_and_closed_token_timeouts_survive_reload() {
let capability = [21u8; 32]; let capability = [21u8; 32];
let mut initial_config = runtime_config(capability, WebCarrier::Https); let mut initial_config = runtime_config(capability, WebCarrier::Https);
initial_config.web.timeouts.body_secs = 3; 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; initial_config.web.timeouts.bootstrap_lifetime_secs = 5;
let generation = test_runtime_generation(1, initial_config); let generation = test_runtime_generation(1, initial_config);
let active_runtime = Arc::new(ArcSwap::from(Arc::clone(&generation))); 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); let mut replacement_config = runtime_config(capability, WebCarrier::Https);
replacement_config.web.timeouts.body_secs = 1; 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; replacement_config.web.timeouts.bootstrap_lifetime_secs = 1;
let replacement = test_runtime_generation(2, replacement_config); let replacement = test_runtime_generation(2, replacement_config);
active_runtime.store(Arc::clone(&replacement)); 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!(retry_headers.starts_with(b"HTTP/1.1 200"));
assert_eq!(response_header(retry_headers, "x-session-token"), session); 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!( 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" "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_sessions().await;
replacement.stop_background_tasks().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;
}
+34 -3
View File
@@ -254,6 +254,7 @@ pub(super) async fn handle(
runtime: Arc<WebProcessRuntime>, runtime: Arc<WebProcessRuntime>,
vhost: Arc<WebRuntimeVhost>, vhost: Arc<WebRuntimeVhost>,
) -> HttpResponse { ) -> HttpResponse {
let request_deadline = super::request_deadline(&request);
let Some(parsed) = parse_upgrade(&request) else { let Some(parsed) = parse_upgrade(&request) else {
return serve_decoy(request, vhost, true, &runtime).await; return serve_decoy(request, vhost, true, &runtime).await;
}; };
@@ -284,7 +285,16 @@ pub(super) async fn handle(
}, },
}; };
let timeouts = session.timeouts().clone(); 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( .admit_websocket(
session.profile_key(), session.profile_key(),
session.trace_session_id(), session.trace_session_id(),
@@ -296,11 +306,17 @@ pub(super) async fn handle(
Duration::from_secs(timeouts.websocket_eviction_secs), Duration::from_secs(timeouts.websocket_eviction_secs),
session.carrier_cancellation(), session.carrier_cancellation(),
) )
.await .await;
{ drop(admission_lease);
let connection = match admitted {
Ok(connection) => connection, Ok(connection) => connection,
Err(_) => return serve_decoy(request, vhost, true, &runtime).await, 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() if let Some(reservation) = probe_reservation.as_mut()
&& reservation.bind(connection.id()).is_err() && reservation.bind(connection.id()).is_err()
{ {
@@ -323,6 +339,20 @@ pub(super) async fn handle(
trace.bind_identity(session.trace_identity()); trace.bind_identity(session.trace_identity());
trace.register_redaction(parsed.protocol.as_bytes()); 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 on_upgrade = hyper::upgrade::on(&mut request);
let protocol = parsed.protocol; let protocol = parsed.protocol;
let accept = parsed.accept; let accept = parsed.accept;
@@ -331,6 +361,7 @@ pub(super) async fn handle(
runtime.spawn_auxiliary(async move { runtime.spawn_auxiliary(async move {
driver::run_upgraded( driver::run_upgraded(
on_upgrade, on_upgrade,
upgrade_deadline,
driver_runtime, driver_runtime,
driver_session, driver_session,
connection, connection,
+17 -10
View File
@@ -8,6 +8,7 @@ use tokio_tungstenite::tungstenite::protocol::{Message, Role, WebSocketConfig};
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
use super::ConnectionIo; use super::ConnectionIo;
use crate::web::http::activity::UpgradeDeadlineLease;
use crate::web::manager::{WebProcessRuntime, WebSocketBudgetLease, WebSocketConnection}; use crate::web::manager::{WebProcessRuntime, WebSocketBudgetLease, WebSocketConnection};
use crate::web::session::{WebSession, WebSocketLaneReservation, WebSocketProbeReservation}; use crate::web::session::{WebSession, WebSocketLaneReservation, WebSocketProbeReservation};
use crate::web::trace::{TraceDirection, TraceWebSocketContext}; 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 io::{flush, process_multiplex, read_message, record_message, reserve_data, send};
use lane::run_lane; use lane::run_lane;
// Upgrade ownership remains explicit across cancellation and reservation boundaries.
#[allow(clippy::too_many_arguments)]
pub(super) async fn run_upgraded( pub(super) async fn run_upgraded(
on_upgrade: hyper::upgrade::OnUpgrade, on_upgrade: hyper::upgrade::OnUpgrade,
upgrade_deadline: Option<UpgradeDeadlineLease>,
runtime: Arc<WebProcessRuntime>, runtime: Arc<WebProcessRuntime>,
session: Arc<WebSession>, session: Arc<WebSession>,
connection: WebSocketConnection, connection: WebSocketConnection,
@@ -34,13 +38,15 @@ pub(super) async fn run_upgraded(
) { ) {
let cancellation = connection.cancellation(); let cancellation = connection.cancellation();
let timeouts = session.timeouts().clone(); 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! { let upgraded = tokio::select! {
_ = cancellation.cancelled() => return, _ = cancellation.cancelled() => return,
result = tokio::time::timeout( result = tokio::time::timeout_at(deadline, on_upgrade) => result,
Duration::from_secs(timeouts.websocket_upgrade_secs),
on_upgrade,
) => result,
}; };
drop(upgrade_deadline);
let Ok(Ok(upgraded)) = upgraded else { let Ok(Ok(upgraded)) = upgraded else {
return; return;
}; };
@@ -95,8 +101,7 @@ pub(super) async fn run_upgraded(
_ = tokio::time::timeout(eviction, socket.close(None)) => {} _ = tokio::time::timeout(eviction, socket.close(None)) => {}
} }
if let Some(reservation) = lane_reservation { if let Some(reservation) = lane_reservation {
session.close_websocket_lane(reservation.lane_id()); session.close_websocket_lane(reservation);
drop(reservation);
} else if !acknowledge_commit || session.is_carrier_committed() { } else if !acknowledge_commit || session.is_carrier_committed() {
session.close(); session.close();
} }
@@ -192,10 +197,12 @@ async fn run_multiplex(
session.close(); session.close();
return Err(()); return Err(());
} }
} else if acknowledge_commit && sequence > 1 && progressed { } else if acknowledge_commit
if !session.websocket_peer_after_commit_ack(connection.id()) { && sequence > 1
return Err(()); && progressed
} && !session.websocket_peer_after_commit_ack(connection.id())
{
return Err(());
} }
if !active && progressed { if !active && progressed {
if !connection.mark_active() { if !connection.mark_active() {
+8 -5
View File
@@ -12,6 +12,7 @@ use crate::web::session::{WebSession, WebSocketLaneReservation};
use crate::web::trace::{TraceDirection, TraceWebSocketContext}; use crate::web::trace::{TraceDirection, TraceWebSocketContext};
#[allow(clippy::too_many_arguments)] #[allow(clippy::too_many_arguments)]
/// Drives one exact WebSocket lane until its isolated failure boundary closes.
pub(super) async fn run_lane( pub(super) async fn run_lane(
socket: &mut CarrierSocket, socket: &mut CarrierSocket,
runtime: &Arc<WebProcessRuntime>, runtime: &Arc<WebProcessRuntime>,
@@ -35,7 +36,7 @@ pub(super) async fn run_lane(
let maximum_message = session.limits().carrier_batch_bytes; let maximum_message = session.limits().carrier_batch_bytes;
let mut active = false; let mut active = false;
loop { 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); tokio::pin!(down);
let event = tokio::select! { let event = tokio::select! {
_ = cancellation.cancelled() => return Err(()), _ = cancellation.cancelled() => return Err(()),
@@ -106,10 +107,12 @@ pub(super) async fn run_lane(
session.close(); session.close();
return Err(()); return Err(());
} }
} else if acknowledge_commit && sequence > 1 && progressed { } else if acknowledge_commit
if !session.websocket_peer_after_commit_ack(connection.id()) { && sequence > 1
return Err(()); && progressed
} && !session.websocket_peer_after_commit_ack(connection.id())
{
return Err(());
} }
if !active && progressed { if !active && progressed {
if !connection.mark_active() { if !connection.mark_active() {
+13 -2
View File
@@ -33,6 +33,7 @@ mod carrier_outcome;
mod admission; mod admission;
// Shutdown and expiry work remain outside request-path coordination. // Shutdown and expiry work remain outside request-path coordination.
mod lifecycle; mod lifecycle;
pub(crate) use lifecycle::WebShutdownOutcome;
// Queue and WebSocket allocations share one process-owned data-plane budget. // Queue and WebSocket allocations share one process-owned data-plane budget.
mod budget; mod budget;
// WebSocket admission, replacement, and liveness are process-scoped. // WebSocket admission, replacement, and liveness are process-scoped.
@@ -294,13 +295,23 @@ impl WebProcessRuntime {
where where
F: Future<Output = ()> + Send + 'static, F: Future<Output = ()> + Send + 'static,
{ {
if self.shutdown.is_cancelled() {
drop(future);
return;
}
let shutdown = self.shutdown.clone(); let shutdown = self.shutdown.clone();
self.tasks.spawn(async move { let tracked = self.tasks.track_future(async move {
tokio::select! { tokio::select! {
biased;
_ = shutdown.cancelled() => {} _ = shutdown.cancelled() => {}
_ = future => {} _ = future => {}
} }
}); });
if self.shutdown.is_cancelled() {
drop(tracked);
return;
}
drop(tokio::spawn(tracked));
} }
/// Reserves one body reader and its declared bounded body allocation. /// Reserves one body reader and its declared bounded body allocation.
@@ -393,7 +404,7 @@ impl WebProcessRuntime {
.try_reserve_websocket(owner, bytes, WebSocketBudgetClass::Data) .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)] #[allow(clippy::too_many_arguments)]
pub(crate) async fn admit_websocket( pub(crate) async fn admit_websocket(
self: &Arc<Self>, self: &Arc<Self>,
+23 -11
View File
@@ -58,6 +58,21 @@ pub(crate) struct WebDataBudgetSnapshot {
pub(crate) high_water_bytes: usize, 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<ProfileKey, usize>,
}
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 { impl WebDataBudget {
pub(super) fn new(limits: WebLimitsConfig) -> Arc<Self> { pub(super) fn new(limits: WebLimitsConfig) -> Arc<Self> {
Arc::new(Self { Arc::new(Self {
@@ -220,16 +235,10 @@ impl WebDataBudget {
self.pressured.store(true, Ordering::Release); self.pressured.store(true, Ordering::Release);
} }
pub(super) fn owner_usage(&self, owner: ProfileKey) -> usize { pub(super) fn fairness_snapshot(
self.state &self,
.lock() additional_owner: Option<ProfileKey>,
.owner_bytes ) -> WebSocketFairnessSnapshot {
.get(&owner)
.copied()
.unwrap_or(0)
}
pub(super) fn fair_share(&self, additional_owner: Option<ProfileKey>) -> usize {
let state = self.state.lock(); let state = self.state.lock();
let mut owners = state.owner_bytes.len(); let mut owners = state.owner_bytes.len();
if additional_owner.is_some_and(|owner| !state.owner_bytes.contains_key(&owner)) { 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_bytes_global,
self.limits.websocket_admission_watermark_pct, 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 { pub(super) fn snapshot(&self) -> WebDataBudgetSnapshot {
+277 -31
View File
@@ -2,13 +2,30 @@ use std::net::IpAddr;
use std::sync::atomic::Ordering; use std::sync::atomic::Ordering;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use tracing::info; use tokio::time::Instant as TokioInstant;
use tracing::{info, warn};
use super::state::{ use super::state::{
decrement_map, remember_closed_token_locked, remove_bootstrap_locked, remove_expired_locked, decrement_map, remember_closed_token_locked, remove_bootstrap_locked, remove_expired_locked,
}; };
use super::{ProfileKey, TokenHash, WebProcessRuntime}; 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<WebProcessRuntime>,
sessions: Vec<std::sync::Arc<crate::web::session::WebSession>>,
started: TokioInstant,
}
impl WebProcessRuntime { impl WebProcessRuntime {
/// Removes one closed session and retains a bounded host-bound replay marker. /// Removes one closed session and retains a bounded host-bound replay marker.
pub(crate) fn session_finished( pub(crate) fn session_finished(
@@ -49,11 +66,20 @@ impl WebProcessRuntime {
self.sessions_closed.fetch_add(1, Ordering::Relaxed); self.sessions_closed.fetch_add(1, Ordering::Relaxed);
} }
/// Stops issuance, closes all sessions, and joins bounded child work. /// Closes every WEB authority gate before any graceful wait begins.
pub(crate) async fn shutdown(&self) { pub(crate) fn begin_shutdown(self: &std::sync::Arc<Self>) -> WebShutdownDrain {
let started = TokioInstant::now();
self.shutdown.cancel(); self.shutdown.cancel();
self.close_websockets(); self.close_websockets();
self.data_budget.close(); 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 sessions = {
let mut state = self.state.lock(); let mut state = self.state.lock();
state.closed = true; state.closed = true;
@@ -65,6 +91,24 @@ impl WebProcessRuntime {
for session in &sessions { for session in &sessions {
session.close(); 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<Self>,
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<Self>) -> WebShutdownOutcome {
let timeout_secs = self let timeout_secs = self
.active_runtime .active_runtime
.load() .load()
@@ -72,34 +116,11 @@ impl WebProcessRuntime {
.web .web
.timeouts .timeouts
.shutdown_secs; .shutdown_secs;
let waits = async { let now = TokioInstant::now();
for session in sessions { let deadline = now
session.wait().await; .checked_add(Duration::from_secs(timeout_secs))
} .unwrap_or(now);
}; self.shutdown_until(deadline).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"
);
} }
/// Expires credentials and closes idle sessions without holding locks across callbacks. /// 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<S, T>(deadline: TokioInstant, sessions: S, tasks: T) -> WebShutdownOutcome
where
S: std::future::Future<Output = ()>,
T: std::future::Future<Output = ()>,
{
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<AtomicUsize>,
drops: Arc<AtomicUsize>,
}
impl Future for DropProbe {
type Output = ();
fn poll(self: Pin<&mut Self>, _context: &mut Context<'_>) -> Poll<Self::Output> {
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<WebProcessRuntime>,
Arc<crate::maestro::generation::RuntimeGeneration>,
) {
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;
}
}
+21 -50
View File
@@ -9,6 +9,9 @@ use tokio_util::sync::CancellationToken;
use super::{ManagerError, ProfileKey, WebProcessRuntime, WebSocketBudgetLease}; 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. /// One process-owned WebSocket carrier class used for eviction priority.
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub(crate) enum WebSocketKind { pub(crate) enum WebSocketKind {
@@ -25,6 +28,7 @@ struct WebSocketClaimKey {
} }
#[repr(u8)] #[repr(u8)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum WebSocketPhase { enum WebSocketPhase {
Claimed, Claimed,
Upgraded, Upgraded,
@@ -239,6 +243,8 @@ enum TryAdmitError {
Closed, Closed,
} }
// Admission inputs stay explicit so quota and cancellation ownership cannot drift.
#[allow(clippy::too_many_arguments)]
fn try_admit( fn try_admit(
runtime: &Arc<WebProcessRuntime>, runtime: &Arc<WebProcessRuntime>,
owner: ProfileKey, owner: ProfileKey,
@@ -353,8 +359,7 @@ fn select_victim(
excluded_id: Option<u64>, excluded_id: Option<u64>,
claim: bool, claim: bool,
) -> Option<Arc<WebSocketEntry>> { ) -> Option<Arc<WebSocketEntry>> {
let fair_share = runtime.data_budget.fair_share(Some(owner)); let fairness = runtime.data_budget.fairness_snapshot(Some(owner));
let requester_usage = runtime.data_budget.owner_usage(owner);
let now = runtime.websocket_tick(); let now = runtime.websocket_tick();
let mut registry = runtime.websockets.lock(); let mut registry = runtime.websockets.lock();
if claim && registry.evictions_in_flight >= runtime.limits.max_websocket_evictions_in_flight { 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| Some(entry.id) != excluded_id)
.filter(|entry| !entry.closing.load(Ordering::Acquire)) .filter(|entry| !entry.closing.load(Ordering::Acquire))
.filter_map(|entry| { .filter_map(|entry| {
let owner_rank = if entry.session_id == session_id { policy::admission_key(entry, now, owner, session_id, client_ip, &fairness)
0 .map(|key| (key, Arc::clone(entry)))
} else if entry.owner == owner {
1
} else if entry.client_ip == client_ip {
2
} else {
if requester_usage >= fair_share
|| runtime.data_budget.owner_usage(entry.owner) <= fair_share
{
return None;
}
3
};
let priority = entry_priority(entry, now);
Some((
(
owner_rank,
priority,
entry.last_progress_tick.load(Ordering::Acquire),
entry.created_tick,
entry.id,
),
Arc::clone(entry),
))
}) })
.min_by_key(|(key, _)| *key) .min_by_key(|(key, _)| *key)
.map(|(_, entry)| entry)?; .map(|(_, entry)| entry)?;
@@ -405,6 +387,7 @@ fn select_pressure_victim(
now: u64, now: u64,
claim: bool, claim: bool,
) -> Option<Arc<WebSocketEntry>> { ) -> Option<Arc<WebSocketEntry>> {
let fairness = runtime.data_budget.fairness_snapshot(None);
let mut registry = runtime.websockets.lock(); let mut registry = runtime.websockets.lock();
if claim && registry.evictions_in_flight >= runtime.limits.max_websocket_evictions_in_flight { if claim && registry.evictions_in_flight >= runtime.limits.max_websocket_evictions_in_flight {
return None; return None;
@@ -415,12 +398,7 @@ fn select_pressure_victim(
.filter(|entry| !entry.closing.load(Ordering::Acquire)) .filter(|entry| !entry.closing.load(Ordering::Acquire))
.map(|entry| { .map(|entry| {
( (
( policy::pressure_key(entry, now, &fairness),
entry_priority(entry, now),
entry.last_progress_tick.load(Ordering::Acquire),
entry.created_tick,
entry.id,
),
Arc::clone(entry), Arc::clone(entry),
) )
}) })
@@ -438,18 +416,23 @@ fn claim_stale_victims(runtime: &WebProcessRuntime, now: u64) -> Vec<Arc<WebSock
.limits .limits
.max_websocket_evictions_in_flight .max_websocket_evictions_in_flight
.saturating_sub(registry.evictions_in_flight); .saturating_sub(registry.evictions_in_flight);
let candidates = registry let mut candidates = registry
.entries .entries
.values() .values()
.filter(|entry| !entry.closing.load(Ordering::Acquire)) .filter(|entry| !entry.closing.load(Ordering::Acquire))
.filter(|entry| { .filter(|entry| policy::victim_class(entry, now) == policy::VictimClass::Dead)
now.saturating_sub(entry.last_peer_tick.load(Ordering::Acquire)) >= dead_after(entry)
})
.take(available)
.cloned() .cloned()
.collect::<Vec<_>>(); .collect::<Vec<_>>();
candidates.sort_unstable_by_key(|entry| {
(
entry.last_peer_tick.load(Ordering::Acquire),
entry.created_tick,
entry.id,
)
});
candidates candidates
.into_iter() .into_iter()
.take(available)
.filter(|entry| claim_entry(&mut registry, entry, runtime)) .filter(|entry| claim_entry(&mut registry, entry, runtime))
.collect() .collect()
} }
@@ -474,18 +457,6 @@ fn claim_entry(
true 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 { fn dead_after(entry: &WebSocketEntry) -> u64 {
entry.liveness_interval_ms.saturating_mul(2) entry.liveness_interval_ms.saturating_mul(2)
} }
+113
View File
@@ -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<VictimKey> {
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,
)
}
+276 -28
View File
@@ -1,25 +1,39 @@
use super::*; 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 { use super::policy::{VictimClass, admission_key, pressure_key, victim_class};
let phase = if opened { use crate::config::ProxyConfig;
WebSocketPhase::Active use crate::maestro::generation::test_runtime_generation;
} else { use crate::web::manager::budget::WebSocketFairnessSnapshot;
WebSocketPhase::Claimed 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 { WebSocketEntry {
id: 1, id,
owner: [0; 32], owner,
session_id: 1, session_id,
claim: WebSocketClaimKey { claim: WebSocketClaimKey {
session_hash: [0; 32], session_hash: [0; 32],
kind, kind,
}, },
client_ip: "192.0.2.10".parse().unwrap(), client_ip: client_ip.parse().unwrap(),
kind, kind,
liveness_interval_ms: 10, liveness_interval_ms: 10,
created_tick: 1, created_tick: 1,
last_peer_tick: AtomicU64::new(peer_tick), 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), phase: AtomicU8::new(phase as u8),
closing: AtomicBool::new(false), closing: AtomicBool::new(false),
cancel: CancellationToken::new(), cancel: CancellationToken::new(),
@@ -27,25 +41,259 @@ fn entry(kind: WebSocketKind, opened: bool, peer_tick: u64) -> WebSocketEntry {
} }
} }
#[test] fn fairness(fair_share: usize, usages: &[(ProfileKey, usize)]) -> WebSocketFairnessSnapshot {
fn preopen_and_dead_entries_precede_live_lane_and_multiplex_victims() { WebSocketFairnessSnapshot {
let preopen = entry(WebSocketKind::Multiplex, false, 90); fair_share,
let dead = entry(WebSocketKind::Multiplex, true, 1); owner_bytes: usages.iter().copied().collect::<HashMap<_, _>>(),
let lane = entry(WebSocketKind::Lane(7), true, 90); }
let multiplex = entry(WebSocketKind::Multiplex, true, 90);
assert_eq!(entry_priority(&preopen, 100), 0);
assert_eq!(entry_priority(&dead, 100), 0);
assert_eq!(entry_priority(&lane, 100), 1);
assert_eq!(entry_priority(&multiplex, 100), 2);
} }
#[test] #[test]
fn dead_classification_keeps_each_connections_creation_time_interval() { fn preactive_and_dead_are_distinct_lifecycle_classes() {
let short_interval = entry(WebSocketKind::Multiplex, true, 80); let preactive = entry(
let mut long_interval = entry(WebSocketKind::Multiplex, true, 80); 1,
long_interval.liveness_interval_ms = 100; [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!(victim_class(&preactive, 100), VictimClass::PreActive);
assert_eq!(entry_priority(&long_interval, 100), 2); 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::<Vec<_>>();
let connections = entries
.iter()
.map(|entry| WebSocketConnection {
runtime: Arc::downgrade(&runtime),
entry: Arc::clone(entry),
slot: None,
base_budget: None,
})
.collect::<Vec<_>>();
{
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;
} }
+22 -3
View File
@@ -62,6 +62,22 @@ pub(crate) struct StreamIdentity {
pub(crate) instance: u64, 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<u64>,
}
struct StreamState { struct StreamState {
instance: u64, instance: u64,
inbound: VecDeque<InboundChunk>, inbound: VecDeque<InboundChunk>,
@@ -144,7 +160,7 @@ struct SessionState {
carrier_lanes: HashMap<u32, CarrierLane>, carrier_lanes: HashMap<u32, CarrierLane>,
lane_open_waits: usize, lane_open_waits: usize,
next_lane_instance: u64, next_lane_instance: u64,
websocket_lane_reservations: HashMap<u32, u16>, websocket_lane_reservations: HashMap<u32, WebSocketLaneClaim>,
pending_bytes: usize, pending_bytes: usize,
pending_items: usize, pending_items: usize,
pending_control_bytes: usize, pending_control_bytes: usize,
@@ -506,11 +522,14 @@ fn remember_closed(state: &mut SessionState, stream_id: u32, limit: usize) -> Op
evicted evicted
} }
fn insert_carrier_lane(state: &mut SessionState, lane_id: u32) -> Option<u64> { fn insert_carrier_lane(state: &mut SessionState, lane_id: u32) -> Option<CarrierLaneIdentity> {
if state.carrier_lanes.contains_key(&lane_id) {
return None;
}
let instance = state.next_lane_instance; let instance = state.next_lane_instance;
state.next_lane_instance = instance.checked_add(1)?; state.next_lane_instance = instance.checked_add(1)?;
state state
.carrier_lanes .carrier_lanes
.insert(lane_id, CarrierLane::new(instance)); .insert(lane_id, CarrierLane::new(instance));
Some(instance) Some(CarrierLaneIdentity { lane_id, instance })
} }
+4 -2
View File
@@ -29,8 +29,10 @@ fn session() -> (Arc<WebSession>, Arc<WebProcessRuntime>) {
max_streams: 1, max_streams: 1,
max_streams_per_session: 1, max_streams_per_session: 1,
}); });
let mut timeouts = WebTimeoutsConfig::default(); let timeouts = WebTimeoutsConfig {
timeouts.long_poll_secs = 1; long_poll_secs: 1,
..WebTimeoutsConfig::default()
};
let session = WebSession::new( let session = WebSession::new(
Arc::downgrade(&manager), Arc::downgrade(&manager),
[1; 32], [1; 32],
+42 -4
View File
@@ -6,8 +6,8 @@ use tokio::sync::OwnedSemaphorePermit;
use super::lane_downlink::take_lane_down_batch; use super::lane_downlink::take_lane_down_batch;
use super::{ use super::{
PendingClass, PollResult, QUEUE_ITEM_COST, QueuedFrame, SessionState, WebSession, CarrierLaneIdentity, PendingClass, PollResult, QUEUE_ITEM_COST, QueuedFrame, SessionState,
remember_closed, WebSession, remember_closed,
}; };
use crate::web::frame::{self, FrameType}; use crate::web::frame::{self, FrameType};
use crate::web::manager::ManagerError; use crate::web::manager::ManagerError;
@@ -18,15 +18,46 @@ impl WebSession {
&self, &self,
lane_id: u32, lane_id: u32,
cursor: u64, cursor: u64,
) -> Result<PollResult, ManagerError> {
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<PollResult, ManagerError> {
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<u64>,
cursor: u64,
) -> Result<PollResult, ManagerError> { ) -> Result<PollResult, ManagerError> {
if !self.carrier().uses_lanes() || lane_id > frame::MAX_STREAM_ID { if !self.carrier().uses_lanes() || lane_id > frame::MAX_STREAM_ID {
return Err(ManagerError::Protocol); return Err(ManagerError::Protocol);
} }
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 { return Ok(PollResult {
body: Bytes::new(), body: Bytes::new(),
next_cursor: cursor, next_cursor: cursor,
lane_closed: false, lane_closed: expected_instance.is_some(),
}); });
} }
let (instance, epoch, notify, healthy) = { let (instance, epoch, notify, healthy) = {
@@ -43,6 +74,13 @@ impl WebSession {
lane_closed: true, 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 let Some(unacked) = &lane.unacked {
if cursor == unacked.base_cursor { if cursor == unacked.base_cursor {
return Ok(PollResult { return Ok(PollResult {
+5
View File
@@ -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. /// Atomically closes a session only when reconnect grace is still due.
pub(crate) fn close_if_due(&self, now: Instant) -> bool { pub(crate) fn close_if_due(&self, now: Instant) -> bool {
let healthy = { let healthy = {
+2
View File
@@ -166,6 +166,8 @@ impl WebSession {
result result
} }
// Batch application keeps every transactional accumulator explicit.
#[allow(clippy::too_many_arguments)]
pub(super) fn apply_batch_locked( pub(super) fn apply_batch_locked(
self: &Arc<Self>, self: &Arc<Self>,
state: &mut SessionState, state: &mut SessionState,
+226 -65
View File
@@ -4,7 +4,10 @@ use std::time::Instant;
use sha2::{Digest, Sha256}; use sha2::{Digest, Sha256};
use super::uplink::{AppliedProgress, inbound_reservation, validate_batch}; 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::config::WebCarrier;
use crate::web::frame; use crate::web::frame;
use crate::web::manager::ManagerError; 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. /// Pre-OPEN stream quota and synthetic tuple ownership for one WebSocket lane.
pub(crate) struct WebSocketLaneReservation { pub(crate) struct WebSocketLaneReservation {
session: Arc<WebSession>, session: Arc<WebSession>,
lane_id: u32, claim: WebSocketLaneClaim,
peer_port: u16, stream: Option<StreamIdentity>,
transferred: bool, 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. /// Session-wide ownership of the only automatic WebSocket carrier probe.
@@ -57,28 +70,99 @@ impl Drop for WebSocketProbeReservation {
impl WebSocketLaneReservation { impl WebSocketLaneReservation {
/// Returns the logical stream owned by this connection. /// Returns the logical stream owned by this connection.
pub(crate) fn lane_id(&self) -> u32 { pub(crate) fn lane_id(&self) -> u32 {
self.lane_id self.claim.lane.lane_id
} }
fn transfer_to_stream(&mut self) { /// Returns the exact lane incarnation owned by this connection.
let removed = self pub(crate) fn lane_identity(&self) -> CarrierLaneIdentity {
.session self.claim.lane
.state }
.lock()
.websocket_lane_reservations /// Binds this pre-upgrade reservation to one admitted process connection.
.remove(&self.lane_id); pub(crate) fn bind(&mut self, connection_id: u64) -> Result<(), ManagerError> {
if removed == Some(self.peer_port) { if self.phase != WebSocketLaneReservationPhase::Reserved {
self.transferred = true; 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 { impl Drop for WebSocketLaneReservation {
fn drop(&mut self) { fn drop(&mut self) {
if !self.transferred { self.release();
self.session
.release_websocket_lane_reservation(self.lane_id, self.peer_port);
}
} }
} }
@@ -162,9 +246,7 @@ impl WebSession {
); );
return Err(ManagerError::Limit); return Err(ManagerError::Limit);
} }
state.websocket_lane_reservations.insert(lane_id, peer_port); let Some(lane) = insert_carrier_lane(&mut state, lane_id) else {
if insert_carrier_lane(&mut state, lane_id).is_none() {
state.websocket_lane_reservations.remove(&lane_id);
state.active_peer_ports.remove(&peer_port); state.active_peer_ports.remove(&peer_port);
manager.release_stream( manager.release_stream(
self.profile_key, self.profile_key,
@@ -173,14 +255,37 @@ impl WebSession {
peer_port, peer_port,
); );
return Err(ManagerError::Protocol); 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); drop(state);
self.lane_open_notify.notify_waiters(); self.lane_open_notify.notify_waiters();
Ok(WebSocketLaneReservation { Ok(WebSocketLaneReservation {
session: Arc::clone(self), session: Arc::clone(self),
lane_id, claim,
peer_port, stream: None,
transferred: false, phase: WebSocketLaneReservationPhase::Reserved,
}) })
} }
@@ -192,12 +297,16 @@ impl WebSession {
body: &[u8], body: &[u8],
) -> Result<bool, ManagerError> { ) -> Result<bool, ManagerError> {
if !Arc::ptr_eq(self, &reservation.session) if !Arc::ptr_eq(self, &reservation.session)
|| reservation.lane_id == 0 || reservation.lane_id() == 0
|| reservation.lane_id > frame::MAX_STREAM_ID || reservation.lane_id() > frame::MAX_STREAM_ID
|| !matches!(
reservation.phase,
WebSocketLaneReservationPhase::Bound | WebSocketLaneReservationPhase::StreamOwned
)
{ {
return Err(ManagerError::Protocol); 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)?; let frames = frame::parse_all(body, &self.limits).map_err(|_| ManagerError::Protocol)?;
if frames if frames
.iter() .iter()
@@ -216,8 +325,19 @@ impl WebSession {
return Err(ManagerError::Closed); return Err(ManagerError::Closed);
} }
self.ensure_carrier_active_locked(&state)?; self.ensure_carrier_active_locked(&state)?;
if !reservation.transferred if state
&& state.websocket_lane_reservations.get(&lane_id) != Some(&reservation.peer_port) .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); return Err(ManagerError::Closed);
} }
@@ -252,8 +372,8 @@ impl WebSession {
let mut unused_bytes = reserve_bytes; let mut unused_bytes = reserve_bytes;
let mut unused_items = reserve_items; let mut unused_items = reserve_items;
let mut progress = AppliedProgress::default(); let mut progress = AppliedProgress::default();
let mut reserved_open = let mut reserved_open = (reservation.phase == WebSocketLaneReservationPhase::Bound)
(!reservation.transferred).then_some((lane_id, reservation.peer_port)); .then_some((lane_id, reservation.claim.peer_port));
let applied = self.apply_batch_locked( let applied = self.apply_batch_locked(
&mut state, &mut state,
&frames, &frames,
@@ -287,15 +407,18 @@ impl WebSession {
self.finish_carrier_health(); self.finish_carrier_health();
} }
for completion in opened { 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); return Err(ManagerError::Protocol);
} }
reservation.transfer_to_stream(stream)?;
reservation.mark_stream_owned(stream)?;
if !self.spawn_stream(completion, true) { if !self.spawn_stream(completion, true) {
reservation.retain_after_rejected_spawn();
return Err(ManagerError::Limit); return Err(ManagerError::Limit);
} }
reservation.transfer_to_stream();
} }
if !reservation.transferred { if reservation.phase != WebSocketLaneReservationPhase::StreamOwned {
return Err(ManagerError::Protocol); return Err(ManagerError::Protocol);
} }
if let Some(manager) = self.manager.upgrade() { if let Some(manager) = self.manager.upgrade() {
@@ -304,49 +427,87 @@ impl WebSession {
Ok(progressed) Ok(progressed)
} }
/// Ends one failed or disconnected lane without closing its parent session. /// Ends one exact failed or disconnected lane without closing its parent session.
pub(crate) fn close_websocket_lane(&self, lane_id: u32) { pub(crate) fn close_websocket_lane(&self, mut reservation: WebSocketLaneReservation) {
let reserved = { if std::ptr::eq(self, Arc::as_ptr(&reservation.session)) {
let mut state = self.state.lock(); reservation.release();
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);
} }
self.lane_open_notify.notify_waiters();
} }
fn release_websocket_lane_reservation(&self, lane_id: u32, peer_port: u16) { fn release_websocket_lane_claim(
let removed = { &self,
claim: WebSocketLaneClaim,
stream: Option<StreamIdentity>,
stream_owned: bool,
) {
let release_port = {
let mut state = self.state.lock(); let mut state = self.state.lock();
if state.websocket_lane_reservations.get(&lane_id) == Some(&peer_port) { let lane_matches = state
state.websocket_lane_reservations.remove(&lane_id); .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); let release_port = if let Some(stream) = stream {
state.active_peer_ports.remove(&peer_port) 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( manager.release_stream(
self.profile_key, self.profile_key,
self.client_ip, self.client_ip,
self.profile.public_addr, self.profile.public_addr,
peer_port, claim.peer_port,
); );
} }
self.lane_open_notify.notify_waiters();
} }
} }
+285 -3
View File
@@ -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] #[tokio::test]
async fn rejected_open_retains_stream_quota_until_lane_socket_teardown() { async fn rejected_open_retains_stream_quota_until_lane_socket_teardown() {
let runtime = runtime(false); let runtime = runtime(false);
let mut reservation = runtime.session.reserve_websocket_lane(7).unwrap(); let mut reservation = runtime.session.reserve_websocket_lane(7).unwrap();
reservation.bind(1).unwrap();
let open = frame::encode(FrameType::Open, 7, &[]); let open = frame::encode(FrameType::Open, 7, &[]);
assert_eq!( assert_eq!(
@@ -93,6 +131,19 @@ async fn rejected_open_retains_stream_quota_until_lane_socket_teardown() {
.process_websocket_lane(&mut reservation, 1, &open), .process_websocket_lane(&mut reservation, 1, &open),
Err(ManagerError::Limit), 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!( assert!(
runtime runtime
.manager .manager
@@ -105,8 +156,124 @@ async fn rejected_open_retains_stream_quota_until_lane_socket_teardown() {
.is_none() .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); 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 let peer_port = runtime
.manager .manager
.try_acquire_stream( .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() { async fn malformed_lane_message_does_not_close_sibling_session_state() {
let runtime = runtime(true); let runtime = runtime(true);
let mut reservation = runtime.session.reserve_websocket_lane(7).unwrap(); let mut reservation = runtime.session.reserve_websocket_lane(7).unwrap();
reservation.bind(1).unwrap();
let data = frame::encode(FrameType::Data, 7, &[1]); let data = frame::encode(FrameType::Data, 7, &[1]);
assert_eq!( assert_eq!(
@@ -139,8 +307,122 @@ async fn malformed_lane_message_does_not_close_sibling_session_state() {
); );
assert!(!runtime.session.state.lock().closed); assert!(!runtime.session.state.lock().closed);
runtime.session.close_websocket_lane(7); runtime.session.close_websocket_lane(reservation);
drop(reservation);
assert!(runtime.session.reserve_websocket_lane(8).is_ok()); assert!(runtime.session.reserve_websocket_lane(8).is_ok());
runtime.shutdown().await; 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;
}