From d706b3f3ba726cadf15f5baf8e85dc35267cd1b3 Mon Sep 17 00:00:00 2001 From: Alexey <247128645+axkurcom@users.noreply.github.com> Date: Sun, 20 Sep 2026 00:28:52 +0300 Subject: [PATCH] Hardswap Invariants in tests + Quota fixes --- src/api/config_edit.rs | 22 +- src/api/config_store/atomic.rs | 99 +++++- src/api/config_store/persistence.rs | 4 +- src/api/handler/user_routes.rs | 75 +++-- src/api/mod.rs | 28 +- src/api/users.rs | 1 + src/api/users/create.rs | 23 +- src/api/users/lifecycle.rs | 36 +- src/api/users/update.rs | 69 +++- src/maestro/control_plane.rs | 40 +++ src/maestro/orchestrator.rs | 15 +- src/maestro/reload_supervisor.rs | 5 +- src/maestro/runtime_build.rs | 1 + src/maestro/runtime_startup.rs | 1 + src/maestro/runtime_tasks.rs | 12 +- src/metrics/render/me_hardswap.rs | 2 +- src/proxy/authenticated.rs | 36 +- src/proxy/client/authenticated.rs | 9 +- src/proxy/direct_relay.rs | 1 + src/proxy/direct_relay/relay.rs | 4 + src/proxy/middle_relay.rs | 5 +- src/proxy/middle_relay/d2c.rs | 7 +- src/proxy/middle_relay/quota.rs | 6 +- src/proxy/middle_relay/session.rs | 16 +- src/proxy/middle_relay/session/tasks.rs | 5 + src/proxy/relay.rs | 2 + src/proxy/relay/adaptive_copy.rs | 5 +- src/proxy/relay/io.rs | 23 +- src/proxy/shared_state.rs | 31 +- ...ddle_relay_atomic_quota_invariant_tests.rs | 1 + src/proxy/traffic_limiter.rs | 35 +- src/proxy/traffic_limiter/buckets.rs | 34 +- src/proxy/traffic_limiter/lease.rs | 43 ++- src/proxy/traffic_limiter/limiter.rs | 31 +- src/proxy/traffic_limiter/tests.rs | 37 ++- src/proxy/user_admission.rs | 92 +++++- src/proxy/user_admission/tests.rs | 68 +++- src/quota_state.rs | 9 +- src/stats/mod.rs | 7 +- src/stats/quota_store.rs | 312 ++++++++++++++++-- src/stats/replay.rs | 67 +++- src/stats/users.rs | 21 ++ .../middle_proxy/pool/writer_admission.rs | 113 +++---- src/transport/middle_proxy/pool_reinit.rs | 15 + .../middle_proxy/pool_reinit/coordination.rs | 94 +++++- .../middle_proxy/pool_reinit/reconcile.rs | 199 ++++++----- .../middle_proxy/pool_reinit/tests.rs | 124 ++++++- src/transport/middle_proxy/pool_status.rs | 3 +- .../pool_status/hardswap_snapshot.rs | 38 ++- .../middle_proxy/pool_writer/publication.rs | 30 +- .../middle_proxy/pool_writer/replacement.rs | 5 + .../middle_proxy/pool_writer/runtime.rs | 11 +- src/web/http/operator_lifecycle_tests.rs | 29 ++ src/web/http/websocket/driver.rs | 2 + src/web/http/websocket/driver/io.rs | 7 + src/web/http/websocket/driver/lane.rs | 1 + src/web/manager/credentials.rs | 59 ++-- src/web/manager/state.rs | 11 +- src/web/session/backend.rs | 8 + src/web/session/downlink.rs | 49 ++- src/web/session/lane_uplink.rs | 7 + src/web/session/lanes.rs | 120 ++++--- src/web/session/lifecycle.rs | 24 +- src/web/session/negotiation.rs | 3 + src/web/session/uplink.rs | 7 + src/web/session/websocket.rs | 27 +- 66 files changed, 1821 insertions(+), 505 deletions(-) diff --git a/src/api/config_edit.rs b/src/api/config_edit.rs index a431072..b4d937a 100644 --- a/src/api/config_edit.rs +++ b/src/api/config_edit.rs @@ -62,6 +62,21 @@ pub(super) async fn patch_config( expected_revision: Option, reload_request: Option, shared: &ApiShared, +) -> Result { + let shared = shared.clone(); + shared + .clone() + .run_mutation_completion(async move { + patch_config_to_completion(patch_json, expected_revision, reload_request, &shared).await + }) + .await +} + +async fn patch_config_to_completion( + patch_json: Json, + expected_revision: Option, + reload_request: Option, + shared: &ApiShared, ) -> Result { let _guard = shared.mutation_lock.lock().await; let active_config = shared.active_runtime.load_full().config(); @@ -83,7 +98,7 @@ pub(super) async fn patch_config( } else { None }; - write_atomic_if_unchanged( + prepared.response.revision = write_atomic_if_unchanged( prepared.config_path, prepared.expected_revision, prepared.owner_path, @@ -123,8 +138,8 @@ pub(super) async fn apply_patch_to_path( patch_json: &Json, expected_revision: Option, ) -> Result { - let prepared = prepare_patch_to_path(config_path, patch_json, expected_revision).await?; - write_atomic_if_unchanged( + let mut prepared = prepare_patch_to_path(config_path, patch_json, expected_revision).await?; + let revision = write_atomic_if_unchanged( prepared.config_path, prepared.expected_revision, prepared.owner_path, @@ -132,6 +147,7 @@ pub(super) async fn apply_patch_to_path( prepared.owner_contents, ) .await?; + prepared.response.revision = revision; Ok(prepared.response) } diff --git a/src/api/config_store/atomic.rs b/src/api/config_store/atomic.rs index c5e7864..6d546b6 100644 --- a/src/api/config_store/atomic.rs +++ b/src/api/config_store/atomic.rs @@ -31,6 +31,11 @@ struct ExistingTarget { metadata: std::fs::Metadata, } +struct GraphFence<'a> { + config_path: &'a Path, + expected_revision: &'a str, +} + struct ConfigWriteLock { #[cfg(unix)] _file: Flock, @@ -38,9 +43,10 @@ struct ConfigWriteLock { impl ConfigWriteLock { fn acquire(path: &Path) -> std::io::Result { + let path = normalize_path(path); #[cfg(unix)] { - let lock_path = sibling_lock_path(path); + let lock_path = sibling_lock_path(&path); let anchored = AnchoredPath::open_creating_parents(&lock_path, 0o750)?; let descriptor = openat( anchored.parent(), @@ -76,7 +82,7 @@ pub(in crate::api) async fn write_atomic( ) -> Result<(), ApiFailure> { tokio::task::spawn_blocking(move || { let _lock = ConfigWriteLock::acquire(&path)?; - write_atomic_sync(&path, None, &contents) + write_atomic_sync(&path, None, &contents, None).map(|_| ()) }) .await .map_err(|error| ApiFailure::internal(format!("failed to join writer: {error}")))? @@ -90,8 +96,10 @@ pub(in crate::api) async fn write_atomic_if_unchanged( path: PathBuf, expected_contents: String, contents: String, -) -> Result<(), ApiFailure> { +) -> Result { tokio::task::spawn_blocking(move || { + let config_path = normalize_path(&config_path); + let path = normalize_path(&path); // Every API mutation locks the root source so writes to different includes serialize. let _lock = ConfigWriteLock::acquire(&config_path).map_err(AtomicWriteError::Io)?; let graph = ProxyConfig::read_source_graph(&config_path) @@ -99,12 +107,26 @@ pub(in crate::api) async fn write_atomic_if_unchanged( if compute_source_revision(&graph) != expected_revision { return Err(AtomicWriteError::Conflict); } - write_atomic_sync(&path, Some(&expected_contents), &contents).map_err(|error| { + write_atomic_sync( + &path, + Some(&expected_contents), + &contents, + Some(GraphFence { + config_path: &config_path, + expected_revision: &expected_revision, + }), + ) + .map_err(|error| { if error.kind() == std::io::ErrorKind::AlreadyExists { AtomicWriteError::Conflict } else { AtomicWriteError::Io(error) } + })? + .ok_or_else(|| { + AtomicWriteError::Io(std::io::Error::other( + "config graph fence did not produce a committed revision", + )) }) }) .await @@ -137,6 +159,51 @@ fn sibling_lock_path(path: &Path) -> PathBuf { path.parent().unwrap_or_else(|| Path::new(".")).join(name) } +fn normalize_path(path: &Path) -> PathBuf { + let absolute = if path.is_absolute() { + path.to_path_buf() + } else { + std::env::current_dir() + .map(|current| current.join(path)) + .unwrap_or_else(|_| path.to_path_buf()) + }; + let mut normalized = PathBuf::new(); + for component in absolute.components() { + match component { + std::path::Component::CurDir => {} + std::path::Component::ParentDir => { + normalized.pop(); + } + component => normalized.push(component.as_os_str()), + } + } + normalized +} + +fn fenced_post_commit_revision( + fence: GraphFence<'_>, + path: &Path, + contents: &str, +) -> std::io::Result { + let mut graph = ProxyConfig::read_source_graph(fence.config_path) + .map_err(|error| std::io::Error::other(error.to_string()))?; + if compute_source_revision(&graph) != fence.expected_revision { + return Err(std::io::Error::new( + std::io::ErrorKind::AlreadyExists, + "config graph changed during persistence", + )); + } + let path = normalize_path(path); + let Some(owner) = graph.source_contents.get_mut(&path) else { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + "config source owner left the source graph during persistence", + )); + }; + *owner = contents.to_string(); + Ok(compute_source_revision(&graph)) +} + #[cfg(unix)] fn open_existing_target(anchored: &AnchoredPath) -> std::io::Result> { let descriptor = match openat( @@ -210,7 +277,8 @@ fn write_atomic_sync( path: &Path, expected_contents: Option<&str>, contents: &str, -) -> std::io::Result<()> { + graph_fence: Option>, +) -> std::io::Result> { let anchored = AnchoredPath::open_creating_parents(path, 0o750)?; let existing = open_existing_target(&anchored)?; validate_expected_contents(existing.as_ref(), expected_contents)?; @@ -234,10 +302,12 @@ fn write_atomic_sync( .map_err(errno_to_io)?; let write_result = write_and_publish( descriptor, + path, &anchored, &temp_name, existing.as_ref(), contents, + graph_fence, ); if write_result.is_err() { let _ = unlinkat( @@ -252,11 +322,13 @@ fn write_atomic_sync( #[cfg(unix)] fn write_and_publish( descriptor: std::os::fd::OwnedFd, + path: &Path, anchored: &AnchoredPath, temp_name: &str, existing: Option<&ExistingTarget>, contents: &str, -) -> std::io::Result<()> { + graph_fence: Option>, +) -> std::io::Result> { let mut file = File::from(descriptor); if let Some(existing) = existing { use nix::unistd::{Gid, Uid, fchown}; @@ -280,6 +352,9 @@ fn write_and_publish( "config target changed during persistence", )); } + let committed_revision = graph_fence + .map(|fence| fenced_post_commit_revision(fence, path, contents)) + .transpose()?; renameat( anchored.parent(), temp_name, @@ -287,7 +362,8 @@ fn write_and_publish( anchored.name(), ) .map_err(errno_to_io)?; - fsync(anchored.parent()).map_err(errno_to_io) + fsync(anchored.parent()).map_err(errno_to_io)?; + Ok(committed_revision) } #[cfg(not(unix))] @@ -295,7 +371,8 @@ fn write_atomic_sync( path: &Path, expected_contents: Option<&str>, contents: &str, -) -> std::io::Result<()> { + graph_fence: Option>, +) -> std::io::Result> { let parent = path.parent().unwrap_or_else(|| Path::new(".")); std::fs::create_dir_all(parent)?; let existing = open_existing_target(path)?; @@ -310,7 +387,11 @@ fn write_atomic_sync( "config target changed during persistence", )); } - std::fs::rename(temp, path) + let committed_revision = graph_fence + .map(|fence| fenced_post_commit_revision(fence, path, contents)) + .transpose()?; + std::fs::rename(temp, path)?; + Ok(committed_revision) } fn validate_expected_contents( diff --git a/src/api/config_store/persistence.rs b/src/api/config_store/persistence.rs index d19e72e..ea68042 100644 --- a/src/api/config_store/persistence.rs +++ b/src/api/config_store/persistence.rs @@ -159,8 +159,8 @@ pub(in crate::api) async fn save_access_sections_to_disk_if_revision( owner_contents.clone(), ) .await?; - let revision = compute_snapshot_revision(&candidate); - write_atomic_if_unchanged( + let _candidate_revision = compute_snapshot_revision(&candidate); + let revision = write_atomic_if_unchanged( config_path.to_path_buf(), loaded_revision, owner_path, diff --git a/src/api/handler/user_routes.rs b/src/api/handler/user_routes.rs index f265213..234c5b5 100644 --- a/src/api/handler/user_routes.rs +++ b/src/api/handler/user_routes.rs @@ -133,38 +133,55 @@ pub(super) async fn handle( )); } let expected_revision = parse_if_match(req.headers()); - let _mutation_guard = shared.mutation_lock.lock().await; - let (disk_cfg, _) = - load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?; - if !disk_cfg.access.users.contains_key(user) { - return Ok(error_response( - request_id, - ApiFailure::new(StatusCode::NOT_FOUND, "not_found", "User not found"), - )); - } - let configured_users = disk_cfg - .access - .users - .keys() - .cloned() - .collect::>(); - let snapshot = match shared.quota_state.reset_user(&configured_users, user).await { - Ok(snapshot) => snapshot, - Err(error) => { - shared.runtime_events.record( - "api.user.reset_quota.failed", - format!("username={} error={}", user, error), + let completion_shared = shared.as_ref().clone(); + let user_owned = user.to_string(); + let completion = shared + .run_mutation_completion(async move { + let _mutation_guard = completion_shared.mutation_lock.lock().await; + let (disk_cfg, _) = load_config_for_mutation( + &completion_shared.config_path, + expected_revision.as_deref(), + ) + .await?; + if !disk_cfg.access.users.contains_key(&user_owned) { + return Err(ApiFailure::new( + StatusCode::NOT_FOUND, + "not_found", + "User not found", + )); + } + let configured_users = disk_cfg + .access + .users + .keys() + .cloned() + .collect::>(); + let snapshot = completion_shared + .quota_state + .reset_user(&configured_users, &user_owned) + .await + .map_err(|error| { + completion_shared.runtime_events.record( + "api.user.reset_quota.failed", + format!("username={} error={}", user_owned, error), + ); + ApiFailure::internal(format!("Failed to reset user quota: {}", error)) + })?; + completion_shared.runtime_events.record( + "api.user.reset_quota.ok", + format!("username={}", user_owned), ); - return Err(ApiFailure::internal(format!( - "Failed to reset user quota: {}", - error - ))); + let revision = current_revision(&completion_shared.config_path).await?; + Ok((snapshot, revision)) + }) + .await; + let (snapshot, revision) = match completion { + Ok(result) => result, + Err(error) if error.code == "not_found" => { + return Ok(error_response(request_id, error)); } + Err(error) => return Err(error), }; - shared - .runtime_events - .record("api.user.reset_quota.ok", format!("username={}", user)); - let revision = current_revision(&shared.config_path).await?; return Ok(success_response( StatusCode::OK, ResetUserQuotaResponse { diff --git a/src/api/mod.rs b/src/api/mod.rs index 9d9ce68..a17b6cf 100644 --- a/src/api/mod.rs +++ b/src/api/mod.rs @@ -17,7 +17,7 @@ use hyper::service::service_fn; use hyper::{Method, Request, Response, StatusCode}; use subtle::ConstantTimeEq; use tokio::net::TcpListener; -use tokio::sync::{Mutex, RwLock, Semaphore, watch}; +use tokio::sync::{Mutex, RwLock, Semaphore, oneshot, watch}; use tokio::time::timeout; use tracing::{debug, info, warn}; @@ -134,6 +134,7 @@ pub(super) struct ApiShared { pub(super) active_runtime: Arc>, pub(super) web_trace: Arc, pub(super) web_runtime_rx: watch::Receiver, + pub(super) control_plane: ProcessControlPlane, } impl ApiShared { @@ -169,8 +170,32 @@ impl ApiShared { active_runtime: self.active_runtime.clone(), web_trace: self.web_trace.clone(), web_runtime_rx: self.web_runtime_rx.clone(), + control_plane: self.control_plane.clone(), } } + + /// Keeps an accepted mutation alive until persistence and mandatory publication finish. + async fn run_mutation_completion(&self, future: F) -> Result + where + T: Send + 'static, + F: std::future::Future> + Send + 'static, + { + let (result_tx, result_rx) = oneshot::channel(); + self.control_plane + .spawn_completion(async move { + let _ = result_tx.send(future.await); + }) + .map_err(|_| { + ApiFailure::new( + StatusCode::SERVICE_UNAVAILABLE, + "control_plane_shutting_down", + "Control plane is shutting down", + ) + })?; + result_rx.await.map_err(|_| { + ApiFailure::internal("accepted config mutation did not report completion") + })? + } } fn auth_header_matches(actual: &str, expected: &str) -> bool { @@ -364,6 +389,7 @@ pub(crate) async fn serve( active_runtime, web_trace, web_runtime_rx, + control_plane: control_plane.clone(), }); spawn_runtime_watchers( diff --git a/src/api/users.rs b/src/api/users.rs index fc6286f..b16f4bb 100644 --- a/src/api/users.rs +++ b/src/api/users.rs @@ -5,6 +5,7 @@ use hyper::StatusCode; use crate::config::ProxyConfig; use crate::config::RateLimitBps; use crate::ip_tracker::UserIpTracker; +use crate::proxy::user_admission::credential_id_from_hex; use crate::stats::Stats; use super::ApiShared; diff --git a/src/api/users/create.rs b/src/api/users/create.rs index 40a48df..84cf8c3 100644 --- a/src/api/users/create.rs +++ b/src/api/users/create.rs @@ -4,6 +4,20 @@ pub(in crate::api) async fn create_user( body: CreateUserRequest, expected_revision: Option, shared: &ApiShared, +) -> Result<(CreateUserResponse, String), ApiFailure> { + let shared = shared.clone(); + shared + .clone() + .run_mutation_completion(async move { + create_user_to_completion(body, expected_revision, &shared).await + }) + .await +} + +async fn create_user_to_completion( + body: CreateUserRequest, + expected_revision: Option, + shared: &ApiShared, ) -> Result<(CreateUserResponse, String), ApiFailure> { let touches_user_ad_tags = body.user_ad_tag.is_some(); let touches_user_max_tcp_conns = body.max_tcp_conns.is_some(); @@ -41,6 +55,8 @@ pub(in crate::api) async fn create_user( } let expiration = parse_optional_expiration(body.expiration_rfc3339.as_deref())?; + let credential_id = credential_id_from_hex(&secret) + .ok_or_else(|| ApiFailure::internal("validated user secret could not be decoded"))?; let _guard = shared.mutation_lock.lock().await; let (mut cfg, base_revision) = load_config_for_mutation(&shared.config_path, expected_revision.as_deref()).await?; @@ -131,12 +147,11 @@ pub(in crate::api) async fn create_user( .await?; shared .proxy_shared - .stage_user( + .stage_user_credential( &body.username, - &secret, + credential_id, cfg.access.is_user_enabled(&body.username), - ) - .ok_or_else(|| ApiFailure::internal("failed to stage user admission policy"))?; + ); if let Some(limit) = updated_limit { shared diff --git a/src/api/users/lifecycle.rs b/src/api/users/lifecycle.rs index e962fb6..a84ee7e 100644 --- a/src/api/users/lifecycle.rs +++ b/src/api/users/lifecycle.rs @@ -6,6 +6,22 @@ pub(in crate::api) async fn rotate_secret( body: RotateSecretRequest, expected_revision: Option, shared: &ApiShared, +) -> Result<(CreateUserResponse, String), ApiFailure> { + let shared = shared.clone(); + let user = user.to_string(); + shared + .clone() + .run_mutation_completion(async move { + rotate_secret_to_completion(&user, body, expected_revision, &shared).await + }) + .await +} + +async fn rotate_secret_to_completion( + user: &str, + body: RotateSecretRequest, + expected_revision: Option, + shared: &ApiShared, ) -> Result<(CreateUserResponse, String), ApiFailure> { let secret = body.secret.unwrap_or_else(random_user_secret); if !is_valid_user_secret(&secret) { @@ -13,6 +29,8 @@ pub(in crate::api) async fn rotate_secret( "secret must be exactly 32 hex characters", )); } + let credential_id = credential_id_from_hex(&secret) + .ok_or_else(|| ApiFailure::internal("validated user secret could not be decoded"))?; let _guard = shared.mutation_lock.lock().await; let (mut cfg, base_revision) = @@ -38,8 +56,7 @@ pub(in crate::api) async fn rotate_secret( .await?; shared .proxy_shared - .stage_user(user, &secret, cfg.access.is_user_enabled(user)) - .ok_or_else(|| ApiFailure::internal("failed to stage rotated user credential"))?; + .stage_user_credential(user, credential_id, cfg.access.is_user_enabled(user)); drop(_guard); let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); @@ -70,6 +87,21 @@ pub(in crate::api) async fn delete_user( user: &str, expected_revision: Option, shared: &ApiShared, +) -> Result<(String, String), ApiFailure> { + let shared = shared.clone(); + let user = user.to_string(); + shared + .clone() + .run_mutation_completion(async move { + delete_user_to_completion(&user, expected_revision, &shared).await + }) + .await +} + +async fn delete_user_to_completion( + user: &str, + expected_revision: Option, + shared: &ApiShared, ) -> Result<(String, String), ApiFailure> { let _guard = shared.mutation_lock.lock().await; let (mut cfg, base_revision) = diff --git a/src/api/users/update.rs b/src/api/users/update.rs index 1619e10..461156c 100644 --- a/src/api/users/update.rs +++ b/src/api/users/update.rs @@ -5,6 +5,22 @@ pub(in crate::api) async fn patch_user( body: PatchUserRequest, expected_revision: Option, shared: &ApiShared, +) -> Result<(UserInfo, String), ApiFailure> { + let shared = shared.clone(); + let user = user.to_string(); + shared + .clone() + .run_mutation_completion(async move { + patch_user_to_completion(&user, body, expected_revision, &shared).await + }) + .await +} + +async fn patch_user_to_completion( + user: &str, + body: PatchUserRequest, + expected_revision: Option, + shared: &ApiShared, ) -> Result<(UserInfo, String), ApiFailure> { let touches_users = body.secret.is_some(); let touches_user_ad_tags = !matches!(&body.user_ad_tag, Patch::Unchanged); @@ -138,6 +154,19 @@ pub(in crate::api) async fn patch_user( cfg.validate() .map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?; + let staged_credential = if touches_users || touches_user_enabled { + let secret = cfg + .access + .users + .get(user) + .ok_or_else(|| ApiFailure::internal("updated user secret is missing"))?; + Some( + credential_id_from_hex(secret) + .ok_or_else(|| ApiFailure::internal("validated user secret could not be decoded"))?, + ) + } else { + None + }; let mut touched_sections = Vec::new(); if touches_users { @@ -176,16 +205,10 @@ pub(in crate::api) async fn patch_user( ) .await? }; - if touches_users || touches_user_enabled { - let secret = cfg - .access - .users - .get(user) - .ok_or_else(|| ApiFailure::internal("updated user secret is missing"))?; + if let Some(credential_id) = staged_credential { shared .proxy_shared - .stage_user(user, secret, cfg.access.is_user_enabled(user)) - .ok_or_else(|| ApiFailure::internal("failed to stage user admission policy"))?; + .stage_user_credential(user, credential_id, cfg.access.is_user_enabled(user)); } match max_unique_ips_change { Some(Some(limit)) => shared.ip_tracker.set_user_limit(user, limit).await, @@ -216,6 +239,22 @@ pub(in crate::api) async fn set_user_enabled( enabled: bool, expected_revision: Option, shared: &ApiShared, +) -> Result<(UserInfo, String), ApiFailure> { + let shared = shared.clone(); + let user = user.to_string(); + shared + .clone() + .run_mutation_completion(async move { + set_user_enabled_to_completion(&user, enabled, expected_revision, &shared).await + }) + .await +} + +async fn set_user_enabled_to_completion( + user: &str, + enabled: bool, + expected_revision: Option, + shared: &ApiShared, ) -> Result<(UserInfo, String), ApiFailure> { let _guard = shared.mutation_lock.lock().await; let (mut cfg, base_revision) = @@ -237,6 +276,12 @@ pub(in crate::api) async fn set_user_enabled( cfg.validate() .map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?; + let credential_id = cfg + .access + .users + .get(user) + .and_then(|secret| credential_id_from_hex(secret)) + .ok_or_else(|| ApiFailure::internal("validated user secret could not be decoded"))?; let revision = save_access_sections_to_disk_if_revision( &shared.config_path, &cfg, @@ -244,15 +289,9 @@ pub(in crate::api) async fn set_user_enabled( Some(&base_revision), ) .await?; - let secret = cfg - .access - .users - .get(user) - .ok_or_else(|| ApiFailure::internal("updated user secret is missing"))?; shared .proxy_shared - .stage_user(user, secret, enabled) - .ok_or_else(|| ApiFailure::internal("failed to stage user admission policy"))?; + .stage_user_credential(user, credential_id, enabled); drop(_guard); let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); diff --git a/src/maestro/control_plane.rs b/src/maestro/control_plane.rs index 89cda3a..c8f6f42 100644 --- a/src/maestro/control_plane.rs +++ b/src/maestro/control_plane.rs @@ -120,6 +120,19 @@ impl ProcessControlPlane { Ok(()) } + /// Registers work that must finish once accepted, even after shutdown cancellation starts. + pub(crate) fn spawn_completion(&self, future: F) -> Result<(), F> + where + F: Future + Send + 'static, + { + let Some(registration) = self.inner.admission.try_register() else { + return Err(future); + }; + self.inner.tasks.spawn(future); + drop(registration); + Ok(()) + } + /// Closes task admission, cancels all owned work, and joins it within the deadline. pub(crate) async fn shutdown(&self, timeout: Duration) -> bool { let deadline = tokio::time::Instant::now() + timeout; @@ -214,4 +227,31 @@ mod tests { assert!(scope.shutdown(Duration::from_secs(1)).await); } + + #[tokio::test] + async fn shutdown_waits_for_accepted_completion_without_cancelling_it() { + let scope = ProcessControlPlane::new(); + let (release_tx, release_rx) = tokio::sync::oneshot::channel(); + let completed = Arc::new(AtomicBool::new(false)); + let completed_task = completed.clone(); + assert!( + scope + .spawn_completion(async move { + let _ = release_rx.await; + completed_task.store(true, Ordering::Release); + }) + .is_ok() + ); + + let shutdown_scope = scope.clone(); + let shutdown = + tokio::spawn(async move { shutdown_scope.shutdown(Duration::from_secs(1)).await }); + tokio::task::yield_now().await; + assert!(!shutdown.is_finished()); + assert!(!completed.load(Ordering::Acquire)); + + release_tx.send(()).unwrap(); + assert!(shutdown.await.unwrap()); + assert!(completed.load(Ordering::Acquire)); + } } diff --git a/src/maestro/orchestrator.rs b/src/maestro/orchestrator.rs index 8dad10c..fe0344d 100644 --- a/src/maestro/orchestrator.rs +++ b/src/maestro/orchestrator.rs @@ -12,6 +12,7 @@ use crate::network::probe::{decide_network_capabilities, log_probe_result, run_p use crate::proxy::direct_buffer_budget::{DirectBufferBudget, resolve_direct_buffer_hard_limit}; use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController}; use crate::proxy::shared_state::ProxySharedState; +use crate::proxy::user_admission::UserAdmissionAuthority; use crate::startup::{COMPONENT_API_BOOTSTRAP, COMPONENT_NETWORK_PROBE}; use crate::stats::telemetry::TelemetryPolicy; use crate::stats::{QuotaStore, Stats}; @@ -106,9 +107,17 @@ pub(super) async fn run_telemt_core( configured_override_bytes = config.general.direct_relay_buffer_budget_max_bytes, "Direct relay buffer budget initialized" ); - let shared_state = - ProxySharedState::new_with_direct_buffer_budget(direct_buffer_budget.clone()); - shared_state.apply_user_config(&config.access.users, &config.access.user_enabled); + let user_admission = UserAdmissionAuthority::new_with_quota_store(quota_store.clone()); + let shared_state = ProxySharedState::new_with_direct_buffer_budget_and_user_admission( + direct_buffer_budget.clone(), + user_admission, + ); + let _ = shared_state.activate_user_config_source( + 1, + None, + &config.access.users, + &config.access.user_enabled, + ); shared_state.traffic_limiter.apply_policy( config.access.user_rate_limits.clone(), config.access.cidr_rate_limits.clone(), diff --git a/src/maestro/reload_supervisor.rs b/src/maestro/reload_supervisor.rs index a4372e4..725c329 100644 --- a/src/maestro/reload_supervisor.rs +++ b/src/maestro/reload_supervisor.rs @@ -303,8 +303,9 @@ impl ReloadSupervisor { let replaced = { let listener_manager = self.listener_manager.lock().await; let config = new_runtime.config(); - let _ = new_runtime.proxy_shared.apply_user_config_if_epoch( - user_admission_epoch, + let _ = new_runtime.proxy_shared.activate_user_config_source( + new_runtime.id, + Some(user_admission_epoch), &config.access.users, &config.access.user_enabled, ); diff --git a/src/maestro/runtime_build.rs b/src/maestro/runtime_build.rs index 16b401d..a289896 100644 --- a/src/maestro/runtime_build.rs +++ b/src/maestro/runtime_build.rs @@ -192,6 +192,7 @@ pub(crate) async fn prepare_runtime( let max_connections = Arc::new(Semaphore::new(max_connections_limit)); let (config_watcher_activation, config_watcher_activation_rx) = watch::channel(false); let watches = runtime_tasks::spawn_runtime_tasks( + generation_id, &config, config_path, &probe, diff --git a/src/maestro/runtime_startup.rs b/src/maestro/runtime_startup.rs index bfa4485..e7193d1 100644 --- a/src/maestro/runtime_startup.rs +++ b/src/maestro/runtime_startup.rs @@ -229,6 +229,7 @@ pub(super) async fn prepare_runtime( } let runtime_watches = runtime_tasks::spawn_runtime_tasks( + 1, &config, config_path, probe, diff --git a/src/maestro/runtime_tasks.rs b/src/maestro/runtime_tasks.rs index 300f4e7..e46e604 100644 --- a/src/maestro/runtime_tasks.rs +++ b/src/maestro/runtime_tasks.rs @@ -90,6 +90,7 @@ impl RuntimeLogFilter { #[allow(clippy::too_many_arguments)] pub(crate) async fn spawn_runtime_tasks( + generation_id: u64, config: &Arc, config_path: &Path, probe: &NetworkProbe, @@ -288,9 +289,14 @@ pub(crate) async fn spawn_runtime_tasks( break; } let cfg = config_rx_user_enabled.borrow_and_update().clone(); - for (user, cancelled) in shared_user_enabled - .apply_user_config(&cfg.access.users, &cfg.access.user_enabled) - { + let Some(cancelled_users) = shared_user_enabled.apply_user_config_from_source( + generation_id, + &cfg.access.users, + &cfg.access.user_enabled, + ) else { + continue; + }; + for (user, cancelled) in cancelled_users { if cancelled > 0 { info!( user = %user, diff --git a/src/metrics/render/me_hardswap.rs b/src/metrics/render/me_hardswap.rs index e8d19d8..878db57 100644 --- a/src/metrics/render/me_hardswap.rs +++ b/src/metrics/render/me_hardswap.rs @@ -76,7 +76,7 @@ pub(super) fn render( ); let _ = writeln!( out, - "# HELP telemt_me_hardswap_pending_missing_dc_groups Desired DC groups missing pending-generation coverage" + "# HELP telemt_me_hardswap_pending_missing_dc_groups Desired DC-family groups below the pending-generation floor" ); let _ = writeln!( out, diff --git a/src/proxy/authenticated.rs b/src/proxy/authenticated.rs index 26f3ec7..571f9a1 100644 --- a/src/proxy/authenticated.rs +++ b/src/proxy/authenticated.rs @@ -15,7 +15,7 @@ use crate::proxy::middle_relay::{handle_via_middle_proxy, handle_via_middle_prox use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController}; use crate::proxy::shared_state::{ConntrackClosePolicy, ProxySharedState}; use crate::proxy::user_admission::UserIncarnation; -use crate::stats::Stats; +use crate::stats::{Stats, UserQuotaHandle}; use crate::stream::{BufferPool, CryptoReader, CryptoWriter}; use crate::transport::UpstreamManager; use crate::transport::middle_proxy::MePool; @@ -85,6 +85,7 @@ where warn!(user = %user, error = %error, "User admission check failed"); error })?; + let quota_handle = user_reservation.quota_handle(); let route_snapshot = deps.route_runtime.snapshot(); let session_id = deps.rng.u64(); @@ -137,6 +138,7 @@ where session_id, session_cancel.clone(), Arc::clone(&deps.shared), + quota_handle.clone(), ) .await } else { @@ -156,6 +158,7 @@ where session_cancel.clone(), Arc::clone(&deps.shared), ConntrackClosePolicy::Suppress, + quota_handle.clone(), ) .await } @@ -171,6 +174,7 @@ where local_addr, session_cancel.clone(), conntrack_close_policy, + quota_handle.clone(), ) .await } @@ -185,6 +189,7 @@ where local_addr, session_cancel, conntrack_close_policy, + quota_handle, ) .await }; @@ -202,6 +207,7 @@ async fn run_direct( local_addr: SocketAddr, session_cancel: tokio_util::sync::CancellationToken, conntrack_close_policy: ConntrackClosePolicy, + quota_handle: UserQuotaHandle, ) -> Result<()> where R: AsyncRead + Unpin + Send + 'static, @@ -223,6 +229,7 @@ where session_cancel, Arc::clone(&deps.shared), conntrack_close_policy, + quota_handle, ) .await } @@ -235,6 +242,7 @@ pub(crate) struct UserConnectionReservation { user: String, ip: IpAddr, incarnation: UserIncarnation, + quota_handle: UserQuotaHandle, tracks_ip: bool, active: bool, } @@ -248,7 +256,16 @@ impl UserConnectionReservation { ip: IpAddr, tracks_ip: bool, ) -> Self { - Self::new_for_incarnation(stats, ip_tracker, user, ip, 0, tracks_ip) + let quota_handle = stats.current_user_quota_handle(&user); + Self::new_for_incarnation( + stats, + ip_tracker, + user, + ip, + 0, + quota_handle, + tracks_ip, + ) } /// Creates a reservation fenced to one authenticated user incarnation. @@ -258,6 +275,7 @@ impl UserConnectionReservation { user: String, ip: IpAddr, incarnation: UserIncarnation, + quota_handle: UserQuotaHandle, tracks_ip: bool, ) -> Self { Self { @@ -266,11 +284,17 @@ impl UserConnectionReservation { user, ip, incarnation, + quota_handle, tracks_ip, active: true, } } + /// Returns quota ownership pinned to the authenticated user incarnation. + pub(crate) fn quota_handle(&self) -> UserQuotaHandle { + self.quota_handle.clone() + } + /// Releases both admission counters through the asynchronous cleanup path. pub(crate) async fn release(mut self) { if !self.active { @@ -354,8 +378,13 @@ async fn acquire_user_connection_reservation_for_incarnation( user: user.to_string(), }); } + let Some(quota_handle) = stats.quota_handle_for_incarnation(user, incarnation) else { + return Err(ProxyError::UserDisabled { + user: user.to_string(), + }); + }; if let Some(quota) = config.access.user_data_quota.get(user) - && stats.get_user_quota_used(user) >= *quota + && quota_handle.used() >= *quota { return Err(ProxyError::DataQuotaExceeded { user: user.to_string(), @@ -399,6 +428,7 @@ async fn acquire_user_connection_reservation_for_incarnation( user.to_string(), peer_addr.ip(), incarnation, + quota_handle, true, )) } diff --git a/src/proxy/client/authenticated.rs b/src/proxy/client/authenticated.rs index 66f385f..b6a0f94 100644 --- a/src/proxy/client/authenticated.rs +++ b/src/proxy/client/authenticated.rs @@ -38,7 +38,14 @@ impl RunningClientHandler { } else { config }; - let shared = ProxySharedState::new(); + let shared = ProxySharedState::new_with_direct_buffer_budget_and_user_admission( + crate::proxy::direct_buffer_budget::DirectBufferBudget::new( + crate::proxy::direct_buffer_budget::fallback_direct_buffer_hard_limit(), + ), + crate::proxy::user_admission::UserAdmissionAuthority::new_with_quota_store( + stats.quota_store(), + ), + ); shared.apply_user_config(&config.access.users, &config.access.user_enabled); Self::handle_authenticated_static_with_shared( client_reader, diff --git a/src/proxy/direct_relay.rs b/src/proxy/direct_relay.rs index c22b209..45a1b8b 100644 --- a/src/proxy/direct_relay.rs +++ b/src/proxy/direct_relay.rs @@ -27,6 +27,7 @@ use crate::proxy::shared_state::{ ProxySharedState, }; use crate::stats::Stats; +use crate::stats::UserQuotaHandle; use crate::stream::{BufferPool, CryptoReader, CryptoWriter}; use crate::transport::UpstreamManager; #[cfg(unix)] diff --git a/src/proxy/direct_relay/relay.rs b/src/proxy/direct_relay/relay.rs index b70420a..298ff5d 100644 --- a/src/proxy/direct_relay/relay.rs +++ b/src/proxy/direct_relay/relay.rs @@ -59,6 +59,7 @@ where R: AsyncRead + Unpin + Send + 'static, W: AsyncWrite + Unpin + Send + 'static, { + let quota_handle = stats.current_user_quota_handle(&success.user); handle_via_direct_with_shared_and_conntrack( client_reader, client_writer, @@ -75,6 +76,7 @@ where session_cancel, shared, ConntrackClosePolicy::Publish, + quota_handle, ) .await } @@ -96,6 +98,7 @@ pub(crate) async fn handle_via_direct_with_shared_and_conntrack( session_cancel: CancellationToken, shared: Arc, conntrack_close_policy: ConntrackClosePolicy, + quota_handle: UserQuotaHandle, ) -> Result<()> where R: AsyncRead + Unpin + Send + 'static, @@ -171,6 +174,7 @@ where config.server.max_connections, user, Arc::clone(&stats), + quota_handle, config.access.user_data_quota.get(user).copied(), traffic_lease, relay_activity_timeout, diff --git a/src/proxy/middle_relay.rs b/src/proxy/middle_relay.rs index 8182c03..9db3730 100644 --- a/src/proxy/middle_relay.rs +++ b/src/proxy/middle_relay.rs @@ -31,7 +31,8 @@ use crate::proxy::shared_state::{ }; use crate::proxy::traffic_limiter::{RateDirection, TrafficLease, next_refill_delay}; use crate::stats::{ - MeD2cFlushReason, MeD2cQuotaRejectStage, MeD2cWriteMode, QuotaReserveError, Stats, UserStats, + MeD2cFlushReason, MeD2cQuotaRejectStage, MeD2cWriteMode, QuotaReserveError, Stats, + UserQuotaHandle, UserStats, }; use crate::stream::{BufferPool, CryptoReader, CryptoWriter, PooledBuffer}; use crate::transport::middle_proxy::{ConnLease, MePool, MeResponse, proto_flags_for_tag}; @@ -108,6 +109,7 @@ pub(crate) async fn handle_via_middle_proxy( session_id: u64, session_cancel: CancellationToken, shared: Arc, + quota_handle: UserQuotaHandle, ) -> Result<()> where R: AsyncRead + Unpin + Send + 'static, @@ -129,6 +131,7 @@ where session_cancel, shared, ConntrackClosePolicy::Publish, + quota_handle, ) .await } diff --git a/src/proxy/middle_relay/d2c.rs b/src/proxy/middle_relay/d2c.rs index a1906fc..c5f887f 100644 --- a/src/proxy/middle_relay/d2c.rs +++ b/src/proxy/middle_relay/d2c.rs @@ -135,6 +135,7 @@ pub(crate) async fn process_me_writer_response( where W: AsyncWrite + Unpin + Send + 'static, { + let quota_handle = quota_limit.map(|_| stats.current_user_quota_handle(user)); process_me_writer_response_with_traffic_lease( response, client_writer, @@ -144,6 +145,7 @@ where stats, user, quota_user_stats, + quota_handle.as_ref(), quota_limit, quota_soft_overshoot_bytes, None, @@ -165,6 +167,7 @@ pub(crate) async fn process_me_writer_response_with_traffic_lease( stats: &Stats, user: &str, quota_user_stats: Option<&UserStats>, + quota_handle: Option<&UserQuotaHandle>, quota_limit: Option, quota_soft_overshoot_bytes: u64, traffic_lease: Option<&Arc>, @@ -185,10 +188,10 @@ where trace!(conn_id, bytes = data.len(), flags, "ME->C data"); } let data_len = data.len() as u64; - if let (Some(limit), Some(user_stats)) = (quota_limit, quota_user_stats) { + if let (Some(limit), Some(quota_handle)) = (quota_limit, quota_handle) { let soft_limit = quota_soft_cap(limit, quota_soft_overshoot_bytes); match reserve_user_quota_with_yield( - user_stats, data_len, soft_limit, stats, cancel, None, + quota_handle, data_len, soft_limit, stats, cancel, None, ) .await { diff --git a/src/proxy/middle_relay/quota.rs b/src/proxy/middle_relay/quota.rs index 3a04c00..1484f2b 100644 --- a/src/proxy/middle_relay/quota.rs +++ b/src/proxy/middle_relay/quota.rs @@ -12,7 +12,7 @@ pub(super) fn quota_soft_cap(limit: u64, overshoot: u64) -> u64 { } pub(super) async fn reserve_user_quota_with_yield( - user_stats: &UserStats, + quota_handle: &UserQuotaHandle, bytes: u64, limit: u64, stats: &Stats, @@ -23,8 +23,8 @@ pub(super) async fn reserve_user_quota_with_yield( let mut backoff_rounds = 0usize; loop { for _ in 0..QUOTA_RESERVE_SPIN_RETRIES { - match user_stats.quota_try_reserve(bytes, limit) { - Ok(total) => return Ok(total), + match quota_handle.try_reserve(bytes, limit) { + Ok(reservation) => return Ok(reservation.commit()), Err(QuotaReserveError::LimitExceeded) => { return Err(MiddleQuotaReserveError::LimitExceeded); } diff --git a/src/proxy/middle_relay/session.rs b/src/proxy/middle_relay/session.rs index bee6675..89fc26e 100644 --- a/src/proxy/middle_relay/session.rs +++ b/src/proxy/middle_relay/session.rs @@ -61,6 +61,7 @@ pub(crate) async fn handle_via_middle_proxy_with_conntrack( session_cancel: CancellationToken, shared: Arc, conntrack_close_policy: ConntrackClosePolicy, + quota_handle: UserQuotaHandle, ) -> Result<()> where R: AsyncRead + Unpin + Send + 'static, @@ -73,6 +74,7 @@ where let quota_limit = config.access.user_data_quota.get(&user).copied(); let quota_user_stats = quota_limit.map(|_| stats.get_or_create_user_stats_handle(&user)); + let quota_handle = quota_limit.map(|_| quota_handle); let peer = success.peer; let traffic_lease = shared.traffic_limiter.acquire_lease(&user, peer.ip()); let proto_tag = success.proto_tag; @@ -200,6 +202,7 @@ where let rng_clone = rng.clone(); let user_clone = user.clone(); let quota_user_stats_me_writer = quota_user_stats.clone(); + let quota_handle_me_writer = quota_handle.clone(); let traffic_lease_me_writer = traffic_lease.clone(); let flow_cancel_me_writer = flow_cancel.clone(); let last_downstream_activity_ms_clone = last_downstream_activity_ms.clone(); @@ -212,6 +215,7 @@ where rng_clone, user_clone, quota_user_stats_me_writer, + quota_handle_me_writer, quota_limit, traffic_lease_me_writer, flow_cancel_me_writer, @@ -340,11 +344,11 @@ where forensics.bytes_c2me = forensics .bytes_c2me .saturating_add(payload.len() as u64); - if let (Some(limit), Some(user_stats)) = - (quota_limit, quota_user_stats.as_deref()) + if let (Some(limit), Some(quota_handle)) = + (quota_limit, quota_handle.as_ref()) { match reserve_user_quota_with_yield( - user_stats, + quota_handle, payload.len() as u64, limit, stats.as_ref(), @@ -379,7 +383,11 @@ where break; } } - stats.add_user_octets_from_handle(user_stats, payload.len() as u64); + if let Some(user_stats) = quota_user_stats.as_deref() { + stats.add_user_octets_from_handle(user_stats, payload.len() as u64); + } else { + stats.add_user_octets_from(&user, payload.len() as u64); + } } else { stats.add_user_octets_from(&user, payload.len() as u64); } diff --git a/src/proxy/middle_relay/session/tasks.rs b/src/proxy/middle_relay/session/tasks.rs index 56f52dd..b2e747e 100644 --- a/src/proxy/middle_relay/session/tasks.rs +++ b/src/proxy/middle_relay/session/tasks.rs @@ -53,6 +53,7 @@ pub(super) async fn run_me_writer( rng_clone: Arc, user_clone: String, quota_user_stats_me_writer: Option>, + quota_handle_me_writer: Option, quota_limit: Option, traffic_lease_me_writer: Option>, flow_cancel_me_writer: CancellationToken, @@ -105,6 +106,7 @@ where stats_clone.as_ref(), &user_clone, quota_user_stats_me_writer.as_deref(), + quota_handle_me_writer.as_ref(), quota_limit, d2c_flush_policy.quota_soft_overshoot_bytes, traffic_lease_me_writer.as_ref(), @@ -167,6 +169,7 @@ where stats_clone.as_ref(), &user_clone, quota_user_stats_me_writer.as_deref(), + quota_handle_me_writer.as_ref(), quota_limit, d2c_flush_policy.quota_soft_overshoot_bytes, traffic_lease_me_writer.as_ref(), @@ -233,6 +236,7 @@ where stats_clone.as_ref(), &user_clone, quota_user_stats_me_writer.as_deref(), + quota_handle_me_writer.as_ref(), quota_limit, d2c_flush_policy.quota_soft_overshoot_bytes, traffic_lease_me_writer.as_ref(), @@ -304,6 +308,7 @@ where stats_clone.as_ref(), &user_clone, quota_user_stats_me_writer.as_deref(), + quota_handle_me_writer.as_ref(), quota_limit, d2c_flush_policy.quota_soft_overshoot_bytes, traffic_lease_me_writer.as_ref(), diff --git a/src/proxy/relay.rs b/src/proxy/relay.rs index 6695026..0cc8cc5 100644 --- a/src/proxy/relay.rs +++ b/src/proxy/relay.rs @@ -290,6 +290,7 @@ where // ── Combine split halves into bidirectional streams ────────────── let client_combined = CombinedStream::new(client_reader, client_writer); let mut server = CombinedStream::new(server_reader, server_writer); + let quota_handle = stats.current_user_quota_handle(&user_owned); // Wrap client with stats/activity tracking let mut client = StatsIo::new_with_traffic_lease( @@ -297,6 +298,7 @@ where Arc::clone(&counters), Arc::clone(&stats), user_owned.clone(), + quota_handle, traffic_lease, quota_limit, Arc::clone("a_exceeded), diff --git a/src/proxy/relay/adaptive_copy.rs b/src/proxy/relay/adaptive_copy.rs index cae723c..ae6ca2a 100644 --- a/src/proxy/relay/adaptive_copy.rs +++ b/src/proxy/relay/adaptive_copy.rs @@ -19,7 +19,7 @@ use crate::proxy::direct_buffer_budget::{ DIRECT_BASE_C2S_BYTES, DIRECT_BASE_S2C_BYTES, DirectBufferBudget, DirectBufferLease, }; use crate::proxy::traffic_limiter::TrafficLease; -use crate::stats::Stats; +use crate::stats::{Stats, UserQuotaHandle}; use super::WATCHDOG_INTERVAL; use super::io::{SharedCounters, StatsIo, is_quota_io_error}; @@ -141,6 +141,7 @@ pub(crate) async fn relay_direct_adaptive( max_connections: u32, user: &str, stats: Arc, + quota_handle: UserQuotaHandle, quota_limit: Option, traffic_lease: Option>, activity_timeout: Duration, @@ -200,6 +201,7 @@ where Arc::clone(&counters), Arc::clone(&stats), user_owned.clone(), + quota_handle.clone(), traffic_lease.clone(), quota_limit, Arc::clone("a_exceeded), @@ -210,6 +212,7 @@ where Arc::clone(&counters), Arc::clone(&stats), user_owned.clone(), + quota_handle, traffic_lease, quota_limit, Arc::clone("a_exceeded), diff --git a/src/proxy/relay/io.rs b/src/proxy/relay/io.rs index e2d5f78..29171a9 100644 --- a/src/proxy/relay/io.rs +++ b/src/proxy/relay/io.rs @@ -1,5 +1,5 @@ use crate::proxy::traffic_limiter::{RateDirection, TrafficLease, next_refill_delay}; -use crate::stats::{Stats, UserStats}; +use crate::stats::{Stats, UserQuotaHandle, UserStats}; use std::io; use std::pin::Pin; use std::sync::Arc; @@ -40,6 +40,7 @@ pub(super) struct StatsIo { stats: Arc, user: String, user_stats: Arc, + quota_handle: UserQuotaHandle, traffic_lease: Option>, c2s_rate_debt_bytes: u64, c2s_wait: RateWaitState, @@ -71,11 +72,13 @@ impl StatsIo { quota_exceeded: Arc, epoch: Instant, ) -> Self { + let quota_handle = stats.current_user_quota_handle(&user); Self::new_with_traffic_lease( inner, counters, stats, user, + quota_handle, None, quota_limit, quota_exceeded, @@ -88,6 +91,7 @@ impl StatsIo { counters: Arc, stats: Arc, user: String, + quota_handle: UserQuotaHandle, traffic_lease: Option>, quota_limit: Option, quota_exceeded: Arc, @@ -102,6 +106,7 @@ impl StatsIo { stats, user, user_stats, + quota_handle, traffic_lease, c2s_rate_debt_bytes: 0, c2s_wait: RateWaitState::default(), @@ -213,7 +218,7 @@ impl AsyncRead for StatsIo { let mut quota_reservation = None; let mut read_limit = buf.remaining(); if let Some(limit) = this.quota_limit { - let used_before = this.user_stats.quota_used(); + let used_before = this.quota_handle.used(); let remaining = limit.saturating_sub(used_before); if remaining == 0 { this.quota_exceeded.store(true, Ordering::Release); @@ -230,7 +235,7 @@ impl AsyncRead for StatsIo { let mut reserve_rounds = 0usize; while quota_reservation.is_none() { for _ in 0..QUOTA_RESERVE_SPIN_RETRIES { - match this.user_stats.quota_reserve(desired, limit) { + match this.quota_handle.try_reserve(desired, limit) { Ok(reservation) => { quota_reservation = Some(reservation); break; @@ -305,7 +310,7 @@ impl AsyncRead for StatsIo { } } if let Some(limit) = this.quota_limit - && this.user_stats.quota_used() >= limit + && this.quota_handle.used() >= limit { this.quota_exceeded.store(true, Ordering::Release); } @@ -401,7 +406,7 @@ impl AsyncWrite for StatsIo { if !write_buf.is_empty() { let mut reserve_rounds = 0usize; while quota_reservation.is_none() { - let used_before = this.user_stats.quota_used(); + let used_before = this.quota_handle.used(); let remaining = limit.saturating_sub(used_before); if remaining == 0 { this.quota_exceeded.store(true, Ordering::Release); @@ -412,7 +417,7 @@ impl AsyncWrite for StatsIo { let desired = remaining.min(write_buf.len() as u64); let mut saw_contention = false; for _ in 0..QUOTA_RESERVE_SPIN_RETRIES { - match this.user_stats.quota_reserve(desired, limit) { + match this.quota_handle.try_reserve(desired, limit) { Ok(reservation) => { quota_reservation = Some(reservation); write_buf = &write_buf[..desired as usize]; @@ -442,7 +447,7 @@ impl AsyncWrite for StatsIo { } } } else { - let used_before = this.user_stats.quota_used(); + let used_before = this.quota_handle.used(); let remaining = limit.saturating_sub(used_before); if remaining == 0 { this.quota_exceeded.store(true, Ordering::Release); @@ -481,7 +486,7 @@ impl AsyncWrite for StatsIo { if let (Some(limit), Some(remaining)) = (this.quota_limit, remaining_before) { if should_immediate_quota_check(remaining, n_to_charge) { this.quota_bytes_since_check = 0; - if this.user_stats.quota_used() >= limit { + if this.quota_handle.used() >= limit { this.quota_exceeded.store(true, Ordering::Release); } } else { @@ -490,7 +495,7 @@ impl AsyncWrite for StatsIo { let interval = quota_adaptive_interval_bytes(remaining); if this.quota_bytes_since_check >= interval { this.quota_bytes_since_check = 0; - if this.user_stats.quota_used() >= limit { + if this.quota_handle.used() >= limit { this.quota_exceeded.store(true, Ordering::Release); } } diff --git a/src/proxy/shared_state.rs b/src/proxy/shared_state.rs index 2cd9d6e..f19f395 100644 --- a/src/proxy/shared_state.rs +++ b/src/proxy/shared_state.rs @@ -198,15 +198,27 @@ impl ProxySharedState { self.user_admission.apply_config(users, user_enabled) } - /// Applies a candidate user policy only when its captured epoch is current. - pub(crate) fn apply_user_config_if_epoch( + /// Transfers user-policy ownership to one runtime generation. + pub(crate) fn activate_user_config_source( &self, - expected_epoch: u64, + source_generation: u64, + expected_epoch: Option, users: &HashMap, user_enabled: &HashMap, ) -> Option> { self.user_admission - .apply_config_if_epoch(expected_epoch, users, user_enabled) + .activate_config_source(source_generation, expected_epoch, users, user_enabled) + } + + /// Applies an update only from the active runtime generation. + pub(crate) fn apply_user_config_from_source( + &self, + source_generation: u64, + users: &HashMap, + user_enabled: &HashMap, + ) -> Option> { + self.user_admission + .apply_config_from_source(source_generation, users, user_enabled) } /// Applies one persisted user mutation before asynchronous config reload. @@ -219,6 +231,17 @@ impl ProxySharedState { self.user_admission.stage_user(user, secret, enabled) } + /// Applies one prevalidated persisted credential before asynchronous reload. + pub(crate) fn stage_user_credential( + &self, + user: &str, + credential_id: UserCredentialId, + enabled: bool, + ) -> UserMutationResult { + self.user_admission + .stage_user_credential(user, credential_id, enabled) + } + /// Installs a deletion tombstone and cancels every current owner. pub(crate) fn delete_user(&self, user: &str) -> UserMutationResult { self.user_admission.delete_user(user) diff --git a/src/proxy/tests/middle_relay_atomic_quota_invariant_tests.rs b/src/proxy/tests/middle_relay_atomic_quota_invariant_tests.rs index 5cf6ac5..a05c13e 100644 --- a/src/proxy/tests/middle_relay_atomic_quota_invariant_tests.rs +++ b/src/proxy/tests/middle_relay_atomic_quota_invariant_tests.rs @@ -240,6 +240,7 @@ async fn me_writer_data_write_obeys_flow_cancellation() { user, None, None, + None, 0, None, &cancel, diff --git a/src/proxy/traffic_limiter.rs b/src/proxy/traffic_limiter.rs index ce454de..40047af 100644 --- a/src/proxy/traffic_limiter.rs +++ b/src/proxy/traffic_limiter.rs @@ -90,20 +90,20 @@ struct DirectionBucket { struct UserBucket { rates: AtomicRatePair, - up: DirectionBucket, - down: DirectionBucket, + up: Arc, + down: Arc, active_leases: AtomicU64, } #[derive(Default)] struct CidrDirectionBucket { - used: DirectionBucket, - active_users: DirectionBucket, + used: Arc, + active_users: Arc, } #[derive(Default)] struct CidrUserDirectionState { - used: DirectionBucket, + used: Arc, } struct CidrUserShare { @@ -155,17 +155,27 @@ struct ShardedRegistry { mask: usize, } -pub struct TrafficLease { +struct TrafficLeaseBinding { limiter: Arc, + revision: u64, user_bucket: Option>, cidr_bucket: Option>, cidr_user_key: Option, cidr_user_share: Option>, } +pub struct TrafficLease { + limiter: Arc, + user: String, + client_ip: IpAddr, + binding: ArcSwap, + refresh: ParkingMutex<()>, +} + pub struct TrafficLimiter { policy: ArcSwap, policy_update: ParkingMutex<()>, + published_revision: AtomicU64, user_buckets: ShardedRegistry, cidr_buckets: ShardedRegistry, user_scope: ScopeMetrics, @@ -173,17 +183,18 @@ pub struct TrafficLimiter { last_cleanup_epoch_secs: AtomicU64, } -struct DirectionDebit<'a> { - bucket: &'a DirectionBucket, +struct DirectionDebit { + bucket: Arc, epoch: u64, refundable: u64, } /// Refunds uncommitted shaping budget when an I/O attempt is cancelled. #[must_use = "traffic reservations must be settled after the I/O attempt"] -pub(crate) struct TrafficReservation<'a> { +pub(crate) struct TrafficReservation { result: TrafficConsumeResult, - user: Option>, - cidr: Option>, - cidr_user: Option>, + _binding: Arc, + user: Option, + cidr: Option, + cidr_user: Option, } diff --git a/src/proxy/traffic_limiter/buckets.rs b/src/proxy/traffic_limiter/buckets.rs index 748ba53..8213151 100644 --- a/src/proxy/traffic_limiter/buckets.rs +++ b/src/proxy/traffic_limiter/buckets.rs @@ -71,11 +71,11 @@ impl DirectionBucket { } pub(super) fn try_reserve_at( - &self, + self: &Arc, epoch: u64, cap: u64, requested: u64, - ) -> Option> { + ) -> Option { if requested == 0 || cap == 0 || epoch > PACKED_EPOCH_MAX { return None; } @@ -109,7 +109,7 @@ impl DirectionBucket { ) { Ok(_) => { return Some(DirectionDebit { - bucket: self, + bucket: Arc::clone(self), epoch, refundable: grant, }); @@ -144,7 +144,7 @@ impl DirectionBucket { } } -impl DirectionDebit<'_> { +impl DirectionDebit { fn granted(&self) -> u64 { self.refundable } @@ -168,7 +168,7 @@ impl DirectionDebit<'_> { } } -impl Drop for DirectionDebit<'_> { +impl Drop for DirectionDebit { fn drop(&mut self) { self.bucket.refund_at(self.epoch, self.refundable); } @@ -178,8 +178,8 @@ impl UserBucket { pub(super) fn new(revision: u64, limits: RateLimitBps) -> Self { Self { rates: AtomicRatePair::new(revision, limits), - up: DirectionBucket::default(), - down: DirectionBucket::default(), + up: Arc::new(DirectionBucket::default()), + down: Arc::new(DirectionBucket::default()), active_leases: AtomicU64::new(0), } } @@ -192,7 +192,7 @@ impl UserBucket { &self, direction: RateDirection, requested: u64, - ) -> (u64, Option>) { + ) -> (u64, Option) { let cap_bps = self.rates.get(direction); if cap_bps == 0 { return (requested, None); @@ -208,12 +208,12 @@ impl UserBucket { } impl CidrDirectionBucket { - pub(super) fn try_reserve<'a>( - &'a self, - user_state: &'a CidrUserDirectionState, + pub(super) fn try_reserve( + &self, + user_state: &CidrUserDirectionState, cap_epoch: u64, requested: u64, - ) -> (u64, Option>, Option>) { + ) -> (u64, Option, Option) { if requested == 0 || cap_epoch == 0 { return (0, None, None); } @@ -260,7 +260,7 @@ impl CidrDirectionBucket { } impl CidrUserDirectionState { - pub(super) fn ensure_active(&self, epoch: u64, active_users: &DirectionBucket) -> bool { + pub(super) fn ensure_active(&self, epoch: u64, active_users: &Arc) -> bool { if epoch > PACKED_EPOCH_MAX { return false; } @@ -340,12 +340,12 @@ impl CidrBucket { }); } - pub(super) fn try_reserve_for_user<'a>( - &'a self, + pub(super) fn try_reserve_for_user( + &self, direction: RateDirection, - share: &'a CidrUserShare, + share: &CidrUserShare, requested: u64, - ) -> (u64, Option>, Option>) { + ) -> (u64, Option, Option) { let cap_bps = self.rates.get(direction); if cap_bps == 0 { return (requested, None, None); diff --git a/src/proxy/traffic_limiter/lease.rs b/src/proxy/traffic_limiter/lease.rs index 99f3d26..452cdfa 100644 --- a/src/proxy/traffic_limiter/lease.rs +++ b/src/proxy/traffic_limiter/lease.rs @@ -1,12 +1,41 @@ use super::*; impl TrafficLease { + fn current_binding(&self) -> Arc { + let published_revision = self.limiter.published_revision.load(Ordering::Acquire); + let current = self.binding.load_full(); + if current.revision == published_revision { + return current; + } + + let refresh = self.refresh.lock(); + let published_revision = self.limiter.published_revision.load(Ordering::Acquire); + let current = self.binding.load_full(); + if current.revision == published_revision { + return current; + } + let policy_update = self.limiter.policy_update.lock(); + let policy = self.limiter.policy.load_full(); + if current.revision == policy.revision { + return current; + } + let next = self + .limiter + .build_binding(&self.user, self.client_ip, &policy); + self.binding.store(Arc::clone(&next)); + drop(policy_update); + drop(refresh); + self.limiter.maybe_cleanup(); + next + } + /// Reserves shaping budget until the associated I/O result is settled. pub(crate) fn try_reserve( &self, direction: RateDirection, requested: u64, - ) -> TrafficReservation<'_> { + ) -> TrafficReservation { + let binding = self.current_binding(); if requested == 0 { return TrafficReservation { result: TrafficConsumeResult { @@ -14,6 +43,7 @@ impl TrafficLease { blocked_user: false, blocked_cidr: false, }, + _binding: binding, user: None, cidr: None, cidr_user: None, @@ -22,7 +52,7 @@ impl TrafficLease { let mut granted = requested; let mut user_debit = None; - if let Some(user_bucket) = self.user_bucket.as_ref() { + if let Some(user_bucket) = binding.user_bucket.as_ref() { let (user_granted, debit) = user_bucket.try_reserve(direction, granted); user_debit = debit; if user_granted == 0 { @@ -33,6 +63,7 @@ impl TrafficLease { blocked_user: true, blocked_cidr: false, }, + _binding: binding, user: user_debit, cidr: None, cidr_user: None, @@ -44,7 +75,7 @@ impl TrafficLease { let mut cidr_debit = None; let mut cidr_user_debit = None; if let (Some(cidr_bucket), Some(cidr_user_share)) = - (self.cidr_bucket.as_ref(), self.cidr_user_share.as_ref()) + (binding.cidr_bucket.as_ref(), binding.cidr_user_share.as_ref()) { let (cidr_granted, aggregate_debit, share_debit) = cidr_bucket.try_reserve_for_user(direction, cidr_user_share, granted); @@ -63,6 +94,7 @@ impl TrafficLease { blocked_user: false, blocked_cidr: true, }, + _binding: binding, user: user_debit, cidr: cidr_debit, cidr_user: cidr_user_debit, @@ -77,6 +109,7 @@ impl TrafficLease { blocked_user: false, blocked_cidr: false, }, + _binding: binding, user: user_debit, cidr: cidr_debit, cidr_user: cidr_user_debit, @@ -105,7 +138,7 @@ impl TrafficLease { } } -impl TrafficReservation<'_> { +impl TrafficReservation { /// Returns the shaping decision associated with this reservation. pub(crate) fn result(&self) -> TrafficConsumeResult { self.result @@ -126,7 +159,7 @@ impl TrafficReservation<'_> { } } -impl Drop for TrafficLease { +impl Drop for TrafficLeaseBinding { fn drop(&mut self) { if let Some(bucket) = self.user_bucket.as_ref() { decrement_atomic_saturating(&bucket.active_leases, 1); diff --git a/src/proxy/traffic_limiter/limiter.rs b/src/proxy/traffic_limiter/limiter.rs index e52df40..d407602 100644 --- a/src/proxy/traffic_limiter/limiter.rs +++ b/src/proxy/traffic_limiter/limiter.rs @@ -6,6 +6,7 @@ impl TrafficLimiter { Arc::new(Self { policy: ArcSwap::from_pointee(PolicySnapshot::default()), policy_update: ParkingMutex::new(()), + published_revision: AtomicU64::new(0), user_buckets: ShardedRegistry::new(REGISTRY_SHARDS), cidr_buckets: ShardedRegistry::new(REGISTRY_SHARDS), user_scope: ScopeMetrics::default(), @@ -92,6 +93,7 @@ impl TrafficLimiter { cidr_auto_rules_v6, cidr_rule_keys, })); + self.published_revision.store(revision, Ordering::Release); drop(policy_update); self.maybe_cleanup(); @@ -102,7 +104,26 @@ impl TrafficLimiter { user: &str, client_ip: IpAddr, ) -> Option> { + let policy_update = self.policy_update.lock(); let policy = self.policy.load_full(); + let binding = self.build_binding(user, client_ip, &policy); + drop(policy_update); + self.maybe_cleanup(); + Some(Arc::new(TrafficLease { + limiter: Arc::clone(self), + user: user.to_string(), + client_ip, + binding: ArcSwap::from(binding), + refresh: ParkingMutex::new(()), + })) + } + + pub(super) fn build_binding( + self: &Arc, + user: &str, + client_ip: IpAddr, + policy: &PolicySnapshot, + ) -> Arc { let mut user_bucket = None; if let Some(limit) = policy.user_limits.get(user).copied() { let bucket = self.user_buckets.get_or_insert_with( @@ -144,18 +165,14 @@ impl TrafficLimiter { cidr_bucket = Some(bucket); } - if user_bucket.is_none() && cidr_bucket.is_none() { - return None; - } - - self.maybe_cleanup(); - Some(Arc::new(TrafficLease { + Arc::new(TrafficLeaseBinding { limiter: Arc::clone(self), + revision: policy.revision, user_bucket, cidr_bucket, cidr_user_key, cidr_user_share, - })) + }) } pub fn metrics_snapshot(&self) -> TrafficLimiterMetricsSnapshot { diff --git a/src/proxy/traffic_limiter/tests.rs b/src/proxy/traffic_limiter/tests.rs index 8bfe6e3..ad352c1 100644 --- a/src/proxy/traffic_limiter/tests.rs +++ b/src/proxy/traffic_limiter/tests.rs @@ -77,7 +77,7 @@ fn auto_cidr_bucket_key_canonicalizes_network_address() { #[test] fn refund_from_an_old_epoch_does_not_reduce_the_current_epoch() { - let bucket = DirectionBucket::default(); + let bucket = Arc::new(DirectionBucket::default()); let old_debit = bucket.try_reserve_at(7, 100, 80).unwrap(); let current_debit = bucket.try_reserve_at(8, 100, 60).unwrap(); @@ -166,7 +166,7 @@ fn stale_policy_revision_cannot_restore_an_old_rate() { #[test] fn dropped_debit_refunds_only_its_packed_epoch() { - let bucket = DirectionBucket::default(); + let bucket = Arc::new(DirectionBucket::default()); let debit = bucket.try_reserve_at(11, 100, 80).unwrap(); drop(debit); @@ -229,9 +229,10 @@ fn dropped_traffic_reservation_refunds_user_and_cidr_debits() { let epoch = reservation.user.as_ref().unwrap().epoch; drop(reservation); - let user_bucket = lease.user_bucket.as_ref().unwrap(); - let cidr_bucket = lease.cidr_bucket.as_ref().unwrap(); - let cidr_user = lease.cidr_user_share.as_ref().unwrap(); + let binding = lease.binding.load_full(); + let user_bucket = binding.user_bucket.as_ref().unwrap(); + let cidr_bucket = binding.cidr_bucket.as_ref().unwrap(); + let cidr_user = binding.cidr_user_share.as_ref().unwrap(); assert_eq!(user_bucket.down.used_at(epoch), Some(0)); assert_eq!(cidr_bucket.down.used.used_at(epoch), Some(0)); assert_eq!(cidr_user.down.used.used_at(epoch), Some(0)); @@ -252,7 +253,31 @@ fn partial_traffic_settlement_charges_only_committed_bytes() { reservation.settle_written(300); assert_eq!( - lease.user_bucket.as_ref().unwrap().down.used_at(epoch), + lease + .binding + .load_full() + .user_bucket + .as_ref() + .unwrap() + .down + .used_at(epoch), Some(300) ); } + +#[test] +fn active_lease_observes_policy_removal() { + let limiter = TrafficLimiter::new(); + let mut user_limits = HashMap::new(); + user_limits.insert("alice".to_string(), rate(1, 0)); + limiter.apply_policy(user_limits, HashMap::new()); + let lease = limiter + .acquire_lease("alice", "203.0.113.7".parse().unwrap()) + .unwrap(); + assert_eq!(lease.try_consume(RateDirection::Up, 1).granted, 1); + assert_eq!(lease.try_consume(RateDirection::Up, 1).granted, 0); + + limiter.apply_policy(HashMap::new(), HashMap::new()); + + assert_eq!(lease.try_consume(RateDirection::Up, 1).granted, 1); +} diff --git a/src/proxy/user_admission.rs b/src/proxy/user_admission.rs index e4d4139..32a3b3e 100644 --- a/src/proxy/user_admission.rs +++ b/src/proxy/user_admission.rs @@ -6,6 +6,7 @@ use parking_lot::{Mutex, MutexGuard}; use tokio_util::sync::CancellationToken; use crate::crypto::sha256; +use crate::stats::QuotaStore; const REGISTRATION_PENDING: u8 = 0; const REGISTRATION_ACTIVE: u8 = 1; @@ -54,6 +55,8 @@ struct RegisteredOwner { struct UserAdmissionState { initialized: bool, epoch: u64, + active_config_source: Option, + stale_config_source_rejections: u64, next_incarnation: UserIncarnation, next_registration_id: u64, users: HashMap, @@ -97,13 +100,20 @@ pub(crate) struct UserMutationResult { /// Process-owned user authentication and live-owner authority. pub(crate) struct UserAdmissionAuthority { state: Mutex, + quota_store: Arc, } impl UserAdmissionAuthority { /// Creates an uninitialized authority for isolated tests and startup wiring. pub(crate) fn new() -> Arc { + Self::new_with_quota_store(Arc::new(QuotaStore::default())) + } + + /// Creates an authority coupled to the process-scoped quota identity store. + pub(crate) fn new_with_quota_store(quota_store: Arc) -> Arc { Arc::new(Self { state: Mutex::new(UserAdmissionState::default()), + quota_store, }) } @@ -118,23 +128,41 @@ impl UserAdmissionAuthority { users: &HashMap, user_enabled: &HashMap, ) -> Vec<(String, usize)> { - self.apply_config_locked(None, users, user_enabled) + self.activate_config_source(0, None, users, user_enabled) .unwrap_or_default() } - /// Applies a candidate configuration only if no newer authority mutation occurred. - pub(crate) fn apply_config_if_epoch( + /// Transfers configuration ownership to one runtime generation. + pub(crate) fn activate_config_source( &self, - expected_epoch: u64, + source_generation: u64, + expected_epoch: Option, users: &HashMap, user_enabled: &HashMap, ) -> Option> { - self.apply_config_locked(Some(expected_epoch), users, user_enabled) + self.apply_config_locked(source_generation, expected_epoch, true, users, user_enabled) + } + + /// Reconciles an update only while its runtime generation owns configuration. + pub(crate) fn apply_config_from_source( + &self, + source_generation: u64, + users: &HashMap, + user_enabled: &HashMap, + ) -> Option> { + self.apply_config_locked(source_generation, None, false, users, user_enabled) + } + + /// Returns the number of rejected updates from non-owning generations. + pub(crate) fn stale_config_source_rejections(&self) -> u64 { + self.state.lock().stale_config_source_rejections } fn apply_config_locked( &self, + source_generation: u64, expected_epoch: Option, + activate_source: bool, users: &HashMap, user_enabled: &HashMap, ) -> Option> { @@ -154,6 +182,21 @@ impl UserAdmissionAuthority { .collect::>(); let cancellations = { let mut state = self.state.lock(); + if activate_source { + if state + .active_config_source + .is_some_and(|active| source_generation < active) + { + state.stale_config_source_rejections = + state.stale_config_source_rejections.saturating_add(1); + return None; + } + state.active_config_source = Some(source_generation); + } else if state.active_config_source != Some(source_generation) { + state.stale_config_source_rejections = + state.stale_config_source_rejections.saturating_add(1); + return None; + } if expected_epoch.is_some_and(|epoch| state.epoch != epoch) { return None; } @@ -195,6 +238,19 @@ impl UserAdmissionAuthority { if let Some(record) = state.users.get_mut(&user) { record.incarnation = incarnation; } + match (old_effective, new_effective) { + (None, Some(_)) => { + self.quota_store.activate_fresh(&user, incarnation); + } + (Some(_), Some(_)) => { + self.quota_store + .advance_preserving_usage(&user, incarnation); + } + (Some(_), None) => { + self.quota_store.retire_through(&user, incarnation); + } + (None, None) => {} + } } if identity_changed || old_effective.is_some_and(|entry| entry.enabled) @@ -213,13 +269,14 @@ impl UserAdmissionAuthority { changed = true; let incarnation = state.allocate_incarnation(); state.users.insert( - user, + user.clone(), UserRecord { configured: Some(desired), mutation_override: None, incarnation, }, ); + self.quota_store.activate_fresh(&user, incarnation); } if changed { @@ -238,6 +295,16 @@ impl UserAdmissionAuthority { enabled: bool, ) -> Option { let credential_id = credential_id_from_hex(secret)?; + Some(self.stage_user_credential(user, credential_id, enabled)) + } + + /// Applies one already validated credential mutation ahead of runtime reload. + pub(crate) fn stage_user_credential( + &self, + user: &str, + credential_id: UserCredentialId, + enabled: bool, + ) -> UserMutationResult { let desired = EffectiveUser { credential_id, enabled, @@ -262,6 +329,14 @@ impl UserAdmissionAuthority { }); record.mutation_override = Some(UserOverride::Present(desired)); record.incarnation = incarnation; + if identity_changed { + if previous.is_some() { + self.quota_store + .advance_preserving_usage(user, incarnation); + } else { + self.quota_store.activate_fresh(user, incarnation); + } + } state.initialized = true; state.bump_epoch(); let newly_disabled = previous.is_some_and(|entry| entry.enabled) && !enabled; @@ -276,11 +351,11 @@ impl UserAdmissionAuthority { for token in tokens { token.cancel(); } - Some(UserMutationResult { + UserMutationResult { incarnation, cancelled, newly_disabled, - }) + } } /// Installs a deletion tombstone and cancels every owner of the old incarnation. @@ -296,6 +371,7 @@ impl UserAdmissionAuthority { }); record.mutation_override = Some(UserOverride::Deleted); record.incarnation = incarnation; + self.quota_store.retire_through(user, incarnation); state.initialized = true; state.bump_epoch(); ( diff --git a/src/proxy/user_admission/tests.rs b/src/proxy/user_admission/tests.rs index 49bcbb5..0743078 100644 --- a/src/proxy/user_admission/tests.rs +++ b/src/proxy/user_admission/tests.rs @@ -43,18 +43,49 @@ fn stale_credential_cannot_cross_delete_and_recreate() { fn stale_candidate_cannot_overwrite_newer_mutation() { let authority = UserAdmissionAuthority::new(); let secret = "00112233445566778899aabbccddeeff"; - authority.apply_config(&users(secret), &HashMap::new()); + authority + .activate_config_source(1, None, &users(secret), &HashMap::new()) + .unwrap(); let candidate_epoch = authority.epoch(); authority.stage_user("alice", secret, false).unwrap(); assert!( authority - .apply_config_if_epoch(candidate_epoch, &users(secret), &HashMap::new()) + .activate_config_source( + 2, + Some(candidate_epoch), + &users(secret), + &HashMap::new(), + ) .is_none() ); assert!(!authority.is_user_enabled("alice")); } +#[test] +fn stale_generation_snapshot_cannot_reopen_reconciled_user() { + let authority = UserAdmissionAuthority::new(); + let secret = "00112233445566778899aabbccddeeff"; + authority + .activate_config_source(1, None, &users(secret), &HashMap::new()) + .unwrap(); + authority.stage_user("alice", secret, false).unwrap(); + + let disabled = HashMap::from([("alice".to_string(), false)]); + let epoch = authority.epoch(); + authority + .activate_config_source(2, Some(epoch), &users(secret), &disabled) + .unwrap(); + assert!( + authority + .apply_config_from_source(1, &users(secret), &HashMap::new()) + .is_none() + ); + + assert!(!authority.is_user_enabled("alice")); + assert_eq!(authority.stale_config_source_rejections(), 1); +} + #[test] fn registration_dropped_before_publication_cannot_leave_an_owner() { let authority = UserAdmissionAuthority::new(); @@ -71,3 +102,36 @@ fn registration_dropped_before_publication_cannot_leave_an_owner() { assert_eq!(authority.cancel_user_owners("alice"), 0); } + +#[test] +fn quota_identity_follows_credential_rotation_and_recreation() { + let quota_store = Arc::new(QuotaStore::default()); + let authority = UserAdmissionAuthority::new_with_quota_store(quota_store.clone()); + let old_secret = "00112233445566778899aabbccddeeff"; + let new_secret = "ffeeddccbbaa99887766554433221100"; + authority.apply_config(&users(old_secret), &HashMap::new()); + let old_incarnation = authority + .authenticated_incarnation("alice", credential_id_from_hex(old_secret).unwrap()) + .unwrap(); + let old_quota = quota_store + .handle_exact("alice", old_incarnation) + .unwrap(); + old_quota.charge(40); + + let rotated = authority.stage_user("alice", new_secret, true).unwrap(); + let rotated_quota = quota_store + .handle_exact("alice", rotated.incarnation) + .unwrap(); + old_quota.charge(20); + + assert_eq!(rotated_quota.used(), 40); + authority.delete_user("alice"); + let recreated = authority.stage_user("alice", old_secret, true).unwrap(); + assert_eq!( + quota_store + .handle_exact("alice", recreated.incarnation) + .unwrap() + .used(), + 0 + ); +} diff --git a/src/quota_state.rs b/src/quota_state.rs index 5c65b62..1516b20 100644 --- a/src/quota_state.rs +++ b/src/quota_state.rs @@ -116,22 +116,19 @@ impl QuotaStateOwner { wait_for_blocking_io(task).await } - /// Removes a deleted user's persisted and in-memory quota ownership. + /// Removes a deleted user's persisted quota checkpoint. pub(crate) async fn remove_user( &self, configured_users: &BTreeSet, user: &str, ) -> std::io::Result<()> { let guard = Arc::clone(&self.mutation).lock_owned().await; + debug_assert!(!configured_users.contains(user)); let state = self.state_for_users(configured_users, None); let path = self.path.clone(); - let store = Arc::clone(&self.store); - let user = user.to_string(); let task = tokio::task::spawn_blocking(move || { let _guard = guard; - let persisted = write_state_file_blocking(&path, &state); - store.remove(&user); - persisted + write_state_file_blocking(&path, &state) }); wait_for_blocking_io(task).await } diff --git a/src/stats/mod.rs b/src/stats/mod.rs index 90a1c8f..f237e29 100644 --- a/src/stats/mod.rs +++ b/src/stats/mod.rs @@ -21,7 +21,7 @@ use std::sync::Arc; use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, AtomicUsize, Ordering}; use std::time::Instant; -pub(crate) use self::quota_store::{QuotaReservation, QuotaStore}; +pub(crate) use self::quota_store::{QuotaReservation, QuotaStore, UserQuotaHandle}; #[allow(unused_imports)] pub use self::replay::{ReplayChecker, ReplayStats}; use self::telemetry::TelemetryPolicy; @@ -430,6 +430,11 @@ impl Stats { *stats.start_time.write() = Some(Instant::now()); stats } + + #[cfg(test)] + pub(crate) fn quota_store(&self) -> Arc { + Arc::clone(&self.quota_store) + } } #[cfg(test)] diff --git a/src/stats/quota_store.rs b/src/stats/quota_store.rs index 5911fa1..bf1d069 100644 --- a/src/stats/quota_store.rs +++ b/src/stats/quota_store.rs @@ -4,13 +4,38 @@ use std::sync::atomic::{AtomicU64, Ordering}; use arc_swap::ArcSwap; use dashmap::DashMap; +use parking_lot::Mutex; use super::{QuotaReserveError, UserQuotaSnapshot}; +use crate::proxy::user_admission::UserIncarnation; /// Process-scoped per-user quota accounting shared by runtime generations. #[derive(Default)] pub struct QuotaStore { - users: DashMap>, + users: DashMap>, +} + +struct QuotaUserSlot { + state: Mutex, +} + +#[derive(Default)] +struct QuotaSlotState { + high_water: UserIncarnation, + current: Option, + startup_seed: Option, +} + +struct QuotaAccount { + incarnation: UserIncarnation, + counters: Arc, +} + +/// Exact quota ownership pinned to one authenticated user incarnation. +#[derive(Clone)] +pub(crate) struct UserQuotaHandle { + incarnation: UserIncarnation, + counters: Arc, } /// Atomically replaceable quota state for one configured user. @@ -32,30 +57,164 @@ pub(crate) struct QuotaReservation { } impl QuotaStore { - pub(crate) fn user(&self, user: &str) -> Arc { + fn slot(&self, user: &str) -> Arc { if let Some(existing) = self.users.get(user) { return Arc::clone(existing.value()); } Arc::clone( self.users .entry(user.to_string()) - .or_insert_with(|| Arc::new(UserQuotaCounters::default())) + .or_insert_with(|| { + Arc::new(QuotaUserSlot { + state: Mutex::new(QuotaSlotState::default()), + }) + }) .value(), ) } + pub(crate) fn current_or_legacy_handle(&self, user: &str) -> UserQuotaHandle { + let slot = self.slot(user); + let mut state = slot.state.lock(); + if let Some(account) = &state.current { + return UserQuotaHandle { + incarnation: account.incarnation, + counters: Arc::clone(&account.counters), + }; + } + let seed = state.startup_seed.take().unwrap_or(UserQuotaSnapshot { + used_bytes: 0, + last_reset_epoch_secs: 0, + }); + let counters = Arc::new(UserQuotaCounters::from_snapshot(&seed)); + let incarnation = state.high_water; + state.current = Some(QuotaAccount { + incarnation, + counters: Arc::clone(&counters), + }); + UserQuotaHandle { + incarnation, + counters, + } + } + + pub(crate) fn user(&self, user: &str) -> Arc { + self.current_or_legacy_handle(user).counters + } + + /// Returns the quota account owned by the exact current incarnation. + pub(crate) fn handle_exact( + &self, + user: &str, + incarnation: UserIncarnation, + ) -> Option { + let slot = self.users.get(user)?; + let state = slot.state.lock(); + let account = state.current.as_ref()?; + (account.incarnation == incarnation).then(|| UserQuotaHandle { + incarnation, + counters: Arc::clone(&account.counters), + }) + } + + /// Creates a quota account for a new username lifetime without inheriting a retired account. + pub(crate) fn activate_fresh(&self, user: &str, incarnation: UserIncarnation) { + let slot = self.slot(user); + let mut state = slot.state.lock(); + if incarnation <= state.high_water { + return; + } + let snapshot = if state.high_water == 0 { + state + .current + .as_ref() + .map(|account| account.counters.snapshot()) + .or_else(|| state.startup_seed.take()) + } else { + state.startup_seed = None; + None + } + .unwrap_or(UserQuotaSnapshot { + used_bytes: 0, + last_reset_epoch_secs: 0, + }); + state.high_water = state.high_water.max(incarnation); + state.current = Some(QuotaAccount { + incarnation, + counters: Arc::new(UserQuotaCounters::from_snapshot(&snapshot)), + }); + } + + /// Advances a credential incarnation while preserving usage captured at the transition. + pub(crate) fn advance_preserving_usage( + &self, + user: &str, + incarnation: UserIncarnation, + ) { + let slot = self.slot(user); + let mut state = slot.state.lock(); + if incarnation <= state.high_water { + return; + } + let snapshot = state + .current + .as_ref() + .map(|account| account.counters.snapshot()) + .or_else(|| state.startup_seed.take()) + .unwrap_or(UserQuotaSnapshot { + used_bytes: 0, + last_reset_epoch_secs: 0, + }); + state.high_water = incarnation; + state.current = Some(QuotaAccount { + incarnation, + counters: Arc::new(UserQuotaCounters::from_snapshot(&snapshot)), + }); + } + + /// Retires quota ownership without affecting a newer incarnation. + pub(crate) fn retire_through(&self, user: &str, incarnation: UserIncarnation) { + let slot = self.slot(user); + let mut state = slot.state.lock(); + if incarnation < state.high_water { + return; + } + state.high_water = incarnation; + state.current = None; + state.startup_seed = None; + } + pub(crate) fn used(&self, user: &str) -> u64 { - self.users.get(user).map(|state| state.used()).unwrap_or(0) + self.users + .get(user) + .and_then(|slot| { + let state = slot.state.lock(); + state + .current + .as_ref() + .map(|account| account.counters.used()) + .or_else(|| state.startup_seed.as_ref().map(|seed| seed.used_bytes)) + }) + .unwrap_or(0) } pub(crate) fn load(&self, user: &str, used_bytes: u64, last_reset_epoch_secs: u64) { - let state = self.user(user); - state.replace(used_bytes, last_reset_epoch_secs); + let slot = self.slot(user); + let mut state = slot.state.lock(); + let snapshot = UserQuotaSnapshot { + used_bytes, + last_reset_epoch_secs, + }; + if let Some(account) = &state.current { + account.replace_from_snapshot(&snapshot); + } else { + state.startup_seed = Some(snapshot); + } } pub(crate) fn reset(&self, user: &str, now_epoch_secs: u64) -> UserQuotaSnapshot { - let state = self.user(user); - state.replace(0, now_epoch_secs); + let state = self.current_or_legacy_handle(user); + state.counters.replace(0, now_epoch_secs); UserQuotaSnapshot { used_bytes: 0, last_reset_epoch_secs: now_epoch_secs, @@ -63,26 +222,28 @@ impl QuotaStore { } pub(crate) fn remove(&self, user: &str) { - self.users.remove(user); + let slot = self.slot(user); + let mut state = slot.state.lock(); + state.high_water = state.high_water.saturating_add(1); + state.current = None; + state.startup_seed = None; } pub(crate) fn snapshot(&self) -> HashMap { let mut out = HashMap::new(); for entry in self.users.iter() { - let state = entry.value(); - let generation = state.generation.load_full(); - let used_bytes = generation.used_bytes.load(Ordering::Relaxed); - let last_reset_epoch_secs = generation.last_reset_epoch_secs; - if used_bytes == 0 && last_reset_epoch_secs == 0 { + let state = entry.value().state.lock(); + let snapshot = if let Some(account) = state.current.as_ref() { + account.counters.snapshot() + } else if let Some(seed) = state.startup_seed.as_ref() { + seed.clone() + } else { + continue; + }; + if snapshot.used_bytes == 0 && snapshot.last_reset_epoch_secs == 0 { continue; } - out.insert( - entry.key().clone(), - UserQuotaSnapshot { - used_bytes, - last_reset_epoch_secs, - }, - ); + out.insert(entry.key().clone(), snapshot); } out } @@ -100,6 +261,23 @@ impl Default for UserQuotaCounters { } impl UserQuotaCounters { + fn from_snapshot(snapshot: &UserQuotaSnapshot) -> Self { + Self { + generation: ArcSwap::from_pointee(QuotaGeneration { + used_bytes: AtomicU64::new(snapshot.used_bytes), + last_reset_epoch_secs: snapshot.last_reset_epoch_secs, + }), + } + } + + fn snapshot(&self) -> UserQuotaSnapshot { + let generation = self.generation.load_full(); + UserQuotaSnapshot { + used_bytes: generation.used_bytes.load(Ordering::Relaxed), + last_reset_epoch_secs: generation.last_reset_epoch_secs, + } + } + fn replace(&self, used_bytes: u64, last_reset_epoch_secs: u64) { self.generation.store(Arc::new(QuotaGeneration { used_bytes: AtomicU64::new(used_bytes), @@ -150,6 +328,39 @@ impl UserQuotaCounters { } } +impl QuotaAccount { + fn replace_from_snapshot(&self, snapshot: &UserQuotaSnapshot) { + self.counters + .replace(snapshot.used_bytes, snapshot.last_reset_epoch_secs); + } +} + +impl UserQuotaHandle { + /// Returns the immutable incarnation owned by this handle. + pub(crate) fn incarnation(&self) -> UserIncarnation { + self.incarnation + } + + #[inline] + pub(crate) fn used(&self) -> u64 { + self.counters.used() + } + + #[inline] + pub(crate) fn charge(&self, bytes: u64) -> u64 { + self.counters.charge(bytes) + } + + #[inline] + pub(crate) fn try_reserve( + &self, + bytes: u64, + limit: u64, + ) -> Result { + self.counters.try_reserve(bytes, limit) + } +} + impl QuotaReservation { /// Returns the number of bytes held by this reservation. pub(crate) fn reserved_bytes(&self) -> u64 { @@ -255,4 +466,63 @@ mod tests { store.reset("alice", generation); } } + + #[test] + fn retired_incarnation_cannot_charge_recreated_username() { + let store = QuotaStore::default(); + store.activate_fresh("alice", 1); + let retired = store.handle_exact("alice", 1).unwrap(); + retired.charge(40); + store.retire_through("alice", 2); + store.activate_fresh("alice", 3); + + retired.charge(20); + + assert_eq!(retired.used(), 60); + assert_eq!(store.handle_exact("alice", 3).unwrap().used(), 0); + } + + #[test] + fn credential_rotation_preserves_usage_without_sharing_future_charges() { + let store = QuotaStore::default(); + store.activate_fresh("alice", 1); + let old = store.handle_exact("alice", 1).unwrap(); + old.charge(40); + store.advance_preserving_usage("alice", 2); + let current = store.handle_exact("alice", 2).unwrap(); + + old.charge(20); + + assert_eq!(old.used(), 60); + assert_eq!(current.used(), 40); + } + + #[test] + fn stale_retirement_cannot_remove_newer_quota_owner() { + let store = QuotaStore::default(); + store.activate_fresh("alice", 1); + store.retire_through("alice", 2); + store.activate_fresh("alice", 3); + + store.retire_through("alice", 2); + + assert!(store.handle_exact("alice", 3).is_some()); + } + + #[test] + fn old_reservation_refund_does_not_debit_recreated_username() { + let store = QuotaStore::default(); + store.activate_fresh("alice", 1); + let old = store.handle_exact("alice", 1).unwrap(); + let reservation = old.try_reserve(80, 100).unwrap(); + store.retire_through("alice", 2); + store.activate_fresh("alice", 3); + let current = store.handle_exact("alice", 3).unwrap(); + current.charge(50); + + drop(reservation); + + assert_eq!(old.used(), 0); + assert_eq!(current.used(), 50); + } } diff --git a/src/stats/replay.rs b/src/stats/replay.rs index 13b0d8f..a8cc204 100644 --- a/src/stats/replay.rs +++ b/src/stats/replay.rs @@ -73,6 +73,7 @@ pub struct ReplayChecker { checks: AtomicU64, hits: AtomicU64, additions: AtomicU64, + capacity_rejections: AtomicU64, cleanups: AtomicU64, next_claim_token: AtomicU64, } @@ -141,21 +142,25 @@ impl ReplayShard { self.cache.get(key).is_some() || self.pending.contains_key(key) } - fn add_owned(&mut self, key: ReplayKey, now: Instant, window: Duration) { + fn add_owned(&mut self, key: ReplayKey, now: Instant, window: Duration) -> bool { if window.is_zero() { - return; + return true; } self.cleanup(now, window); if self.cache.peek(key.as_slice()).is_some() || self.pending.contains_key(key.as_slice()) { - return; + return true; } - while self.queue.len() >= self.capacity { + while self.cache.len().saturating_add(self.pending.len()) >= self.capacity { + if self.queue.is_empty() { + return false; + } self.evict_queue_front(); } let seq = self.next_seq(); self.cache.put(key.clone(), ReplayEntry { seq }); self.queue.push_back((now, key, seq)); + true } fn claim_owned( @@ -218,7 +223,12 @@ impl TlsReplayClaim<'_> { if !shard.remove_pending(key.as_slice(), self.token) { return false; } - shard.add_owned(key, Instant::now(), self.checker.tls_window); + if !shard.add_owned(key, Instant::now(), self.checker.tls_window) { + self.checker + .capacity_rejections + .fetch_add(1, Ordering::Relaxed); + return false; + } self.checker.additions.fetch_add(1, Ordering::Relaxed); self.reserved = false; true @@ -261,6 +271,7 @@ impl ReplayChecker { checks: AtomicU64::new(0), hits: AtomicU64::new(0), additions: AtomicU64::new(0), + capacity_rejections: AtomicU64::new(0), cleanups: AtomicU64::new(0), next_claim_token: AtomicU64::new(1), } @@ -304,11 +315,15 @@ impl ReplayChecker { let found = shard.check(data, now, window); if found { self.hits.fetch_add(1, Ordering::Relaxed); - } else { - shard.add_owned(owned_key, now, window); - self.additions.fetch_add(1, Ordering::Relaxed); + return true; + } + if shard.add_owned(owned_key, now, window) { + self.additions.fetch_add(1, Ordering::Relaxed); + false + } else { + self.capacity_rejections.fetch_add(1, Ordering::Relaxed); + true } - found } fn check_only_internal( @@ -328,11 +343,14 @@ impl ReplayChecker { } fn add_only(&self, data: &[u8], shards: &[Mutex], window: Duration) { - self.additions.fetch_add(1, Ordering::Relaxed); let idx = self.get_shard_idx(data); let owned_key = ReplayKey::from_slice(data); let mut shard = shards[idx].lock(); - shard.add_owned(owned_key, Instant::now(), window); + if shard.add_owned(owned_key, Instant::now(), window) { + self.additions.fetch_add(1, Ordering::Relaxed); + } else { + self.capacity_rejections.fetch_add(1, Ordering::Relaxed); + } } pub fn check_and_add_handshake(&self, data: &[u8]) -> bool { @@ -394,12 +412,12 @@ impl ReplayChecker { let mut total_queue_len = 0; for shard in &self.handshake_shards { let s = shard.lock(); - total_entries += s.cache.len(); + total_entries += s.len(); total_queue_len += s.queue.len(); } for shard in &self.tls_shards { let s = shard.lock(); - total_entries += s.cache.len(); + total_entries += s.len(); total_queue_len += s.queue.len(); } @@ -409,6 +427,7 @@ impl ReplayChecker { total_checks: self.checks.load(Ordering::Relaxed), total_hits: self.hits.load(Ordering::Relaxed), total_additions: self.additions.load(Ordering::Relaxed), + total_capacity_rejections: self.capacity_rejections.load(Ordering::Relaxed), total_cleanups: self.cleanups.load(Ordering::Relaxed), num_shards: self.handshake_shards.len() + self.tls_shards.len(), window_secs: self.window.as_secs(), @@ -459,6 +478,7 @@ pub struct ReplayStats { pub total_checks: u64, pub total_hits: u64, pub total_additions: u64, + pub total_capacity_rejections: u64, pub total_cleanups: u64, pub num_shards: usize, pub window_secs: u64, @@ -481,3 +501,24 @@ impl ReplayStats { } } } + +#[cfg(test)] +mod capacity_tests { + use super::*; + + #[test] + fn committed_and_pending_entries_share_one_shard_capacity() { + let capacity = NonZeroUsize::new(2).unwrap(); + let mut shard = ReplayShard::new(capacity); + let now = Instant::now(); + let window = Duration::from_secs(60); + assert!(shard.claim_owned(ReplayKey::from_slice(b"pending-a"), now, window, 1)); + assert!(shard.claim_owned(ReplayKey::from_slice(b"pending-b"), now, window, 2)); + + shard.add_owned(ReplayKey::from_slice(b"committed"), now, window); + + assert!(shard.len() <= capacity.get()); + assert!(shard.pending.contains_key(b"pending-a".as_slice())); + assert!(shard.pending.contains_key(b"pending-b".as_slice())); + } +} diff --git a/src/stats/users.rs b/src/stats/users.rs index 2c9d4af..982c2c1 100644 --- a/src/stats/users.rs +++ b/src/stats/users.rs @@ -122,6 +122,27 @@ impl Stats { self.quota_store.used(user) } + /// Returns quota ownership for the exact authenticated user incarnation. + pub(crate) fn quota_handle_for_incarnation( + &self, + user: &str, + incarnation: crate::proxy::user_admission::UserIncarnation, + ) -> Option { + if let Some(handle) = self.quota_store.handle_exact(user, incarnation) { + return Some(handle); + } + if incarnation != 0 { + return None; + } + let handle = self.quota_store.current_or_legacy_handle(user); + (handle.incarnation() == incarnation).then_some(handle) + } + + /// Returns the currently published quota owner for compatibility relay entrypoints. + pub(crate) fn current_user_quota_handle(&self, user: &str) -> UserQuotaHandle { + self.quota_store.current_or_legacy_handle(user) + } + pub fn load_user_quota_state(&self, user: &str, used_bytes: u64, last_reset_epoch_secs: u64) { self.quota_store .load(user, used_bytes, last_reset_epoch_secs); diff --git a/src/transport/middle_proxy/pool/writer_admission.rs b/src/transport/middle_proxy/pool/writer_admission.rs index 91acf18..b11673c 100644 --- a/src/transport/middle_proxy/pool/writer_admission.rs +++ b/src/transport/middle_proxy/pool/writer_admission.rs @@ -100,66 +100,64 @@ impl MePool { contour: WriterContour, intent: WriterOpenIntent, writer_dc: i32, + family: IpFamily, ) -> bool { if intent == WriterOpenIntent::Replacement { return true; } let (active_writers, warm_writers, _) = self.non_draining_writer_counts_by_contour().await; - match contour { - WriterContour::Active => { - let active_cap = self.adaptive_floor_active_cap_configured_total(); - if active_writers < active_cap { - return true; - } - if intent != WriterOpenIntent::Coverage { - return false; - } - - let mut endpoints_len = 0; - let now_epoch = Self::now_epoch_secs(); - let endpoint_snapshot = self.endpoint_snapshot.load(); - if self.family_enabled_for_drain_coverage(IpFamily::V4, now_epoch) { - if let Some(addrs) = endpoint_snapshot.map_v4.get(&writer_dc) { - endpoints_len += addrs.len(); - } - } - if self.family_enabled_for_drain_coverage(IpFamily::V6, now_epoch) { - if let Some(addrs) = endpoint_snapshot.map_v6.get(&writer_dc) { - endpoints_len += addrs.len(); - } - } - - if endpoints_len > 0 { - let base_req = - self.required_writers_for_dc_with_floor_mode(endpoints_len, false); - let active_generation = self.reinit.status.load().active_generation; - let active_for_dc = { - let ws = self.writers.read().await; - ws.iter() - .filter(|w| { - !w.draining.load(std::sync::atomic::Ordering::Relaxed) - && w.writer_dc == writer_dc - && w.generation == active_generation - && matches!( - WriterContour::from_u8( - w.contour.load(std::sync::atomic::Ordering::Relaxed), - ), - WriterContour::Active - ) - }) - .count() - }; - if active_for_dc < base_req { - return true; - } - } - - let coverage_required = self.active_coverage_required_total().await; - active_writers < coverage_required - } - WriterContour::Warm => warm_writers < self.adaptive_floor_warm_cap_configured_total(), - WriterContour::Draining => true, + let live = match contour { + WriterContour::Active => active_writers, + WriterContour::Warm => warm_writers, + WriterContour::Draining => return true, + }; + let configured_cap = match contour { + WriterContour::Active => self.adaptive_floor_active_cap_configured_total(), + WriterContour::Warm => self.adaptive_floor_warm_cap_configured_total(), + WriterContour::Draining => usize::MAX, + }; + if live < configured_cap { + return true; } + if intent != WriterOpenIntent::Coverage { + return false; + } + + let endpoint_snapshot = self.endpoint_snapshot.load(); + let endpoint_count = match family { + IpFamily::V4 => endpoint_snapshot.map_v4.get(&writer_dc), + IpFamily::V6 => endpoint_snapshot.map_v6.get(&writer_dc), + } + .map(Vec::len) + .unwrap_or(0); + if endpoint_count == 0 { + return false; + } + let required = self.required_writers_for_dc_with_floor_mode(endpoint_count, false); + let status = self.reinit.status.load(); + let generation = match contour { + WriterContour::Active => status.active_generation, + WriterContour::Warm => status.pending_hardswap_generation, + WriterContour::Draining => 0, + }; + let family_count = { + let writers = self.writers.read().await; + writers + .iter() + .filter(|writer| { + !writer.draining.load(Ordering::Relaxed) + && writer.writer_dc == writer_dc + && writer.generation == generation + && WriterContour::from_u8(writer.contour.load(Ordering::Relaxed)) == contour + && writer.addr.is_ipv4() == (family == IpFamily::V4) + }) + .count() + }; + if family_count < required { + return true; + } + + live < self.active_coverage_required_total().await } /// Reserves bounded transient capacity for a writer open attempt. @@ -168,6 +166,7 @@ impl MePool { contour: WriterContour, intent: WriterOpenIntent, writer_dc: i32, + family: IpFamily, ) -> Option> { let counter = match contour { WriterContour::Active => &self.writer_connect_active_reserved, @@ -222,7 +221,7 @@ impl MePool { loop { if !self - .can_open_writer_for_contour(contour, intent, writer_dc) + .can_open_writer_for_contour(contour, intent, writer_dc, family) .await { return None; @@ -239,7 +238,9 @@ impl MePool { WriterContour::Warm => self.adaptive_floor_warm_cap_configured_total(), WriterContour::Draining => usize::MAX, }; - if contour == WriterContour::Active && intent == WriterOpenIntent::Coverage { + if intent == WriterOpenIntent::Coverage + && matches!(contour, WriterContour::Active | WriterContour::Warm) + { limit = limit .max(self.active_coverage_required_total().await) .saturating_add( diff --git a/src/transport/middle_proxy/pool_reinit.rs b/src/transport/middle_proxy/pool_reinit.rs index cfaabf9..9c9ff07 100644 --- a/src/transport/middle_proxy/pool_reinit.rs +++ b/src/transport/middle_proxy/pool_reinit.rs @@ -56,20 +56,35 @@ struct ReinitReservation { struct ReinitCommitOutcome { coverage_ratio: f32, missing_dc: Vec, + missing_groups: Vec, stale_writer_ids: Vec, force_close_writer_ids: Vec, } +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] +struct DcFamilyGroup { + dc: i32, + family: IpFamily, +} + +struct HardswapCoverage { + ratio: f32, + missing_groups: Vec, + writer_deficit: usize, +} + #[derive(Debug)] enum ReinitCommitFailure { Superseded, Coverage { coverage_ratio: f32, missing_dc: Vec, + missing_groups: Vec, }, Redundancy { coverage_ratio: f32, missing_dc: Vec, + missing_groups: Vec, }, } diff --git a/src/transport/middle_proxy/pool_reinit/coordination.rs b/src/transport/middle_proxy/pool_reinit/coordination.rs index ab8e3c6..736c798 100644 --- a/src/transport/middle_proxy/pool_reinit/coordination.rs +++ b/src/transport/middle_proxy/pool_reinit/coordination.rs @@ -138,22 +138,35 @@ impl MePool { } }) .map(|writer| (writer.writer_dc, writer.addr)) - .collect::>(); - let (coverage_ratio, missing_dc) = - Self::coverage_ratio(desired_by_dc, &authoritative_writer_addrs); + .collect::>(); + let (coverage_ratio, missing_dc, missing_groups) = if attempt.hardswap { + let coverage = self.hardswap_coverage(desired_by_dc, &authoritative_writer_addrs); + let missing_dc = Self::missing_group_dcs(&coverage.missing_groups); + (coverage.ratio, missing_dc, coverage.missing_groups) + } else { + let authoritative_writer_addrs = authoritative_writer_addrs + .iter() + .copied() + .collect::>(); + let (coverage_ratio, missing_dc) = + Self::coverage_ratio(desired_by_dc, &authoritative_writer_addrs); + (coverage_ratio, missing_dc, Vec::new()) + }; if coverage_ratio < min_ratio { return Err(ReinitCommitFailure::Coverage { coverage_ratio, missing_dc, + missing_groups, }); } if attempt.hardswap - && !missing_dc.is_empty() + && !missing_groups.is_empty() && self.bind_stale_mode() == MeBindStaleMode::Never { return Err(ReinitCommitFailure::Redundancy { coverage_ratio, missing_dc, + missing_groups, }); } if !commit_reinit_state( @@ -183,7 +196,7 @@ impl MePool { .iter() .flat_map(|(dc, endpoints)| endpoints.iter().copied().map(|addr| (*dc, addr))) .collect::>(); - let missing_dc_set = missing_dc.iter().copied().collect::>(); + let missing_group_set = missing_groups.iter().copied().collect::>(); let mut stale_writer_ids = Vec::::new(); let mut force_close_writer_ids = Vec::::new(); for writer in writers.iter() { @@ -199,8 +212,15 @@ impl MePool { continue; } - let preserve_fallback = attempt.hardswap - && missing_dc_set.contains(&writer.writer_dc); + let writer_group = DcFamilyGroup { + dc: writer.writer_dc, + family: if writer.addr.is_ipv4() { + IpFamily::V4 + } else { + IpFamily::V6 + }, + }; + let preserve_fallback = attempt.hardswap && missing_group_set.contains(&writer_group); if !preserve_fallback && attempt.hardswap { registry_registration.retire(writer.id); } @@ -225,6 +245,7 @@ impl MePool { Ok(ReinitCommitOutcome { coverage_ratio, missing_dc, + missing_groups, stale_writer_ids, force_close_writer_ids, }) @@ -265,6 +286,65 @@ impl MePool { (ratio, missing_dc) } + /// Evaluates full writer-floor coverage independently for every DC and address family. + pub(super) fn hardswap_coverage( + &self, + desired_by_dc: &HashMap>, + writer_addrs: &[(i32, SocketAddr)], + ) -> HardswapCoverage { + let mut covered = 0usize; + let mut total = 0usize; + let mut writer_deficit = 0usize; + let mut missing_groups = Vec::new(); + for (dc, endpoints) in desired_by_dc { + for family in [IpFamily::V4, IpFamily::V6] { + let endpoint_count = endpoints + .iter() + .filter(|endpoint| endpoint.is_ipv4() == (family == IpFamily::V4)) + .count(); + if endpoint_count == 0 { + continue; + } + total = total.saturating_add(1); + let required = self.required_writers_for_dc(endpoint_count); + let alive = writer_addrs + .iter() + .filter(|(writer_dc, endpoint)| { + *writer_dc == *dc + && endpoint.is_ipv4() == (family == IpFamily::V4) + && endpoints.contains(endpoint) + }) + .count(); + if alive >= required { + covered = covered.saturating_add(1); + } else { + writer_deficit = + writer_deficit.saturating_add(required.saturating_sub(alive)); + missing_groups.push(DcFamilyGroup { dc: *dc, family }); + } + } + } + missing_groups.sort_unstable_by_key(|group| { + (group.dc, matches!(group.family, IpFamily::V6)) + }); + HardswapCoverage { + ratio: if total == 0 { + 1.0 + } else { + (covered as f32) / (total as f32) + }, + missing_groups, + writer_deficit, + } + } + + fn missing_group_dcs(groups: &[DcFamilyGroup]) -> Vec { + let mut dcs = groups.iter().map(|group| group.dc).collect::>(); + dcs.sort_unstable(); + dcs.dedup(); + dcs + } + /// Restores at least one active writer for every enabled desired DC group. pub async fn reconcile_connections(self: &Arc, rng: &SecureRandom) { let endpoint_snapshot = self.endpoint_snapshot.load_full(); diff --git a/src/transport/middle_proxy/pool_reinit/reconcile.rs b/src/transport/middle_proxy/pool_reinit/reconcile.rs index 8eac110..6ac3d1a 100644 --- a/src/transport/middle_proxy/pool_reinit/reconcile.rs +++ b/src/transport/middle_proxy/pool_reinit/reconcile.rs @@ -15,102 +15,118 @@ impl MePool { let total_passes = 1 + extra_passes; for (dc, endpoints) in desired_by_dc { - if endpoints.is_empty() { - continue; - } - - let mut endpoint_list: Vec = endpoints.iter().copied().collect(); - endpoint_list.sort_unstable(); - let required = self.required_writers_for_dc(endpoint_list.len()); - let mut completed = false; - let mut last_fresh_count = self - .fresh_writer_count_for_dc_endpoints(generation, *dc, endpoints) - .await; - - for pass_idx in 0..total_passes { - if last_fresh_count >= required { - completed = true; - break; + for family in [IpFamily::V4, IpFamily::V6] { + let family_endpoints = endpoints + .iter() + .copied() + .filter(|endpoint| endpoint.is_ipv4() == (family == IpFamily::V4)) + .collect::>(); + if family_endpoints.is_empty() { + continue; } - let missing = required.saturating_sub(last_fresh_count); - debug!( - dc = *dc, - pass = pass_idx + 1, - total_passes, - fresh_count = last_fresh_count, - required, - missing, - endpoint_count = endpoint_list.len(), - "ME hardswap warmup pass started" - ); + let mut endpoint_list = family_endpoints.iter().copied().collect::>(); + endpoint_list.sort_unstable(); + let required = self.required_writers_for_dc(endpoint_list.len()); + let mut completed = false; + let mut last_fresh_count = self + .fresh_writer_count_for_dc_endpoints(generation, *dc, &family_endpoints) + .await; - for attempt_idx in 0..missing { - let delay_ms = self.hardswap_warmup_connect_delay_ms(); - tokio::time::sleep(Duration::from_millis(delay_ms)).await; + for pass_idx in 0..total_passes { + if last_fresh_count >= required { + completed = true; + break; + } - let connected = self - .connect_endpoints_round_robin_with_generation_contour( - *dc, - &endpoint_list, - rng, + let missing = required.saturating_sub(last_fresh_count); + debug!( + dc = *dc, + family = ?family, + pass = pass_idx + 1, + total_passes, + fresh_count = last_fresh_count, + required, + missing, + endpoint_count = endpoint_list.len(), + "ME hardswap family warmup pass started" + ); + + for attempt_idx in 0..missing { + let delay_ms = self.hardswap_warmup_connect_delay_ms(); + tokio::time::sleep(Duration::from_millis(delay_ms)).await; + + let connected = self + .connect_endpoints_round_robin_with_generation_contour( + *dc, + &endpoint_list, + rng, + generation, + WriterContour::Warm, + WriterOpenIntent::Coverage, + ) + .await; + debug!( + dc = *dc, + family = ?family, + pass = pass_idx + 1, + total_passes, + attempt = attempt_idx + 1, + delay_ms, + connected, + "ME hardswap family warmup connect attempt finished" + ); + } + + last_fresh_count = self + .fresh_writer_count_for_dc_endpoints( generation, - WriterContour::Warm, - WriterOpenIntent::Normal, + *dc, + &family_endpoints, ) .await; - debug!( - dc = *dc, - pass = pass_idx + 1, - total_passes, - attempt = attempt_idx + 1, - delay_ms, - connected, - "ME hardswap warmup connect attempt finished" - ); + if last_fresh_count >= required { + completed = true; + info!( + dc = *dc, + family = ?family, + pass = pass_idx + 1, + total_passes, + fresh_count = last_fresh_count, + required, + "ME hardswap writer floor reached for DC family" + ); + break; + } + + if pass_idx + 1 < total_passes { + let backoff_ms = self.hardswap_warmup_backoff_ms(pass_idx); + debug!( + dc = *dc, + family = ?family, + pass = pass_idx + 1, + total_passes, + fresh_count = last_fresh_count, + required, + backoff_ms, + "ME hardswap family warmup incomplete, delaying next pass" + ); + tokio::time::sleep(Duration::from_millis(backoff_ms)).await; + } } - last_fresh_count = self - .fresh_writer_count_for_dc_endpoints(generation, *dc, endpoints) - .await; - if last_fresh_count >= required { - completed = true; - info!( + if !completed { + warn!( dc = *dc, - pass = pass_idx + 1, - total_passes, + family = ?family, fresh_count = last_fresh_count, required, - "ME hardswap warmup floor reached for DC" - ); - break; - } - - if pass_idx + 1 < total_passes { - let backoff_ms = self.hardswap_warmup_backoff_ms(pass_idx); - debug!( - dc = *dc, - pass = pass_idx + 1, + endpoint_count = endpoint_list.len(), total_passes, - fresh_count = last_fresh_count, - required, - backoff_ms, - "ME hardswap warmup pass incomplete, delaying next pass" + "ME warmup stopped below the required DC-family writer floor" ); - tokio::time::sleep(Duration::from_millis(backoff_ms)).await; } } - - if !completed { - warn!( - dc = *dc, - fresh_count = last_fresh_count, - required, - endpoint_count = endpoint_list.len(), - total_passes, - "ME warmup stopped: unable to reach required writer floor for DC" - ); - } } } @@ -220,27 +236,27 @@ impl MePool { } if hardswap { - let fresh_writer_addrs: HashSet<(i32, SocketAddr)> = writers + let fresh_writer_addrs: Vec<(i32, SocketAddr)> = writers .iter() .filter(|w| !w.draining.load(Ordering::Relaxed)) .filter(|w| w.generation == generation) .map(|w| (w.writer_dc, w.addr)) .collect(); - let (fresh_coverage_ratio, fresh_missing_dc) = - Self::coverage_ratio(&desired_by_dc, &fresh_writer_addrs); - if fresh_coverage_ratio < min_ratio { + let fresh_coverage = self.hardswap_coverage(&desired_by_dc, &fresh_writer_addrs); + if fresh_coverage.ratio < min_ratio { self.set_last_drain_gate( false, - fresh_missing_dc.is_empty(), + fresh_coverage.missing_groups.is_empty(), MeDrainGateReason::CoverageQuorum, now_epoch_secs, ); warn!( previous_generation, generation, - fresh_coverage_ratio = format_args!("{fresh_coverage_ratio:.3}"), - missing_dc = ?fresh_missing_dc, - "ME hardswap pending: fresh generation DC coverage incomplete" + fresh_coverage_ratio = format_args!("{:.3}", fresh_coverage.ratio), + writer_deficit = fresh_coverage.writer_deficit, + missing_groups = ?fresh_coverage.missing_groups, + "ME hardswap pending: fresh generation DC-family floors incomplete" ); return false; } @@ -263,6 +279,7 @@ impl MePool { Err(ReinitCommitFailure::Coverage { coverage_ratio, missing_dc, + missing_groups, }) => { self.set_last_drain_gate( false, @@ -276,6 +293,7 @@ impl MePool { coverage_ratio = format_args!("{coverage_ratio:.3}"), min_ratio = format_args!("{min_ratio:.3}"), missing_dc = ?missing_dc, + missing_groups = ?missing_groups, "ME reinit coverage changed before commit; keeping current generation" ); return false; @@ -283,6 +301,7 @@ impl MePool { Err(ReinitCommitFailure::Redundancy { coverage_ratio, missing_dc, + missing_groups, }) => { self.set_last_drain_gate( true, @@ -296,6 +315,7 @@ impl MePool { coverage_ratio = format_args!("{coverage_ratio:.3}"), min_ratio = format_args!("{min_ratio:.3}"), missing_dc = ?missing_dc, + missing_groups = ?missing_groups, "ME hardswap weighted quorum requires stale-binding fallback" ); return false; @@ -310,9 +330,10 @@ impl MePool { if !outcome.missing_dc.is_empty() { warn!( missing_dc = ?outcome.missing_dc, + missing_groups = ?outcome.missing_groups, coverage_ratio = format_args!("{:.3}", outcome.coverage_ratio), min_ratio = format_args!("{min_ratio:.3}"), - "ME reinit committed with bounded stale fallback for uncovered DC groups" + "ME reinit committed with bounded stale fallback for uncovered DC-family groups" ); } diff --git a/src/transport/middle_proxy/pool_reinit/tests.rs b/src/transport/middle_proxy/pool_reinit/tests.rs index 8f19001..25cb58f 100644 --- a/src/transport/middle_proxy/pool_reinit/tests.rs +++ b/src/transport/middle_proxy/pool_reinit/tests.rs @@ -1,5 +1,5 @@ use std::collections::{HashMap, HashSet}; -use std::net::{IpAddr, Ipv4Addr, SocketAddr}; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; use std::sync::Arc; use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU32, AtomicU64, Ordering}; use std::time::Instant; @@ -7,7 +7,7 @@ use std::time::Instant; use tokio::sync::mpsc; use tokio_util::sync::CancellationToken; -use super::{MePool, ReinitCommitFailure, commit_reinit_state}; +use super::{DcFamilyGroup, MePool, ReinitCommitFailure, commit_reinit_state}; use crate::config::MeBindStaleMode; use crate::transport::middle_proxy::codec::WriterCommand; use crate::transport::middle_proxy::pool::{ @@ -19,6 +19,13 @@ fn addr(octet: u8, port: u16) -> SocketAddr { SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, octet)), port) } +fn addr_v6(segment: u16, port: u16) -> SocketAddr { + SocketAddr::new( + IpAddr::V6(Ipv6Addr::new(0x2001, 0xdb8, 0, 0, 0, 0, 0, segment)), + port, + ) +} + async fn insert_writer( pool: &Arc, writer_id: u64, @@ -56,6 +63,32 @@ async fn insert_writer( writer } +async fn insert_writer_floor( + pool: &Arc, + first_writer_id: u64, + writer_dc: i32, + endpoint: SocketAddr, + generation: u64, + contour: WriterContour, +) -> Vec { + let required = pool.required_writers_for_dc(1); + let mut writers = Vec::with_capacity(required); + for offset in 0..required { + writers.push( + insert_writer( + pool, + first_writer_id + offset as u64, + writer_dc, + endpoint, + generation, + contour, + ) + .await, + ); + } + writers +} + fn desired_two_dcs() -> HashMap> { HashMap::from([ (1, HashSet::from([addr(1, 2001)])), @@ -181,7 +214,7 @@ async fn partial_hardswap_is_rejected_when_stale_binding_is_disabled() { let reservation = pool .reserve_reinit_attempt(true, map_hash, endpoint_revision, 100) .expect("endpoint revision must remain current"); - insert_writer( + insert_writer_floor( &pool, 201, 1, @@ -232,7 +265,7 @@ async fn partial_hardswap_preserves_fallback_only_for_missing_dc() { let reservation = pool .reserve_reinit_attempt(true, map_hash, endpoint_revision, 100) .expect("endpoint revision must remain current"); - let fresh_dc1 = insert_writer( + let fresh_dc1 = insert_writer_floor( &pool, 401, 1, @@ -255,11 +288,72 @@ async fn partial_hardswap_preserves_fallback_only_for_missing_dc() { assert!(old_dc2.draining.load(Ordering::Acquire)); assert!(old_dc2.allow_drain_fallback.load(Ordering::Acquire)); assert_eq!( - WriterContour::from_u8(fresh_dc1.contour.load(Ordering::Acquire)), + WriterContour::from_u8(fresh_dc1[0].contour.load(Ordering::Acquire)), WriterContour::Active ); } +#[tokio::test] +async fn partial_hardswap_preserves_fallback_only_for_underfloor_family() { + let pool = make_pool().await; + pool.binding_policy + .me_bind_stale_mode + .store(MeBindStaleMode::Ttl.as_u8(), Ordering::Release); + let v4 = addr(1, 2001); + let v6 = addr_v6(1, 2001); + let desired_by_dc = HashMap::from([(1, HashSet::from([v4, v6]))]); + let active_generation = pool.current_generation(); + let old_v4 = insert_writer( + &pool, + 451, + 1, + v4, + active_generation, + WriterContour::Active, + ) + .await; + let old_v6 = insert_writer( + &pool, + 452, + 1, + v6, + active_generation, + WriterContour::Active, + ) + .await; + let map_hash = MePool::desired_map_hash(&desired_by_dc); + let endpoint_revision = pool.endpoint_snapshot.load().revision; + let reservation = pool + .reserve_reinit_attempt(true, map_hash, endpoint_revision, 100) + .expect("endpoint revision must remain current"); + insert_writer_floor( + &pool, + 461, + 1, + v4, + reservation.attempt.generation, + WriterContour::Warm, + ) + .await; + + let outcome = pool + .commit_reinit_attempt(&reservation.attempt, &desired_by_dc, 0.5) + .await + .expect("one covered family satisfies the configured weighted quorum"); + + assert_eq!( + outcome.missing_groups, + vec![DcFamilyGroup { + dc: 1, + family: crate::network::IpFamily::V6, + }] + ); + assert!(outcome.force_close_writer_ids.contains(&old_v4.id)); + assert!(!outcome.force_close_writer_ids.contains(&old_v6.id)); + assert!(!old_v4.allow_drain_fallback.load(Ordering::Acquire)); + assert!(old_v6.allow_drain_fallback.load(Ordering::Acquire)); +} + #[tokio::test] async fn complete_hardswap_promotes_fresh_generation_and_retires_old_writers() { let pool = make_pool().await; @@ -288,7 +382,7 @@ async fn complete_hardswap_promotes_fresh_generation_and_retires_old_writers() { let reservation = pool .reserve_reinit_attempt(true, map_hash, endpoint_revision, 100) .expect("endpoint revision must remain current"); - let fresh_dc1 = insert_writer( + let fresh_dc1 = insert_writer_floor( &pool, 601, 1, @@ -297,9 +391,9 @@ async fn complete_hardswap_promotes_fresh_generation_and_retires_old_writers() { WriterContour::Warm, ) .await; - let fresh_dc2 = insert_writer( + let fresh_dc2 = insert_writer_floor( &pool, - 602, + 611, 2, addr(2, 2002), reservation.attempt.generation, @@ -318,11 +412,11 @@ async fn complete_hardswap_promotes_fresh_generation_and_retires_old_writers() { assert!(old_dc1.draining.load(Ordering::Acquire)); assert!(old_dc2.draining.load(Ordering::Acquire)); assert_eq!( - WriterContour::from_u8(fresh_dc1.contour.load(Ordering::Acquire)), + WriterContour::from_u8(fresh_dc1[0].contour.load(Ordering::Acquire)), WriterContour::Active ); assert_eq!( - WriterContour::from_u8(fresh_dc2.contour.load(Ordering::Acquire)), + WriterContour::from_u8(fresh_dc2[0].contour.load(Ordering::Acquire)), WriterContour::Active ); } @@ -442,10 +536,8 @@ async fn generation_role_reconciliation_promotes_active_warm_and_drains_orphans( ); assert!(orphan_warm.draining.load(Ordering::Acquire)); assert!(!orphan_warm.allow_drain_fallback.load(Ordering::Acquire)); - assert_eq!( - pool.api_hardswap_snapshot() - .await - .orphan_warm_writers_current, - 0 - ); + let snapshot = pool.api_hardswap_snapshot().await; + assert_eq!(snapshot.orphan_warm_writers_current, 0); + assert_eq!(snapshot.pending_writer_deficit, 2); + assert_eq!(snapshot.pending_missing_dc_groups, 1); } diff --git a/src/transport/middle_proxy/pool_status.rs b/src/transport/middle_proxy/pool_status.rs index 0e3cecc..d19a859 100644 --- a/src/transport/middle_proxy/pool_status.rs +++ b/src/transport/middle_proxy/pool_status.rs @@ -6,6 +6,7 @@ use std::time::Instant; use super::pool::{MePool, ReinitStatusSnapshot, WriterContour}; use crate::config::{MeBindStaleMode, MeFloorMode, MeSocksKdfPolicy}; +use crate::network::IpFamily; use crate::transport::upstream::IpPreference; // ME writer and DC coverage snapshots. @@ -105,7 +106,7 @@ pub(crate) struct MeApiRuntimeSnapshot { pub pending_writers_current: usize, /// Number of writers still required to reach the pending generation floor. pub pending_writer_deficit: usize, - /// Number of desired DC groups without pending-generation coverage. + /// Number of desired DC-family groups below their pending-generation floor. pub pending_missing_dc_groups: usize, /// Whether the pending generation targets the current desired endpoint map. pub pending_map_current: Option, diff --git a/src/transport/middle_proxy/pool_status/hardswap_snapshot.rs b/src/transport/middle_proxy/pool_status/hardswap_snapshot.rs index 770791b..26aa8cd 100644 --- a/src/transport/middle_proxy/pool_status/hardswap_snapshot.rs +++ b/src/transport/middle_proxy/pool_status/hardswap_snapshot.rs @@ -11,7 +11,7 @@ pub(crate) struct MeApiHardswapSnapshot { pub pending_writers_current: usize, /// Number of writers still required to reach the pending generation floor. pub pending_writer_deficit: usize, - /// Number of desired DC groups without a pending-generation writer. + /// Number of desired DC-family groups below their pending-generation writer floor. pub pending_missing_dc_groups: usize, /// Whether the pending generation targets the current desired endpoint map. pub pending_map_current: Option, @@ -41,7 +41,7 @@ impl MePool { let pending_generation = reinit.pending_hardswap_generation; let pending = pending_generation != 0; let mut pending_writers_current = 0usize; - let mut pending_by_dc = HashMap::::new(); + let mut pending_by_group = HashMap::<(i32, IpFamily), usize>::new(); let mut orphan_warm_writers_current = 0usize; for writer in writers.iter() { @@ -60,7 +60,14 @@ impl MePool { .is_some_and(|endpoints| endpoints.contains(&writer.addr)) { pending_writers_current = pending_writers_current.saturating_add(1); - *pending_by_dc.entry(writer.writer_dc).or_insert(0) += 1; + let family = if writer.addr.is_ipv4() { + IpFamily::V4 + } else { + IpFamily::V6 + }; + *pending_by_group + .entry((writer.writer_dc, family)) + .or_insert(0) += 1; } } @@ -68,15 +75,22 @@ impl MePool { let mut pending_missing_dc_groups = 0usize; if pending { for (dc, endpoints) in &desired_by_dc { - if endpoints.is_empty() { - continue; - } - let alive = pending_by_dc.get(dc).copied().unwrap_or(0); - let required = self.required_writers_for_dc(endpoints.len()); - pending_writer_deficit = pending_writer_deficit - .saturating_add(required.saturating_sub(alive)); - if alive == 0 { - pending_missing_dc_groups = pending_missing_dc_groups.saturating_add(1); + for family in [IpFamily::V4, IpFamily::V6] { + let endpoint_count = endpoints + .iter() + .filter(|endpoint| endpoint.is_ipv4() == (family == IpFamily::V4)) + .count(); + if endpoint_count == 0 { + continue; + } + let alive = pending_by_group.get(&(*dc, family)).copied().unwrap_or(0); + let required = self.required_writers_for_dc(endpoint_count); + let deficit = required.saturating_sub(alive); + pending_writer_deficit = pending_writer_deficit.saturating_add(deficit); + if deficit > 0 { + pending_missing_dc_groups = + pending_missing_dc_groups.saturating_add(1); + } } } } diff --git a/src/transport/middle_proxy/pool_writer/publication.rs b/src/transport/middle_proxy/pool_writer/publication.rs index ef1a45c..b68b940 100644 --- a/src/transport/middle_proxy/pool_writer/publication.rs +++ b/src/transport/middle_proxy/pool_writer/publication.rs @@ -102,28 +102,12 @@ impl MePool { )); }; let required = match contour { - WriterContour::Active => self.required_writers_for_dc( + WriterContour::Active | WriterContour::Warm => self.required_writers_for_dc( endpoints .iter() .filter(|endpoint| endpoint.is_ipv4() == writer.addr.is_ipv4()) .count(), ), - WriterContour::Warm => self.required_writers_for_dc( - endpoints - .iter() - .filter(|endpoint| { - let endpoint_family = if endpoint.is_ipv4() { - crate::network::IpFamily::V4 - } else { - crate::network::IpFamily::V6 - }; - self.family_enabled_for_drain_coverage( - endpoint_family, - now_epoch_secs, - ) - }) - .count(), - ), WriterContour::Draining => 0, }; let current = writers @@ -134,18 +118,8 @@ impl MePool { && candidate.generation == writer.generation && WriterContour::from_u8(candidate.contour.load(Ordering::Acquire)) == contour - && (contour == WriterContour::Warm - || candidate.addr.is_ipv4() == writer.addr.is_ipv4()) + && candidate.addr.is_ipv4() == writer.addr.is_ipv4() && endpoints.contains(&candidate.addr) - && (contour != WriterContour::Warm - || self.family_enabled_for_drain_coverage( - if candidate.addr.is_ipv4() { - crate::network::IpFamily::V4 - } else { - crate::network::IpFamily::V6 - }, - now_epoch_secs, - )) }) .count(); if current >= required { diff --git a/src/transport/middle_proxy/pool_writer/replacement.rs b/src/transport/middle_proxy/pool_writer/replacement.rs index c1a3fb3..ff34bda 100644 --- a/src/transport/middle_proxy/pool_writer/replacement.rs +++ b/src/transport/middle_proxy/pool_writer/replacement.rs @@ -223,6 +223,11 @@ mod tests { WriterContour::Active, WriterOpenIntent::Replacement, writer_dc, + if addr.is_ipv4() { + crate::network::IpFamily::V4 + } else { + crate::network::IpFamily::V6 + }, ) .await .expect("replacement open must be admitted"); diff --git a/src/transport/middle_proxy/pool_writer/runtime.rs b/src/transport/middle_proxy/pool_writer/runtime.rs index 382f844..412b2cd 100644 --- a/src/transport/middle_proxy/pool_writer/runtime.rs +++ b/src/transport/middle_proxy/pool_writer/runtime.rs @@ -92,7 +92,16 @@ impl MePool { intent: WriterOpenIntent, ) -> Result> { let Some(writer_open_reservation) = self - .reserve_writer_open(contour, intent, writer_dc) + .reserve_writer_open( + contour, + intent, + writer_dc, + if addr.is_ipv4() { + crate::network::IpFamily::V4 + } else { + crate::network::IpFamily::V6 + }, + ) .await else { return Err(ProxyError::Proxy(format!( diff --git a/src/web/http/operator_lifecycle_tests.rs b/src/web/http/operator_lifecycle_tests.rs index e579d59..a07c866 100644 --- a/src/web/http/operator_lifecycle_tests.rs +++ b/src/web/http/operator_lifecycle_tests.rs @@ -90,6 +90,35 @@ async fn stop_runtime( generation.stop_background_tasks().await; } +#[tokio::test] +async fn user_revocation_interrupts_live_session_before_periodic_cleanup() { + let (runtime, generation, listener) = live_runtime().await; + let bootstrap = issue_bootstrap(&runtime); + let (_, token) = create_session(&listener, &runtime, &bootstrap).await; + let session_hash = token_hash(&token); + let session = runtime + .get_session(session_hash, "proxy.example.com") + .unwrap(); + let polling = Arc::clone(&session); + let poll = tokio::spawn(async move { polling.poll_down(0).await }); + tokio::task::yield_now().await; + + let mutation = generation.proxy_shared.delete_user("alice"); + + assert!(mutation.cancelled >= 1); + assert!(matches!( + tokio::time::timeout(Duration::from_millis(250), poll) + .await + .unwrap() + .unwrap(), + Err(ManagerError::Closed) + )); + assert!(runtime + .get_session(session_hash, "proxy.example.com") + .is_err()); + stop_runtime(runtime, generation).await; +} + #[tokio::test] async fn pause_preserves_decoy_retry_and_exact_session_replay() { let (runtime, generation, listener) = live_runtime().await; diff --git a/src/web/http/websocket/driver.rs b/src/web/http/websocket/driver.rs index f34df62..b936ea3 100644 --- a/src/web/http/websocket/driver.rs +++ b/src/web/http/websocket/driver.rs @@ -45,6 +45,7 @@ pub(super) async fn run_upgraded( UpgradeDeadlineLease::deadline, ); let upgraded = tokio::select! { + biased; _ = cancellation.cancelled() => return, result = tokio::time::timeout_at(deadline, on_upgrade) => result, }; @@ -137,6 +138,7 @@ async fn run_multiplex( let down = session.poll_down_websocket(cursor); tokio::pin!(down); let event = tokio::select! { + biased; _ = cancellation.cancelled() => return Err(()), _ = tokio::time::sleep_until(open_deadline.into()), if !active => return Err(()), _ = tokio::time::sleep_until(next_ping.into()) => DriverEvent::Liveness, diff --git a/src/web/http/websocket/driver/io.rs b/src/web/http/websocket/driver/io.rs index 20223e8..4f4633f 100644 --- a/src/web/http/websocket/driver/io.rs +++ b/src/web/http/websocket/driver/io.rs @@ -20,6 +20,7 @@ pub(super) async fn read_message( backpressure_timeout: Duration, ) -> Result<(Message, Option), ()> { tokio::select! { + biased; _ = cancellation.cancelled() => return Err(()), ready = socket.get_ref().readable() => ready.map_err(|_| ())?, } @@ -28,6 +29,7 @@ pub(super) async fn read_message( Some(reserve_data(runtime, owner, maximum, cancellation, backpressure_timeout).await?); } let message = tokio::select! { + biased; _ = cancellation.cancelled() => return Err(()), message = socket.next() => message.ok_or(())?.map_err(|_| ())?, }; @@ -59,6 +61,7 @@ pub(super) async fn reserve_data( return Ok(budget); } tokio::select! { + biased; _ = cancellation.cancelled() => return Err(()), _ = notified => {} } @@ -114,6 +117,7 @@ pub(super) async fn process_lane( Err(_) => return Err(()), } tokio::select! { + biased; _ = cancellation.cancelled() => return Err(()), _ = notified => {} } @@ -147,6 +151,7 @@ where Err(_) => return Err(()), } tokio::select! { + biased; _ = cancellation.cancelled() => return Err(()), _ = notified => {} } @@ -163,6 +168,7 @@ pub(super) async fn send( timeout: Duration, ) -> Result<(), ()> { tokio::select! { + biased; _ = cancellation.cancelled() => Err(()), result = tokio::time::timeout(timeout, socket.send(message)) => { result.map_err(|_| ())?.map_err(|_| ()) @@ -176,6 +182,7 @@ pub(super) async fn flush( timeout: Duration, ) -> Result<(), ()> { tokio::select! { + biased; _ = cancellation.cancelled() => Err(()), result = tokio::time::timeout(timeout, socket.flush()) => { result.map_err(|_| ())?.map_err(|_| ()) diff --git a/src/web/http/websocket/driver/lane.rs b/src/web/http/websocket/driver/lane.rs index e635bab..8780d0e 100644 --- a/src/web/http/websocket/driver/lane.rs +++ b/src/web/http/websocket/driver/lane.rs @@ -39,6 +39,7 @@ pub(super) async fn run_lane( let down = session.poll_down_websocket_lane(reservation.lane_identity(), cursor); tokio::pin!(down); let event = tokio::select! { + biased; _ = cancellation.cancelled() => return Err(()), _ = tokio::time::sleep_until(open_deadline.into()), if !active => return Err(()), _ = tokio::time::sleep_until(next_ping.into()) => DriverEvent::Liveness, diff --git a/src/web/manager/credentials.rs b/src/web/manager/credentials.rs index 8d9fd86..f2de283 100644 --- a/src/web/manager/credentials.rs +++ b/src/web/manager/credentials.rs @@ -116,19 +116,9 @@ impl WebProcessRuntime { .record_rejection(WebRejectionReason::BootstrapCapacity); return Err(ManagerError::Limit); } - if !allow_rate( - &mut state.bootstrap_rate, - now, - self.limits.new_bootstraps_per_minute, - self.limits.new_bootstraps_burst, - ) { - self.record_limit_hit(); - self.telemetry - .record_rejection(WebRejectionReason::BootstrapRate); - return Err(ManagerError::Limit); - } - if state.bootstraps.len() >= self.limits.max_bootstraps_global - && !evict_oldest_unused_bootstrap(&mut state) + let global_capacity_full = state.bootstraps.len() >= self.limits.max_bootstraps_global; + if global_capacity_full + && !state.bootstraps.values().any(|bootstrap| !bootstrap.used) { self.record_limit_hit(); self.telemetry @@ -152,6 +142,23 @@ impl WebProcessRuntime { let Some(user_registration) = user_publication.take_registration() else { return Err(ManagerError::Closed); }; + if !allow_rate( + &mut state.bootstrap_rate, + now, + self.limits.new_bootstraps_per_minute, + self.limits.new_bootstraps_burst, + ) { + self.record_limit_hit(); + self.telemetry + .record_rejection(WebRejectionReason::BootstrapRate); + return Err(ManagerError::Limit); + } + let evicted_bootstrap = global_capacity_full + .then(|| evict_oldest_unused_bootstrap(&mut state)) + .flatten(); + if global_capacity_full && evicted_bootstrap.is_none() { + return Err(ManagerError::Limit); + } let trace_session_id = self.trace.next_session_id(); let bridge_diagnostics_enabled = config.web.debug.bridge_diagnostics_enabled(); let (user_agent, user_agent_id) = bounded_user_agent(user_agent); @@ -196,6 +203,7 @@ impl WebProcessRuntime { *state.bootstraps_per_ip.entry(client_ip).or_insert(0) += 1; user_publication.commit(); drop(state); + drop(evicted_bootstrap); if recovery { self.telemetry .record_bridge_recovery(WebBridgeRecoveryEvent::BootstrapIssued); @@ -243,7 +251,11 @@ impl WebProcessRuntime { .lock() .bootstraps .get(&hash) - .filter(|entry| entry.profile.host == host && now <= entry.expires_at) + .filter(|entry| { + entry.profile.host == host + && now <= entry.expires_at + && !entry.user_registration.is_cancelled() + }) .map(|entry| { ( entry.trace_session_id, @@ -263,20 +275,23 @@ impl WebProcessRuntime { host: &str, ) -> std::result::Result, ManagerError> { let state = self.state.lock(); - if let Some(session) = state + let session = state .sessions .get(&hash) .cloned() - .filter(|session| session.matches_host(host)) - { - return Ok(session); - } + .filter(|session| session.matches_host(host)); let retired_carrier = state .closed_tokens .get(&hash) .filter(|closed| closed.host == host) .map(|closed| closed.carrier); drop(state); + if let Some(session) = session { + if session.close_if_cancelled() { + return Err(ManagerError::Closed); + } + return Ok(session); + } if let Some(carrier) = retired_carrier { self.telemetry.record_session_observation( carrier, @@ -294,14 +309,16 @@ impl WebProcessRuntime { profile: &WebRuntimeProfile, ) -> Option> { let expected_profile = profile_key(profile); - self.state + let session = self + .state .lock() .sessions .get(&hash) .filter(|session| { session.matches_host(host) && session.profile_key() == expected_profile }) - .cloned() + .cloned(); + session.filter(|session| !session.close_if_cancelled()) } /// Closes a live token and accepts bounded tombstone retries. diff --git a/src/web/manager/state.rs b/src/web/manager/state.rs index 47a4f43..ca16949 100644 --- a/src/web/manager/state.rs +++ b/src/web/manager/state.rs @@ -307,8 +307,8 @@ pub(super) fn allow_rate(state: &mut RateState, now: Instant, per_minute: u32, b true } -/// Evicts the oldest unused bootstrap while preserving used retry state. -pub(super) fn evict_oldest_unused_bootstrap(state: &mut ManagerState) -> bool { +/// Detaches the oldest unused bootstrap while preserving used retry state. +pub(super) fn evict_oldest_unused_bootstrap(state: &mut ManagerState) -> Option { let Some(hash) = state .bootstraps .iter() @@ -316,10 +316,11 @@ pub(super) fn evict_oldest_unused_bootstrap(state: &mut ManagerState) -> bool { .min_by_key(|(_, bootstrap)| bootstrap.issued_at) .map(|(hash, _)| *hash) else { - return false; + return None; }; - remove_bootstrap_locked(state, hash); - true + let bootstrap = state.bootstraps.remove(&hash)?; + decrement_map(&mut state.bootstraps_per_ip, &bootstrap.issuance_ip); + Some(bootstrap) } /// Removes expired bootstrap and closed-token entries while the manager lock is held. diff --git a/src/web/session/backend.rs b/src/web/session/backend.rs index 173f75c..004800c 100644 --- a/src/web/session/backend.rs +++ b/src/web/session/backend.rs @@ -20,6 +20,13 @@ impl WebSession { completion: StreamCompletion, retain_reservation_on_reject: bool, ) -> bool { + if self.close_if_cancelled() { + completion + .retain_rejected + .store(retain_reservation_on_reject, Ordering::Release); + drop(completion); + return false; + } let stream = completion.stream; let peer_port = completion.peer_port; let Some(manager) = self.manager.upgrade() else { @@ -60,6 +67,7 @@ impl WebSession { ); let logical_stream = WebLogicalStream::new(Arc::clone(&session), stream); tokio::select! { + biased; _ = cancel.cancelled() => {} _ = run_stream( Arc::clone(&session), diff --git a/src/web/session/downlink.rs b/src/web/session/downlink.rs index 6447610..b59b9ab 100644 --- a/src/web/session/downlink.rs +++ b/src/web/session/downlink.rs @@ -26,9 +26,14 @@ impl WebSession { if !self.carrier().is_multiplexed() { return Err(ManagerError::Protocol); } + if self.close_if_cancelled() { + return Err(ManagerError::Closed); + } let (epoch, healthy) = { let mut state = self.state.lock(); - if state.closed { + if state.closed || self.cancel.is_cancelled() { + drop(state); + self.close_if_cancelled(); return Err(ManagerError::Closed); } if let Some(unacked) = &state.unacked { @@ -92,6 +97,11 @@ impl WebSession { notified.as_mut().enable(); { let mut state = self.state.lock(); + if self.cancel.is_cancelled() { + drop(state); + self.close_if_cancelled(); + return Err(ManagerError::Closed); + } if state.down_epoch != epoch { return Ok(PollResult { body: Bytes::new(), @@ -129,18 +139,33 @@ impl WebSession { notified.await; } }; - match tokio::time::timeout(deadline, poll).await { - Ok(result) => result, - Err(_) => { - let mut state = self.state.lock(); - if state.down_epoch == epoch { - state.activity.touch_progress(Instant::now()); + tokio::select! { + biased; + _ = self.cancel.cancelled() => { + self.close_if_cancelled(); + Err(ManagerError::Closed) + } + result = tokio::time::timeout(deadline, poll) => match result { + Ok(result) => result, + Err(_) => { + if self.close_if_cancelled() { + return Err(ManagerError::Closed); + } + let mut state = self.state.lock(); + if self.cancel.is_cancelled() { + drop(state); + self.close_if_cancelled(); + return Err(ManagerError::Closed); + } + if state.down_epoch == epoch { + state.activity.touch_progress(Instant::now()); + } + Ok(PollResult { + body: Bytes::new(), + next_cursor: cursor, + lane_closed: false, + }) } - Ok(PollResult { - body: Bytes::new(), - next_cursor: cursor, - lane_closed: false, - }) } } } diff --git a/src/web/session/lane_uplink.rs b/src/web/session/lane_uplink.rs index 092f007..e167411 100644 --- a/src/web/session/lane_uplink.rs +++ b/src/web/session/lane_uplink.rs @@ -22,6 +22,9 @@ impl WebSession { if self.carrier() != WebCarrier::HttpsLanes || lane_id > frame::MAX_STREAM_ID { return Err(ManagerError::Protocol); } + if self.close_if_cancelled() { + return Err(ManagerError::Closed); + } let frames = match frame::parse_all(body, &self.limits) { Ok(frames) => frames, Err(_) => { @@ -175,6 +178,10 @@ impl WebSession { if matches!(result, Err(ManagerError::Backpressure)) { return result; } + if matches!(result, Err(ManagerError::Closed)) && self.close_if_cancelled() { + drop(opened); + return result; + } if result.is_err() { self.close(SessionCloseReason::Protocol); drop(opened); diff --git a/src/web/session/lanes.rs b/src/web/session/lanes.rs index 5cfa72e..295a3bb 100644 --- a/src/web/session/lanes.rs +++ b/src/web/session/lanes.rs @@ -42,9 +42,14 @@ impl WebSession { if !self.carrier().uses_lanes() || lane_id > frame::MAX_STREAM_ID { return Err(ManagerError::Protocol); } + if self.close_if_cancelled() { + return Err(ManagerError::Closed); + } let lane_ready = if let Some(expected_instance) = expected_instance { let state = self.state.lock(); - if state.closed { + if state.closed || self.cancel.is_cancelled() { + drop(state); + self.close_if_cancelled(); return Err(ManagerError::Closed); } state @@ -63,7 +68,9 @@ impl WebSession { } let (instance, epoch, notify, healthy) = { let mut state = self.state.lock(); - if state.closed { + if state.closed || self.cancel.is_cancelled() { + drop(state); + self.close_if_cancelled(); return Err(ManagerError::Closed); } let (acknowledged, replay) = { @@ -175,7 +182,9 @@ impl WebSession { notified.as_mut().enable(); { let mut state = self.state.lock(); - if state.closed { + if state.closed || self.cancel.is_cancelled() { + drop(state); + self.close_if_cancelled(); return Err(ManagerError::Closed); } let carrier_health_eligible = lane_id != 0 @@ -245,55 +254,72 @@ impl WebSession { notified.await; } }; - match tokio::time::timeout(deadline, poll).await { - Ok(result) => result, - Err(_) => { - let mut state = self.state.lock(); - if state.closed { - return Err(ManagerError::Closed); - } - if !state.carrier_lanes.contains_key(&lane_id) { - return Ok(PollResult { - body: Bytes::new(), - next_cursor: cursor, - lane_closed: true, - }); - } - if lane_id != 0 - && !state.streams.contains_key(&lane_id) - && state.closed_streams.contains(&lane_id) - { - return Ok(PollResult { - body: Bytes::new(), - next_cursor: cursor, - lane_closed: true, - }); - } - if let Some(lane) = state.carrier_lanes.get(&lane_id) { - if lane.instance != instance { + tokio::select! { + biased; + _ = self.cancel.cancelled() => { + self.close_if_cancelled(); + Err(ManagerError::Closed) + } + result = tokio::time::timeout(deadline, poll) => match result { + Ok(result) => result, + Err(_) => { + if self.close_if_cancelled() { + return Err(ManagerError::Closed); + } + let mut state = self.state.lock(); + if state.closed || self.cancel.is_cancelled() { + drop(state); + self.close_if_cancelled(); + return Err(ManagerError::Closed); + } + if !state.carrier_lanes.contains_key(&lane_id) { return Ok(PollResult { body: Bytes::new(), next_cursor: cursor, lane_closed: true, }); } - if lane.down_epoch == epoch { - state.activity.touch_progress(Instant::now()); + if lane_id != 0 + && !state.streams.contains_key(&lane_id) + && state.closed_streams.contains(&lane_id) + { + return Ok(PollResult { + body: Bytes::new(), + next_cursor: cursor, + lane_closed: true, + }); } + if let Some(lane) = state.carrier_lanes.get(&lane_id) { + if lane.instance != instance { + return Ok(PollResult { + body: Bytes::new(), + next_cursor: cursor, + lane_closed: true, + }); + } + if lane.down_epoch == epoch { + state.activity.touch_progress(Instant::now()); + } + } + Ok(PollResult { + body: Bytes::new(), + next_cursor: cursor, + lane_closed: false, + }) } - Ok(PollResult { - body: Bytes::new(), - next_cursor: cursor, - lane_closed: false, - }) } } } async fn wait_for_lane_open(&self, lane_id: u32, cursor: u64) -> Result { + if self.close_if_cancelled() { + return Err(ManagerError::Closed); + } let wait = { let mut state = self.state.lock(); - if state.closed { + if state.closed || self.cancel.is_cancelled() { + drop(state); + self.close_if_cancelled(); return Err(ManagerError::Closed); } if state.carrier_lanes.contains_key(&lane_id) { @@ -334,7 +360,9 @@ impl WebSession { notified.as_mut().enable(); { let state = self.state.lock(); - if state.closed { + if state.closed || self.cancel.is_cancelled() { + drop(state); + self.close_if_cancelled(); return Err(ManagerError::Closed); } if state.carrier_lanes.contains_key(&lane_id) @@ -346,14 +374,24 @@ impl WebSession { } notified.await; } - }) - .await; + }); + let opened = tokio::select! { + biased; + _ = self.cancel.cancelled() => { + drop(wait); + self.close_if_cancelled(); + return Err(ManagerError::Closed); + } + opened = opened => opened, + }; drop(wait); match opened { Ok(result) => result, Err(_) => { let state = self.state.lock(); - if state.closed { + if state.closed || self.cancel.is_cancelled() { + drop(state); + self.close_if_cancelled(); Err(ManagerError::Closed) } else { Ok(state.carrier_lanes.contains_key(&lane_id) diff --git a/src/web/session/lifecycle.rs b/src/web/session/lifecycle.rs index f58986c..fb9a80a 100644 --- a/src/web/session/lifecycle.rs +++ b/src/web/session/lifecycle.rs @@ -121,6 +121,15 @@ impl CarrierSupersedeCompletion<'_> { } impl WebSession { + /// Closes a bearer as soon as its process-owned user registration is revoked. + pub(crate) fn close_if_cancelled(&self) -> bool { + if !self.cancel.is_cancelled() { + return false; + } + self.close(SessionCloseReason::UserDisabled); + true + } + /// Closes carrier state while relay tasks retain their admission until exit. pub(crate) fn close(&self, reason: SessionCloseReason) -> SessionCloseOutcome { let mut state = self.state.lock(); @@ -221,19 +230,8 @@ impl WebSession { /// Atomically closes a session only when reconnect grace is still due. pub(crate) fn close_if_due(&self, now: Instant) -> bool { - if self.cancel.is_cancelled() { - let released = { - let mut state = self.state.lock(); - if state.closed || state.close_requested.is_some() { - None - } else { - Some(self.release_on_close_locked(&mut state, SessionCloseReason::UserDisabled)) - } - }; - if let Some(released) = released { - self.finish_close(released); - return true; - } + if self.close_if_cancelled() { + return true; } let healthy = { let mut state = self.state.lock(); diff --git a/src/web/session/negotiation.rs b/src/web/session/negotiation.rs index 3e1de9d..9a2e539 100644 --- a/src/web/session/negotiation.rs +++ b/src/web/session/negotiation.rs @@ -34,6 +34,9 @@ impl WebSession { &self, state: &SessionState, ) -> Result<(), crate::web::manager::ManagerError> { + if self.cancel.is_cancelled() { + return Err(crate::web::manager::ManagerError::Closed); + } if state.negotiation_phase == SessionNegotiationPhase::Uncommitted && self .carrier_deadline_at diff --git a/src/web/session/uplink.rs b/src/web/session/uplink.rs index 5c87e52..9b058d5 100644 --- a/src/web/session/uplink.rs +++ b/src/web/session/uplink.rs @@ -69,6 +69,9 @@ impl WebSession { if !self.carrier().is_multiplexed() { return Err(ManagerError::Protocol); } + if self.close_if_cancelled() { + return Err(ManagerError::Closed); + } if self .up_active .compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire) @@ -157,6 +160,10 @@ impl WebSession { if matches!(result, Err(ManagerError::Backpressure)) { return result; } + if matches!(result, Err(ManagerError::Closed)) && self.close_if_cancelled() { + drop(opened); + return result; + } if result.is_err() { self.close(SessionCloseReason::Protocol); drop(opened); diff --git a/src/web/session/websocket.rs b/src/web/session/websocket.rs index c0d3c72..b26193d 100644 --- a/src/web/session/websocket.rs +++ b/src/web/session/websocket.rs @@ -39,8 +39,12 @@ pub(crate) struct WebSocketProbeReservation { impl WebSocketProbeReservation { /// Binds the admitted process connection to the future commit acknowledgement. pub(crate) fn bind(&mut self, owner: u64) -> Result<(), ManagerError> { + if self.session.close_if_cancelled() { + return Err(ManagerError::Closed); + } let mut state = self.session.state.lock(); if state.closed + || self.session.cancel.is_cancelled() || !state.websocket_probe_claimed || state.websocket_commit_ack_owner.is_some() { @@ -85,8 +89,12 @@ impl WebSocketLaneReservation { if self.phase != WebSocketLaneReservationPhase::Reserved { return Err(ManagerError::Concurrent); } + if self.session.close_if_cancelled() { + return Err(ManagerError::Closed); + } let mut state = self.session.state.lock(); if state.closed + || self.session.cancel.is_cancelled() || state .carrier_lanes .get(&self.claim.lane.lane_id) @@ -113,8 +121,12 @@ impl WebSocketLaneReservation { { return Err(ManagerError::Protocol); } + if self.session.close_if_cancelled() { + return Err(ManagerError::Closed); + } let mut state = self.session.state.lock(); - if state + if self.session.cancel.is_cancelled() + || state .carrier_lanes .get(&self.claim.lane.lane_id) .is_none_or(|lane| lane.instance != self.claim.lane.instance) @@ -174,8 +186,11 @@ impl WebSession { self: &Arc, acknowledge_commit: bool, ) -> Result, ManagerError> { + if self.close_if_cancelled() { + return Err(ManagerError::Closed); + } let mut state = self.state.lock(); - if state.closed { + if state.closed || self.cancel.is_cancelled() { return Err(ManagerError::Closed); } self.ensure_carrier_active_locked(&state)?; @@ -216,8 +231,11 @@ impl WebSession { { return Err(ManagerError::Protocol); } + if self.close_if_cancelled() { + return Err(ManagerError::Closed); + } let mut state = self.state.lock(); - if state.closed { + if state.closed || self.cancel.is_cancelled() { return Err(ManagerError::Closed); } if state.active_peer_ports.len() >= self.profile.max_streams_per_session @@ -306,6 +324,9 @@ impl WebSession { { return Err(ManagerError::Protocol); } + if self.close_if_cancelled() { + return Err(ManagerError::Closed); + } let lane_id = reservation.lane_id(); let frames = frame::parse_all(body, &self.limits).map_err(|_| ManagerError::Protocol)?; if frames