From 0ac236955ae16f29aa38d6afc8b1449b10735350 Mon Sep 17 00:00:00 2001 From: Alexey <247128645+axkurcom@users.noreply.github.com> Date: Mon, 14 Sep 2026 20:19:48 +0300 Subject: [PATCH] Proxy Shared User Drafts --- src/api/handler/fixed_routes.rs | 17 - src/api/handler/user_routes.rs | 30 +- src/api/mod.rs | 1 - src/api/users/create.rs | 10 +- src/api/users/lifecycle.rs | 12 +- src/api/users/update.rs | 22 +- src/config/load.rs | 20 + src/config/load/runtime_auth.rs | 13 + src/config/load/runtime_web.rs | 1 + src/config/types/web/runtime.rs | 2 + src/ip_tracker.rs | 20 +- src/ip_tracker/admission.rs | 47 ++ src/ip_tracker/cleanup.rs | 72 ++- src/ip_tracker/snapshot.rs | 18 + src/maestro/orchestrator.rs | 2 +- src/maestro/reload_supervisor.rs | 8 + src/maestro/reload_supervisor_tests.rs | 2 + src/maestro/runtime_build.rs | 13 +- src/maestro/runtime_tasks.rs | 4 +- src/proxy/authenticated.rs | 83 ++- src/proxy/mod.rs | 1 + src/proxy/shared_state.rs | 265 ++++---- src/proxy/user_admission.rs | 598 ++++++++++++++++++ src/transport/middle_proxy/health/family.rs | 7 +- src/transport/middle_proxy/ping.rs | 5 +- src/transport/middle_proxy/pool.rs | 18 +- .../middle_proxy/pool/construction.rs | 15 +- src/transport/middle_proxy/pool/routing.rs | 82 ++- .../middle_proxy/pool/writer_admission.rs | 12 +- src/transport/middle_proxy/pool_config.rs | 67 +- src/transport/middle_proxy/pool_init.rs | 2 +- src/transport/middle_proxy/pool_refill.rs | 12 +- src/transport/middle_proxy/pool_reinit.rs | 16 +- .../middle_proxy/pool_reinit/coordination.rs | 53 +- .../middle_proxy/pool_reinit/reconcile.rs | 16 +- .../middle_proxy/pool_reinit/tests.rs | 94 ++- .../pool_status/status_snapshot.rs | 15 +- .../middle_proxy/pool_writer/publication.rs | 7 +- .../middle_proxy/pool_writer/replacement.rs | 29 +- src/transport/middle_proxy/send.rs | 9 +- src/transport/middle_proxy/send/recovery.rs | 3 +- src/transport/middle_proxy/send/selection.rs | 3 +- .../tests/pool_refill_security_tests.rs | 14 +- .../tests/pool_writer_publication_tests.rs | 14 +- .../tests/send_adversarial_tests.rs | 18 +- src/web/http/tests.rs | 1 + src/web/manager/credentials.rs | 28 +- src/web/manager/lifecycle.rs | 5 +- src/web/manager/session_creation.rs | 16 +- .../manager/session_creation/replacement.rs | 18 +- src/web/manager/state.rs | 4 + src/web/session.rs | 10 +- src/web/session/backend_tests.rs | 2 + src/web/session/downlink_tests.rs | 2 + src/web/session/lanes/tests.rs | 2 + src/web/session/lifecycle.rs | 20 +- src/web/session/negotiation.rs | 2 + src/web/session/uplink_tests.rs | 2 + src/web/session/websocket/tests.rs | 2 + 59 files changed, 1485 insertions(+), 401 deletions(-) create mode 100644 src/proxy/user_admission.rs diff --git a/src/api/handler/fixed_routes.rs b/src/api/handler/fixed_routes.rs index 96ae7c6..e6e55bb 100644 --- a/src/api/handler/fixed_routes.rs +++ b/src/api/handler/fixed_routes.rs @@ -21,7 +21,6 @@ pub(super) async fn create_user_route( } let expected_revision = parse_if_match(req.headers()); let body = read_json::(req.into_body(), body_limit).await?; - let requested_enabled = body.enabled; let result = create_user(body, expected_revision, shared).await; let (mut data, revision) = match result { Ok(ok) => ok, @@ -34,22 +33,6 @@ pub(super) async fn create_user_route( }; let runtime_cfg = config_rx.borrow().clone(); data.user.in_runtime = runtime_cfg.access.users.contains_key(&data.user.username); - if let Some(enabled) = requested_enabled { - let (_, cancelled) = shared - .proxy_shared - .set_user_enabled(&data.user.username, enabled); - if !enabled { - if cancelled > 0 { - shared.runtime_events.record( - "api.user.disable.runtime", - format!( - "username={} cancelled_sessions={}", - data.user.username, cancelled - ), - ); - } - } - } shared.runtime_events.record( "api.user.create.ok", format!("username={}", data.user.username), diff --git a/src/api/handler/user_routes.rs b/src/api/handler/user_routes.rs index 72ec55d..ca50e0d 100644 --- a/src/api/handler/user_routes.rs +++ b/src/api/handler/user_routes.rs @@ -61,7 +61,6 @@ pub(super) async fn handle( }; let runtime_cfg = config_rx.borrow().clone(); data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); - shared.proxy_shared.set_user_enabled(base_user, true); shared .runtime_events .record("api.user.enable.ok", format!("username={}", base_user)); @@ -104,13 +103,9 @@ pub(super) async fn handle( }; let runtime_cfg = config_rx.borrow().clone(); data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); - let (newly_disabled, cancelled) = shared.proxy_shared.set_user_enabled(base_user, false); shared.runtime_events.record( "api.user.disable.ok", - format!( - "username={} newly_disabled={} cancelled_sessions={}", - base_user, newly_disabled, cancelled - ), + format!("username={}", base_user), ); let status = if data.in_runtime { StatusCode::OK @@ -270,11 +265,6 @@ pub(super) async fn handle( } let expected_revision = parse_if_match(req.headers()); let body = read_json::(req.into_body(), body_limit).await?; - let enabled_update = match &body.enabled { - Patch::Unchanged => None, - Patch::Remove => Some(true), - Patch::Set(enabled) => Some(*enabled), - }; let result = patch_user(user, body, expected_revision, shared).await; let (mut data, revision) = match result { Ok(ok) => ok, @@ -288,20 +278,6 @@ pub(super) async fn handle( }; let runtime_cfg = config_rx.borrow().clone(); data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); - if let Some(enabled) = enabled_update { - let (_, cancelled) = shared - .proxy_shared - .set_user_enabled(&data.username, enabled); - if !enabled { - shared.runtime_events.record( - "api.user.disable.runtime", - format!( - "username={} cancelled_sessions={}", - data.username, cancelled - ), - ); - } - } shared .runtime_events .record("api.user.patch.ok", format!("username={}", data.username)); @@ -335,11 +311,9 @@ pub(super) async fn handle( return Err(error); } }; - shared.proxy_shared.set_user_enabled(&deleted_user, true); - let cancelled = shared.proxy_shared.cancel_user_sessions(&deleted_user); shared.runtime_events.record( "api.user.delete.ok", - format!("username={} cancelled_sessions={}", deleted_user, cancelled), + format!("username={}", deleted_user), ); let runtime_cfg = config_rx.borrow().clone(); let in_runtime = runtime_cfg.access.users.contains_key(&deleted_user); diff --git a/src/api/mod.rs b/src/api/mod.rs index e4aad93..edf288a 100644 --- a/src/api/mod.rs +++ b/src/api/mod.rs @@ -70,7 +70,6 @@ use model::{ PatchUserRequest, ResetUserQuotaResponse, RotateSecretRequest, SummaryData, UserActiveIps, is_valid_username, }; -use patch::Patch; use runtime_edge::{ EdgeConnectionsCacheEntry, build_runtime_connections_summary_data, build_runtime_events_recent_data, build_runtime_tls_fingerprints_data, diff --git a/src/api/users/create.rs b/src/api/users/create.rs index a46ed59..3ee3d4a 100644 --- a/src/api/users/create.rs +++ b/src/api/users/create.rs @@ -124,7 +124,14 @@ pub(in crate::api) async fn create_user( let revision = save_access_sections_to_disk(&shared.config_path, &cfg, &touched_sections).await?; - drop(_guard); + shared + .proxy_shared + .stage_user( + &body.username, + &secret, + 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 @@ -132,6 +139,7 @@ pub(in crate::api) async fn create_user( .set_user_limit(&body.username, limit) .await; } + drop(_guard); let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); let users = users_from_config( diff --git a/src/api/users/lifecycle.rs b/src/api/users/lifecycle.rs index a68850e..546b53e 100644 --- a/src/api/users/lifecycle.rs +++ b/src/api/users/lifecycle.rs @@ -31,6 +31,10 @@ pub(in crate::api) async fn rotate_secret( .map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?; let revision = save_access_sections_to_disk(&shared.config_path, &cfg, &[AccessSection::Users]).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"))?; drop(_guard); let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); @@ -109,6 +113,7 @@ pub(in crate::api) async fn delete_user( .map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?; let revision = save_access_sections_to_disk(&shared.config_path, &cfg, &touched_sections).await?; + let deleted_incarnation = shared.proxy_shared.delete_user(user).incarnation; let configured_users = cfg.access.users.keys().cloned().collect(); if let Err(error) = shared .quota_state @@ -121,9 +126,12 @@ pub(in crate::api) async fn delete_user( "Deleted user quota checkpoint cleanup will be reconciled on restart" ); } - drop(_guard); shared.ip_tracker.remove_user_limit(user).await; - shared.ip_tracker.clear_user_ips(user).await; + shared + .ip_tracker + .clear_user_ips_if_not_newer(user, deleted_incarnation) + .await; + drop(_guard); Ok((user.to_string(), revision)) } diff --git a/src/api/users/update.rs b/src/api/users/update.rs index bfd1af2..4ff090c 100644 --- a/src/api/users/update.rs +++ b/src/api/users/update.rs @@ -170,12 +170,23 @@ pub(in crate::api) async fn patch_user( } else { save_access_sections_to_disk(&shared.config_path, &cfg, &touched_sections).await? }; - drop(_guard); + if touches_users || touches_user_enabled { + let secret = cfg + .access + .users + .get(user) + .ok_or_else(|| ApiFailure::internal("updated user secret is missing"))?; + shared + .proxy_shared + .stage_user(user, secret, cfg.access.is_user_enabled(user)) + .ok_or_else(|| ApiFailure::internal("failed to stage user admission policy"))?; + } match max_unique_ips_change { Some(Some(limit)) => shared.ip_tracker.set_user_limit(user, limit).await, Some(None) => shared.ip_tracker.remove_user_limit(user).await, None => {} } + drop(_guard); let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); let users = users_from_config( &cfg, @@ -223,6 +234,15 @@ pub(in crate::api) async fn set_user_enabled( let revision = save_access_sections_to_disk(&shared.config_path, &cfg, &[AccessSection::UserEnabled]) .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"))?; drop(_guard); let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); diff --git a/src/config/load.rs b/src/config/load.rs index e6560a2..4b19197 100644 --- a/src/config/load.rs +++ b/src/config/load.rs @@ -9,6 +9,7 @@ use rand::RngExt; use serde::{Deserialize, Serialize}; use tracing::warn; +use crate::crypto::sha256; use crate::error::{ProxyError, Result}; use super::defaults::*; @@ -230,6 +231,25 @@ impl ProxyConfig { self.runtime_user_auth.as_deref() } + /// Returns the credential identity frozen into this runtime snapshot. + pub(crate) fn runtime_user_credential_id(&self, user: &str) -> Option<[u8; 16]> { + self.runtime_user_auth() + .and_then(|snapshot| snapshot.credential_id_by_name(user)) + .or_else(|| { + self.access + .users + .get(user) + .and_then(|secret| hex::decode(secret).ok()) + .and_then(|secret| <[u8; 16]>::try_from(secret).ok()) + .map(|secret| { + let digest = sha256(&secret); + let mut credential_id = [0; 16]; + credential_id.copy_from_slice(&digest[..16]); + credential_id + }) + }) + } + /// Validates cross-field configuration invariants after deserialization. pub fn validate(&self) -> Result<()> { if self.access.users.is_empty() { diff --git a/src/config/load/runtime_auth.rs b/src/config/load/runtime_auth.rs index 01d1c4b..3fc914b 100644 --- a/src/config/load/runtime_auth.rs +++ b/src/config/load/runtime_auth.rs @@ -2,6 +2,7 @@ use std::collections::HashMap; use std::collections::hash_map::DefaultHasher; use std::hash::Hasher; +use crate::crypto::sha256; use crate::error::{ProxyError, Result}; const ACCESS_SECRET_BYTES: usize = 16; @@ -19,6 +20,8 @@ pub(crate) struct UserAuthSnapshot { pub(crate) struct UserAuthEntry { pub(crate) user: String, pub(crate) secret: [u8; ACCESS_SECRET_BYTES], + /// Stable secret identity used by process-wide admission fencing. + pub(crate) credential_id: [u8; 16], } impl UserAuthSnapshot { @@ -46,9 +49,13 @@ impl UserAuthSnapshot { let mut secret = [0u8; ACCESS_SECRET_BYTES]; secret.copy_from_slice(&decoded); + let digest = sha256(&secret); + let mut credential_id = [0; 16]; + credential_id.copy_from_slice(&digest[..16]); entries.push(UserAuthEntry { user: user.clone(), secret, + credential_id, }); by_name.insert(user.clone(), user_id); sni_index @@ -88,6 +95,12 @@ impl UserAuthSnapshot { self.entries.get(idx) } + pub(crate) fn credential_id_by_name(&self, user: &str) -> Option<[u8; 16]> { + self.user_id_by_name(user) + .and_then(|user_id| self.entry_by_id(user_id)) + .map(|entry| entry.credential_id) + } + pub(crate) fn sni_candidates(&self, sni: &str) -> Option<&[u32]> { self.sni_index .get(&Self::sni_lookup_hash(sni)) diff --git a/src/config/load/runtime_web.rs b/src/config/load/runtime_web.rs index 159a5fa..d99ea57 100644 --- a/src/config/load/runtime_web.rs +++ b/src/config/load/runtime_web.rs @@ -63,6 +63,7 @@ pub(super) fn rebuild(config: &mut ProxyConfig) -> Result<()> { host: vhost.host.clone(), public_addr: vhost.public_addr, user: profile.user.clone(), + credential_id: auth_entry.credential_id, secret_mode: profile.secret_mode, carrier: config.web.carrier, carrier_negotiation_enabled: config.web.carrier_negotiation_enabled(), diff --git a/src/config/types/web/runtime.rs b/src/config/types/web/runtime.rs index 0485ee2..3922889 100644 --- a/src/config/types/web/runtime.rs +++ b/src/config/types/web/runtime.rs @@ -35,6 +35,8 @@ pub(crate) struct WebRuntimeProfile { pub(crate) public_addr: SocketAddr, /// Exact access user authenticated by logical streams. pub(crate) user: String, + /// Stable credential identity used by process-wide admission fencing. + pub(crate) credential_id: [u8; 16], /// Client secret representation and inner protocol policy. pub(crate) secret_mode: WebSecretMode, /// Sole carrier or final fallback frozen into the issued bridge policy. diff --git a/src/ip_tracker.rs b/src/ip_tracker.rs index a54410f..990a405 100644 --- a/src/ip_tracker.rs +++ b/src/ip_tracker.rs @@ -15,6 +15,7 @@ use arc_swap::ArcSwap; use tokio::sync::{Mutex as AsyncMutex, RwLock}; use crate::config::UserMaxUniqueIpsMode; +use crate::proxy::user_admission::UserIncarnation; const CLEANUP_DRAIN_BATCH_LIMIT: usize = 1024; const MAX_ACTIVE_IP_ENTRIES: u64 = 131_072; @@ -32,11 +33,12 @@ mod tests; struct UserIpShard { active_ips: HashMap>, recent_ips: HashMap>, + incarnations: HashMap, } #[derive(Debug, Default)] struct CleanupShard { - queue: Mutex>>, + queue: Mutex>>, } #[derive(Debug, Clone)] @@ -194,19 +196,19 @@ impl UserIpTracker { } pub(super) fn pop_one_cleanup( - queue: &mut HashMap>, - ) -> Option<(String, IpAddr, usize)> { - let user = queue.keys().next().cloned()?; - let ip = queue.get(&user)?.keys().next().copied()?; - let count = queue.get_mut(&user)?.remove(&ip)?; + queue: &mut HashMap<(String, UserIncarnation), HashMap>, + ) -> Option<(String, UserIncarnation, IpAddr, usize)> { + let owner = queue.keys().next().cloned()?; + let ip = queue.get(&owner)?.keys().next().copied()?; + let count = queue.get_mut(&owner)?.remove(&ip)?; let remove_user = queue - .get(&user) + .get(&owner) .map(|user_queue| user_queue.is_empty()) .unwrap_or(false); if remove_user { - queue.remove(&user); + queue.remove(&owner); } - Some((user, ip, count)) + Some((owner.0, owner.1, ip, count)) } #[cfg(test)] diff --git a/src/ip_tracker/admission.rs b/src/ip_tracker/admission.rs index df3bd98..26e477b 100644 --- a/src/ip_tracker/admission.rs +++ b/src/ip_tracker/admission.rs @@ -59,6 +59,16 @@ impl UserIpTracker { } pub async fn check_and_add(&self, username: &str, ip: IpAddr) -> Result<(), String> { + self.check_and_add_for_incarnation(username, 0, ip).await + } + + /// Reserves an IP slot for one exact user incarnation. + pub(crate) async fn check_and_add_for_incarnation( + &self, + username: &str, + incarnation: UserIncarnation, + ip: IpAddr, + ) -> Result<(), String> { self.drain_cleanup_for_user(username).await; self.maybe_compact_empty_users().await; let policy = self.limit_policy.load(); @@ -69,6 +79,30 @@ impl UserIpTracker { let shard_idx = Self::shard_idx(username); let mut shard = self.shards[shard_idx].write().await; + if let Some(current) = shard.incarnations.get(username).copied() { + if current > incarnation { + return Err(format!( + "IP tracker rejected stale user incarnation for '{username}'" + )); + } + if current < incarnation { + let removed_active = shard + .active_ips + .remove(username) + .map(|ips| ips.len()) + .unwrap_or(0); + let removed_recent = shard + .recent_ips + .remove(username) + .map(|ips| ips.len()) + .unwrap_or(0); + Self::decrement_counter(&self.active_entry_count, removed_active); + Self::decrement_counter(&self.recent_entry_count, removed_recent); + shard.incarnations.insert(username.to_string(), incarnation); + } + } else { + shard.incarnations.insert(username.to_string(), incarnation); + } let user_active = shard.active_ips.entry(username.to_string()).or_default(); let active_contains_ip = user_active.contains_key(&ip); let active_len = user_active.len(); @@ -174,9 +208,22 @@ impl UserIpTracker { } pub async fn remove_ip(&self, username: &str, ip: IpAddr) { + self.remove_ip_for_incarnation(username, 0, ip).await; + } + + /// Releases an IP slot only from the incarnation that acquired it. + pub(crate) async fn remove_ip_for_incarnation( + &self, + username: &str, + incarnation: UserIncarnation, + ip: IpAddr, + ) { self.maybe_compact_empty_users().await; let shard_idx = Self::shard_idx(username); let mut shard = self.shards[shard_idx].write().await; + if shard.incarnations.get(username).copied() != Some(incarnation) { + return; + } let mut removed_active_entries = 0usize; if let Some(user_ips) = shard.active_ips.get_mut(username) { if let Some(count) = user_ips.get_mut(&ip) { diff --git a/src/ip_tracker/cleanup.rs b/src/ip_tracker/cleanup.rs index 86b4f8b..a884009 100644 --- a/src/ip_tracker/cleanup.rs +++ b/src/ip_tracker/cleanup.rs @@ -3,12 +3,22 @@ use super::*; impl UserIpTracker { /// Queues a deferred active IP cleanup for a later async drain. pub fn enqueue_cleanup(&self, user: String, ip: IpAddr) { + self.enqueue_cleanup_for_incarnation(user, 0, ip); + } + + /// Queues cleanup for the exact user incarnation that owns the reservation. + pub(crate) fn enqueue_cleanup_for_incarnation( + &self, + user: String, + incarnation: UserIncarnation, + ip: IpAddr, + ) { self.observe_cleanup_poison_for_tests(); let shard_idx = Self::shard_idx(&user); let cleanup_shard = &self.cleanup_shards[shard_idx]; match cleanup_shard.queue.lock() { Ok(mut queue) => { - let user_queue = queue.entry(user).or_default(); + let user_queue = queue.entry((user, incarnation)).or_default(); let count = user_queue.entry(ip).or_insert(0); if *count == 0 { self.cleanup_queue_len.fetch_add(1, Ordering::Relaxed); @@ -19,7 +29,7 @@ impl UserIpTracker { } Err(poisoned) => { let mut queue = poisoned.into_inner(); - let user_queue = queue.entry(user.clone()).or_default(); + let user_queue = queue.entry((user.clone(), incarnation)).or_default(); let count = user_queue.entry(ip).or_insert(0); if *count == 0 { self.cleanup_queue_len.fetch_add(1, Ordering::Relaxed); @@ -65,10 +75,10 @@ impl UserIpTracker { let shard_idx = Self::shard_idx(user); let cleanup_shard = &self.cleanup_shards[shard_idx]; let to_remove = match cleanup_shard.queue.lock() { - Ok(mut queue) => queue.remove(user).unwrap_or_default(), + Ok(mut queue) => drain_user_cleanup(&mut queue, user), Err(poisoned) => { let mut queue = poisoned.into_inner(); - let drained = queue.remove(user).unwrap_or_default(); + let drained = drain_user_cleanup(&mut queue, user); cleanup_shard.queue.clear_poison(); drained } @@ -76,14 +86,23 @@ impl UserIpTracker { if to_remove.is_empty() { return; } + let removed_queue_entries = to_remove + .iter() + .map(|(_, ips)| ips.len()) + .sum::(); self.cleanup_queue_len - .fetch_sub(to_remove.len() as u64, Ordering::Relaxed); + .fetch_sub(removed_queue_entries as u64, Ordering::Relaxed); let mut shard = self.shards[shard_idx].write().await; let mut removed_active_entries = 0usize; - for (ip, pending_count) in to_remove { - removed_active_entries = removed_active_entries.saturating_add( - Self::apply_active_cleanup(&mut shard.active_ips, user, ip, pending_count), - ); + for (incarnation, ips) in to_remove { + if shard.incarnations.get(user).copied() != Some(incarnation) { + continue; + } + for (ip, pending_count) in ips { + removed_active_entries = removed_active_entries.saturating_add( + Self::apply_active_cleanup(&mut shard.active_ips, user, ip, pending_count), + ); + } } Self::decrement_counter(&self.active_entry_count, removed_active_entries); } @@ -103,11 +122,13 @@ impl UserIpTracker { let mut drained = HashMap::with_capacity(queue.len().min(CLEANUP_DRAIN_BATCH_LIMIT)); for _ in 0..CLEANUP_DRAIN_BATCH_LIMIT { - let Some((user, ip, count)) = Self::pop_one_cleanup(&mut queue) else { + let Some((user, incarnation, ip, count)) = + Self::pop_one_cleanup(&mut queue) + else { break; }; self.cleanup_queue_len.fetch_sub(1, Ordering::Relaxed); - drained.insert((user, ip), count); + drained.insert((user, incarnation, ip), count); } drained } @@ -120,11 +141,13 @@ impl UserIpTracker { let mut drained = HashMap::with_capacity(queue.len().min(CLEANUP_DRAIN_BATCH_LIMIT)); for _ in 0..CLEANUP_DRAIN_BATCH_LIMIT { - let Some((user, ip, count)) = Self::pop_one_cleanup(&mut queue) else { + let Some((user, incarnation, ip, count)) = + Self::pop_one_cleanup(&mut queue) + else { break; }; self.cleanup_queue_len.fetch_sub(1, Ordering::Relaxed); - drained.insert((user, ip), count); + drained.insert((user, incarnation, ip), count); } cleanup_shard.queue.clear_poison(); drained @@ -138,7 +161,10 @@ impl UserIpTracker { let mut shard = self.shards[shard_idx].write().await; let mut removed_active_entries = 0usize; - for ((user, ip), pending_count) in to_remove { + for ((user, incarnation, ip), pending_count) in to_remove { + if shard.incarnations.get(&user).copied() != Some(incarnation) { + continue; + } removed_active_entries = removed_active_entries.saturating_add( Self::apply_active_cleanup(&mut shard.active_ips, &user, ip, pending_count), ); @@ -146,3 +172,21 @@ impl UserIpTracker { Self::decrement_counter(&self.active_entry_count, removed_active_entries); } } + +fn drain_user_cleanup( + queue: &mut HashMap<(String, UserIncarnation), HashMap>, + user: &str, +) -> Vec<(UserIncarnation, HashMap)> { + let owners = queue + .keys() + .filter(|(queued_user, _)| queued_user == user) + .cloned() + .collect::>(); + owners + .into_iter() + .filter_map(|owner| { + let incarnation = owner.1; + queue.remove(&owner).map(|ips| (incarnation, ips)) + }) + .collect() +} diff --git a/src/ip_tracker/snapshot.rs b/src/ip_tracker/snapshot.rs index f696c4e..816e01b 100644 --- a/src/ip_tracker/snapshot.rs +++ b/src/ip_tracker/snapshot.rs @@ -228,8 +228,25 @@ impl UserIpTracker { } pub async fn clear_user_ips(&self, username: &str) { + self.clear_user_ips_if_not_newer(username, 0).await; + } + + /// Clears state while advancing the username fence to a newer incarnation. + pub(crate) async fn clear_user_ips_if_not_newer( + &self, + username: &str, + incarnation: UserIncarnation, + ) { let shard_idx = Self::shard_idx(username); let mut shard = self.shards[shard_idx].write().await; + if shard + .incarnations + .get(username) + .is_some_and(|current| *current > incarnation) + { + return; + } + shard.incarnations.insert(username.to_string(), incarnation); let removed_active_entries = shard .active_ips .remove(username) @@ -250,6 +267,7 @@ impl UserIpTracker { let mut shard = shard_lock.write().await; shard.active_ips.clear(); shard.recent_ips.clear(); + shard.incarnations.clear(); } self.active_entry_count.store(0, Ordering::Relaxed); self.recent_entry_count.store(0, Ordering::Relaxed); diff --git a/src/maestro/orchestrator.rs b/src/maestro/orchestrator.rs index cf0fc61..8dad10c 100644 --- a/src/maestro/orchestrator.rs +++ b/src/maestro/orchestrator.rs @@ -108,7 +108,7 @@ pub(super) async fn run_telemt_core( ); let shared_state = ProxySharedState::new_with_direct_buffer_budget(direct_buffer_budget.clone()); - shared_state.apply_user_enabled_config(&config.access.user_enabled); + shared_state.apply_user_config(&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 d186745..a4372e4 100644 --- a/src/maestro/reload_supervisor.rs +++ b/src/maestro/reload_supervisor.rs @@ -177,6 +177,7 @@ impl ReloadSupervisor { self.quota_store.clone(), self.runtime_log_filter.clone(), self.tls_full_cert_budget.clone(), + old_runtime.proxy_shared.user_admission(), ) .await { @@ -277,6 +278,7 @@ impl ReloadSupervisor { generation: new_runtime, detected_ips, config_watcher_activation, + user_admission_epoch, } = prepared; let pending_listener_transition = if let Some(listener_transition) = listener_transition { match self @@ -300,6 +302,12 @@ 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, + &config.access.users, + &config.access.user_enabled, + ); old_runtime.stop_accepting_sessions(); listener_manager.activate_runtime_generation(new_runtime.clone()) }; diff --git a/src/maestro/reload_supervisor_tests.rs b/src/maestro/reload_supervisor_tests.rs index 27b6e4b..57f6adf 100644 --- a/src/maestro/reload_supervisor_tests.rs +++ b/src/maestro/reload_supervisor_tests.rs @@ -23,10 +23,12 @@ fn runtime_log_filter() -> RuntimeLogFilter { fn prepared_runtime(generation: Arc) -> PreparedRuntime { let (config_watcher_activation, _activation_rx) = watch::channel(false); + let user_admission_epoch = generation.proxy_shared.user_admission().epoch(); PreparedRuntime { generation, detected_ips: (None, None), config_watcher_activation, + user_admission_epoch, } } diff --git a/src/maestro/runtime_build.rs b/src/maestro/runtime_build.rs index 638478f..16b401d 100644 --- a/src/maestro/runtime_build.rs +++ b/src/maestro/runtime_build.rs @@ -16,6 +16,7 @@ use crate::proxy::direct_buffer_budget::{ }; use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController}; use crate::proxy::shared_state::ProxySharedState; +use crate::proxy::user_admission::UserAdmissionAuthority; use crate::startup::StartupTracker; use crate::stats::beobachten::BeobachtenStore; use crate::stats::telemetry::TelemetryPolicy; @@ -39,6 +40,8 @@ pub(crate) struct PreparedRuntime { pub(crate) detected_ips: (Option, Option), /// Gate opened only after the candidate becomes the active generation. pub(crate) config_watcher_activation: watch::Sender, + /// User-authority epoch captured before candidate construction. + pub(crate) user_admission_epoch: u64, } pub(crate) async fn prepare_runtime( @@ -48,7 +51,9 @@ pub(crate) async fn prepare_runtime( quota_store: Arc, runtime_log_filter: RuntimeLogFilter, tls_full_cert_budget: Arc, + user_admission: Arc, ) -> Result { + let user_admission_epoch = user_admission.epoch(); config .validate_web_decoy_listener_separation() .map_err(|error| error.to_string())?; @@ -92,9 +97,10 @@ pub(crate) async fn prepare_runtime( let hard_limit = resolve_direct_buffer_hard_limit(config.general.direct_relay_buffer_budget_max_bytes).await; let direct_buffer_budget = DirectBufferBudget::new(hard_limit); - let proxy_shared = - ProxySharedState::new_with_direct_buffer_budget(direct_buffer_budget.clone()); - proxy_shared.apply_user_enabled_config(&config.access.user_enabled); + let proxy_shared = ProxySharedState::new_with_direct_buffer_budget_and_user_admission( + direct_buffer_budget.clone(), + user_admission, + ); proxy_shared.traffic_limiter.apply_policy( config.access.user_rate_limits.clone(), config.access.cidr_rate_limits.clone(), @@ -311,6 +317,7 @@ pub(crate) async fn prepare_runtime( Ok(PreparedRuntime { generation, config_watcher_activation, + user_admission_epoch, detected_ips: ( probe.detected_ipv4.map(IpAddr::V4), probe.detected_ipv6.map(IpAddr::V6), diff --git a/src/maestro/runtime_tasks.rs b/src/maestro/runtime_tasks.rs index c8a302a..300f4e7 100644 --- a/src/maestro/runtime_tasks.rs +++ b/src/maestro/runtime_tasks.rs @@ -288,8 +288,8 @@ 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_enabled_config(&cfg.access.user_enabled) + for (user, cancelled) in shared_user_enabled + .apply_user_config(&cfg.access.users, &cfg.access.user_enabled) { if cancelled > 0 { info!( diff --git a/src/proxy/authenticated.rs b/src/proxy/authenticated.rs index 77e0c9d..26f3ec7 100644 --- a/src/proxy/authenticated.rs +++ b/src/proxy/authenticated.rs @@ -14,6 +14,7 @@ use crate::proxy::handshake::HandshakeSuccess; use crate::proxy::middle_relay::{handle_via_middle_proxy, handle_via_middle_proxy_with_conntrack}; 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::stream::{BufferPool, CryptoReader, CryptoWriter}; use crate::transport::UpstreamManager; @@ -59,13 +60,21 @@ where W: AsyncWrite + Unpin + Send + 'static, { let user = success.user.clone(); - if !deps.shared.is_user_enabled(&user) { + let Some(credential_id) = deps.config.runtime_user_credential_id(&user) else { + warn!(user = %user, "Authenticated user is absent from the runtime credential snapshot"); + return Err(ProxyError::UserDisabled { user }); + }; + let Some(user_incarnation) = deps + .shared + .authenticated_user_incarnation(&user, credential_id) + else { warn!(user = %user, "Disabled user rejected"); return Err(ProxyError::UserDisabled { user }); - } + }; - let user_reservation = acquire_user_connection_reservation( + let user_reservation = acquire_user_connection_reservation_for_incarnation( &user, + user_incarnation, &deps.config, Arc::clone(&deps.stats), peer_addr, @@ -79,11 +88,20 @@ where let route_snapshot = deps.route_runtime.snapshot(); let session_id = deps.rng.u64(); - let Some(user_session) = deps.shared.register_user_session(&user, session_id) else { + let Some(user_session) = deps + .shared + .register_authenticated_user_session(&user, credential_id) + else { user_reservation.release_deferred(); warn!(user = %user, "Disabled user rejected during final admission"); return Err(ProxyError::UserDisabled { user }); }; + if user_session.incarnation() != user_incarnation { + drop(user_session); + user_reservation.release_deferred(); + warn!(user = %user, "User incarnation changed during admission"); + return Err(ProxyError::UserDisabled { user }); + } let session_cancel = user_session.token(); let selected_me_pool = if deps.config.general.use_middle_proxy && matches!(route_snapshot.mode, RelayRouteMode::Middle) @@ -216,6 +234,7 @@ pub(crate) struct UserConnectionReservation { ip_tracker: Arc, user: String, ip: IpAddr, + incarnation: UserIncarnation, tracks_ip: bool, active: bool, } @@ -228,12 +247,25 @@ impl UserConnectionReservation { user: String, ip: IpAddr, tracks_ip: bool, + ) -> Self { + Self::new_for_incarnation(stats, ip_tracker, user, ip, 0, tracks_ip) + } + + /// Creates a reservation fenced to one authenticated user incarnation. + pub(crate) fn new_for_incarnation( + stats: Arc, + ip_tracker: Arc, + user: String, + ip: IpAddr, + incarnation: UserIncarnation, + tracks_ip: bool, ) -> Self { Self { stats, ip_tracker, user, ip, + incarnation, tracks_ip, active: true, } @@ -246,7 +278,9 @@ impl UserConnectionReservation { } self.active = false; if self.tracks_ip { - self.ip_tracker.remove_ip(&self.user, self.ip).await; + self.ip_tracker + .remove_ip_for_incarnation(&self.user, self.incarnation, self.ip) + .await; } self.stats.decrement_user_curr_connects(&self.user); } @@ -259,7 +293,11 @@ impl UserConnectionReservation { self.active = false; self.stats.decrement_user_curr_connects(&self.user); if self.tracks_ip { - self.ip_tracker.enqueue_cleanup(self.user.clone(), self.ip); + self.ip_tracker.enqueue_cleanup_for_incarnation( + self.user.clone(), + self.incarnation, + self.ip, + ); } } } @@ -273,7 +311,11 @@ impl Drop for UserConnectionReservation { self.stats.increment_session_drop_fallback_total(); self.stats.decrement_user_curr_connects(&self.user); if self.tracks_ip { - self.ip_tracker.enqueue_cleanup(self.user.clone(), self.ip); + self.ip_tracker.enqueue_cleanup_for_incarnation( + self.user.clone(), + self.incarnation, + self.ip, + ); } } } @@ -285,6 +327,25 @@ pub(crate) async fn acquire_user_connection_reservation( stats: Arc, peer_addr: SocketAddr, ip_tracker: Arc, +) -> Result { + acquire_user_connection_reservation_for_incarnation( + user, + 0, + config, + stats, + peer_addr, + ip_tracker, + ) + .await +} + +async fn acquire_user_connection_reservation_for_incarnation( + user: &str, + incarnation: UserIncarnation, + config: &ProxyConfig, + stats: Arc, + peer_addr: SocketAddr, + ip_tracker: Arc, ) -> Result { if let Some(expiration) = config.access.user_expirations.get(user) && chrono::Utc::now() > *expiration @@ -316,7 +377,10 @@ pub(crate) async fn acquire_user_connection_reservation( }); } - if let Err(reason) = ip_tracker.check_and_add(user, peer_addr.ip()).await { + if let Err(reason) = ip_tracker + .check_and_add_for_incarnation(user, incarnation, peer_addr.ip()) + .await + { stats.decrement_user_curr_connects(user); warn!( user = %user, @@ -329,11 +393,12 @@ pub(crate) async fn acquire_user_connection_reservation( }); } - Ok(UserConnectionReservation::new( + Ok(UserConnectionReservation::new_for_incarnation( stats, ip_tracker, user.to_string(), peer_addr.ip(), + incarnation, true, )) } diff --git a/src/proxy/mod.rs b/src/proxy/mod.rs index b1b71bf..7908e12 100644 --- a/src/proxy/mod.rs +++ b/src/proxy/mod.rs @@ -73,6 +73,7 @@ pub mod route_mode; pub mod session_eviction; pub mod shared_state; pub mod traffic_limiter; +pub(crate) mod user_admission; pub use client::ClientHandler; #[allow(unused_imports)] diff --git a/src/proxy/shared_state.rs b/src/proxy/shared_state.rs index 9623b0c..706388d 100644 --- a/src/proxy/shared_state.rs +++ b/src/proxy/shared_state.rs @@ -6,14 +6,16 @@ use std::sync::{Arc, Mutex}; use std::time::Instant; use dashmap::DashMap; -use parking_lot::Mutex as ParkingMutex; use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc}; -use tokio_util::sync::CancellationToken; use crate::proxy::direct_buffer_budget::{DirectBufferBudget, fallback_direct_buffer_hard_limit}; use crate::proxy::handshake::{AuthProbeSaturationState, AuthProbeState}; use crate::proxy::middle_relay::{DesyncDedupRotationState, RelayIdleCandidateRegistry}; use crate::proxy::traffic_limiter::TrafficLimiter; +use crate::proxy::user_admission::{ + UserAdmissionAuthority, UserAdmissionPublication, UserCredentialId, UserIncarnation, + UserMutationResult, UserSessionRegistration, +}; const HANDSHAKE_RECENT_USER_RING_LEN: usize = 64; const MASKING_FALLBACK_MAX_CONCURRENT: usize = 512; @@ -76,57 +78,17 @@ pub(crate) struct MiddleRelaySharedState { pub(crate) relay_idle_mark_seq: AtomicU64, } -#[derive(Default)] -struct UserAdmissionState { - disabled_users: HashSet, - sessions_by_user: HashMap>, -} - pub(crate) struct ProxySharedState { pub(crate) handshake: HandshakeSharedState, pub(crate) middle_relay: MiddleRelaySharedState, pub(crate) traffic_limiter: Arc, pub(crate) direct_buffer_budget: Arc, - user_admission: ParkingMutex, + user_admission: Arc, pub(crate) conntrack_pressure_active: AtomicBool, pub(crate) conntrack_close_tx: Mutex>>, masking_fallback_permits: Arc, } -#[must_use = "registered user sessions must be kept alive until relay completion"] -pub(crate) struct UserSessionRegistration { - token: CancellationToken, - _guard: UserSessionGuard, -} - -impl UserSessionRegistration { - pub(crate) fn token(&self) -> CancellationToken { - self.token.clone() - } -} - -struct UserSessionGuard { - shared: Arc, - key: (String, u64), -} - -impl Drop for UserSessionGuard { - fn drop(&mut self) { - let mut admission = self.shared.user_admission.lock(); - let remove_user = admission - .sessions_by_user - .get_mut(&self.key.0) - .map(|sessions| { - sessions.remove(&self.key.1); - sessions.is_empty() - }) - .unwrap_or(false); - if remove_user { - admission.sessions_by_user.remove(&self.key.0); - } - } -} - impl ProxySharedState { pub(crate) fn new() -> Arc { Self::new_with_direct_buffer_budget(DirectBufferBudget::new( @@ -137,6 +99,17 @@ impl ProxySharedState { /// Creates process state with the startup-resolved Direct buffer envelope. pub(crate) fn new_with_direct_buffer_budget( direct_buffer_budget: Arc, + ) -> Arc { + Self::new_with_direct_buffer_budget_and_user_admission( + direct_buffer_budget, + UserAdmissionAuthority::new(), + ) + } + + /// Creates generation state around one process-owned user authority. + pub(crate) fn new_with_direct_buffer_budget_and_user_admission( + direct_buffer_budget: Arc, + user_admission: Arc, ) -> Arc { Arc::new(Self { handshake: HandshakeSharedState { @@ -167,7 +140,7 @@ impl ProxySharedState { }, traffic_limiter: TrafficLimiter::new(), direct_buffer_budget, - user_admission: ParkingMutex::new(UserAdmissionState::default()), + user_admission, conntrack_pressure_active: AtomicBool::new(false), conntrack_close_tx: Mutex::new(None), masking_fallback_permits: Arc::new(Semaphore::new(MASKING_FALLBACK_MAX_CONCURRENT)), @@ -183,106 +156,91 @@ impl ProxySharedState { } pub(crate) fn is_user_enabled(&self, user: &str) -> bool { - !self.user_admission.lock().disabled_users.contains(user) + self.user_admission.is_user_enabled(user) } - pub(crate) fn set_user_enabled(&self, user: &str, enabled: bool) -> (bool, usize) { - let (newly_disabled, tokens) = { - let mut admission = self.user_admission.lock(); - if enabled { - admission.disabled_users.remove(user); - (false, Vec::new()) - } else { - let newly_disabled = admission.disabled_users.insert(user.to_string()); - let tokens = admission - .sessions_by_user - .get(user) - .map(|sessions| sessions.values().cloned().collect()) - .unwrap_or_default(); - (newly_disabled, tokens) - } - }; - for token in &tokens { - token.cancel(); - } - (newly_disabled, tokens.len()) + /// Returns the process authority shared by every runtime generation. + pub(crate) fn user_admission(&self) -> Arc { + Arc::clone(&self.user_admission) } - pub(crate) fn apply_user_enabled_config( + /// Reconciles the complete user authentication policy from configuration. + pub(crate) fn apply_user_config( &self, + users: &HashMap, user_enabled: &HashMap, ) -> Vec<(String, usize)> { - let desired_disabled = user_enabled - .iter() - .filter_map(|(user, enabled)| (!*enabled).then_some(user.clone())) - .collect::>(); - let cancellations = { - let mut admission = self.user_admission.lock(); - let newly_disabled = desired_disabled - .difference(&admission.disabled_users) - .cloned() - .collect::>(); - admission.disabled_users = desired_disabled; - newly_disabled - .into_iter() - .map(|user| { - let tokens = admission - .sessions_by_user - .get(&user) - .map(|sessions| sessions.values().cloned().collect()) - .unwrap_or_default(); - (user, tokens) - }) - .collect::)>>() - }; - cancellations - .into_iter() - .map(|(user, tokens)| { - for token in &tokens { - token.cancel(); - } - (user, tokens.len()) - }) - .collect() + 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( + &self, + expected_epoch: u64, + users: &HashMap, + user_enabled: &HashMap, + ) -> Option> { + self.user_admission + .apply_config_if_epoch(expected_epoch, users, user_enabled) + } + + /// Applies one persisted user mutation before asynchronous config reload. + pub(crate) fn stage_user( + &self, + user: &str, + secret: &str, + enabled: bool, + ) -> Option { + self.user_admission.stage_user(user, secret, 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) + } + + /// Returns the current incarnation for an exact authenticated credential. + pub(crate) fn authenticated_user_incarnation( + &self, + user: &str, + credential_id: UserCredentialId, + ) -> Option { + self.user_admission + .authenticated_incarnation(user, credential_id) + } + + /// Starts an atomic publication boundary for an authenticated owner. + pub(crate) fn claim_authenticated_user( + self: &Arc, + user: &str, + credential_id: UserCredentialId, + ) -> Option> { + self.user_admission + .claim_authenticated(user, credential_id) } pub(crate) fn register_user_session( self: &Arc, user: &str, - session_id: u64, + _session_id: u64, ) -> Option { - let token = CancellationToken::new(); - let key = (user.to_string(), session_id); - let mut admission = self.user_admission.lock(); - if admission.disabled_users.contains(user) { - return None; - } - admission - .sessions_by_user - .entry(key.0.clone()) - .or_default() - .insert(session_id, token.clone()); - Some(UserSessionRegistration { - token, - _guard: UserSessionGuard { - shared: Arc::clone(self), - key, - }, - }) + self.user_admission.register_legacy(user) + } + + /// Registers a relay session against the exact credential that authenticated it. + pub(crate) fn register_authenticated_user_session( + self: &Arc, + user: &str, + credential_id: UserCredentialId, + ) -> Option { + let mut publication = self.claim_authenticated_user(user, credential_id)?; + let registration = publication.take_registration()?; + publication.commit(); + Some(registration) } pub(crate) fn cancel_user_sessions(&self, user: &str) -> usize { - let tokens: Vec = self - .user_admission - .lock() - .sessions_by_user - .get(user) - .map(|sessions| sessions.values().cloned().collect()) - .unwrap_or_default(); - for token in &tokens { - token.cancel(); - } - tokens.len() + self.user_admission.cancel_user_owners(user) } pub(crate) fn set_conntrack_close_sender(&self, tx: mpsc::Sender) { @@ -350,31 +308,53 @@ impl ProxySharedState { mod tests { use super::*; + const ALICE_SECRET: &str = "00112233445566778899aabbccddeeff"; + + fn configured_shared() -> Arc { + let shared = ProxySharedState::new(); + let users = HashMap::from([ + ("alice".to_string(), ALICE_SECRET.to_string()), + ( + "bob".to_string(), + "ffeeddccbbaa99887766554433221100".to_string(), + ), + ]); + shared.apply_user_config(&users, &HashMap::new()); + shared + } + #[test] fn user_enabled_config_sync_tracks_disabled_overrides() { - let shared = ProxySharedState::new(); + let shared = configured_shared(); assert!(shared.is_user_enabled("alice")); + let users = HashMap::from([ + ("alice".to_string(), ALICE_SECRET.to_string()), + ( + "bob".to_string(), + "ffeeddccbbaa99887766554433221100".to_string(), + ), + ]); let mut user_enabled = HashMap::new(); user_enabled.insert("alice".to_string(), false); user_enabled.insert("bob".to_string(), true); - let mut newly_disabled = shared.apply_user_enabled_config(&user_enabled); + let mut newly_disabled = shared.apply_user_config(&users, &user_enabled); newly_disabled.sort(); assert_eq!(newly_disabled, vec![("alice".to_string(), 0)]); assert!(!shared.is_user_enabled("alice")); assert!(shared.is_user_enabled("bob")); - assert!(shared.apply_user_enabled_config(&user_enabled).is_empty()); + assert!(shared.apply_user_config(&users, &user_enabled).is_empty()); user_enabled.clear(); - assert!(shared.apply_user_enabled_config(&user_enabled).is_empty()); + assert!(shared.apply_user_config(&users, &user_enabled).is_empty()); assert!(shared.is_user_enabled("alice")); } #[test] fn cancel_user_sessions_cancels_only_registered_matching_user() { - let shared = ProxySharedState::new(); + let shared = configured_shared(); let alice_1 = shared.register_user_session("alice", 1).unwrap(); let alice_2 = shared.register_user_session("alice", 2).unwrap(); let bob = shared.register_user_session("bob", 1).unwrap(); @@ -392,9 +372,11 @@ mod tests { #[test] fn disabled_user_cannot_register_after_the_cancellation_snapshot() { - let shared = ProxySharedState::new(); + let shared = configured_shared(); - assert_eq!(shared.set_user_enabled("alice", false), (true, 0)); + let result = shared.stage_user("alice", ALICE_SECRET, false).unwrap(); + assert!(result.newly_disabled); + assert_eq!(result.cancelled, 0); assert_eq!(shared.cancel_user_sessions("alice"), 0); let late = shared.register_user_session("alice", 1); @@ -406,15 +388,19 @@ mod tests { #[test] fn disabling_user_cancels_existing_sessions_before_return() { - let shared = ProxySharedState::new(); + let shared = configured_shared(); let registration = shared.register_user_session("alice", 1).unwrap(); let token = registration.token(); - assert_eq!(shared.set_user_enabled("alice", false), (true, 1)); + let result = shared.stage_user("alice", ALICE_SECRET, false).unwrap(); + assert!(result.newly_disabled); + assert_eq!(result.cancelled, 1); assert!(token.is_cancelled()); assert!(shared.register_user_session("alice", 2).is_none()); - assert_eq!(shared.set_user_enabled("alice", true), (false, 0)); + let result = shared.stage_user("alice", ALICE_SECRET, true).unwrap(); + assert!(!result.newly_disabled); + assert_eq!(result.cancelled, 0); assert!(shared.register_user_session("alice", 3).is_some()); } @@ -439,7 +425,8 @@ mod tests { for session_id in 0..ITERATIONS as u64 { let user = format!("user-{session_id}"); barrier.wait(); - shared.set_user_enabled(&user, false); + let secret = "00112233445566778899aabbccddeeff"; + shared.stage_user(&user, secret, false).unwrap(); } for registration in register.join().unwrap().into_iter().flatten() { diff --git a/src/proxy/user_admission.rs b/src/proxy/user_admission.rs new file mode 100644 index 0000000..8efd9fb --- /dev/null +++ b/src/proxy/user_admission.rs @@ -0,0 +1,598 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; + +use parking_lot::{Mutex, MutexGuard}; +use tokio_util::sync::CancellationToken; + +use crate::crypto::sha256; + +/// Stable secret identity used to fence authentication across runtime generations. +pub(crate) type UserCredentialId = [u8; 16]; + +/// Monotonic identity of one configured username lifetime. +pub(crate) type UserIncarnation = u64; + +#[derive(Clone, Copy, PartialEq, Eq)] +struct EffectiveUser { + credential_id: UserCredentialId, + enabled: bool, +} + +#[derive(Clone, Copy, PartialEq, Eq)] +enum UserOverride { + Present(EffectiveUser), + Deleted, +} + +struct UserRecord { + configured: Option, + mutation_override: Option, + incarnation: UserIncarnation, +} + +impl UserRecord { + fn effective(&self) -> Option { + match self.mutation_override { + Some(UserOverride::Present(user)) => Some(user), + Some(UserOverride::Deleted) => None, + None => self.configured, + } + } +} + +struct RegisteredOwner { + token: CancellationToken, + incarnation: UserIncarnation, +} + +#[derive(Default)] +struct UserAdmissionState { + initialized: bool, + epoch: u64, + next_incarnation: UserIncarnation, + next_registration_id: u64, + users: HashMap, + owners_by_user: HashMap>, +} + +impl UserAdmissionState { + fn allocate_incarnation(&mut self) -> UserIncarnation { + self.next_incarnation = self.next_incarnation.checked_add(1).unwrap_or(u64::MAX); + self.next_incarnation + } + + fn allocate_registration_id(&mut self) -> Option { + let next = self.next_registration_id.checked_add(1)?; + self.next_registration_id = next; + Some(next) + } + + fn bump_epoch(&mut self) { + self.epoch = self.epoch.checked_add(1).unwrap_or(u64::MAX); + } + + fn owner_tokens(&self, user: &str) -> Vec { + self.owners_by_user + .get(user) + .map(|owners| owners.values().map(|owner| owner.token.clone()).collect()) + .unwrap_or_default() + } +} + +/// Result of one durable user mutation applied to the process admission authority. +pub(crate) struct UserMutationResult { + /// Incarnation invalidated or created by the mutation. + pub(crate) incarnation: UserIncarnation, + /// Number of live owners cancelled by the mutation. + pub(crate) cancelled: usize, + /// Whether the effective enabled state changed from enabled to disabled. + pub(crate) newly_disabled: bool, +} + +/// Process-owned user authentication and live-owner authority. +pub(crate) struct UserAdmissionAuthority { + state: Mutex, +} + +impl UserAdmissionAuthority { + /// Creates an uninitialized authority for isolated tests and startup wiring. + pub(crate) fn new() -> Arc { + Arc::new(Self { + state: Mutex::new(UserAdmissionState::default()), + }) + } + + /// Returns the mutation epoch used to reject stale candidate configuration. + pub(crate) fn epoch(&self) -> u64 { + self.state.lock().epoch + } + + /// Reconciles the complete configured user set into the process authority. + pub(crate) fn apply_config( + &self, + users: &HashMap, + user_enabled: &HashMap, + ) -> Vec<(String, usize)> { + self.apply_config_locked(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( + &self, + expected_epoch: u64, + users: &HashMap, + user_enabled: &HashMap, + ) -> Option> { + self.apply_config_locked(Some(expected_epoch), users, user_enabled) + } + + fn apply_config_locked( + &self, + expected_epoch: Option, + users: &HashMap, + user_enabled: &HashMap, + ) -> Option> { + let configured = users + .iter() + .filter_map(|(user, secret)| { + credential_id_from_hex(secret).map(|credential_id| { + ( + user.clone(), + EffectiveUser { + credential_id, + enabled: user_enabled.get(user).copied().unwrap_or(true), + }, + ) + }) + }) + .collect::>(); + let cancellations = { + let mut state = self.state.lock(); + if expected_epoch.is_some_and(|epoch| state.epoch != epoch) { + return None; + } + + let mut changed = !state.initialized; + state.initialized = true; + let existing_users = state.users.keys().cloned().collect::>(); + let mut cancellations = Vec::new(); + + for user in existing_users { + let desired = configured.get(&user).copied(); + let old_effective = state.users.get(&user).and_then(UserRecord::effective); + let override_matches = state.users.get(&user).is_some_and(|record| { + matches!( + (record.mutation_override, desired), + (Some(UserOverride::Present(current)), Some(next)) if current == next + ) || matches!( + (record.mutation_override, desired), + (Some(UserOverride::Deleted), None) + ) + }); + + if let Some(record) = state.users.get_mut(&user) { + if record.configured != desired || override_matches { + changed = true; + } + record.configured = desired; + if override_matches { + record.mutation_override = None; + } + } + + let new_effective = state.users.get(&user).and_then(UserRecord::effective); + if old_effective != new_effective { + let identity_changed = old_effective.map(|entry| entry.credential_id) + != new_effective.map(|entry| entry.credential_id); + if identity_changed { + let incarnation = state.allocate_incarnation(); + if let Some(record) = state.users.get_mut(&user) { + record.incarnation = incarnation; + } + } + if identity_changed + || old_effective.is_some_and(|entry| entry.enabled) + && new_effective.is_none_or(|entry| !entry.enabled) + { + let tokens = state.owner_tokens(&user); + cancellations.push((user, tokens)); + } + } + } + + for (user, desired) in configured { + if state.users.contains_key(&user) { + continue; + } + changed = true; + let incarnation = state.allocate_incarnation(); + state.users.insert( + user, + UserRecord { + configured: Some(desired), + mutation_override: None, + incarnation, + }, + ); + } + + if changed { + state.bump_epoch(); + } + cancellations + }; + Some(cancel_owners(cancellations)) + } + + /// Applies one persisted user value ahead of asynchronous runtime reload. + pub(crate) fn stage_user( + &self, + user: &str, + secret: &str, + enabled: bool, + ) -> Option { + let credential_id = credential_id_from_hex(secret)?; + let desired = EffectiveUser { + credential_id, + enabled, + }; + let (incarnation, newly_disabled, tokens) = { + let mut state = self.state.lock(); + let previous = state.users.get(user).and_then(UserRecord::effective); + let identity_changed = previous.map(|entry| entry.credential_id) != Some(credential_id); + let incarnation = if identity_changed { + state.allocate_incarnation() + } else { + state + .users + .get(user) + .map(|record| record.incarnation) + .unwrap_or_else(|| state.allocate_incarnation()) + }; + let record = state.users.entry(user.to_string()).or_insert(UserRecord { + configured: None, + mutation_override: None, + incarnation, + }); + record.mutation_override = Some(UserOverride::Present(desired)); + record.incarnation = incarnation; + state.initialized = true; + state.bump_epoch(); + let newly_disabled = previous.is_some_and(|entry| entry.enabled) && !enabled; + let tokens = if identity_changed || !enabled { + state.owner_tokens(user) + } else { + Vec::new() + }; + (incarnation, newly_disabled, tokens) + }; + let cancelled = tokens.len(); + for token in tokens { + token.cancel(); + } + Some(UserMutationResult { + incarnation, + cancelled, + newly_disabled, + }) + } + + /// Installs a deletion tombstone and cancels every owner of the old incarnation. + pub(crate) fn delete_user(&self, user: &str) -> UserMutationResult { + let (incarnation, newly_disabled, tokens) = { + let mut state = self.state.lock(); + let previous = state.users.get(user).and_then(UserRecord::effective); + let incarnation = state.allocate_incarnation(); + let record = state.users.entry(user.to_string()).or_insert(UserRecord { + configured: None, + mutation_override: None, + incarnation, + }); + record.mutation_override = Some(UserOverride::Deleted); + record.incarnation = incarnation; + state.initialized = true; + state.bump_epoch(); + ( + incarnation, + previous.is_some_and(|entry| entry.enabled), + state.owner_tokens(user), + ) + }; + let cancelled = tokens.len(); + for token in tokens { + token.cancel(); + } + UserMutationResult { + incarnation, + cancelled, + newly_disabled, + } + } + + /// Returns whether the effective process policy currently enables a user. + pub(crate) fn is_user_enabled(&self, user: &str) -> bool { + let state = self.state.lock(); + if !state.initialized { + return true; + } + state + .users + .get(user) + .and_then(UserRecord::effective) + .is_some_and(|entry| entry.enabled) + } + + /// Returns the authenticated incarnation for an exact current credential. + pub(crate) fn authenticated_incarnation( + &self, + user: &str, + credential_id: UserCredentialId, + ) -> Option { + let state = self.state.lock(); + if !state.initialized { + return Some(0); + } + let record = state.users.get(user)?; + let effective = record.effective()?; + (effective.enabled && effective.credential_id == credential_id).then_some(record.incarnation) + } + + /// Starts a short publication critical section for one authenticated owner. + pub(crate) fn claim_authenticated( + self: &Arc, + user: &str, + credential_id: UserCredentialId, + ) -> Option> { + let mut state = self.state.lock(); + let incarnation = if state.initialized { + let record = state.users.get(user)?; + let effective = record.effective()?; + if !effective.enabled || effective.credential_id != credential_id { + return None; + } + record.incarnation + } else { + 0 + }; + let registration_id = state.allocate_registration_id()?; + let token = CancellationToken::new(); + let active = Arc::new(AtomicBool::new(false)); + Some(UserAdmissionPublication { + state, + authority: Arc::clone(self), + user: user.to_string(), + registration_id, + incarnation, + token, + active, + registration_taken: false, + }) + } + + /// Registers a legacy owner when no credential snapshot is available. + pub(crate) fn register_legacy( + self: &Arc, + user: &str, + ) -> Option { + let credential_id = { + let state = self.state.lock(); + if !state.initialized { + [0; 16] + } else { + state.users.get(user)?.effective()?.credential_id + } + }; + let mut publication = self.claim_authenticated(user, credential_id)?; + let registration = publication.take_registration()?; + publication.commit(); + Some(registration) + } + + /// Cancels all current owners without changing admission policy. + pub(crate) fn cancel_user_owners(&self, user: &str) -> usize { + let tokens = self.state.lock().owner_tokens(user); + let count = tokens.len(); + for token in tokens { + token.cancel(); + } + count + } + + fn unregister(&self, user: &str, registration_id: u64, incarnation: UserIncarnation) { + let mut state = self.state.lock(); + let remove_user = state + .owners_by_user + .get_mut(user) + .map(|owners| { + if owners + .get(®istration_id) + .is_some_and(|owner| owner.incarnation == incarnation) + { + owners.remove(®istration_id); + } + owners.is_empty() + }) + .unwrap_or(false); + if remove_user { + state.owners_by_user.remove(user); + } + } +} + +/// Authority lock retained until the caller publishes its owned object. +pub(crate) struct UserAdmissionPublication<'a> { + state: MutexGuard<'a, UserAdmissionState>, + authority: Arc, + user: String, + registration_id: u64, + incarnation: UserIncarnation, + token: CancellationToken, + active: Arc, + registration_taken: bool, +} + +impl UserAdmissionPublication<'_> { + /// Moves the registered owner out while retaining the authority lock. + pub(crate) fn take_registration(&mut self) -> Option { + if self.registration_taken { + return None; + } + self.registration_taken = true; + Some(UserSessionRegistration { + authority: Arc::clone(&self.authority), + user: self.user.clone(), + registration_id: self.registration_id, + incarnation: self.incarnation, + token: self.token.clone(), + active: Arc::clone(&self.active), + }) + } + + /// Commits the owner record after the caller publishes its lifecycle object. + pub(crate) fn commit(mut self) { + if !self.registration_taken { + return; + } + self.state + .owners_by_user + .entry(self.user.clone()) + .or_default() + .insert( + self.registration_id, + RegisteredOwner { + token: self.token.clone(), + incarnation: self.incarnation, + }, + ); + self.active.store(true, Ordering::Release); + } +} + +/// RAII ownership registered against one user incarnation. +#[must_use = "registered user ownership must be retained until lifecycle completion"] +pub(crate) struct UserSessionRegistration { + authority: Arc, + user: String, + registration_id: u64, + incarnation: UserIncarnation, + token: CancellationToken, + active: Arc, +} + +impl UserSessionRegistration { + /// Returns the cancellation signal for revocation or credential replacement. + pub(crate) fn token(&self) -> CancellationToken { + self.token.clone() + } + + /// Returns the immutable user incarnation owned by this registration. + pub(crate) fn incarnation(&self) -> UserIncarnation { + self.incarnation + } + + /// Returns whether revocation has cancelled this ownership. + pub(crate) fn is_cancelled(&self) -> bool { + self.token.is_cancelled() + } +} + +impl Drop for UserSessionRegistration { + fn drop(&mut self) { + if self.active.swap(false, Ordering::AcqRel) { + self.authority + .unregister(&self.user, self.registration_id, self.incarnation); + } + } +} + +/// Derives the stable credential identity from one decoded MTProxy secret. +pub(crate) fn credential_id(secret: &[u8; 16]) -> UserCredentialId { + let digest = sha256(secret); + let mut id = [0; 16]; + id.copy_from_slice(&digest[..16]); + id +} + +/// Decodes one configured secret and derives its credential identity. +pub(crate) fn credential_id_from_hex(secret: &str) -> Option { + let decoded = hex::decode(secret).ok()?; + let secret: [u8; 16] = decoded.try_into().ok()?; + Some(credential_id(&secret)) +} + +fn cancel_owners( + cancellations: Vec<(String, Vec)>, +) -> Vec<(String, usize)> { + cancellations + .into_iter() + .map(|(user, tokens)| { + let count = tokens.len(); + for token in tokens { + token.cancel(); + } + (user, count) + }) + .collect() +} + +#[cfg(test)] +mod tests { + use super::*; + + fn users(secret: &str) -> HashMap { + HashMap::from([("alice".to_string(), secret.to_string())]) + } + + #[test] + fn shared_authority_rejects_registration_through_an_old_generation() { + let authority = UserAdmissionAuthority::new(); + let secret = "00112233445566778899aabbccddeeff"; + authority.apply_config(&users(secret), &HashMap::new()); + let credential = credential_id_from_hex(secret).unwrap(); + + assert!(authority.claim_authenticated("alice", credential).is_some()); + authority.stage_user("alice", secret, false).unwrap(); + + assert!(authority.claim_authenticated("alice", credential).is_none()); + } + + #[test] + fn stale_credential_cannot_cross_delete_and_recreate() { + let authority = UserAdmissionAuthority::new(); + let old_secret = "00112233445566778899aabbccddeeff"; + let new_secret = "ffeeddccbbaa99887766554433221100"; + authority.apply_config(&users(old_secret), &HashMap::new()); + let old_credential = credential_id_from_hex(old_secret).unwrap(); + let old_incarnation = authority + .authenticated_incarnation("alice", old_credential) + .unwrap(); + + authority.delete_user("alice"); + let recreated = authority.stage_user("alice", new_secret, true).unwrap(); + + assert!(recreated.incarnation > old_incarnation); + assert!( + authority + .authenticated_incarnation("alice", old_credential) + .is_none() + ); + } + + #[test] + fn stale_candidate_cannot_overwrite_newer_mutation() { + let authority = UserAdmissionAuthority::new(); + let secret = "00112233445566778899aabbccddeeff"; + authority.apply_config(&users(secret), &HashMap::new()); + let candidate_epoch = authority.epoch(); + authority.stage_user("alice", secret, false).unwrap(); + + assert!( + authority + .apply_config_if_epoch(candidate_epoch, &users(secret), &HashMap::new()) + .is_none() + ); + assert!(!authority.is_user_enabled("alice")); + } +} diff --git a/src/transport/middle_proxy/health/family.rs b/src/transport/middle_proxy/health/family.rs index b03677c..0bb5d72 100644 --- a/src/transport/middle_proxy/health/family.rs +++ b/src/transport/middle_proxy/health/family.rs @@ -25,9 +25,10 @@ pub(super) async fn check_family( let mut family_degraded = false; let mut dc_endpoints = HashMap::>::new(); + let endpoint_snapshot = pool.endpoint_snapshot.load(); let map_guard = match family { - IpFamily::V4 => pool.proxy_map_v4.read().await, - IpFamily::V6 => pool.proxy_map_v6.read().await, + IpFamily::V4 => &endpoint_snapshot.map_v4, + IpFamily::V6 => &endpoint_snapshot.map_v6, }; for (dc, addrs) in map_guard.iter() { let entry = dc_endpoints.entry(*dc).or_default(); @@ -35,7 +36,7 @@ pub(super) async fn check_family( entry.push(SocketAddr::new(ip, port)); } } - drop(map_guard); + drop(endpoint_snapshot); for endpoints in dc_endpoints.values_mut() { endpoints.sort_unstable(); endpoints.dedup(); diff --git a/src/transport/middle_proxy/ping.rs b/src/transport/middle_proxy/ping.rs index 85888bd..8a86204 100644 --- a/src/transport/middle_proxy/ping.rs +++ b/src/transport/middle_proxy/ping.rs @@ -329,13 +329,14 @@ mod tests { pub async fn run_me_ping(pool: &Arc, rng: &SecureRandom) -> Vec { let mut reports = Vec::new(); + let endpoint_snapshot = pool.endpoint_snapshot.load_full(); let v4_map = if pool.decision.ipv4_me { - pool.proxy_map_v4.read().await.clone() + endpoint_snapshot.map_v4.clone() } else { HashMap::new() }; let v6_map = if pool.decision.ipv6_me { - pool.proxy_map_v6.read().await.clone() + endpoint_snapshot.map_v6.clone() } else { HashMap::new() }; diff --git a/src/transport/middle_proxy/pool.rs b/src/transport/middle_proxy/pool.rs index 093de8c..8c551cc 100644 --- a/src/transport/middle_proxy/pool.rs +++ b/src/transport/middle_proxy/pool.rs @@ -266,7 +266,17 @@ pub struct RoutingCore { pub(super) writers: Arc, pub(super) rr: AtomicU64, pub(super) writer_epoch: watch::Sender, - pub(super) preferred_endpoints_by_dc: ArcSwap>>, + pub(super) endpoint_snapshot: ArcSwap, +} + +/// Immutable endpoint routing authority published as one coherent revision. +#[derive(Clone, Debug)] +pub(super) struct EndpointSnapshot { + pub(super) revision: u64, + pub(super) map_v4: HashMap>, + pub(super) map_v6: HashMap>, + pub(super) endpoint_dc_map: HashMap>, + pub(super) preferred_endpoints_by_dc: HashMap>, } pub(super) struct ReinitCore { @@ -302,12 +312,14 @@ pub(super) struct ReinitPendingState { pub(super) generation: u64, pub(super) started_at_epoch_secs: u64, pub(super) map_hash: u64, + pub(super) endpoint_revision: u64, } #[derive(Clone, Copy)] pub(super) struct ReinitAttemptState { pub(super) generation: u64, pub(super) map_hash: u64, + pub(super) endpoint_revision: u64, pub(super) hardswap: bool, pub(super) committed: bool, } @@ -316,6 +328,7 @@ pub(super) struct ReinitCoordinatorState { pub(super) next_attempt_id: u64, pub(super) active_generation: u64, pub(super) desired_map_hash: u64, + pub(super) endpoint_revision: u64, pub(super) pending: Option, pub(super) attempts: HashMap, } @@ -475,9 +488,6 @@ pub struct MePool { pub(super) rng: Arc, pub(super) proxy_tag: Option>, pub(super) proxy_secret: Arc>, - pub(super) proxy_map_v4: Arc>>>, - pub(super) proxy_map_v6: Arc>>>, - pub(super) endpoint_dc_map: Arc>>>, pub(super) default_dc: AtomicI32, pub(super) next_writer_id: AtomicU64, pub(super) writer_connect_active_reserved: AtomicUsize, diff --git a/src/transport/middle_proxy/pool/construction.rs b/src/transport/middle_proxy/pool/construction.rs index a14a086..e66ddc1 100644 --- a/src/transport/middle_proxy/pool/construction.rs +++ b/src/transport/middle_proxy/pool/construction.rs @@ -120,9 +120,12 @@ impl MePool { me_route_inline_recovery_wait_ms: u64, me_connection_cleanup_capacity: usize, ) -> Arc { - let endpoint_dc_map = Self::build_endpoint_dc_map_from_maps(&proxy_map_v4, &proxy_map_v6); - let preferred_endpoints_by_dc = - Self::build_preferred_endpoints_by_dc(&decision, &proxy_map_v4, &proxy_map_v6); + let endpoint_snapshot = Self::build_endpoint_snapshot( + &decision, + proxy_map_v4, + proxy_map_v6, + 1, + ); let registry = Arc::new(ConnRegistry::with_route_and_cleanup_capacity( me_route_channel_capacity, me_connection_cleanup_capacity, @@ -149,7 +152,7 @@ impl MePool { writers: Arc::new(WritersState::new()), rr: AtomicU64::new(0), writer_epoch, - preferred_endpoints_by_dc: ArcSwap::from_pointee(preferred_endpoints_by_dc), + endpoint_snapshot: ArcSwap::from_pointee(endpoint_snapshot), }), reinit: Arc::new(ReinitCore { generation: AtomicU64::new(1), @@ -164,6 +167,7 @@ impl MePool { next_attempt_id: 1, active_generation: 1, desired_map_hash: 0, + endpoint_revision: 1, pending: None, attempts: HashMap::new(), }), @@ -391,9 +395,6 @@ impl MePool { })), stats, pool_size: 2, - proxy_map_v4: Arc::new(RwLock::new(proxy_map_v4)), - proxy_map_v6: Arc::new(RwLock::new(proxy_map_v6)), - endpoint_dc_map: Arc::new(RwLock::new(endpoint_dc_map)), default_dc: AtomicI32::new(default_dc.unwrap_or(2)), next_writer_id: AtomicU64::new(1), writer_connect_active_reserved: AtomicUsize::new(0), diff --git a/src/transport/middle_proxy/pool/routing.rs b/src/transport/middle_proxy/pool/routing.rs index c999730..50e9616 100644 --- a/src/transport/middle_proxy/pool/routing.rs +++ b/src/transport/middle_proxy/pool/routing.rs @@ -76,16 +76,23 @@ impl MePool { &self, dc: i32, ) -> bool { + let snapshot = self.endpoint_snapshot.load(); if self.decision.ipv4_me { - let map = self.proxy_map_v4.read().await; - if map.get(&dc).is_some_and(|endpoints| !endpoints.is_empty()) { + if snapshot + .map_v4 + .get(&dc) + .is_some_and(|endpoints| !endpoints.is_empty()) + { return true; } } if self.decision.ipv6_me { - let map = self.proxy_map_v6.read().await; - if map.get(&dc).is_some_and(|endpoints| !endpoints.is_empty()) { + if snapshot + .map_v6 + .get(&dc) + .is_some_and(|endpoints| !endpoints.is_empty()) + { return true; } } @@ -112,7 +119,12 @@ impl MePool { &self, addr: SocketAddr, ) -> i32 { - if let Some(cached) = self.endpoint_dc_map.read().await.get(&addr).copied() + if let Some(cached) = self + .endpoint_snapshot + .load() + .endpoint_dc_map + .get(&addr) + .copied() && let Some(dc) = cached { return dc; @@ -125,9 +137,45 @@ impl MePool { &self, family: IpFamily, ) -> HashMap> { + let snapshot = self.endpoint_snapshot.load(); match family { - IpFamily::V4 => self.proxy_map_v4.read().await.clone(), - IpFamily::V6 => self.proxy_map_v6.read().await.clone(), + IpFamily::V4 => snapshot.map_v4.clone(), + IpFamily::V6 => snapshot.map_v6.clone(), + } + } + + pub(in crate::transport::middle_proxy) fn build_endpoint_snapshot( + decision: &NetworkDecision, + mut map_v4: HashMap>, + mut map_v6: HashMap>, + revision: u64, + ) -> EndpointSnapshot { + Self::mirror_negative_dcs(&mut map_v4); + Self::mirror_negative_dcs(&mut map_v6); + let endpoint_dc_map = Self::build_endpoint_dc_map_from_maps(&map_v4, &map_v6); + let preferred_endpoints_by_dc = + Self::build_preferred_endpoints_by_dc(decision, &map_v4, &map_v6); + EndpointSnapshot { + revision, + map_v4, + map_v6, + endpoint_dc_map, + preferred_endpoints_by_dc, + } + } + + fn mirror_negative_dcs(map: &mut HashMap>) { + let positive_dcs = map + .keys() + .copied() + .filter(|dc| *dc > 0) + .collect::>(); + for dc in positive_dcs { + if !map.contains_key(&-dc) + && let Some(endpoints) = map.get(&dc).cloned() + { + map.insert(-dc, endpoints); + } } } @@ -224,17 +272,11 @@ impl MePool { endpoint_dc_map } - pub(in crate::transport::middle_proxy) async fn rebuild_endpoint_dc_map(&self) { - let map_v4 = self.proxy_map_v4.read().await.clone(); - let map_v6 = self.proxy_map_v6.read().await.clone(); - let rebuilt = Self::build_endpoint_dc_map_from_maps(&map_v4, &map_v6); - let preferred = Self::build_preferred_endpoints_by_dc(&self.decision, &map_v4, &map_v6); - *self.endpoint_dc_map.write().await = rebuilt; - self.preferred_endpoints_by_dc.store(Arc::new(preferred)); + pub(in crate::transport::middle_proxy) async fn prune_endpoint_runtime_state(&self) { let configured_endpoints = self + .endpoint_snapshot + .load() .endpoint_dc_map - .read() - .await .keys() .copied() .collect::>(); @@ -253,8 +295,12 @@ impl MePool { &self, dc: i32, ) -> Vec { - let guard = self.preferred_endpoints_by_dc.load(); - guard.get(&dc).cloned().unwrap_or_default() + self.endpoint_snapshot + .load() + .preferred_endpoints_by_dc + .get(&dc) + .cloned() + .unwrap_or_default() } pub(in crate::transport::middle_proxy) fn health_interval_unhealthy(&self) -> Duration { diff --git a/src/transport/middle_proxy/pool/writer_admission.rs b/src/transport/middle_proxy/pool/writer_admission.rs index 0041c3d..91acf18 100644 --- a/src/transport/middle_proxy/pool/writer_admission.rs +++ b/src/transport/middle_proxy/pool/writer_admission.rs @@ -65,10 +65,10 @@ impl MePool { pub(in crate::transport::middle_proxy) async fn active_coverage_required_total(&self) -> usize { let now_epoch_secs = Self::now_epoch_secs(); let mut required_total = 0usize; + let endpoint_snapshot = self.endpoint_snapshot.load(); if self.family_enabled_for_drain_coverage(IpFamily::V4, now_epoch_secs) { - let map = self.proxy_map_v4.read().await; - for addrs in map.values() { + for addrs in endpoint_snapshot.map_v4.values() { let mut endpoints = HashSet::::new(); for (ip, port) in addrs.iter().copied() { endpoints.insert(SocketAddr::new(ip, port)); @@ -80,8 +80,7 @@ impl MePool { } if self.family_enabled_for_drain_coverage(IpFamily::V6, now_epoch_secs) { - let map = self.proxy_map_v6.read().await; - for addrs in map.values() { + for addrs in endpoint_snapshot.map_v6.values() { let mut endpoints = HashSet::::new(); for (ip, port) in addrs.iter().copied() { endpoints.insert(SocketAddr::new(ip, port)); @@ -118,13 +117,14 @@ impl MePool { 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) = self.proxy_map_v4.read().await.get(&writer_dc) { + 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) = self.proxy_map_v6.read().await.get(&writer_dc) { + if let Some(addrs) = endpoint_snapshot.map_v6.get(&writer_dc) { endpoints_len += addrs.len(); } } diff --git a/src/transport/middle_proxy/pool_config.rs b/src/transport/middle_proxy/pool_config.rs index 73695b7..280f45b 100644 --- a/src/transport/middle_proxy/pool_config.rs +++ b/src/transport/middle_proxy/pool_config.rs @@ -31,48 +31,35 @@ impl MePool { return SnapshotApplyOutcome::RejectedEmpty; } - let mut changed = false; - { - let mut guard = self.proxy_map_v4.write().await; - if !new_v4.is_empty() && *guard != new_v4 { - *guard = new_v4; - changed = true; + let changed = { + // Endpoint publication and reinit commit share this barrier. + let mut coordinator = self.reinit.coordinator.lock(); + let current = self.endpoint_snapshot.load_full(); + let map_v4 = if new_v4.is_empty() { + current.map_v4.clone() + } else { + new_v4 + }; + let map_v6 = match new_v6 { + Some(map) if !map.is_empty() => map, + _ => current.map_v6.clone(), + }; + let candidate = Self::build_endpoint_snapshot( + &self.decision, + map_v4, + map_v6, + current.revision.saturating_add(1), + ); + if candidate.map_v4 == current.map_v4 && candidate.map_v6 == current.map_v6 { + false + } else { + coordinator.endpoint_revision = candidate.revision; + self.endpoint_snapshot.store(Arc::new(candidate)); + true } - } - if let Some(v6) = new_v6 { - let mut guard = self.proxy_map_v6.write().await; - if !v6.is_empty() && *guard != v6 { - *guard = v6; - changed = true; - } - } - // Ensure negative DC entries mirror positives when absent (Telegram convention). - { - let mut guard = self.proxy_map_v4.write().await; - let keys: Vec = guard.keys().cloned().collect(); - for k in keys.iter().cloned().filter(|k| *k > 0) { - if !guard.contains_key(&-k) - && let Some(addrs) = guard.get(&k).cloned() - { - guard.insert(-k, addrs); - changed = true; - } - } - } - { - let mut guard = self.proxy_map_v6.write().await; - let keys: Vec = guard.keys().cloned().collect(); - for k in keys.iter().cloned().filter(|k| *k > 0) { - if !guard.contains_key(&-k) - && let Some(addrs) = guard.get(&k).cloned() - { - guard.insert(-k, addrs); - changed = true; - } - } - } + }; if changed { - self.rebuild_endpoint_dc_map().await; + self.prune_endpoint_runtime_state().await; self.notify_writer_epoch(); } if changed { diff --git a/src/transport/middle_proxy/pool_init.rs b/src/transport/middle_proxy/pool_init.rs index eaa0476..4a17298 100644 --- a/src/transport/middle_proxy/pool_init.rs +++ b/src/transport/middle_proxy/pool_init.rs @@ -19,7 +19,7 @@ impl MePool { .me_reconnect_max_concurrent_per_dc .max(1) as usize; let ks = self.key_selector().await; - let me_servers = self.proxy_map_v4.read().await.len(); + let me_servers = self.endpoint_snapshot.load().map_v4.len(); let secret_len = self.proxy_secret.read().await.secret.len(); info!( me_servers, diff --git a/src/transport/middle_proxy/pool_refill.rs b/src/transport/middle_proxy/pool_refill.rs index c6f1747..60ec7e6 100644 --- a/src/transport/middle_proxy/pool_refill.rs +++ b/src/transport/middle_proxy/pool_refill.rs @@ -70,9 +70,9 @@ impl Drop for RefillRunGuard { impl MePool { pub(super) async fn sweep_endpoint_quarantine(&self) { let configured = self + .endpoint_snapshot + .load() .endpoint_dc_map - .read() - .await .keys() .copied() .collect::>(); @@ -266,9 +266,10 @@ impl MePool { if !self.family_enabled_for_drain_coverage(target.family, now_epoch_secs) { return Vec::new(); } + let snapshot = self.endpoint_snapshot.load(); let map = match target.family { - IpFamily::V4 => self.proxy_map_v4.read().await, - IpFamily::V6 => self.proxy_map_v6.read().await, + IpFamily::V4 => &snapshot.map_v4, + IpFamily::V6 => &snapshot.map_v6, }; let mut endpoints = map .get(&target.dc) @@ -294,8 +295,9 @@ impl MePool { }; role_is_authoritative && self - .preferred_endpoints_by_dc + .endpoint_snapshot .load() + .preferred_endpoints_by_dc .get(&target.dc) .is_some_and(|endpoints| { endpoints.iter().any(|endpoint| match target.family { diff --git a/src/transport/middle_proxy/pool_reinit.rs b/src/transport/middle_proxy/pool_reinit.rs index f0a3db5..cfaabf9 100644 --- a/src/transport/middle_proxy/pool_reinit.rs +++ b/src/transport/middle_proxy/pool_reinit.rs @@ -15,8 +15,8 @@ use crate::config::MeBindStaleMode; use crate::network::IpFamily; use super::pool::{ - MeDrainGateReason, MePool, ReinitAttemptState, ReinitCoordinatorState, ReinitCore, - ReinitPendingState, ReinitStatusSnapshot, WriterContour, WriterOpenIntent, + EndpointSnapshot, MeDrainGateReason, MePool, ReinitAttemptState, ReinitCoordinatorState, + ReinitCore, ReinitPendingState, ReinitStatusSnapshot, WriterContour, WriterOpenIntent, }; // Reinitialization admission, generation state, and coverage checks. @@ -34,6 +34,7 @@ struct ReinitAttemptGuard { generation: u64, previous_generation: u64, map_hash: u64, + endpoint_revision: u64, hardswap: bool, } @@ -119,17 +120,24 @@ fn commit_reinit_state( attempt_id: u64, generation: u64, map_hash: u64, + endpoint_revision: u64, hardswap: bool, ) -> bool { let Some(record) = state.attempts.get(&attempt_id).copied() else { return false; }; - if record.map_hash != state.desired_map_hash || record.map_hash != map_hash { + if record.map_hash != state.desired_map_hash + || record.map_hash != map_hash + || record.endpoint_revision != endpoint_revision + || record.endpoint_revision != state.endpoint_revision + { return false; } if hardswap { let pending_matches = state.pending.is_some_and(|pending| { - pending.generation == generation && pending.map_hash == map_hash + pending.generation == generation + && pending.map_hash == map_hash + && pending.endpoint_revision == endpoint_revision }); if !pending_matches || generation < state.active_generation { return false; diff --git a/src/transport/middle_proxy/pool_reinit/coordination.rs b/src/transport/middle_proxy/pool_reinit/coordination.rs index a81b672..4b27b19 100644 --- a/src/transport/middle_proxy/pool_reinit/coordination.rs +++ b/src/transport/middle_proxy/pool_reinit/coordination.rs @@ -27,9 +27,13 @@ impl MePool { self: &Arc, hardswap: bool, map_hash: u64, + endpoint_revision: u64, now_epoch_secs: u64, - ) -> ReinitReservation { + ) -> Option { let mut state = self.reinit.coordinator.lock(); + if state.endpoint_revision != endpoint_revision { + return None; + } state.desired_map_hash = map_hash; let previous_generation = state.active_generation; let mut pending_reused = false; @@ -43,6 +47,7 @@ impl MePool { && pending_age_secs > ME_HARDSWAP_PENDING_TTL_SECS; pending.generation >= previous_generation && pending.map_hash == map_hash + && pending.endpoint_revision == endpoint_revision && !pending_expired }); if let Some(pending) = reusable { @@ -54,6 +59,7 @@ impl MePool { generation, started_at_epoch_secs: now_epoch_secs, map_hash, + endpoint_revision, }); generation } @@ -69,24 +75,26 @@ impl MePool { ReinitAttemptState { generation, map_hash, + endpoint_revision, hardswap, committed: false, }, ); publish_reinit_state(self.reinit.as_ref(), &state); - ReinitReservation { + Some(ReinitReservation { attempt: ReinitAttemptGuard { reinit: Arc::clone(&self.reinit), attempt_id, generation, previous_generation, map_hash, + endpoint_revision, hardswap, }, pending_reused, pending_expired, pending_age_secs, - } + }) } /// Revalidates coverage and commits generation ownership under the publication barrier. @@ -105,10 +113,13 @@ impl MePool { if record.generation != attempt.generation || record.map_hash != state.desired_map_hash || record.map_hash != attempt.map_hash + || record.endpoint_revision != attempt.endpoint_revision + || record.endpoint_revision != state.endpoint_revision || (attempt.hardswap && !state.pending.is_some_and(|pending| { pending.generation == attempt.generation && pending.map_hash == attempt.map_hash + && pending.endpoint_revision == attempt.endpoint_revision })) { return Err(ReinitCommitFailure::Superseded); @@ -150,6 +161,7 @@ impl MePool { attempt.attempt_id, attempt.generation, attempt.map_hash, + attempt.endpoint_revision, attempt.hardswap, ) { return Err(ReinitCommitFailure::Superseded); @@ -255,9 +267,13 @@ impl MePool { /// 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(); for family in self.family_order() { - let map = self.proxy_map_for_family(family).await; - for (dc, addrs) in &map { + let map = match family { + IpFamily::V4 => &endpoint_snapshot.map_v4, + IpFamily::V6 => &endpoint_snapshot.map_v6, + }; + for (dc, addrs) in map { let dc_addrs: Vec = addrs .iter() .map(|(ip, port)| SocketAddr::new(*ip, *port)) @@ -286,26 +302,32 @@ impl MePool { /// Returns the currently authoritative endpoint set for drain and coverage decisions. pub(in crate::transport::middle_proxy) async fn desired_dc_endpoints( &self, + ) -> HashMap> { + let endpoint_snapshot = self.endpoint_snapshot.load_full(); + self.desired_dc_endpoints_from_snapshot(&endpoint_snapshot) + } + + pub(super) fn desired_dc_endpoints_from_snapshot( + &self, + endpoint_snapshot: &EndpointSnapshot, ) -> HashMap> { let now_epoch_secs = Self::now_epoch_secs(); let mut out: HashMap> = HashMap::new(); if self.family_enabled_for_drain_coverage(IpFamily::V4, now_epoch_secs) { - let map_v4 = self.proxy_map_v4.read().await.clone(); - for (dc, addrs) in map_v4 { - let entry = out.entry(dc).or_default(); + for (dc, addrs) in &endpoint_snapshot.map_v4 { + let entry = out.entry(*dc).or_default(); for (ip, port) in addrs { - entry.insert(SocketAddr::new(ip, port)); + entry.insert(SocketAddr::new(*ip, *port)); } } } if self.family_enabled_for_drain_coverage(IpFamily::V6, now_epoch_secs) { - let map_v6 = self.proxy_map_v6.read().await.clone(); - for (dc, addrs) in map_v6 { - let entry = out.entry(dc).or_default(); + for (dc, addrs) in &endpoint_snapshot.map_v6 { + let entry = out.entry(*dc).or_default(); for (ip, port) in addrs { - entry.insert(SocketAddr::new(ip, port)); + entry.insert(SocketAddr::new(*ip, *port)); } } } @@ -320,7 +342,8 @@ impl MePool { let state = self.reinit.coordinator.lock(); let active_generation = state.active_generation; let pending_generation = state.pending.map(|pending| pending.generation); - let preferred = self.preferred_endpoints_by_dc.load(); + let endpoint_snapshot = self.endpoint_snapshot.load(); + let preferred = &endpoint_snapshot.preferred_endpoints_by_dc; let now_epoch_secs = Self::now_epoch_secs(); let mut changed = 0usize; @@ -362,7 +385,7 @@ impl MePool { self.apply_writer_draining_state(writer, self.force_close_timeout(), false); changed = changed.saturating_add(1); } - drop(preferred); + drop(endpoint_snapshot); drop(state); drop(registry_registration); drop(writers); diff --git a/src/transport/middle_proxy/pool_reinit/reconcile.rs b/src/transport/middle_proxy/pool_reinit/reconcile.rs index b2b5ee4..8eac110 100644 --- a/src/transport/middle_proxy/pool_reinit/reconcile.rs +++ b/src/transport/middle_proxy/pool_reinit/reconcile.rs @@ -118,7 +118,8 @@ impl MePool { self: &Arc, rng: &SecureRandom, ) -> bool { - let desired_by_dc = self.desired_dc_endpoints().await; + let endpoint_snapshot = self.endpoint_snapshot.load_full(); + let desired_by_dc = self.desired_dc_endpoints_from_snapshot(&endpoint_snapshot); let now_epoch_secs = Self::now_epoch_secs(); let v4_suppressed = self.is_family_temporarily_suppressed(IpFamily::V4, now_epoch_secs); let v6_suppressed = self.is_family_temporarily_suppressed(IpFamily::V6, now_epoch_secs); @@ -137,7 +138,18 @@ impl MePool { let desired_map_hash = Self::desired_map_hash(&desired_by_dc); let hardswap = self.reinit.hardswap.load(Ordering::Relaxed); - let reservation = self.reserve_reinit_attempt(hardswap, desired_map_hash, now_epoch_secs); + let Some(reservation) = self.reserve_reinit_attempt( + hardswap, + desired_map_hash, + endpoint_snapshot.revision, + now_epoch_secs, + ) else { + debug!( + endpoint_revision = endpoint_snapshot.revision, + "ME reinit snapshot superseded before reservation" + ); + return false; + }; let attempt = reservation.attempt; let previous_generation = attempt.previous_generation; let generation = attempt.generation; diff --git a/src/transport/middle_proxy/pool_reinit/tests.rs b/src/transport/middle_proxy/pool_reinit/tests.rs index 4f500e1..8f19001 100644 --- a/src/transport/middle_proxy/pool_reinit/tests.rs +++ b/src/transport/middle_proxy/pool_reinit/tests.rs @@ -115,10 +115,12 @@ fn stale_concurrent_attempt_cannot_regress_active_generation() { next_attempt_id: 3, active_generation: 1, desired_map_hash: 22, + endpoint_revision: 7, pending: Some(ReinitPendingState { generation: 3, started_at_epoch_secs: 1, map_hash: 22, + endpoint_revision: 7, }), attempts: HashMap::from([ ( @@ -126,6 +128,7 @@ fn stale_concurrent_attempt_cannot_regress_active_generation() { ReinitAttemptState { generation: 2, map_hash: 11, + endpoint_revision: 6, hardswap: true, committed: false, }, @@ -135,6 +138,7 @@ fn stale_concurrent_attempt_cannot_regress_active_generation() { ReinitAttemptState { generation: 3, map_hash: 22, + endpoint_revision: 7, hardswap: true, committed: false, }, @@ -142,9 +146,9 @@ fn stale_concurrent_attempt_cannot_regress_active_generation() { ]), }; - assert!(commit_reinit_state(&mut state, 2, 3, 22, true)); + assert!(commit_reinit_state(&mut state, 2, 3, 22, 7, true)); assert_eq!(state.active_generation, 3); - assert!(!commit_reinit_state(&mut state, 1, 2, 11, true)); + assert!(!commit_reinit_state(&mut state, 1, 2, 11, 6, true)); assert_eq!(state.active_generation, 3); assert!(state.pending.is_none()); } @@ -173,7 +177,10 @@ async fn partial_hardswap_is_rejected_when_stale_binding_is_disabled() { ) .await; let map_hash = MePool::desired_map_hash(&desired_by_dc); - let reservation = pool.reserve_reinit_attempt(true, map_hash, 100); + 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( &pool, 201, @@ -221,7 +228,10 @@ async fn partial_hardswap_preserves_fallback_only_for_missing_dc() { ) .await; let map_hash = MePool::desired_map_hash(&desired_by_dc); - let reservation = pool.reserve_reinit_attempt(true, map_hash, 100); + 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"); let fresh_dc1 = insert_writer( &pool, 401, @@ -274,7 +284,10 @@ async fn complete_hardswap_promotes_fresh_generation_and_retires_old_writers() { ) .await; let map_hash = MePool::desired_map_hash(&desired_by_dc); - let reservation = pool.reserve_reinit_attempt(true, map_hash, 100); + 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"); let fresh_dc1 = insert_writer( &pool, 601, @@ -314,15 +327,80 @@ async fn complete_hardswap_promotes_fresh_generation_and_retires_old_writers() { ); } +#[tokio::test] +async fn endpoint_revision_change_supersedes_hardswap_before_writer_drain() { + let pool = make_pool().await; + let old_endpoint = addr(1, 2001); + pool.update_proxy_maps( + HashMap::from([(1, vec![(old_endpoint.ip(), old_endpoint.port())])]), + None, + ) + .await; + let active_generation = pool.current_generation(); + let old_writer = insert_writer( + &pool, + 651, + 1, + old_endpoint, + active_generation, + WriterContour::Active, + ) + .await; + let desired_by_dc = HashMap::from([(1, HashSet::from([old_endpoint]))]); + 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"); + let fresh_writer = insert_writer( + &pool, + 652, + 1, + old_endpoint, + reservation.attempt.generation, + WriterContour::Warm, + ) + .await; + + let replacement_endpoint = addr(3, 2003); + pool.update_proxy_maps( + HashMap::from([( + 1, + vec![(replacement_endpoint.ip(), replacement_endpoint.port())], + )]), + None, + ) + .await; + + let result = pool + .commit_reinit_attempt(&reservation.attempt, &desired_by_dc, 1.0) + .await; + + assert!(matches!(result, Err(ReinitCommitFailure::Superseded))); + assert_eq!(pool.current_generation(), active_generation); + assert!(!old_writer.draining.load(Ordering::Acquire)); + assert!(!fresh_writer.draining.load(Ordering::Acquire)); + assert_eq!( + WriterContour::from_u8(fresh_writer.contour.load(Ordering::Acquire)), + WriterContour::Warm + ); +} + #[tokio::test] async fn generation_role_reconciliation_promotes_active_warm_and_drains_orphans() { let pool = make_pool().await; let endpoint = addr(1, 2001); - pool.preferred_endpoints_by_dc - .store(Arc::new(HashMap::from([(1, vec![endpoint])]))); + pool.update_proxy_maps( + HashMap::from([(1, vec![(endpoint.ip(), endpoint.port())])]), + None, + ) + .await; let desired_by_dc = HashMap::from([(1, HashSet::from([endpoint]))]); let map_hash = MePool::desired_map_hash(&desired_by_dc); - let reservation = pool.reserve_reinit_attempt(true, map_hash, 100); + 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"); let active_warm = insert_writer( &pool, 701, diff --git a/src/transport/middle_proxy/pool_status/status_snapshot.rs b/src/transport/middle_proxy/pool_status/status_snapshot.rs index 4dc6d86..a246b44 100644 --- a/src/transport/middle_proxy/pool_status/status_snapshot.rs +++ b/src/transport/middle_proxy/pool_status/status_snapshot.rs @@ -3,12 +3,13 @@ use super::*; impl MePool { pub(crate) async fn admission_ready_conditional_cast(&self) -> bool { let mut endpoints_by_dc = BTreeMap::>::new(); + let endpoint_snapshot = self.endpoint_snapshot.load_full(); if self.decision.ipv4_me { - let map = self.proxy_map_v4.read().await.clone(); + let map = endpoint_snapshot.map_v4.clone(); extend_signed_endpoints(&mut endpoints_by_dc, map); } if self.decision.ipv6_me { - let map = self.proxy_map_v6.read().await.clone(); + let map = endpoint_snapshot.map_v6.clone(); extend_signed_endpoints(&mut endpoints_by_dc, map); } @@ -45,12 +46,13 @@ impl MePool { #[allow(dead_code)] pub(crate) async fn admission_ready_full_floor(&self) -> bool { let mut endpoints_by_dc = BTreeMap::>::new(); + let endpoint_snapshot = self.endpoint_snapshot.load_full(); if self.decision.ipv4_me { - let map = self.proxy_map_v4.read().await.clone(); + let map = endpoint_snapshot.map_v4.clone(); extend_signed_endpoints(&mut endpoints_by_dc, map); } if self.decision.ipv6_me { - let map = self.proxy_map_v6.read().await.clone(); + let map = endpoint_snapshot.map_v6.clone(); extend_signed_endpoints(&mut endpoints_by_dc, map); } @@ -106,12 +108,13 @@ impl MePool { .load(Ordering::Relaxed); let mut endpoints_by_dc = BTreeMap::>::new(); + let endpoint_snapshot = self.endpoint_snapshot.load_full(); if self.decision.ipv4_me { - let map = self.proxy_map_v4.read().await.clone(); + let map = endpoint_snapshot.map_v4.clone(); extend_signed_endpoints(&mut endpoints_by_dc, map); } if self.decision.ipv6_me { - let map = self.proxy_map_v6.read().await.clone(); + let map = endpoint_snapshot.map_v6.clone(); extend_signed_endpoints(&mut endpoints_by_dc, map); } diff --git a/src/transport/middle_proxy/pool_writer/publication.rs b/src/transport/middle_proxy/pool_writer/publication.rs index 2027e33..ef1a45c 100644 --- a/src/transport/middle_proxy/pool_writer/publication.rs +++ b/src/transport/middle_proxy/pool_writer/publication.rs @@ -41,8 +41,9 @@ impl MePool { coordinator: &crate::transport::middle_proxy::pool::ReinitCoordinatorState, ) -> Result { let endpoint_is_current = self - .preferred_endpoints_by_dc + .endpoint_snapshot .load() + .preferred_endpoints_by_dc .get(&writer.writer_dc) .is_some_and(|endpoints| endpoints.contains(&writer.addr)); if !endpoint_is_current { @@ -61,6 +62,7 @@ impl MePool { && coordinator.pending.is_some_and(|pending| { pending.generation == writer.generation && pending.map_hash == coordinator.desired_map_hash + && pending.endpoint_revision == coordinator.endpoint_revision }) { return Ok(WriterContour::Warm); @@ -92,7 +94,8 @@ impl MePool { if intent == WriterOpenIntent::Replacement || contour == WriterContour::Draining { return Ok(()); } - let preferred = self.preferred_endpoints_by_dc.load(); + let endpoint_snapshot = self.endpoint_snapshot.load(); + let preferred = &endpoint_snapshot.preferred_endpoints_by_dc; let Some(endpoints) = preferred.get(&writer.writer_dc) else { return Err(ProxyError::Proxy( "ME writer target changed before publication".into(), diff --git a/src/transport/middle_proxy/pool_writer/replacement.rs b/src/transport/middle_proxy/pool_writer/replacement.rs index 85a9927..c1a3fb3 100644 --- a/src/transport/middle_proxy/pool_writer/replacement.rs +++ b/src/transport/middle_proxy/pool_writer/replacement.rs @@ -85,7 +85,8 @@ impl MePool { "ME floor rebalance lost active-generation authority".into(), )); } - let preferred = self.preferred_endpoints_by_dc.load(); + let endpoint_snapshot = self.endpoint_snapshot.load(); + let preferred = &endpoint_snapshot.preferred_endpoints_by_dc; let donor_count = writers .iter() .filter(|candidate| { @@ -265,8 +266,11 @@ mod tests { async fn replacement_commit_publishes_successor_before_draining_victim() { let pool = make_pool().await; let addr = endpoint(1); - pool.preferred_endpoints_by_dc - .store(Arc::new(HashMap::from([(2, vec![addr])]))); + pool.update_proxy_maps( + HashMap::from([(2, vec![(addr.ip(), addr.port())])]), + None, + ) + .await; let victim = install_writer(&pool, 1001, 2, addr).await; let expected_role = WriterRole::from_writer(&victim); let mut reservation = pool @@ -305,8 +309,11 @@ mod tests { async fn cancelled_replacement_waiting_for_publication_restores_all_reservations() { let pool = make_pool().await; let addr = endpoint(2); - pool.preferred_endpoints_by_dc - .store(Arc::new(HashMap::from([(2, vec![addr])]))); + pool.update_proxy_maps( + HashMap::from([(2, vec![(addr.ip(), addr.port())])]), + None, + ) + .await; let victim = install_writer(&pool, 2001, 2, addr).await; let expected_role = WriterRole::from_writer(&victim); let mut reservation = pool @@ -342,10 +349,14 @@ mod tests { let pool = make_pool().await; let donor_addr = endpoint(3); let receiver_addr = endpoint(4); - pool.preferred_endpoints_by_dc.store(Arc::new(HashMap::from([ - (1, vec![donor_addr]), - (2, vec![receiver_addr]), - ]))); + pool.update_proxy_maps( + HashMap::from([ + (1, vec![(donor_addr.ip(), donor_addr.port())]), + (2, vec![(receiver_addr.ip(), receiver_addr.port())]), + ]), + None, + ) + .await; let victim = install_writer(&pool, 3001, 1, donor_addr).await; let expected_role = WriterRole::from_writer(&victim); let mut reservation = pool diff --git a/src/transport/middle_proxy/send.rs b/src/transport/middle_proxy/send.rs index 88e73b0..cec8795 100644 --- a/src/transport/middle_proxy/send.rs +++ b/src/transport/middle_proxy/send.rs @@ -348,8 +348,10 @@ impl MePool { for _ in 0..self.route_runtime.me_route_inline_recovery_attempts.max(1) { - let preferred = self.preferred_endpoints_by_dc.load_full(); - for (dc, addrs) in preferred.iter() { + let endpoint_snapshot = self.endpoint_snapshot.load_full(); + for (dc, addrs) in + &endpoint_snapshot.preferred_endpoints_by_dc + { for addr in addrs { let _ = self .connect_one_for_dc(*addr, *dc, self.rng.as_ref()) @@ -470,8 +472,9 @@ impl MePool { } emergency_attempts += 1; let mut endpoints = self - .preferred_endpoints_by_dc + .endpoint_snapshot .load() + .preferred_endpoints_by_dc .get(&routed_dc) .cloned() .unwrap_or_default(); diff --git a/src/transport/middle_proxy/send/recovery.rs b/src/transport/middle_proxy/send/recovery.rs index ab38b5e..da95a3c 100644 --- a/src/transport/middle_proxy/send/recovery.rs +++ b/src/transport/middle_proxy/send/recovery.rs @@ -88,7 +88,8 @@ impl MePool { pub(super) async fn trigger_async_recovery_global(self: &Arc) { self.stats.increment_me_async_recovery_trigger_total(); - let preferred = self.preferred_endpoints_by_dc.load(); + let endpoint_snapshot = self.endpoint_snapshot.load(); + let preferred = &endpoint_snapshot.preferred_endpoints_by_dc; let mut triggered = 0usize; for (dc, addrs) in preferred.iter() { for addr in addrs { diff --git a/src/transport/middle_proxy/send/selection.rs b/src/transport/middle_proxy/send/selection.rs index 837062c..5c75b13 100644 --- a/src/transport/middle_proxy/send/selection.rs +++ b/src/transport/middle_proxy/send/selection.rs @@ -15,7 +15,8 @@ impl MePool { routed_dc: i32, include_warm: bool, ) -> Vec { - let preferred_snapshot = self.preferred_endpoints_by_dc.load(); + let endpoint_snapshot = self.endpoint_snapshot.load(); + let preferred_snapshot = &endpoint_snapshot.preferred_endpoints_by_dc; let mut out = Vec::new(); if let Some(preferred) = preferred_snapshot .get(&routed_dc) diff --git a/src/transport/middle_proxy/tests/pool_refill_security_tests.rs b/src/transport/middle_proxy/tests/pool_refill_security_tests.rs index 0cae5c5..db9245d 100644 --- a/src/transport/middle_proxy/tests/pool_refill_security_tests.rs +++ b/src/transport/middle_proxy/tests/pool_refill_security_tests.rs @@ -183,8 +183,18 @@ async fn refill_preserves_bounded_pending_cardinality_and_cleans_up_before_first let first = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 31, 0, 21)), 443); let second = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 31, 0, 22)), 443); let latest = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 31, 0, 23)), 443); - pool.preferred_endpoints_by_dc - .store(Arc::new(HashMap::from([(2, vec![first, second, latest])]))); + pool.update_proxy_maps( + HashMap::from([( + 2, + vec![ + (first.ip(), first.port()), + (second.ip(), second.port()), + (latest.ip(), latest.port()), + ], + )]), + None, + ) + .await; pool.trigger_immediate_refill_for_dc(first, 2); pool.trigger_immediate_refill_for_dc(second, 2); diff --git a/src/transport/middle_proxy/tests/pool_writer_publication_tests.rs b/src/transport/middle_proxy/tests/pool_writer_publication_tests.rs index 4135df1..84def94 100644 --- a/src/transport/middle_proxy/tests/pool_writer_publication_tests.rs +++ b/src/transport/middle_proxy/tests/pool_writer_publication_tests.rs @@ -43,8 +43,11 @@ fn unregistered_writer( async fn normal_warm_publication_cannot_race_past_the_dc_floor() { let pool = make_pool().await; let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); - pool.preferred_endpoints_by_dc - .store(Arc::new(std::collections::HashMap::from([(2, vec![addr])]))); + pool.update_proxy_maps( + std::collections::HashMap::from([(2, vec![(addr.ip(), addr.port())])]), + None, + ) + .await; let generation = 2; let writers = (1..=3) .map(|writer_id| { @@ -86,8 +89,11 @@ async fn normal_warm_publication_cannot_race_past_the_dc_floor() { async fn normal_active_publication_cannot_race_past_the_family_floor() { let pool = make_pool().await; let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); - pool.preferred_endpoints_by_dc - .store(Arc::new(std::collections::HashMap::from([(2, vec![addr])]))); + pool.update_proxy_maps( + std::collections::HashMap::from([(2, vec![(addr.ip(), addr.port())])]), + None, + ) + .await; let generation = pool.current_generation(); let writers = (1..=3) .map(|writer_id| { diff --git a/src/transport/middle_proxy/tests/send_adversarial_tests.rs b/src/transport/middle_proxy/tests/send_adversarial_tests.rs index 8f78562..a06f415 100644 --- a/src/transport/middle_proxy/tests/send_adversarial_tests.rs +++ b/src/transport/middle_proxy/tests/send_adversarial_tests.rs @@ -154,13 +154,11 @@ async fn insert_writer( }; pool.writers.write().await.push(writer); - { - let mut map = pool.proxy_map_v4.write().await; - map.entry(writer_dc) - .or_insert_with(Vec::new) - .push((addr.ip(), addr.port())); - } - pool.rebuild_endpoint_dc_map().await; + let mut map = pool.endpoint_snapshot.load().map_v4.clone(); + map.entry(writer_dc) + .or_insert_with(Vec::new) + .push((addr.ip(), addr.port())); + pool.update_proxy_maps(map, None).await; if register_in_registry { pool.registry .register_writer(writer_id, tx, byte_budget) @@ -241,11 +239,7 @@ async fn send_proxy_req_uses_live_same_dc_writer_while_preferred_endpoint_refill assert!(pool.admission_ready_conditional_cast().await); assert_eq!( - pool.preferred_endpoints_by_dc - .load() - .get(&2) - .cloned() - .unwrap_or_default(), + pool.preferred_endpoints_for_dc(2).await, vec![new_addr] ); diff --git a/src/web/http/tests.rs b/src/web/http/tests.rs index ca686f1..ea6bf6b 100644 --- a/src/web/http/tests.rs +++ b/src/web/http/tests.rs @@ -123,6 +123,7 @@ fn runtime_config_with_carriers_and_deadlines( carriers: Arc::clone(&carriers), carrier_negotiation_deadlines_secs, capability, + credential_id: [0; 16], key_fingerprint: "0000000000000000".to_string(), max_sessions: 4, max_streams: 16, diff --git a/src/web/manager/credentials.rs b/src/web/manager/credentials.rs index bc2caec..e00df38 100644 --- a/src/web/manager/credentials.rs +++ b/src/web/manager/credentials.rs @@ -89,11 +89,6 @@ impl WebProcessRuntime { .record_rejection(WebRejectionReason::ConfigDisabled); return Err(ManagerError::Closed); } - if !generation.proxy_shared.is_user_enabled(&profile.user) { - self.telemetry - .record_rejection(WebRejectionReason::UserDisabled); - return Err(ManagerError::Closed); - } let _operator_admission = self.try_operator_admission()?; let now = Instant::now(); let mut state = self.state.lock(); @@ -146,12 +141,25 @@ impl WebProcessRuntime { .record_rejection(WebRejectionReason::BootstrapCapacity); return Err(ManagerError::Limit); }; + let Some(mut user_publication) = generation + .proxy_shared + .claim_authenticated_user(&profile.user, profile.credential_id) + else { + self.telemetry + .record_rejection(WebRejectionReason::UserDisabled); + return Err(ManagerError::Closed); + }; + let Some(user_registration) = user_publication.take_registration() else { + return Err(ManagerError::Closed); + }; 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); + let issued_profile = Arc::clone(&profile); state.bootstraps.insert( hash, Bootstrap { + user_registration, expires_at: now + Duration::from_secs(config.web.timeouts.bootstrap_lifetime_secs), issued_at: now, issuance_ip: client_ip, @@ -186,11 +194,7 @@ impl WebProcessRuntime { }, ); *state.bootstraps_per_ip.entry(client_ip).or_insert(0) += 1; - let profile = state - .bootstraps - .get(&hash) - .map(|entry| Arc::clone(&entry.profile)) - .ok_or(ManagerError::Closed)?; + user_publication.commit(); drop(state); if recovery { self.telemetry @@ -202,7 +206,7 @@ impl WebProcessRuntime { Some(client_ip), crate::web::trace::TraceIdentity::from_optional_profile( Some(trace_session_id), - &profile, + &issued_profile, ), crate::web::trace::TraceLifecycleEvent::BridgeIssued, None, @@ -216,7 +220,7 @@ impl WebProcessRuntime { self.trace.record_profile_lifecycle( client_ip, Some(trace_session_id), - &profile, + &issued_profile, crate::web::trace::TraceLifecycleEvent::BridgeIssued, None, None, diff --git a/src/web/manager/lifecycle.rs b/src/web/manager/lifecycle.rs index 6ce297c..b1b9d06 100644 --- a/src/web/manager/lifecycle.rs +++ b/src/web/manager/lifecycle.rs @@ -30,14 +30,17 @@ pub(crate) struct WebShutdownDrain { } impl WebProcessRuntime { - /// Applies learning policy before publishing one new runtime generation. + /// Applies issuance and learning policy before publishing one new generation. pub(crate) fn activate_generation( &self, generation: Arc, ) -> Arc { let config = generation.config(); let (replaced, detached) = { + // Manager state precedes learning in the request-path lock order. + let mut state = self.state.lock(); let mut learning = self.learning.lock(); + state.apply_issuance_policy(generation.id, config.web.enabled); let outcome = learning.apply_policy( Instant::now(), generation.id, diff --git a/src/web/manager/session_creation.rs b/src/web/manager/session_creation.rs index 817dbde..696778f 100644 --- a/src/web/manager/session_creation.rs +++ b/src/web/manager/session_creation.rs @@ -69,7 +69,10 @@ impl WebProcessRuntime { let Some(entry) = state.bootstraps.get(&bootstrap_hash) else { return Err(ManagerError::Authentication); }; - if entry.profile.host != host || now > entry.expires_at { + if entry.profile.host != host + || now > entry.expires_at + || entry.user_registration.is_cancelled() + { return Err(ManagerError::Authentication); } if entry.used { @@ -301,6 +304,15 @@ impl WebProcessRuntime { return Err(ManagerError::Protocol); }; let _operator_admission = self.try_operator_admission()?; + let Some(mut user_publication) = generation + .proxy_shared + .claim_authenticated_user(&profile.user, profile.credential_id) + else { + return Err(ManagerError::Closed); + }; + let Some(user_registration) = user_publication.take_registration() else { + return Err(ManagerError::Closed); + }; if !admit_initial(self, &mut state, now, client_ip, profile_key, &profile) { return Err(ManagerError::Limit); } @@ -338,6 +350,7 @@ impl WebProcessRuntime { recovery, self.limits.clone(), issued_timeouts.clone(), + Some(user_registration), ); state.sessions.insert(session_hash, Arc::clone(&session)); *state.sessions_per_ip.entry(client_ip).or_insert(0) += 1; @@ -402,6 +415,7 @@ impl WebProcessRuntime { user_agent_id, }, ); + user_publication.commit(); drop(state); self.telemetry .record_carrier_selection(carrier, learning_disposition); diff --git a/src/web/manager/session_creation/replacement.rs b/src/web/manager/session_creation/replacement.rs index 1af2cef..865974b 100644 --- a/src/web/manager/session_creation/replacement.rs +++ b/src/web/manager/session_creation/replacement.rs @@ -43,14 +43,24 @@ impl WebProcessRuntime { if !valid || state.closed || !state.issuance_enabled - || !generation - .proxy_shared - .is_user_enabled(&replacement.profile.user) { drop(state); self.cancel_replacement(bootstrap_hash, &replacement.old_session); return Err(ManagerError::Closed); } + let Some(mut user_publication) = generation.proxy_shared.claim_authenticated_user( + &replacement.profile.user, + replacement.profile.credential_id, + ) else { + drop(state); + self.cancel_replacement(bootstrap_hash, &replacement.old_session); + return Err(ManagerError::Closed); + }; + let Some(user_registration) = user_publication.take_registration() else { + drop(state); + self.cancel_replacement(bootstrap_hash, &replacement.old_session); + return Err(ManagerError::Closed); + }; let Some((session_token, session_hash)) = new_unique_token(&generation, &state) else { self.record_limit_hit(); self.telemetry @@ -86,6 +96,7 @@ impl WebProcessRuntime { replacement.recovery, self.limits.clone(), replacement.old_session.timeouts().clone(), + Some(user_registration), ); let Some(supersede) = replacement.old_session.prepare_carrier_supersede() else { drop(state); @@ -146,6 +157,7 @@ impl WebProcessRuntime { index.bootstrap_hash = bootstrap_hash; index.attempt = replacement.attempt; } + user_publication.commit(); let identity = session.trace_identity(); let old_identity = replacement.old_session.trace_identity(); drop(state); diff --git a/src/web/manager/state.rs b/src/web/manager/state.rs index 6838d98..47a4f43 100644 --- a/src/web/manager/state.rs +++ b/src/web/manager/state.rs @@ -10,6 +10,7 @@ use zeroize::Zeroizing; use super::{CarrierRequest, ProfileKey, TOKEN_BYTES, TokenHash}; use crate::config::{WebCarrier, WebRuntimeConfig, WebRuntimeProfile, WebTimeoutsConfig}; use crate::maestro::generation::RuntimeGeneration; +use crate::proxy::user_admission::UserSessionRegistration; use crate::web::session::WebSession; use crate::web::telemetry::WebCarrierSelectionDisposition; @@ -34,6 +35,8 @@ impl CarrierChainPhase { /// One issued bootstrap and optional idempotent session-creation replay state. pub(super) struct Bootstrap { + /// User authority ownership retained for the credential lifetime. + pub(super) user_registration: UserSessionRegistration, /// Credential and replay-state expiry deadline. pub(super) expires_at: Instant, /// Stable ordering point used for bounded eviction. @@ -273,6 +276,7 @@ pub(super) fn matching_profile( profile.host == expected.host && profile.public_addr == expected.public_addr && profile.user == expected.user + && profile.credential_id == expected.credential_id && profile.secret_mode == expected.secret_mode && profile.carrier == expected.carrier && profile.carrier_negotiation_enabled == expected.carrier_negotiation_enabled diff --git a/src/web/session.rs b/src/web/session.rs index dd93eb2..e200e61 100644 --- a/src/web/session.rs +++ b/src/web/session.rs @@ -17,6 +17,7 @@ use crate::web::frame::{self, FrameType}; use crate::web::manager::{ CarrierClientClass, CarrierLearningContext, ProfileKey, TokenHash, WebProcessRuntime, }; +use crate::proxy::user_admission::UserSessionRegistration; // Backend tasks own generation admission and authenticated MTProxy relay lifetimes. mod backend; @@ -213,6 +214,7 @@ pub(crate) struct WebSession { created_at: Instant, limits: WebLimitsConfig, timeouts: WebTimeoutsConfig, + _user_registration: Option, state: Mutex, carrier_health_publication: AtomicU8, close_complete: AtomicBool, @@ -257,8 +259,13 @@ impl WebSession { recovery: bool, limits: WebLimitsConfig, timeouts: WebTimeoutsConfig, + user_registration: Option, ) -> Arc { let created_at = Instant::now(); + let cancel = user_registration + .as_ref() + .map(UserSessionRegistration::token) + .unwrap_or_default(); let mut carrier_lanes = HashMap::new(); let mut next_lane_instance = 1; if selected_carrier == WebCarrier::HttpsLanes { @@ -283,6 +290,7 @@ impl WebSession { created_at, limits, timeouts, + _user_registration: user_registration, state: Mutex::new(SessionState { streams: HashMap::new(), closing_streams: HashMap::new(), @@ -328,7 +336,7 @@ impl WebSession { close_notify: Notify::new(), down_notify: Arc::new(Notify::new()), lane_open_notify: Arc::new(Notify::new()), - cancel: CancellationToken::new(), + cancel, tasks_live: AtomicUsize::new(0), tasks_done: Arc::new(Notify::new()), resident: Arc::new(resident::ResidentCounters::default()), diff --git a/src/web/session/backend_tests.rs b/src/web/session/backend_tests.rs index 6704c7f..3f04878 100644 --- a/src/web/session/backend_tests.rs +++ b/src/web/session/backend_tests.rs @@ -75,6 +75,7 @@ fn test_runtime_with_dc( carriers: Arc::from([carrier]), carrier_negotiation_deadlines_secs: [3, 5, 8, 12], capability: [7; 32], + credential_id: [0; 16], key_fingerprint: "0000000000000000".to_string(), max_sessions: 4, max_streams: 16, @@ -134,6 +135,7 @@ fn test_runtime_with_dc( false, limits, timeouts, + None, ); TestRuntime { session, diff --git a/src/web/session/downlink_tests.rs b/src/web/session/downlink_tests.rs index 29bdb58..c444b0f 100644 --- a/src/web/session/downlink_tests.rs +++ b/src/web/session/downlink_tests.rs @@ -24,6 +24,7 @@ fn session() -> (Arc, Arc) { carriers: Arc::from([WebCarrier::Https]), carrier_negotiation_deadlines_secs: [3, 5, 8, 12], capability: [0; 32], + credential_id: [0; 16], key_fingerprint: "0000000000000000".to_string(), max_sessions: 1, max_streams: 1, @@ -50,6 +51,7 @@ fn session() -> (Arc, Arc) { false, WebLimitsConfig::default(), timeouts, + None, ); (session, manager) } diff --git a/src/web/session/lanes/tests.rs b/src/web/session/lanes/tests.rs index a836009..020ea5f 100644 --- a/src/web/session/lanes/tests.rs +++ b/src/web/session/lanes/tests.rs @@ -39,6 +39,7 @@ fn new_session_with_automatic( carriers: Arc::from([WebCarrier::HttpsLanes]), carrier_negotiation_deadlines_secs: [3, 5, 8, 12], capability: [0; 32], + credential_id: [0; 16], key_fingerprint: "0000000000000000".to_string(), max_sessions: 1, max_streams: 2, @@ -65,6 +66,7 @@ fn new_session_with_automatic( false, limits, WebTimeoutsConfig::default(), + None, ) } diff --git a/src/web/session/lifecycle.rs b/src/web/session/lifecycle.rs index f3f21ca..f58986c 100644 --- a/src/web/session/lifecycle.rs +++ b/src/web/session/lifecycle.rs @@ -25,6 +25,8 @@ pub(crate) enum SessionCloseReason { WebSocketEnded, /// An authenticated control-plane request selected this session. ApiClose, + /// Process user authority revoked or replaced the authenticated credential. + UserDisabled, /// A graceful operator drain reached its force-close deadline. OperatorForce, /// Terminal process shutdown closed all remaining sessions. @@ -33,7 +35,7 @@ pub(crate) enum SessionCloseReason { impl SessionCloseReason { /// Complete fixed reason set in stable API and metric order. - pub(crate) const ALL: [Self; 11] = [ + pub(crate) const ALL: [Self; 12] = [ Self::ClientDelete, Self::BridgeRecovery, Self::PeerIdle, @@ -43,6 +45,7 @@ impl SessionCloseReason { Self::Backpressure, Self::WebSocketEnded, Self::ApiClose, + Self::UserDisabled, Self::OperatorForce, Self::RuntimeShutdown, ]; @@ -59,6 +62,7 @@ impl SessionCloseReason { Self::Backpressure => "backpressure", Self::WebSocketEnded => "websocket_ended", Self::ApiClose => "api_close", + Self::UserDisabled => "user_disabled", Self::OperatorForce => "operator_force", Self::RuntimeShutdown => "runtime_shutdown", } @@ -217,6 +221,20 @@ 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; + } + } let healthy = { let mut state = self.state.lock(); self.carrier_health_ready_locked(&mut state, now) diff --git a/src/web/session/negotiation.rs b/src/web/session/negotiation.rs index dd2a96d..3e1de9d 100644 --- a/src/web/session/negotiation.rs +++ b/src/web/session/negotiation.rs @@ -320,6 +320,7 @@ mod tests { carriers: Arc::from([carrier]), carrier_negotiation_deadlines_secs: [3, 5, 8, 12], capability: [0; 32], + credential_id: [0; 16], key_fingerprint: "0000000000000000".to_string(), max_sessions: 1, max_streams: 1, @@ -342,6 +343,7 @@ mod tests { false, WebLimitsConfig::default(), WebTimeoutsConfig::default(), + None, ) } diff --git a/src/web/session/uplink_tests.rs b/src/web/session/uplink_tests.rs index ed09655..69a46b6 100644 --- a/src/web/session/uplink_tests.rs +++ b/src/web/session/uplink_tests.rs @@ -22,6 +22,7 @@ fn session_with_automatic(automatic: bool) -> Arc { carriers: Arc::from([WebCarrier::Https]), carrier_negotiation_deadlines_secs: [3, 5, 8, 12], capability: [0; 32], + credential_id: [0; 16], key_fingerprint: "0000000000000000".to_string(), max_sessions: 1, max_streams: 1, @@ -48,6 +49,7 @@ fn session_with_automatic(automatic: bool) -> Arc { false, WebLimitsConfig::default(), WebTimeoutsConfig::default(), + None, ) } diff --git a/src/web/session/websocket/tests.rs b/src/web/session/websocket/tests.rs index f35c043..7a4900c 100644 --- a/src/web/session/websocket/tests.rs +++ b/src/web/session/websocket/tests.rs @@ -39,6 +39,7 @@ fn runtime(admission: bool) -> TestRuntime { carriers: Arc::from([WebCarrier::WebsocketLanes]), carrier_negotiation_deadlines_secs: [3, 5, 8, 12], capability: [7; 32], + credential_id: [0; 16], key_fingerprint: "0000000000000000".to_string(), max_sessions: 2, max_streams: 1, @@ -75,6 +76,7 @@ fn runtime(admission: bool) -> TestRuntime { false, limits, timeouts, + None, ); TestRuntime { session,