Proxy Shared User Drafts

This commit is contained in:
Alexey
2026-09-14 20:19:48 +03:00
parent b37f1ebdeb
commit 0ac236955a
59 changed files with 1485 additions and 401 deletions
-17
View File
@@ -21,7 +21,6 @@ pub(super) async fn create_user_route(
} }
let expected_revision = parse_if_match(req.headers()); let expected_revision = parse_if_match(req.headers());
let body = read_json::<CreateUserRequest>(req.into_body(), body_limit).await?; let body = read_json::<CreateUserRequest>(req.into_body(), body_limit).await?;
let requested_enabled = body.enabled;
let result = create_user(body, expected_revision, shared).await; let result = create_user(body, expected_revision, shared).await;
let (mut data, revision) = match result { let (mut data, revision) = match result {
Ok(ok) => ok, Ok(ok) => ok,
@@ -34,22 +33,6 @@ pub(super) async fn create_user_route(
}; };
let runtime_cfg = config_rx.borrow().clone(); let runtime_cfg = config_rx.borrow().clone();
data.user.in_runtime = runtime_cfg.access.users.contains_key(&data.user.username); 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( shared.runtime_events.record(
"api.user.create.ok", "api.user.create.ok",
format!("username={}", data.user.username), format!("username={}", data.user.username),
+2 -28
View File
@@ -61,7 +61,6 @@ pub(super) async fn handle(
}; };
let runtime_cfg = config_rx.borrow().clone(); let runtime_cfg = config_rx.borrow().clone();
data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); data.in_runtime = runtime_cfg.access.users.contains_key(&data.username);
shared.proxy_shared.set_user_enabled(base_user, true);
shared shared
.runtime_events .runtime_events
.record("api.user.enable.ok", format!("username={}", base_user)); .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(); let runtime_cfg = config_rx.borrow().clone();
data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); 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( shared.runtime_events.record(
"api.user.disable.ok", "api.user.disable.ok",
format!( format!("username={}", base_user),
"username={} newly_disabled={} cancelled_sessions={}",
base_user, newly_disabled, cancelled
),
); );
let status = if data.in_runtime { let status = if data.in_runtime {
StatusCode::OK StatusCode::OK
@@ -270,11 +265,6 @@ pub(super) async fn handle(
} }
let expected_revision = parse_if_match(req.headers()); let expected_revision = parse_if_match(req.headers());
let body = read_json::<PatchUserRequest>(req.into_body(), body_limit).await?; let body = read_json::<PatchUserRequest>(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 result = patch_user(user, body, expected_revision, shared).await;
let (mut data, revision) = match result { let (mut data, revision) = match result {
Ok(ok) => ok, Ok(ok) => ok,
@@ -288,20 +278,6 @@ pub(super) async fn handle(
}; };
let runtime_cfg = config_rx.borrow().clone(); let runtime_cfg = config_rx.borrow().clone();
data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); 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 shared
.runtime_events .runtime_events
.record("api.user.patch.ok", format!("username={}", data.username)); .record("api.user.patch.ok", format!("username={}", data.username));
@@ -335,11 +311,9 @@ pub(super) async fn handle(
return Err(error); 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( shared.runtime_events.record(
"api.user.delete.ok", "api.user.delete.ok",
format!("username={} cancelled_sessions={}", deleted_user, cancelled), format!("username={}", deleted_user),
); );
let runtime_cfg = config_rx.borrow().clone(); let runtime_cfg = config_rx.borrow().clone();
let in_runtime = runtime_cfg.access.users.contains_key(&deleted_user); let in_runtime = runtime_cfg.access.users.contains_key(&deleted_user);
-1
View File
@@ -70,7 +70,6 @@ use model::{
PatchUserRequest, ResetUserQuotaResponse, RotateSecretRequest, SummaryData, UserActiveIps, PatchUserRequest, ResetUserQuotaResponse, RotateSecretRequest, SummaryData, UserActiveIps,
is_valid_username, is_valid_username,
}; };
use patch::Patch;
use runtime_edge::{ use runtime_edge::{
EdgeConnectionsCacheEntry, build_runtime_connections_summary_data, EdgeConnectionsCacheEntry, build_runtime_connections_summary_data,
build_runtime_events_recent_data, build_runtime_tls_fingerprints_data, build_runtime_events_recent_data, build_runtime_tls_fingerprints_data,
+9 -1
View File
@@ -124,7 +124,14 @@ pub(in crate::api) async fn create_user(
let revision = let revision =
save_access_sections_to_disk(&shared.config_path, &cfg, &touched_sections).await?; 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 { if let Some(limit) = updated_limit {
shared shared
@@ -132,6 +139,7 @@ pub(in crate::api) async fn create_user(
.set_user_limit(&body.username, limit) .set_user_limit(&body.username, limit)
.await; .await;
} }
drop(_guard);
let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips();
let users = users_from_config( let users = users_from_config(
+10 -2
View File
@@ -31,6 +31,10 @@ pub(in crate::api) async fn rotate_secret(
.map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?; .map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?;
let revision = let revision =
save_access_sections_to_disk(&shared.config_path, &cfg, &[AccessSection::Users]).await?; 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); drop(_guard);
let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); 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)))?; .map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?;
let revision = let revision =
save_access_sections_to_disk(&shared.config_path, &cfg, &touched_sections).await?; 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(); let configured_users = cfg.access.users.keys().cloned().collect();
if let Err(error) = shared if let Err(error) = shared
.quota_state .quota_state
@@ -121,9 +126,12 @@ pub(in crate::api) async fn delete_user(
"Deleted user quota checkpoint cleanup will be reconciled on restart" "Deleted user quota checkpoint cleanup will be reconciled on restart"
); );
} }
drop(_guard);
shared.ip_tracker.remove_user_limit(user).await; 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)) Ok((user.to_string(), revision))
} }
+21 -1
View File
@@ -170,12 +170,23 @@ pub(in crate::api) async fn patch_user(
} else { } else {
save_access_sections_to_disk(&shared.config_path, &cfg, &touched_sections).await? 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 { match max_unique_ips_change {
Some(Some(limit)) => shared.ip_tracker.set_user_limit(user, limit).await, Some(Some(limit)) => shared.ip_tracker.set_user_limit(user, limit).await,
Some(None) => shared.ip_tracker.remove_user_limit(user).await, Some(None) => shared.ip_tracker.remove_user_limit(user).await,
None => {} None => {}
} }
drop(_guard);
let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips();
let users = users_from_config( let users = users_from_config(
&cfg, &cfg,
@@ -223,6 +234,15 @@ pub(in crate::api) async fn set_user_enabled(
let revision = let revision =
save_access_sections_to_disk(&shared.config_path, &cfg, &[AccessSection::UserEnabled]) save_access_sections_to_disk(&shared.config_path, &cfg, &[AccessSection::UserEnabled])
.await?; .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); drop(_guard);
let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips(); let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips();
+20
View File
@@ -9,6 +9,7 @@ use rand::RngExt;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use tracing::warn; use tracing::warn;
use crate::crypto::sha256;
use crate::error::{ProxyError, Result}; use crate::error::{ProxyError, Result};
use super::defaults::*; use super::defaults::*;
@@ -230,6 +231,25 @@ impl ProxyConfig {
self.runtime_user_auth.as_deref() 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. /// Validates cross-field configuration invariants after deserialization.
pub fn validate(&self) -> Result<()> { pub fn validate(&self) -> Result<()> {
if self.access.users.is_empty() { if self.access.users.is_empty() {
+13
View File
@@ -2,6 +2,7 @@ use std::collections::HashMap;
use std::collections::hash_map::DefaultHasher; use std::collections::hash_map::DefaultHasher;
use std::hash::Hasher; use std::hash::Hasher;
use crate::crypto::sha256;
use crate::error::{ProxyError, Result}; use crate::error::{ProxyError, Result};
const ACCESS_SECRET_BYTES: usize = 16; const ACCESS_SECRET_BYTES: usize = 16;
@@ -19,6 +20,8 @@ pub(crate) struct UserAuthSnapshot {
pub(crate) struct UserAuthEntry { pub(crate) struct UserAuthEntry {
pub(crate) user: String, pub(crate) user: String,
pub(crate) secret: [u8; ACCESS_SECRET_BYTES], pub(crate) secret: [u8; ACCESS_SECRET_BYTES],
/// Stable secret identity used by process-wide admission fencing.
pub(crate) credential_id: [u8; 16],
} }
impl UserAuthSnapshot { impl UserAuthSnapshot {
@@ -46,9 +49,13 @@ impl UserAuthSnapshot {
let mut secret = [0u8; ACCESS_SECRET_BYTES]; let mut secret = [0u8; ACCESS_SECRET_BYTES];
secret.copy_from_slice(&decoded); 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 { entries.push(UserAuthEntry {
user: user.clone(), user: user.clone(),
secret, secret,
credential_id,
}); });
by_name.insert(user.clone(), user_id); by_name.insert(user.clone(), user_id);
sni_index sni_index
@@ -88,6 +95,12 @@ impl UserAuthSnapshot {
self.entries.get(idx) 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]> { pub(crate) fn sni_candidates(&self, sni: &str) -> Option<&[u32]> {
self.sni_index self.sni_index
.get(&Self::sni_lookup_hash(sni)) .get(&Self::sni_lookup_hash(sni))
+1
View File
@@ -63,6 +63,7 @@ pub(super) fn rebuild(config: &mut ProxyConfig) -> Result<()> {
host: vhost.host.clone(), host: vhost.host.clone(),
public_addr: vhost.public_addr, public_addr: vhost.public_addr,
user: profile.user.clone(), user: profile.user.clone(),
credential_id: auth_entry.credential_id,
secret_mode: profile.secret_mode, secret_mode: profile.secret_mode,
carrier: config.web.carrier, carrier: config.web.carrier,
carrier_negotiation_enabled: config.web.carrier_negotiation_enabled(), carrier_negotiation_enabled: config.web.carrier_negotiation_enabled(),
+2
View File
@@ -35,6 +35,8 @@ pub(crate) struct WebRuntimeProfile {
pub(crate) public_addr: SocketAddr, pub(crate) public_addr: SocketAddr,
/// Exact access user authenticated by logical streams. /// Exact access user authenticated by logical streams.
pub(crate) user: String, 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. /// Client secret representation and inner protocol policy.
pub(crate) secret_mode: WebSecretMode, pub(crate) secret_mode: WebSecretMode,
/// Sole carrier or final fallback frozen into the issued bridge policy. /// Sole carrier or final fallback frozen into the issued bridge policy.
+11 -9
View File
@@ -15,6 +15,7 @@ use arc_swap::ArcSwap;
use tokio::sync::{Mutex as AsyncMutex, RwLock}; use tokio::sync::{Mutex as AsyncMutex, RwLock};
use crate::config::UserMaxUniqueIpsMode; use crate::config::UserMaxUniqueIpsMode;
use crate::proxy::user_admission::UserIncarnation;
const CLEANUP_DRAIN_BATCH_LIMIT: usize = 1024; const CLEANUP_DRAIN_BATCH_LIMIT: usize = 1024;
const MAX_ACTIVE_IP_ENTRIES: u64 = 131_072; const MAX_ACTIVE_IP_ENTRIES: u64 = 131_072;
@@ -32,11 +33,12 @@ mod tests;
struct UserIpShard { struct UserIpShard {
active_ips: HashMap<String, HashMap<IpAddr, usize>>, active_ips: HashMap<String, HashMap<IpAddr, usize>>,
recent_ips: HashMap<String, HashMap<IpAddr, Instant>>, recent_ips: HashMap<String, HashMap<IpAddr, Instant>>,
incarnations: HashMap<String, UserIncarnation>,
} }
#[derive(Debug, Default)] #[derive(Debug, Default)]
struct CleanupShard { struct CleanupShard {
queue: Mutex<HashMap<String, HashMap<IpAddr, usize>>>, queue: Mutex<HashMap<(String, UserIncarnation), HashMap<IpAddr, usize>>>,
} }
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
@@ -194,19 +196,19 @@ impl UserIpTracker {
} }
pub(super) fn pop_one_cleanup( pub(super) fn pop_one_cleanup(
queue: &mut HashMap<String, HashMap<IpAddr, usize>>, queue: &mut HashMap<(String, UserIncarnation), HashMap<IpAddr, usize>>,
) -> Option<(String, IpAddr, usize)> { ) -> Option<(String, UserIncarnation, IpAddr, usize)> {
let user = queue.keys().next().cloned()?; let owner = queue.keys().next().cloned()?;
let ip = queue.get(&user)?.keys().next().copied()?; let ip = queue.get(&owner)?.keys().next().copied()?;
let count = queue.get_mut(&user)?.remove(&ip)?; let count = queue.get_mut(&owner)?.remove(&ip)?;
let remove_user = queue let remove_user = queue
.get(&user) .get(&owner)
.map(|user_queue| user_queue.is_empty()) .map(|user_queue| user_queue.is_empty())
.unwrap_or(false); .unwrap_or(false);
if remove_user { if remove_user {
queue.remove(&user); queue.remove(&owner);
} }
Some((user, ip, count)) Some((owner.0, owner.1, ip, count))
} }
#[cfg(test)] #[cfg(test)]
+47
View File
@@ -59,6 +59,16 @@ impl UserIpTracker {
} }
pub async fn check_and_add(&self, username: &str, ip: IpAddr) -> Result<(), String> { 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.drain_cleanup_for_user(username).await;
self.maybe_compact_empty_users().await; self.maybe_compact_empty_users().await;
let policy = self.limit_policy.load(); let policy = self.limit_policy.load();
@@ -69,6 +79,30 @@ impl UserIpTracker {
let shard_idx = Self::shard_idx(username); let shard_idx = Self::shard_idx(username);
let mut shard = self.shards[shard_idx].write().await; 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 user_active = shard.active_ips.entry(username.to_string()).or_default();
let active_contains_ip = user_active.contains_key(&ip); let active_contains_ip = user_active.contains_key(&ip);
let active_len = user_active.len(); let active_len = user_active.len();
@@ -174,9 +208,22 @@ impl UserIpTracker {
} }
pub async fn remove_ip(&self, username: &str, ip: IpAddr) { 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; self.maybe_compact_empty_users().await;
let shard_idx = Self::shard_idx(username); let shard_idx = Self::shard_idx(username);
let mut shard = self.shards[shard_idx].write().await; let mut shard = self.shards[shard_idx].write().await;
if shard.incarnations.get(username).copied() != Some(incarnation) {
return;
}
let mut removed_active_entries = 0usize; let mut removed_active_entries = 0usize;
if let Some(user_ips) = shard.active_ips.get_mut(username) { if let Some(user_ips) = shard.active_ips.get_mut(username) {
if let Some(count) = user_ips.get_mut(&ip) { if let Some(count) = user_ips.get_mut(&ip) {
+58 -14
View File
@@ -3,12 +3,22 @@ use super::*;
impl UserIpTracker { impl UserIpTracker {
/// Queues a deferred active IP cleanup for a later async drain. /// Queues a deferred active IP cleanup for a later async drain.
pub fn enqueue_cleanup(&self, user: String, ip: IpAddr) { 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(); self.observe_cleanup_poison_for_tests();
let shard_idx = Self::shard_idx(&user); let shard_idx = Self::shard_idx(&user);
let cleanup_shard = &self.cleanup_shards[shard_idx]; let cleanup_shard = &self.cleanup_shards[shard_idx];
match cleanup_shard.queue.lock() { match cleanup_shard.queue.lock() {
Ok(mut queue) => { 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); let count = user_queue.entry(ip).or_insert(0);
if *count == 0 { if *count == 0 {
self.cleanup_queue_len.fetch_add(1, Ordering::Relaxed); self.cleanup_queue_len.fetch_add(1, Ordering::Relaxed);
@@ -19,7 +29,7 @@ impl UserIpTracker {
} }
Err(poisoned) => { Err(poisoned) => {
let mut queue = poisoned.into_inner(); 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); let count = user_queue.entry(ip).or_insert(0);
if *count == 0 { if *count == 0 {
self.cleanup_queue_len.fetch_add(1, Ordering::Relaxed); self.cleanup_queue_len.fetch_add(1, Ordering::Relaxed);
@@ -65,10 +75,10 @@ impl UserIpTracker {
let shard_idx = Self::shard_idx(user); let shard_idx = Self::shard_idx(user);
let cleanup_shard = &self.cleanup_shards[shard_idx]; let cleanup_shard = &self.cleanup_shards[shard_idx];
let to_remove = match cleanup_shard.queue.lock() { 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) => { Err(poisoned) => {
let mut queue = poisoned.into_inner(); 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(); cleanup_shard.queue.clear_poison();
drained drained
} }
@@ -76,14 +86,23 @@ impl UserIpTracker {
if to_remove.is_empty() { if to_remove.is_empty() {
return; return;
} }
let removed_queue_entries = to_remove
.iter()
.map(|(_, ips)| ips.len())
.sum::<usize>();
self.cleanup_queue_len 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 shard = self.shards[shard_idx].write().await;
let mut removed_active_entries = 0usize; let mut removed_active_entries = 0usize;
for (ip, pending_count) in to_remove { for (incarnation, ips) in to_remove {
removed_active_entries = removed_active_entries.saturating_add( if shard.incarnations.get(user).copied() != Some(incarnation) {
Self::apply_active_cleanup(&mut shard.active_ips, user, ip, pending_count), 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); Self::decrement_counter(&self.active_entry_count, removed_active_entries);
} }
@@ -103,11 +122,13 @@ impl UserIpTracker {
let mut drained = let mut drained =
HashMap::with_capacity(queue.len().min(CLEANUP_DRAIN_BATCH_LIMIT)); HashMap::with_capacity(queue.len().min(CLEANUP_DRAIN_BATCH_LIMIT));
for _ in 0..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; break;
}; };
self.cleanup_queue_len.fetch_sub(1, Ordering::Relaxed); self.cleanup_queue_len.fetch_sub(1, Ordering::Relaxed);
drained.insert((user, ip), count); drained.insert((user, incarnation, ip), count);
} }
drained drained
} }
@@ -120,11 +141,13 @@ impl UserIpTracker {
let mut drained = let mut drained =
HashMap::with_capacity(queue.len().min(CLEANUP_DRAIN_BATCH_LIMIT)); HashMap::with_capacity(queue.len().min(CLEANUP_DRAIN_BATCH_LIMIT));
for _ in 0..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; break;
}; };
self.cleanup_queue_len.fetch_sub(1, Ordering::Relaxed); 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(); cleanup_shard.queue.clear_poison();
drained drained
@@ -138,7 +161,10 @@ impl UserIpTracker {
let mut shard = self.shards[shard_idx].write().await; let mut shard = self.shards[shard_idx].write().await;
let mut removed_active_entries = 0usize; 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( removed_active_entries = removed_active_entries.saturating_add(
Self::apply_active_cleanup(&mut shard.active_ips, &user, ip, pending_count), 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); Self::decrement_counter(&self.active_entry_count, removed_active_entries);
} }
} }
fn drain_user_cleanup(
queue: &mut HashMap<(String, UserIncarnation), HashMap<IpAddr, usize>>,
user: &str,
) -> Vec<(UserIncarnation, HashMap<IpAddr, usize>)> {
let owners = queue
.keys()
.filter(|(queued_user, _)| queued_user == user)
.cloned()
.collect::<Vec<_>>();
owners
.into_iter()
.filter_map(|owner| {
let incarnation = owner.1;
queue.remove(&owner).map(|ips| (incarnation, ips))
})
.collect()
}
+18
View File
@@ -228,8 +228,25 @@ impl UserIpTracker {
} }
pub async fn clear_user_ips(&self, username: &str) { 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 shard_idx = Self::shard_idx(username);
let mut shard = self.shards[shard_idx].write().await; 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 let removed_active_entries = shard
.active_ips .active_ips
.remove(username) .remove(username)
@@ -250,6 +267,7 @@ impl UserIpTracker {
let mut shard = shard_lock.write().await; let mut shard = shard_lock.write().await;
shard.active_ips.clear(); shard.active_ips.clear();
shard.recent_ips.clear(); shard.recent_ips.clear();
shard.incarnations.clear();
} }
self.active_entry_count.store(0, Ordering::Relaxed); self.active_entry_count.store(0, Ordering::Relaxed);
self.recent_entry_count.store(0, Ordering::Relaxed); self.recent_entry_count.store(0, Ordering::Relaxed);
+1 -1
View File
@@ -108,7 +108,7 @@ pub(super) async fn run_telemt_core(
); );
let shared_state = let shared_state =
ProxySharedState::new_with_direct_buffer_budget(direct_buffer_budget.clone()); 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( shared_state.traffic_limiter.apply_policy(
config.access.user_rate_limits.clone(), config.access.user_rate_limits.clone(),
config.access.cidr_rate_limits.clone(), config.access.cidr_rate_limits.clone(),
+8
View File
@@ -177,6 +177,7 @@ impl ReloadSupervisor {
self.quota_store.clone(), self.quota_store.clone(),
self.runtime_log_filter.clone(), self.runtime_log_filter.clone(),
self.tls_full_cert_budget.clone(), self.tls_full_cert_budget.clone(),
old_runtime.proxy_shared.user_admission(),
) )
.await .await
{ {
@@ -277,6 +278,7 @@ impl ReloadSupervisor {
generation: new_runtime, generation: new_runtime,
detected_ips, detected_ips,
config_watcher_activation, config_watcher_activation,
user_admission_epoch,
} = prepared; } = prepared;
let pending_listener_transition = if let Some(listener_transition) = listener_transition { let pending_listener_transition = if let Some(listener_transition) = listener_transition {
match self match self
@@ -300,6 +302,12 @@ impl ReloadSupervisor {
}; };
let replaced = { let replaced = {
let listener_manager = self.listener_manager.lock().await; 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(); old_runtime.stop_accepting_sessions();
listener_manager.activate_runtime_generation(new_runtime.clone()) listener_manager.activate_runtime_generation(new_runtime.clone())
}; };
+2
View File
@@ -23,10 +23,12 @@ fn runtime_log_filter() -> RuntimeLogFilter {
fn prepared_runtime(generation: Arc<RuntimeGeneration>) -> PreparedRuntime { fn prepared_runtime(generation: Arc<RuntimeGeneration>) -> PreparedRuntime {
let (config_watcher_activation, _activation_rx) = watch::channel(false); let (config_watcher_activation, _activation_rx) = watch::channel(false);
let user_admission_epoch = generation.proxy_shared.user_admission().epoch();
PreparedRuntime { PreparedRuntime {
generation, generation,
detected_ips: (None, None), detected_ips: (None, None),
config_watcher_activation, config_watcher_activation,
user_admission_epoch,
} }
} }
+10 -3
View File
@@ -16,6 +16,7 @@ use crate::proxy::direct_buffer_budget::{
}; };
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController}; use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::shared_state::ProxySharedState; use crate::proxy::shared_state::ProxySharedState;
use crate::proxy::user_admission::UserAdmissionAuthority;
use crate::startup::StartupTracker; use crate::startup::StartupTracker;
use crate::stats::beobachten::BeobachtenStore; use crate::stats::beobachten::BeobachtenStore;
use crate::stats::telemetry::TelemetryPolicy; use crate::stats::telemetry::TelemetryPolicy;
@@ -39,6 +40,8 @@ pub(crate) struct PreparedRuntime {
pub(crate) detected_ips: (Option<IpAddr>, Option<IpAddr>), pub(crate) detected_ips: (Option<IpAddr>, Option<IpAddr>),
/// Gate opened only after the candidate becomes the active generation. /// Gate opened only after the candidate becomes the active generation.
pub(crate) config_watcher_activation: watch::Sender<bool>, pub(crate) config_watcher_activation: watch::Sender<bool>,
/// User-authority epoch captured before candidate construction.
pub(crate) user_admission_epoch: u64,
} }
pub(crate) async fn prepare_runtime( pub(crate) async fn prepare_runtime(
@@ -48,7 +51,9 @@ pub(crate) async fn prepare_runtime(
quota_store: Arc<QuotaStore>, quota_store: Arc<QuotaStore>,
runtime_log_filter: RuntimeLogFilter, runtime_log_filter: RuntimeLogFilter,
tls_full_cert_budget: Arc<TlsFullCertBudget>, tls_full_cert_budget: Arc<TlsFullCertBudget>,
user_admission: Arc<UserAdmissionAuthority>,
) -> Result<PreparedRuntime, String> { ) -> Result<PreparedRuntime, String> {
let user_admission_epoch = user_admission.epoch();
config config
.validate_web_decoy_listener_separation() .validate_web_decoy_listener_separation()
.map_err(|error| error.to_string())?; .map_err(|error| error.to_string())?;
@@ -92,9 +97,10 @@ pub(crate) async fn prepare_runtime(
let hard_limit = let hard_limit =
resolve_direct_buffer_hard_limit(config.general.direct_relay_buffer_budget_max_bytes).await; resolve_direct_buffer_hard_limit(config.general.direct_relay_buffer_budget_max_bytes).await;
let direct_buffer_budget = DirectBufferBudget::new(hard_limit); let direct_buffer_budget = DirectBufferBudget::new(hard_limit);
let proxy_shared = let proxy_shared = ProxySharedState::new_with_direct_buffer_budget_and_user_admission(
ProxySharedState::new_with_direct_buffer_budget(direct_buffer_budget.clone()); direct_buffer_budget.clone(),
proxy_shared.apply_user_enabled_config(&config.access.user_enabled); user_admission,
);
proxy_shared.traffic_limiter.apply_policy( proxy_shared.traffic_limiter.apply_policy(
config.access.user_rate_limits.clone(), config.access.user_rate_limits.clone(),
config.access.cidr_rate_limits.clone(), config.access.cidr_rate_limits.clone(),
@@ -311,6 +317,7 @@ pub(crate) async fn prepare_runtime(
Ok(PreparedRuntime { Ok(PreparedRuntime {
generation, generation,
config_watcher_activation, config_watcher_activation,
user_admission_epoch,
detected_ips: ( detected_ips: (
probe.detected_ipv4.map(IpAddr::V4), probe.detected_ipv4.map(IpAddr::V4),
probe.detected_ipv6.map(IpAddr::V6), probe.detected_ipv6.map(IpAddr::V6),
+2 -2
View File
@@ -288,8 +288,8 @@ pub(crate) async fn spawn_runtime_tasks(
break; break;
} }
let cfg = config_rx_user_enabled.borrow_and_update().clone(); let cfg = config_rx_user_enabled.borrow_and_update().clone();
for (user, cancelled) in for (user, cancelled) in shared_user_enabled
shared_user_enabled.apply_user_enabled_config(&cfg.access.user_enabled) .apply_user_config(&cfg.access.users, &cfg.access.user_enabled)
{ {
if cancelled > 0 { if cancelled > 0 {
info!( info!(
+74 -9
View File
@@ -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::middle_relay::{handle_via_middle_proxy, handle_via_middle_proxy_with_conntrack};
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController}; use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::shared_state::{ConntrackClosePolicy, ProxySharedState}; use crate::proxy::shared_state::{ConntrackClosePolicy, ProxySharedState};
use crate::proxy::user_admission::UserIncarnation;
use crate::stats::Stats; use crate::stats::Stats;
use crate::stream::{BufferPool, CryptoReader, CryptoWriter}; use crate::stream::{BufferPool, CryptoReader, CryptoWriter};
use crate::transport::UpstreamManager; use crate::transport::UpstreamManager;
@@ -59,13 +60,21 @@ where
W: AsyncWrite + Unpin + Send + 'static, W: AsyncWrite + Unpin + Send + 'static,
{ {
let user = success.user.clone(); 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"); warn!(user = %user, "Disabled user rejected");
return Err(ProxyError::UserDisabled { user }); return Err(ProxyError::UserDisabled { user });
} };
let user_reservation = acquire_user_connection_reservation( let user_reservation = acquire_user_connection_reservation_for_incarnation(
&user, &user,
user_incarnation,
&deps.config, &deps.config,
Arc::clone(&deps.stats), Arc::clone(&deps.stats),
peer_addr, peer_addr,
@@ -79,11 +88,20 @@ where
let route_snapshot = deps.route_runtime.snapshot(); let route_snapshot = deps.route_runtime.snapshot();
let session_id = deps.rng.u64(); 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(); user_reservation.release_deferred();
warn!(user = %user, "Disabled user rejected during final admission"); warn!(user = %user, "Disabled user rejected during final admission");
return Err(ProxyError::UserDisabled { user }); 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 session_cancel = user_session.token();
let selected_me_pool = if deps.config.general.use_middle_proxy let selected_me_pool = if deps.config.general.use_middle_proxy
&& matches!(route_snapshot.mode, RelayRouteMode::Middle) && matches!(route_snapshot.mode, RelayRouteMode::Middle)
@@ -216,6 +234,7 @@ pub(crate) struct UserConnectionReservation {
ip_tracker: Arc<UserIpTracker>, ip_tracker: Arc<UserIpTracker>,
user: String, user: String,
ip: IpAddr, ip: IpAddr,
incarnation: UserIncarnation,
tracks_ip: bool, tracks_ip: bool,
active: bool, active: bool,
} }
@@ -228,12 +247,25 @@ impl UserConnectionReservation {
user: String, user: String,
ip: IpAddr, ip: IpAddr,
tracks_ip: bool, 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<Stats>,
ip_tracker: Arc<UserIpTracker>,
user: String,
ip: IpAddr,
incarnation: UserIncarnation,
tracks_ip: bool,
) -> Self { ) -> Self {
Self { Self {
stats, stats,
ip_tracker, ip_tracker,
user, user,
ip, ip,
incarnation,
tracks_ip, tracks_ip,
active: true, active: true,
} }
@@ -246,7 +278,9 @@ impl UserConnectionReservation {
} }
self.active = false; self.active = false;
if self.tracks_ip { 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); self.stats.decrement_user_curr_connects(&self.user);
} }
@@ -259,7 +293,11 @@ impl UserConnectionReservation {
self.active = false; self.active = false;
self.stats.decrement_user_curr_connects(&self.user); self.stats.decrement_user_curr_connects(&self.user);
if self.tracks_ip { 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.increment_session_drop_fallback_total();
self.stats.decrement_user_curr_connects(&self.user); self.stats.decrement_user_curr_connects(&self.user);
if self.tracks_ip { 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<Stats>, stats: Arc<Stats>,
peer_addr: SocketAddr, peer_addr: SocketAddr,
ip_tracker: Arc<UserIpTracker>, ip_tracker: Arc<UserIpTracker>,
) -> Result<UserConnectionReservation> {
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<Stats>,
peer_addr: SocketAddr,
ip_tracker: Arc<UserIpTracker>,
) -> Result<UserConnectionReservation> { ) -> Result<UserConnectionReservation> {
if let Some(expiration) = config.access.user_expirations.get(user) if let Some(expiration) = config.access.user_expirations.get(user)
&& chrono::Utc::now() > *expiration && 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); stats.decrement_user_curr_connects(user);
warn!( warn!(
user = %user, user = %user,
@@ -329,11 +393,12 @@ pub(crate) async fn acquire_user_connection_reservation(
}); });
} }
Ok(UserConnectionReservation::new( Ok(UserConnectionReservation::new_for_incarnation(
stats, stats,
ip_tracker, ip_tracker,
user.to_string(), user.to_string(),
peer_addr.ip(), peer_addr.ip(),
incarnation,
true, true,
)) ))
} }
+1
View File
@@ -73,6 +73,7 @@ pub mod route_mode;
pub mod session_eviction; pub mod session_eviction;
pub mod shared_state; pub mod shared_state;
pub mod traffic_limiter; pub mod traffic_limiter;
pub(crate) mod user_admission;
pub use client::ClientHandler; pub use client::ClientHandler;
#[allow(unused_imports)] #[allow(unused_imports)]
+126 -139
View File
@@ -6,14 +6,16 @@ use std::sync::{Arc, Mutex};
use std::time::Instant; use std::time::Instant;
use dashmap::DashMap; use dashmap::DashMap;
use parking_lot::Mutex as ParkingMutex;
use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc}; 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::direct_buffer_budget::{DirectBufferBudget, fallback_direct_buffer_hard_limit};
use crate::proxy::handshake::{AuthProbeSaturationState, AuthProbeState}; use crate::proxy::handshake::{AuthProbeSaturationState, AuthProbeState};
use crate::proxy::middle_relay::{DesyncDedupRotationState, RelayIdleCandidateRegistry}; use crate::proxy::middle_relay::{DesyncDedupRotationState, RelayIdleCandidateRegistry};
use crate::proxy::traffic_limiter::TrafficLimiter; 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 HANDSHAKE_RECENT_USER_RING_LEN: usize = 64;
const MASKING_FALLBACK_MAX_CONCURRENT: usize = 512; const MASKING_FALLBACK_MAX_CONCURRENT: usize = 512;
@@ -76,57 +78,17 @@ pub(crate) struct MiddleRelaySharedState {
pub(crate) relay_idle_mark_seq: AtomicU64, pub(crate) relay_idle_mark_seq: AtomicU64,
} }
#[derive(Default)]
struct UserAdmissionState {
disabled_users: HashSet<String>,
sessions_by_user: HashMap<String, HashMap<u64, CancellationToken>>,
}
pub(crate) struct ProxySharedState { pub(crate) struct ProxySharedState {
pub(crate) handshake: HandshakeSharedState, pub(crate) handshake: HandshakeSharedState,
pub(crate) middle_relay: MiddleRelaySharedState, pub(crate) middle_relay: MiddleRelaySharedState,
pub(crate) traffic_limiter: Arc<TrafficLimiter>, pub(crate) traffic_limiter: Arc<TrafficLimiter>,
pub(crate) direct_buffer_budget: Arc<DirectBufferBudget>, pub(crate) direct_buffer_budget: Arc<DirectBufferBudget>,
user_admission: ParkingMutex<UserAdmissionState>, user_admission: Arc<UserAdmissionAuthority>,
pub(crate) conntrack_pressure_active: AtomicBool, pub(crate) conntrack_pressure_active: AtomicBool,
pub(crate) conntrack_close_tx: Mutex<Option<mpsc::Sender<ConntrackCloseEvent>>>, pub(crate) conntrack_close_tx: Mutex<Option<mpsc::Sender<ConntrackCloseEvent>>>,
masking_fallback_permits: Arc<Semaphore>, masking_fallback_permits: Arc<Semaphore>,
} }
#[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<ProxySharedState>,
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 { impl ProxySharedState {
pub(crate) fn new() -> Arc<Self> { pub(crate) fn new() -> Arc<Self> {
Self::new_with_direct_buffer_budget(DirectBufferBudget::new( Self::new_with_direct_buffer_budget(DirectBufferBudget::new(
@@ -137,6 +99,17 @@ impl ProxySharedState {
/// Creates process state with the startup-resolved Direct buffer envelope. /// Creates process state with the startup-resolved Direct buffer envelope.
pub(crate) fn new_with_direct_buffer_budget( pub(crate) fn new_with_direct_buffer_budget(
direct_buffer_budget: Arc<DirectBufferBudget>, direct_buffer_budget: Arc<DirectBufferBudget>,
) -> Arc<Self> {
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<DirectBufferBudget>,
user_admission: Arc<UserAdmissionAuthority>,
) -> Arc<Self> { ) -> Arc<Self> {
Arc::new(Self { Arc::new(Self {
handshake: HandshakeSharedState { handshake: HandshakeSharedState {
@@ -167,7 +140,7 @@ impl ProxySharedState {
}, },
traffic_limiter: TrafficLimiter::new(), traffic_limiter: TrafficLimiter::new(),
direct_buffer_budget, direct_buffer_budget,
user_admission: ParkingMutex::new(UserAdmissionState::default()), user_admission,
conntrack_pressure_active: AtomicBool::new(false), conntrack_pressure_active: AtomicBool::new(false),
conntrack_close_tx: Mutex::new(None), conntrack_close_tx: Mutex::new(None),
masking_fallback_permits: Arc::new(Semaphore::new(MASKING_FALLBACK_MAX_CONCURRENT)), 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 { 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) { /// Returns the process authority shared by every runtime generation.
let (newly_disabled, tokens) = { pub(crate) fn user_admission(&self) -> Arc<UserAdmissionAuthority> {
let mut admission = self.user_admission.lock(); Arc::clone(&self.user_admission)
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())
} }
pub(crate) fn apply_user_enabled_config( /// Reconciles the complete user authentication policy from configuration.
pub(crate) fn apply_user_config(
&self, &self,
users: &HashMap<String, String>,
user_enabled: &HashMap<String, bool>, user_enabled: &HashMap<String, bool>,
) -> Vec<(String, usize)> { ) -> Vec<(String, usize)> {
let desired_disabled = user_enabled self.user_admission.apply_config(users, user_enabled)
.iter() }
.filter_map(|(user, enabled)| (!*enabled).then_some(user.clone()))
.collect::<HashSet<_>>(); /// Applies a candidate user policy only when its captured epoch is current.
let cancellations = { pub(crate) fn apply_user_config_if_epoch(
let mut admission = self.user_admission.lock(); &self,
let newly_disabled = desired_disabled expected_epoch: u64,
.difference(&admission.disabled_users) users: &HashMap<String, String>,
.cloned() user_enabled: &HashMap<String, bool>,
.collect::<Vec<_>>(); ) -> Option<Vec<(String, usize)>> {
admission.disabled_users = desired_disabled; self.user_admission
newly_disabled .apply_config_if_epoch(expected_epoch, users, user_enabled)
.into_iter() }
.map(|user| {
let tokens = admission /// Applies one persisted user mutation before asynchronous config reload.
.sessions_by_user pub(crate) fn stage_user(
.get(&user) &self,
.map(|sessions| sessions.values().cloned().collect()) user: &str,
.unwrap_or_default(); secret: &str,
(user, tokens) enabled: bool,
}) ) -> Option<UserMutationResult> {
.collect::<Vec<(String, Vec<CancellationToken>)>>() self.user_admission.stage_user(user, secret, enabled)
}; }
cancellations
.into_iter() /// Installs a deletion tombstone and cancels every current owner.
.map(|(user, tokens)| { pub(crate) fn delete_user(&self, user: &str) -> UserMutationResult {
for token in &tokens { self.user_admission.delete_user(user)
token.cancel(); }
}
(user, tokens.len()) /// Returns the current incarnation for an exact authenticated credential.
}) pub(crate) fn authenticated_user_incarnation(
.collect() &self,
user: &str,
credential_id: UserCredentialId,
) -> Option<UserIncarnation> {
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<Self>,
user: &str,
credential_id: UserCredentialId,
) -> Option<UserAdmissionPublication<'_>> {
self.user_admission
.claim_authenticated(user, credential_id)
} }
pub(crate) fn register_user_session( pub(crate) fn register_user_session(
self: &Arc<Self>, self: &Arc<Self>,
user: &str, user: &str,
session_id: u64, _session_id: u64,
) -> Option<UserSessionRegistration> { ) -> Option<UserSessionRegistration> {
let token = CancellationToken::new(); self.user_admission.register_legacy(user)
let key = (user.to_string(), session_id); }
let mut admission = self.user_admission.lock();
if admission.disabled_users.contains(user) { /// Registers a relay session against the exact credential that authenticated it.
return None; pub(crate) fn register_authenticated_user_session(
} self: &Arc<Self>,
admission user: &str,
.sessions_by_user credential_id: UserCredentialId,
.entry(key.0.clone()) ) -> Option<UserSessionRegistration> {
.or_default() let mut publication = self.claim_authenticated_user(user, credential_id)?;
.insert(session_id, token.clone()); let registration = publication.take_registration()?;
Some(UserSessionRegistration { publication.commit();
token, Some(registration)
_guard: UserSessionGuard {
shared: Arc::clone(self),
key,
},
})
} }
pub(crate) fn cancel_user_sessions(&self, user: &str) -> usize { pub(crate) fn cancel_user_sessions(&self, user: &str) -> usize {
let tokens: Vec<CancellationToken> = self self.user_admission.cancel_user_owners(user)
.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()
} }
pub(crate) fn set_conntrack_close_sender(&self, tx: mpsc::Sender<ConntrackCloseEvent>) { pub(crate) fn set_conntrack_close_sender(&self, tx: mpsc::Sender<ConntrackCloseEvent>) {
@@ -350,31 +308,53 @@ impl ProxySharedState {
mod tests { mod tests {
use super::*; use super::*;
const ALICE_SECRET: &str = "00112233445566778899aabbccddeeff";
fn configured_shared() -> Arc<ProxySharedState> {
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] #[test]
fn user_enabled_config_sync_tracks_disabled_overrides() { fn user_enabled_config_sync_tracks_disabled_overrides() {
let shared = ProxySharedState::new(); let shared = configured_shared();
assert!(shared.is_user_enabled("alice")); 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(); let mut user_enabled = HashMap::new();
user_enabled.insert("alice".to_string(), false); user_enabled.insert("alice".to_string(), false);
user_enabled.insert("bob".to_string(), true); 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(); newly_disabled.sort();
assert_eq!(newly_disabled, vec![("alice".to_string(), 0)]); assert_eq!(newly_disabled, vec![("alice".to_string(), 0)]);
assert!(!shared.is_user_enabled("alice")); assert!(!shared.is_user_enabled("alice"));
assert!(shared.is_user_enabled("bob")); 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(); 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")); assert!(shared.is_user_enabled("alice"));
} }
#[test] #[test]
fn cancel_user_sessions_cancels_only_registered_matching_user() { 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_1 = shared.register_user_session("alice", 1).unwrap();
let alice_2 = shared.register_user_session("alice", 2).unwrap(); let alice_2 = shared.register_user_session("alice", 2).unwrap();
let bob = shared.register_user_session("bob", 1).unwrap(); let bob = shared.register_user_session("bob", 1).unwrap();
@@ -392,9 +372,11 @@ mod tests {
#[test] #[test]
fn disabled_user_cannot_register_after_the_cancellation_snapshot() { 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); assert_eq!(shared.cancel_user_sessions("alice"), 0);
let late = shared.register_user_session("alice", 1); let late = shared.register_user_session("alice", 1);
@@ -406,15 +388,19 @@ mod tests {
#[test] #[test]
fn disabling_user_cancels_existing_sessions_before_return() { 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 registration = shared.register_user_session("alice", 1).unwrap();
let token = registration.token(); 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!(token.is_cancelled());
assert!(shared.register_user_session("alice", 2).is_none()); 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()); assert!(shared.register_user_session("alice", 3).is_some());
} }
@@ -439,7 +425,8 @@ mod tests {
for session_id in 0..ITERATIONS as u64 { for session_id in 0..ITERATIONS as u64 {
let user = format!("user-{session_id}"); let user = format!("user-{session_id}");
barrier.wait(); 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() { for registration in register.join().unwrap().into_iter().flatten() {
+598
View File
@@ -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<EffectiveUser>,
mutation_override: Option<UserOverride>,
incarnation: UserIncarnation,
}
impl UserRecord {
fn effective(&self) -> Option<EffectiveUser> {
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<String, UserRecord>,
owners_by_user: HashMap<String, HashMap<u64, RegisteredOwner>>,
}
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<u64> {
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<CancellationToken> {
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<UserAdmissionState>,
}
impl UserAdmissionAuthority {
/// Creates an uninitialized authority for isolated tests and startup wiring.
pub(crate) fn new() -> Arc<Self> {
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<String, String>,
user_enabled: &HashMap<String, bool>,
) -> 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<String, String>,
user_enabled: &HashMap<String, bool>,
) -> Option<Vec<(String, usize)>> {
self.apply_config_locked(Some(expected_epoch), users, user_enabled)
}
fn apply_config_locked(
&self,
expected_epoch: Option<u64>,
users: &HashMap<String, String>,
user_enabled: &HashMap<String, bool>,
) -> Option<Vec<(String, usize)>> {
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::<HashMap<_, _>>();
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::<Vec<_>>();
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<UserMutationResult> {
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<UserIncarnation> {
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<Self>,
user: &str,
credential_id: UserCredentialId,
) -> Option<UserAdmissionPublication<'_>> {
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<Self>,
user: &str,
) -> Option<UserSessionRegistration> {
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(&registration_id)
.is_some_and(|owner| owner.incarnation == incarnation)
{
owners.remove(&registration_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<UserAdmissionAuthority>,
user: String,
registration_id: u64,
incarnation: UserIncarnation,
token: CancellationToken,
active: Arc<AtomicBool>,
registration_taken: bool,
}
impl UserAdmissionPublication<'_> {
/// Moves the registered owner out while retaining the authority lock.
pub(crate) fn take_registration(&mut self) -> Option<UserSessionRegistration> {
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<UserAdmissionAuthority>,
user: String,
registration_id: u64,
incarnation: UserIncarnation,
token: CancellationToken,
active: Arc<AtomicBool>,
}
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<UserCredentialId> {
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<CancellationToken>)>,
) -> 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<String, String> {
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"));
}
}
+4 -3
View File
@@ -25,9 +25,10 @@ pub(super) async fn check_family(
let mut family_degraded = false; let mut family_degraded = false;
let mut dc_endpoints = HashMap::<i32, Vec<SocketAddr>>::new(); let mut dc_endpoints = HashMap::<i32, Vec<SocketAddr>>::new();
let endpoint_snapshot = pool.endpoint_snapshot.load();
let map_guard = match family { let map_guard = match family {
IpFamily::V4 => pool.proxy_map_v4.read().await, IpFamily::V4 => &endpoint_snapshot.map_v4,
IpFamily::V6 => pool.proxy_map_v6.read().await, IpFamily::V6 => &endpoint_snapshot.map_v6,
}; };
for (dc, addrs) in map_guard.iter() { for (dc, addrs) in map_guard.iter() {
let entry = dc_endpoints.entry(*dc).or_default(); let entry = dc_endpoints.entry(*dc).or_default();
@@ -35,7 +36,7 @@ pub(super) async fn check_family(
entry.push(SocketAddr::new(ip, port)); entry.push(SocketAddr::new(ip, port));
} }
} }
drop(map_guard); drop(endpoint_snapshot);
for endpoints in dc_endpoints.values_mut() { for endpoints in dc_endpoints.values_mut() {
endpoints.sort_unstable(); endpoints.sort_unstable();
endpoints.dedup(); endpoints.dedup();
+3 -2
View File
@@ -329,13 +329,14 @@ mod tests {
pub async fn run_me_ping(pool: &Arc<MePool>, rng: &SecureRandom) -> Vec<MePingReport> { pub async fn run_me_ping(pool: &Arc<MePool>, rng: &SecureRandom) -> Vec<MePingReport> {
let mut reports = Vec::new(); let mut reports = Vec::new();
let endpoint_snapshot = pool.endpoint_snapshot.load_full();
let v4_map = if pool.decision.ipv4_me { let v4_map = if pool.decision.ipv4_me {
pool.proxy_map_v4.read().await.clone() endpoint_snapshot.map_v4.clone()
} else { } else {
HashMap::new() HashMap::new()
}; };
let v6_map = if pool.decision.ipv6_me { let v6_map = if pool.decision.ipv6_me {
pool.proxy_map_v6.read().await.clone() endpoint_snapshot.map_v6.clone()
} else { } else {
HashMap::new() HashMap::new()
}; };
+14 -4
View File
@@ -266,7 +266,17 @@ pub struct RoutingCore {
pub(super) writers: Arc<WritersState>, pub(super) writers: Arc<WritersState>,
pub(super) rr: AtomicU64, pub(super) rr: AtomicU64,
pub(super) writer_epoch: watch::Sender<u64>, pub(super) writer_epoch: watch::Sender<u64>,
pub(super) preferred_endpoints_by_dc: ArcSwap<HashMap<i32, Vec<SocketAddr>>>, pub(super) endpoint_snapshot: ArcSwap<EndpointSnapshot>,
}
/// 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<i32, Vec<(IpAddr, u16)>>,
pub(super) map_v6: HashMap<i32, Vec<(IpAddr, u16)>>,
pub(super) endpoint_dc_map: HashMap<SocketAddr, Option<i32>>,
pub(super) preferred_endpoints_by_dc: HashMap<i32, Vec<SocketAddr>>,
} }
pub(super) struct ReinitCore { pub(super) struct ReinitCore {
@@ -302,12 +312,14 @@ pub(super) struct ReinitPendingState {
pub(super) generation: u64, pub(super) generation: u64,
pub(super) started_at_epoch_secs: u64, pub(super) started_at_epoch_secs: u64,
pub(super) map_hash: u64, pub(super) map_hash: u64,
pub(super) endpoint_revision: u64,
} }
#[derive(Clone, Copy)] #[derive(Clone, Copy)]
pub(super) struct ReinitAttemptState { pub(super) struct ReinitAttemptState {
pub(super) generation: u64, pub(super) generation: u64,
pub(super) map_hash: u64, pub(super) map_hash: u64,
pub(super) endpoint_revision: u64,
pub(super) hardswap: bool, pub(super) hardswap: bool,
pub(super) committed: bool, pub(super) committed: bool,
} }
@@ -316,6 +328,7 @@ pub(super) struct ReinitCoordinatorState {
pub(super) next_attempt_id: u64, pub(super) next_attempt_id: u64,
pub(super) active_generation: u64, pub(super) active_generation: u64,
pub(super) desired_map_hash: u64, pub(super) desired_map_hash: u64,
pub(super) endpoint_revision: u64,
pub(super) pending: Option<ReinitPendingState>, pub(super) pending: Option<ReinitPendingState>,
pub(super) attempts: HashMap<u64, ReinitAttemptState>, pub(super) attempts: HashMap<u64, ReinitAttemptState>,
} }
@@ -475,9 +488,6 @@ pub struct MePool {
pub(super) rng: Arc<SecureRandom>, pub(super) rng: Arc<SecureRandom>,
pub(super) proxy_tag: Option<Vec<u8>>, pub(super) proxy_tag: Option<Vec<u8>>,
pub(super) proxy_secret: Arc<RwLock<SecretSnapshot>>, pub(super) proxy_secret: Arc<RwLock<SecretSnapshot>>,
pub(super) proxy_map_v4: Arc<RwLock<HashMap<i32, Vec<(IpAddr, u16)>>>>,
pub(super) proxy_map_v6: Arc<RwLock<HashMap<i32, Vec<(IpAddr, u16)>>>>,
pub(super) endpoint_dc_map: Arc<RwLock<HashMap<SocketAddr, Option<i32>>>>,
pub(super) default_dc: AtomicI32, pub(super) default_dc: AtomicI32,
pub(super) next_writer_id: AtomicU64, pub(super) next_writer_id: AtomicU64,
pub(super) writer_connect_active_reserved: AtomicUsize, pub(super) writer_connect_active_reserved: AtomicUsize,
@@ -120,9 +120,12 @@ impl MePool {
me_route_inline_recovery_wait_ms: u64, me_route_inline_recovery_wait_ms: u64,
me_connection_cleanup_capacity: usize, me_connection_cleanup_capacity: usize,
) -> Arc<Self> { ) -> Arc<Self> {
let endpoint_dc_map = Self::build_endpoint_dc_map_from_maps(&proxy_map_v4, &proxy_map_v6); let endpoint_snapshot = Self::build_endpoint_snapshot(
let preferred_endpoints_by_dc = &decision,
Self::build_preferred_endpoints_by_dc(&decision, &proxy_map_v4, &proxy_map_v6); proxy_map_v4,
proxy_map_v6,
1,
);
let registry = Arc::new(ConnRegistry::with_route_and_cleanup_capacity( let registry = Arc::new(ConnRegistry::with_route_and_cleanup_capacity(
me_route_channel_capacity, me_route_channel_capacity,
me_connection_cleanup_capacity, me_connection_cleanup_capacity,
@@ -149,7 +152,7 @@ impl MePool {
writers: Arc::new(WritersState::new()), writers: Arc::new(WritersState::new()),
rr: AtomicU64::new(0), rr: AtomicU64::new(0),
writer_epoch, writer_epoch,
preferred_endpoints_by_dc: ArcSwap::from_pointee(preferred_endpoints_by_dc), endpoint_snapshot: ArcSwap::from_pointee(endpoint_snapshot),
}), }),
reinit: Arc::new(ReinitCore { reinit: Arc::new(ReinitCore {
generation: AtomicU64::new(1), generation: AtomicU64::new(1),
@@ -164,6 +167,7 @@ impl MePool {
next_attempt_id: 1, next_attempt_id: 1,
active_generation: 1, active_generation: 1,
desired_map_hash: 0, desired_map_hash: 0,
endpoint_revision: 1,
pending: None, pending: None,
attempts: HashMap::new(), attempts: HashMap::new(),
}), }),
@@ -391,9 +395,6 @@ impl MePool {
})), })),
stats, stats,
pool_size: 2, 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)), default_dc: AtomicI32::new(default_dc.unwrap_or(2)),
next_writer_id: AtomicU64::new(1), next_writer_id: AtomicU64::new(1),
writer_connect_active_reserved: AtomicUsize::new(0), writer_connect_active_reserved: AtomicUsize::new(0),
+64 -18
View File
@@ -76,16 +76,23 @@ impl MePool {
&self, &self,
dc: i32, dc: i32,
) -> bool { ) -> bool {
let snapshot = self.endpoint_snapshot.load();
if self.decision.ipv4_me { if self.decision.ipv4_me {
let map = self.proxy_map_v4.read().await; if snapshot
if map.get(&dc).is_some_and(|endpoints| !endpoints.is_empty()) { .map_v4
.get(&dc)
.is_some_and(|endpoints| !endpoints.is_empty())
{
return true; return true;
} }
} }
if self.decision.ipv6_me { if self.decision.ipv6_me {
let map = self.proxy_map_v6.read().await; if snapshot
if map.get(&dc).is_some_and(|endpoints| !endpoints.is_empty()) { .map_v6
.get(&dc)
.is_some_and(|endpoints| !endpoints.is_empty())
{
return true; return true;
} }
} }
@@ -112,7 +119,12 @@ impl MePool {
&self, &self,
addr: SocketAddr, addr: SocketAddr,
) -> i32 { ) -> 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 && let Some(dc) = cached
{ {
return dc; return dc;
@@ -125,9 +137,45 @@ impl MePool {
&self, &self,
family: IpFamily, family: IpFamily,
) -> HashMap<i32, Vec<(IpAddr, u16)>> { ) -> HashMap<i32, Vec<(IpAddr, u16)>> {
let snapshot = self.endpoint_snapshot.load();
match family { match family {
IpFamily::V4 => self.proxy_map_v4.read().await.clone(), IpFamily::V4 => snapshot.map_v4.clone(),
IpFamily::V6 => self.proxy_map_v6.read().await.clone(), IpFamily::V6 => snapshot.map_v6.clone(),
}
}
pub(in crate::transport::middle_proxy) fn build_endpoint_snapshot(
decision: &NetworkDecision,
mut map_v4: HashMap<i32, Vec<(IpAddr, u16)>>,
mut map_v6: HashMap<i32, Vec<(IpAddr, u16)>>,
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<i32, Vec<(IpAddr, u16)>>) {
let positive_dcs = map
.keys()
.copied()
.filter(|dc| *dc > 0)
.collect::<Vec<_>>();
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 endpoint_dc_map
} }
pub(in crate::transport::middle_proxy) async fn rebuild_endpoint_dc_map(&self) { pub(in crate::transport::middle_proxy) async fn prune_endpoint_runtime_state(&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));
let configured_endpoints = self let configured_endpoints = self
.endpoint_snapshot
.load()
.endpoint_dc_map .endpoint_dc_map
.read()
.await
.keys() .keys()
.copied() .copied()
.collect::<HashSet<SocketAddr>>(); .collect::<HashSet<SocketAddr>>();
@@ -253,8 +295,12 @@ impl MePool {
&self, &self,
dc: i32, dc: i32,
) -> Vec<SocketAddr> { ) -> Vec<SocketAddr> {
let guard = self.preferred_endpoints_by_dc.load(); self.endpoint_snapshot
guard.get(&dc).cloned().unwrap_or_default() .load()
.preferred_endpoints_by_dc
.get(&dc)
.cloned()
.unwrap_or_default()
} }
pub(in crate::transport::middle_proxy) fn health_interval_unhealthy(&self) -> Duration { pub(in crate::transport::middle_proxy) fn health_interval_unhealthy(&self) -> Duration {
@@ -65,10 +65,10 @@ impl MePool {
pub(in crate::transport::middle_proxy) async fn active_coverage_required_total(&self) -> usize { pub(in crate::transport::middle_proxy) async fn active_coverage_required_total(&self) -> usize {
let now_epoch_secs = Self::now_epoch_secs(); let now_epoch_secs = Self::now_epoch_secs();
let mut required_total = 0usize; let mut required_total = 0usize;
let endpoint_snapshot = self.endpoint_snapshot.load();
if self.family_enabled_for_drain_coverage(IpFamily::V4, now_epoch_secs) { if self.family_enabled_for_drain_coverage(IpFamily::V4, now_epoch_secs) {
let map = self.proxy_map_v4.read().await; for addrs in endpoint_snapshot.map_v4.values() {
for addrs in map.values() {
let mut endpoints = HashSet::<SocketAddr>::new(); let mut endpoints = HashSet::<SocketAddr>::new();
for (ip, port) in addrs.iter().copied() { for (ip, port) in addrs.iter().copied() {
endpoints.insert(SocketAddr::new(ip, port)); endpoints.insert(SocketAddr::new(ip, port));
@@ -80,8 +80,7 @@ impl MePool {
} }
if self.family_enabled_for_drain_coverage(IpFamily::V6, now_epoch_secs) { if self.family_enabled_for_drain_coverage(IpFamily::V6, now_epoch_secs) {
let map = self.proxy_map_v6.read().await; for addrs in endpoint_snapshot.map_v6.values() {
for addrs in map.values() {
let mut endpoints = HashSet::<SocketAddr>::new(); let mut endpoints = HashSet::<SocketAddr>::new();
for (ip, port) in addrs.iter().copied() { for (ip, port) in addrs.iter().copied() {
endpoints.insert(SocketAddr::new(ip, port)); endpoints.insert(SocketAddr::new(ip, port));
@@ -118,13 +117,14 @@ impl MePool {
let mut endpoints_len = 0; let mut endpoints_len = 0;
let now_epoch = Self::now_epoch_secs(); 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 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(); endpoints_len += addrs.len();
} }
} }
if self.family_enabled_for_drain_coverage(IpFamily::V6, now_epoch) { 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(); endpoints_len += addrs.len();
} }
} }
+27 -40
View File
@@ -31,48 +31,35 @@ impl MePool {
return SnapshotApplyOutcome::RejectedEmpty; return SnapshotApplyOutcome::RejectedEmpty;
} }
let mut changed = false; let changed = {
{ // Endpoint publication and reinit commit share this barrier.
let mut guard = self.proxy_map_v4.write().await; let mut coordinator = self.reinit.coordinator.lock();
if !new_v4.is_empty() && *guard != new_v4 { let current = self.endpoint_snapshot.load_full();
*guard = new_v4; let map_v4 = if new_v4.is_empty() {
changed = true; 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<i32> = 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<i32> = 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 { if changed {
self.rebuild_endpoint_dc_map().await; self.prune_endpoint_runtime_state().await;
self.notify_writer_epoch(); self.notify_writer_epoch();
} }
if changed { if changed {
+1 -1
View File
@@ -19,7 +19,7 @@ impl MePool {
.me_reconnect_max_concurrent_per_dc .me_reconnect_max_concurrent_per_dc
.max(1) as usize; .max(1) as usize;
let ks = self.key_selector().await; 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(); let secret_len = self.proxy_secret.read().await.secret.len();
info!( info!(
me_servers, me_servers,
+7 -5
View File
@@ -70,9 +70,9 @@ impl Drop for RefillRunGuard {
impl MePool { impl MePool {
pub(super) async fn sweep_endpoint_quarantine(&self) { pub(super) async fn sweep_endpoint_quarantine(&self) {
let configured = self let configured = self
.endpoint_snapshot
.load()
.endpoint_dc_map .endpoint_dc_map
.read()
.await
.keys() .keys()
.copied() .copied()
.collect::<HashSet<SocketAddr>>(); .collect::<HashSet<SocketAddr>>();
@@ -266,9 +266,10 @@ impl MePool {
if !self.family_enabled_for_drain_coverage(target.family, now_epoch_secs) { if !self.family_enabled_for_drain_coverage(target.family, now_epoch_secs) {
return Vec::new(); return Vec::new();
} }
let snapshot = self.endpoint_snapshot.load();
let map = match target.family { let map = match target.family {
IpFamily::V4 => self.proxy_map_v4.read().await, IpFamily::V4 => &snapshot.map_v4,
IpFamily::V6 => self.proxy_map_v6.read().await, IpFamily::V6 => &snapshot.map_v6,
}; };
let mut endpoints = map let mut endpoints = map
.get(&target.dc) .get(&target.dc)
@@ -294,8 +295,9 @@ impl MePool {
}; };
role_is_authoritative role_is_authoritative
&& self && self
.preferred_endpoints_by_dc .endpoint_snapshot
.load() .load()
.preferred_endpoints_by_dc
.get(&target.dc) .get(&target.dc)
.is_some_and(|endpoints| { .is_some_and(|endpoints| {
endpoints.iter().any(|endpoint| match target.family { endpoints.iter().any(|endpoint| match target.family {
+12 -4
View File
@@ -15,8 +15,8 @@ use crate::config::MeBindStaleMode;
use crate::network::IpFamily; use crate::network::IpFamily;
use super::pool::{ use super::pool::{
MeDrainGateReason, MePool, ReinitAttemptState, ReinitCoordinatorState, ReinitCore, EndpointSnapshot, MeDrainGateReason, MePool, ReinitAttemptState, ReinitCoordinatorState,
ReinitPendingState, ReinitStatusSnapshot, WriterContour, WriterOpenIntent, ReinitCore, ReinitPendingState, ReinitStatusSnapshot, WriterContour, WriterOpenIntent,
}; };
// Reinitialization admission, generation state, and coverage checks. // Reinitialization admission, generation state, and coverage checks.
@@ -34,6 +34,7 @@ struct ReinitAttemptGuard {
generation: u64, generation: u64,
previous_generation: u64, previous_generation: u64,
map_hash: u64, map_hash: u64,
endpoint_revision: u64,
hardswap: bool, hardswap: bool,
} }
@@ -119,17 +120,24 @@ fn commit_reinit_state(
attempt_id: u64, attempt_id: u64,
generation: u64, generation: u64,
map_hash: u64, map_hash: u64,
endpoint_revision: u64,
hardswap: bool, hardswap: bool,
) -> bool { ) -> bool {
let Some(record) = state.attempts.get(&attempt_id).copied() else { let Some(record) = state.attempts.get(&attempt_id).copied() else {
return false; 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; return false;
} }
if hardswap { if hardswap {
let pending_matches = state.pending.is_some_and(|pending| { 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 { if !pending_matches || generation < state.active_generation {
return false; return false;
@@ -27,9 +27,13 @@ impl MePool {
self: &Arc<Self>, self: &Arc<Self>,
hardswap: bool, hardswap: bool,
map_hash: u64, map_hash: u64,
endpoint_revision: u64,
now_epoch_secs: u64, now_epoch_secs: u64,
) -> ReinitReservation { ) -> Option<ReinitReservation> {
let mut state = self.reinit.coordinator.lock(); let mut state = self.reinit.coordinator.lock();
if state.endpoint_revision != endpoint_revision {
return None;
}
state.desired_map_hash = map_hash; state.desired_map_hash = map_hash;
let previous_generation = state.active_generation; let previous_generation = state.active_generation;
let mut pending_reused = false; let mut pending_reused = false;
@@ -43,6 +47,7 @@ impl MePool {
&& pending_age_secs > ME_HARDSWAP_PENDING_TTL_SECS; && pending_age_secs > ME_HARDSWAP_PENDING_TTL_SECS;
pending.generation >= previous_generation pending.generation >= previous_generation
&& pending.map_hash == map_hash && pending.map_hash == map_hash
&& pending.endpoint_revision == endpoint_revision
&& !pending_expired && !pending_expired
}); });
if let Some(pending) = reusable { if let Some(pending) = reusable {
@@ -54,6 +59,7 @@ impl MePool {
generation, generation,
started_at_epoch_secs: now_epoch_secs, started_at_epoch_secs: now_epoch_secs,
map_hash, map_hash,
endpoint_revision,
}); });
generation generation
} }
@@ -69,24 +75,26 @@ impl MePool {
ReinitAttemptState { ReinitAttemptState {
generation, generation,
map_hash, map_hash,
endpoint_revision,
hardswap, hardswap,
committed: false, committed: false,
}, },
); );
publish_reinit_state(self.reinit.as_ref(), &state); publish_reinit_state(self.reinit.as_ref(), &state);
ReinitReservation { Some(ReinitReservation {
attempt: ReinitAttemptGuard { attempt: ReinitAttemptGuard {
reinit: Arc::clone(&self.reinit), reinit: Arc::clone(&self.reinit),
attempt_id, attempt_id,
generation, generation,
previous_generation, previous_generation,
map_hash, map_hash,
endpoint_revision,
hardswap, hardswap,
}, },
pending_reused, pending_reused,
pending_expired, pending_expired,
pending_age_secs, pending_age_secs,
} })
} }
/// Revalidates coverage and commits generation ownership under the publication barrier. /// Revalidates coverage and commits generation ownership under the publication barrier.
@@ -105,10 +113,13 @@ impl MePool {
if record.generation != attempt.generation if record.generation != attempt.generation
|| record.map_hash != state.desired_map_hash || record.map_hash != state.desired_map_hash
|| record.map_hash != attempt.map_hash || record.map_hash != attempt.map_hash
|| record.endpoint_revision != attempt.endpoint_revision
|| record.endpoint_revision != state.endpoint_revision
|| (attempt.hardswap || (attempt.hardswap
&& !state.pending.is_some_and(|pending| { && !state.pending.is_some_and(|pending| {
pending.generation == attempt.generation pending.generation == attempt.generation
&& pending.map_hash == attempt.map_hash && pending.map_hash == attempt.map_hash
&& pending.endpoint_revision == attempt.endpoint_revision
})) }))
{ {
return Err(ReinitCommitFailure::Superseded); return Err(ReinitCommitFailure::Superseded);
@@ -150,6 +161,7 @@ impl MePool {
attempt.attempt_id, attempt.attempt_id,
attempt.generation, attempt.generation,
attempt.map_hash, attempt.map_hash,
attempt.endpoint_revision,
attempt.hardswap, attempt.hardswap,
) { ) {
return Err(ReinitCommitFailure::Superseded); return Err(ReinitCommitFailure::Superseded);
@@ -255,9 +267,13 @@ impl MePool {
/// Restores at least one active writer for every enabled desired DC group. /// Restores at least one active writer for every enabled desired DC group.
pub async fn reconcile_connections(self: &Arc<Self>, rng: &SecureRandom) { pub async fn reconcile_connections(self: &Arc<Self>, rng: &SecureRandom) {
let endpoint_snapshot = self.endpoint_snapshot.load_full();
for family in self.family_order() { for family in self.family_order() {
let map = self.proxy_map_for_family(family).await; let map = match family {
for (dc, addrs) in &map { IpFamily::V4 => &endpoint_snapshot.map_v4,
IpFamily::V6 => &endpoint_snapshot.map_v6,
};
for (dc, addrs) in map {
let dc_addrs: Vec<SocketAddr> = addrs let dc_addrs: Vec<SocketAddr> = addrs
.iter() .iter()
.map(|(ip, port)| SocketAddr::new(*ip, *port)) .map(|(ip, port)| SocketAddr::new(*ip, *port))
@@ -286,26 +302,32 @@ impl MePool {
/// Returns the currently authoritative endpoint set for drain and coverage decisions. /// Returns the currently authoritative endpoint set for drain and coverage decisions.
pub(in crate::transport::middle_proxy) async fn desired_dc_endpoints( pub(in crate::transport::middle_proxy) async fn desired_dc_endpoints(
&self, &self,
) -> HashMap<i32, HashSet<SocketAddr>> {
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<i32, HashSet<SocketAddr>> { ) -> HashMap<i32, HashSet<SocketAddr>> {
let now_epoch_secs = Self::now_epoch_secs(); let now_epoch_secs = Self::now_epoch_secs();
let mut out: HashMap<i32, HashSet<SocketAddr>> = HashMap::new(); let mut out: HashMap<i32, HashSet<SocketAddr>> = HashMap::new();
if self.family_enabled_for_drain_coverage(IpFamily::V4, now_epoch_secs) { 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 &endpoint_snapshot.map_v4 {
for (dc, addrs) in map_v4 { let entry = out.entry(*dc).or_default();
let entry = out.entry(dc).or_default();
for (ip, port) in addrs { 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) { 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 &endpoint_snapshot.map_v6 {
for (dc, addrs) in map_v6 { let entry = out.entry(*dc).or_default();
let entry = out.entry(dc).or_default();
for (ip, port) in addrs { 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 state = self.reinit.coordinator.lock();
let active_generation = state.active_generation; let active_generation = state.active_generation;
let pending_generation = state.pending.map(|pending| pending.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 now_epoch_secs = Self::now_epoch_secs();
let mut changed = 0usize; let mut changed = 0usize;
@@ -362,7 +385,7 @@ impl MePool {
self.apply_writer_draining_state(writer, self.force_close_timeout(), false); self.apply_writer_draining_state(writer, self.force_close_timeout(), false);
changed = changed.saturating_add(1); changed = changed.saturating_add(1);
} }
drop(preferred); drop(endpoint_snapshot);
drop(state); drop(state);
drop(registry_registration); drop(registry_registration);
drop(writers); drop(writers);
@@ -118,7 +118,8 @@ impl MePool {
self: &Arc<Self>, self: &Arc<Self>,
rng: &SecureRandom, rng: &SecureRandom,
) -> bool { ) -> 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 now_epoch_secs = Self::now_epoch_secs();
let v4_suppressed = self.is_family_temporarily_suppressed(IpFamily::V4, 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); 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 desired_map_hash = Self::desired_map_hash(&desired_by_dc);
let hardswap = self.reinit.hardswap.load(Ordering::Relaxed); 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 attempt = reservation.attempt;
let previous_generation = attempt.previous_generation; let previous_generation = attempt.previous_generation;
let generation = attempt.generation; let generation = attempt.generation;
@@ -115,10 +115,12 @@ fn stale_concurrent_attempt_cannot_regress_active_generation() {
next_attempt_id: 3, next_attempt_id: 3,
active_generation: 1, active_generation: 1,
desired_map_hash: 22, desired_map_hash: 22,
endpoint_revision: 7,
pending: Some(ReinitPendingState { pending: Some(ReinitPendingState {
generation: 3, generation: 3,
started_at_epoch_secs: 1, started_at_epoch_secs: 1,
map_hash: 22, map_hash: 22,
endpoint_revision: 7,
}), }),
attempts: HashMap::from([ attempts: HashMap::from([
( (
@@ -126,6 +128,7 @@ fn stale_concurrent_attempt_cannot_regress_active_generation() {
ReinitAttemptState { ReinitAttemptState {
generation: 2, generation: 2,
map_hash: 11, map_hash: 11,
endpoint_revision: 6,
hardswap: true, hardswap: true,
committed: false, committed: false,
}, },
@@ -135,6 +138,7 @@ fn stale_concurrent_attempt_cannot_regress_active_generation() {
ReinitAttemptState { ReinitAttemptState {
generation: 3, generation: 3,
map_hash: 22, map_hash: 22,
endpoint_revision: 7,
hardswap: true, hardswap: true,
committed: false, 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_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_eq!(state.active_generation, 3);
assert!(state.pending.is_none()); assert!(state.pending.is_none());
} }
@@ -173,7 +177,10 @@ async fn partial_hardswap_is_rejected_when_stale_binding_is_disabled() {
) )
.await; .await;
let map_hash = MePool::desired_map_hash(&desired_by_dc); 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( insert_writer(
&pool, &pool,
201, 201,
@@ -221,7 +228,10 @@ async fn partial_hardswap_preserves_fallback_only_for_missing_dc() {
) )
.await; .await;
let map_hash = MePool::desired_map_hash(&desired_by_dc); 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( let fresh_dc1 = insert_writer(
&pool, &pool,
401, 401,
@@ -274,7 +284,10 @@ async fn complete_hardswap_promotes_fresh_generation_and_retires_old_writers() {
) )
.await; .await;
let map_hash = MePool::desired_map_hash(&desired_by_dc); 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( let fresh_dc1 = insert_writer(
&pool, &pool,
601, 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] #[tokio::test]
async fn generation_role_reconciliation_promotes_active_warm_and_drains_orphans() { async fn generation_role_reconciliation_promotes_active_warm_and_drains_orphans() {
let pool = make_pool().await; let pool = make_pool().await;
let endpoint = addr(1, 2001); let endpoint = addr(1, 2001);
pool.preferred_endpoints_by_dc pool.update_proxy_maps(
.store(Arc::new(HashMap::from([(1, vec![endpoint])]))); HashMap::from([(1, vec![(endpoint.ip(), endpoint.port())])]),
None,
)
.await;
let desired_by_dc = HashMap::from([(1, HashSet::from([endpoint]))]); let desired_by_dc = HashMap::from([(1, HashSet::from([endpoint]))]);
let map_hash = MePool::desired_map_hash(&desired_by_dc); 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( let active_warm = insert_writer(
&pool, &pool,
701, 701,
@@ -3,12 +3,13 @@ use super::*;
impl MePool { impl MePool {
pub(crate) async fn admission_ready_conditional_cast(&self) -> bool { pub(crate) async fn admission_ready_conditional_cast(&self) -> bool {
let mut endpoints_by_dc = BTreeMap::<i16, BTreeSet<SocketAddr>>::new(); let mut endpoints_by_dc = BTreeMap::<i16, BTreeSet<SocketAddr>>::new();
let endpoint_snapshot = self.endpoint_snapshot.load_full();
if self.decision.ipv4_me { 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); extend_signed_endpoints(&mut endpoints_by_dc, map);
} }
if self.decision.ipv6_me { 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); extend_signed_endpoints(&mut endpoints_by_dc, map);
} }
@@ -45,12 +46,13 @@ impl MePool {
#[allow(dead_code)] #[allow(dead_code)]
pub(crate) async fn admission_ready_full_floor(&self) -> bool { pub(crate) async fn admission_ready_full_floor(&self) -> bool {
let mut endpoints_by_dc = BTreeMap::<i16, BTreeSet<SocketAddr>>::new(); let mut endpoints_by_dc = BTreeMap::<i16, BTreeSet<SocketAddr>>::new();
let endpoint_snapshot = self.endpoint_snapshot.load_full();
if self.decision.ipv4_me { 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); extend_signed_endpoints(&mut endpoints_by_dc, map);
} }
if self.decision.ipv6_me { 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); extend_signed_endpoints(&mut endpoints_by_dc, map);
} }
@@ -106,12 +108,13 @@ impl MePool {
.load(Ordering::Relaxed); .load(Ordering::Relaxed);
let mut endpoints_by_dc = BTreeMap::<i16, BTreeSet<SocketAddr>>::new(); let mut endpoints_by_dc = BTreeMap::<i16, BTreeSet<SocketAddr>>::new();
let endpoint_snapshot = self.endpoint_snapshot.load_full();
if self.decision.ipv4_me { 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); extend_signed_endpoints(&mut endpoints_by_dc, map);
} }
if self.decision.ipv6_me { 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); extend_signed_endpoints(&mut endpoints_by_dc, map);
} }
@@ -41,8 +41,9 @@ impl MePool {
coordinator: &crate::transport::middle_proxy::pool::ReinitCoordinatorState, coordinator: &crate::transport::middle_proxy::pool::ReinitCoordinatorState,
) -> Result<WriterContour> { ) -> Result<WriterContour> {
let endpoint_is_current = self let endpoint_is_current = self
.preferred_endpoints_by_dc .endpoint_snapshot
.load() .load()
.preferred_endpoints_by_dc
.get(&writer.writer_dc) .get(&writer.writer_dc)
.is_some_and(|endpoints| endpoints.contains(&writer.addr)); .is_some_and(|endpoints| endpoints.contains(&writer.addr));
if !endpoint_is_current { if !endpoint_is_current {
@@ -61,6 +62,7 @@ impl MePool {
&& coordinator.pending.is_some_and(|pending| { && coordinator.pending.is_some_and(|pending| {
pending.generation == writer.generation pending.generation == writer.generation
&& pending.map_hash == coordinator.desired_map_hash && pending.map_hash == coordinator.desired_map_hash
&& pending.endpoint_revision == coordinator.endpoint_revision
}) })
{ {
return Ok(WriterContour::Warm); return Ok(WriterContour::Warm);
@@ -92,7 +94,8 @@ impl MePool {
if intent == WriterOpenIntent::Replacement || contour == WriterContour::Draining { if intent == WriterOpenIntent::Replacement || contour == WriterContour::Draining {
return Ok(()); 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 { let Some(endpoints) = preferred.get(&writer.writer_dc) else {
return Err(ProxyError::Proxy( return Err(ProxyError::Proxy(
"ME writer target changed before publication".into(), "ME writer target changed before publication".into(),
@@ -85,7 +85,8 @@ impl MePool {
"ME floor rebalance lost active-generation authority".into(), "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 let donor_count = writers
.iter() .iter()
.filter(|candidate| { .filter(|candidate| {
@@ -265,8 +266,11 @@ mod tests {
async fn replacement_commit_publishes_successor_before_draining_victim() { async fn replacement_commit_publishes_successor_before_draining_victim() {
let pool = make_pool().await; let pool = make_pool().await;
let addr = endpoint(1); let addr = endpoint(1);
pool.preferred_endpoints_by_dc pool.update_proxy_maps(
.store(Arc::new(HashMap::from([(2, vec![addr])]))); HashMap::from([(2, vec![(addr.ip(), addr.port())])]),
None,
)
.await;
let victim = install_writer(&pool, 1001, 2, addr).await; let victim = install_writer(&pool, 1001, 2, addr).await;
let expected_role = WriterRole::from_writer(&victim); let expected_role = WriterRole::from_writer(&victim);
let mut reservation = pool let mut reservation = pool
@@ -305,8 +309,11 @@ mod tests {
async fn cancelled_replacement_waiting_for_publication_restores_all_reservations() { async fn cancelled_replacement_waiting_for_publication_restores_all_reservations() {
let pool = make_pool().await; let pool = make_pool().await;
let addr = endpoint(2); let addr = endpoint(2);
pool.preferred_endpoints_by_dc pool.update_proxy_maps(
.store(Arc::new(HashMap::from([(2, vec![addr])]))); HashMap::from([(2, vec![(addr.ip(), addr.port())])]),
None,
)
.await;
let victim = install_writer(&pool, 2001, 2, addr).await; let victim = install_writer(&pool, 2001, 2, addr).await;
let expected_role = WriterRole::from_writer(&victim); let expected_role = WriterRole::from_writer(&victim);
let mut reservation = pool let mut reservation = pool
@@ -342,10 +349,14 @@ mod tests {
let pool = make_pool().await; let pool = make_pool().await;
let donor_addr = endpoint(3); let donor_addr = endpoint(3);
let receiver_addr = endpoint(4); let receiver_addr = endpoint(4);
pool.preferred_endpoints_by_dc.store(Arc::new(HashMap::from([ pool.update_proxy_maps(
(1, vec![donor_addr]), HashMap::from([
(2, vec![receiver_addr]), (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 victim = install_writer(&pool, 3001, 1, donor_addr).await;
let expected_role = WriterRole::from_writer(&victim); let expected_role = WriterRole::from_writer(&victim);
let mut reservation = pool let mut reservation = pool
+6 -3
View File
@@ -348,8 +348,10 @@ impl MePool {
for _ in for _ in
0..self.route_runtime.me_route_inline_recovery_attempts.max(1) 0..self.route_runtime.me_route_inline_recovery_attempts.max(1)
{ {
let preferred = self.preferred_endpoints_by_dc.load_full(); let endpoint_snapshot = self.endpoint_snapshot.load_full();
for (dc, addrs) in preferred.iter() { for (dc, addrs) in
&endpoint_snapshot.preferred_endpoints_by_dc
{
for addr in addrs { for addr in addrs {
let _ = self let _ = self
.connect_one_for_dc(*addr, *dc, self.rng.as_ref()) .connect_one_for_dc(*addr, *dc, self.rng.as_ref())
@@ -470,8 +472,9 @@ impl MePool {
} }
emergency_attempts += 1; emergency_attempts += 1;
let mut endpoints = self let mut endpoints = self
.preferred_endpoints_by_dc .endpoint_snapshot
.load() .load()
.preferred_endpoints_by_dc
.get(&routed_dc) .get(&routed_dc)
.cloned() .cloned()
.unwrap_or_default(); .unwrap_or_default();
+2 -1
View File
@@ -88,7 +88,8 @@ impl MePool {
pub(super) async fn trigger_async_recovery_global(self: &Arc<Self>) { pub(super) async fn trigger_async_recovery_global(self: &Arc<Self>) {
self.stats.increment_me_async_recovery_trigger_total(); 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; let mut triggered = 0usize;
for (dc, addrs) in preferred.iter() { for (dc, addrs) in preferred.iter() {
for addr in addrs { for addr in addrs {
+2 -1
View File
@@ -15,7 +15,8 @@ impl MePool {
routed_dc: i32, routed_dc: i32,
include_warm: bool, include_warm: bool,
) -> Vec<usize> { ) -> Vec<usize> {
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(); let mut out = Vec::new();
if let Some(preferred) = preferred_snapshot if let Some(preferred) = preferred_snapshot
.get(&routed_dc) .get(&routed_dc)
@@ -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 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 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); let latest = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 31, 0, 23)), 443);
pool.preferred_endpoints_by_dc pool.update_proxy_maps(
.store(Arc::new(HashMap::from([(2, vec![first, second, latest])]))); 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(first, 2);
pool.trigger_immediate_refill_for_dc(second, 2); pool.trigger_immediate_refill_for_dc(second, 2);
@@ -43,8 +43,11 @@ fn unregistered_writer(
async fn normal_warm_publication_cannot_race_past_the_dc_floor() { async fn normal_warm_publication_cannot_race_past_the_dc_floor() {
let pool = make_pool().await; let pool = make_pool().await;
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443);
pool.preferred_endpoints_by_dc pool.update_proxy_maps(
.store(Arc::new(std::collections::HashMap::from([(2, vec![addr])]))); std::collections::HashMap::from([(2, vec![(addr.ip(), addr.port())])]),
None,
)
.await;
let generation = 2; let generation = 2;
let writers = (1..=3) let writers = (1..=3)
.map(|writer_id| { .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() { async fn normal_active_publication_cannot_race_past_the_family_floor() {
let pool = make_pool().await; let pool = make_pool().await;
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443);
pool.preferred_endpoints_by_dc pool.update_proxy_maps(
.store(Arc::new(std::collections::HashMap::from([(2, vec![addr])]))); std::collections::HashMap::from([(2, vec![(addr.ip(), addr.port())])]),
None,
)
.await;
let generation = pool.current_generation(); let generation = pool.current_generation();
let writers = (1..=3) let writers = (1..=3)
.map(|writer_id| { .map(|writer_id| {
@@ -154,13 +154,11 @@ async fn insert_writer(
}; };
pool.writers.write().await.push(writer); pool.writers.write().await.push(writer);
{ let mut map = pool.endpoint_snapshot.load().map_v4.clone();
let mut map = pool.proxy_map_v4.write().await; map.entry(writer_dc)
map.entry(writer_dc) .or_insert_with(Vec::new)
.or_insert_with(Vec::new) .push((addr.ip(), addr.port()));
.push((addr.ip(), addr.port())); pool.update_proxy_maps(map, None).await;
}
pool.rebuild_endpoint_dc_map().await;
if register_in_registry { if register_in_registry {
pool.registry pool.registry
.register_writer(writer_id, tx, byte_budget) .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!(pool.admission_ready_conditional_cast().await);
assert_eq!( assert_eq!(
pool.preferred_endpoints_by_dc pool.preferred_endpoints_for_dc(2).await,
.load()
.get(&2)
.cloned()
.unwrap_or_default(),
vec![new_addr] vec![new_addr]
); );
+1
View File
@@ -123,6 +123,7 @@ fn runtime_config_with_carriers_and_deadlines(
carriers: Arc::clone(&carriers), carriers: Arc::clone(&carriers),
carrier_negotiation_deadlines_secs, carrier_negotiation_deadlines_secs,
capability, capability,
credential_id: [0; 16],
key_fingerprint: "0000000000000000".to_string(), key_fingerprint: "0000000000000000".to_string(),
max_sessions: 4, max_sessions: 4,
max_streams: 16, max_streams: 16,
+16 -12
View File
@@ -89,11 +89,6 @@ impl WebProcessRuntime {
.record_rejection(WebRejectionReason::ConfigDisabled); .record_rejection(WebRejectionReason::ConfigDisabled);
return Err(ManagerError::Closed); 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 _operator_admission = self.try_operator_admission()?;
let now = Instant::now(); let now = Instant::now();
let mut state = self.state.lock(); let mut state = self.state.lock();
@@ -146,12 +141,25 @@ impl WebProcessRuntime {
.record_rejection(WebRejectionReason::BootstrapCapacity); .record_rejection(WebRejectionReason::BootstrapCapacity);
return Err(ManagerError::Limit); 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 trace_session_id = self.trace.next_session_id();
let bridge_diagnostics_enabled = config.web.debug.bridge_diagnostics_enabled(); let bridge_diagnostics_enabled = config.web.debug.bridge_diagnostics_enabled();
let (user_agent, user_agent_id) = bounded_user_agent(user_agent); let (user_agent, user_agent_id) = bounded_user_agent(user_agent);
let issued_profile = Arc::clone(&profile);
state.bootstraps.insert( state.bootstraps.insert(
hash, hash,
Bootstrap { Bootstrap {
user_registration,
expires_at: now + Duration::from_secs(config.web.timeouts.bootstrap_lifetime_secs), expires_at: now + Duration::from_secs(config.web.timeouts.bootstrap_lifetime_secs),
issued_at: now, issued_at: now,
issuance_ip: client_ip, issuance_ip: client_ip,
@@ -186,11 +194,7 @@ impl WebProcessRuntime {
}, },
); );
*state.bootstraps_per_ip.entry(client_ip).or_insert(0) += 1; *state.bootstraps_per_ip.entry(client_ip).or_insert(0) += 1;
let profile = state user_publication.commit();
.bootstraps
.get(&hash)
.map(|entry| Arc::clone(&entry.profile))
.ok_or(ManagerError::Closed)?;
drop(state); drop(state);
if recovery { if recovery {
self.telemetry self.telemetry
@@ -202,7 +206,7 @@ impl WebProcessRuntime {
Some(client_ip), Some(client_ip),
crate::web::trace::TraceIdentity::from_optional_profile( crate::web::trace::TraceIdentity::from_optional_profile(
Some(trace_session_id), Some(trace_session_id),
&profile, &issued_profile,
), ),
crate::web::trace::TraceLifecycleEvent::BridgeIssued, crate::web::trace::TraceLifecycleEvent::BridgeIssued,
None, None,
@@ -216,7 +220,7 @@ impl WebProcessRuntime {
self.trace.record_profile_lifecycle( self.trace.record_profile_lifecycle(
client_ip, client_ip,
Some(trace_session_id), Some(trace_session_id),
&profile, &issued_profile,
crate::web::trace::TraceLifecycleEvent::BridgeIssued, crate::web::trace::TraceLifecycleEvent::BridgeIssued,
None, None,
None, None,
+4 -1
View File
@@ -30,14 +30,17 @@ pub(crate) struct WebShutdownDrain {
} }
impl WebProcessRuntime { 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( pub(crate) fn activate_generation(
&self, &self,
generation: Arc<RuntimeGeneration>, generation: Arc<RuntimeGeneration>,
) -> Arc<RuntimeGeneration> { ) -> Arc<RuntimeGeneration> {
let config = generation.config(); let config = generation.config();
let (replaced, detached) = { 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(); let mut learning = self.learning.lock();
state.apply_issuance_policy(generation.id, config.web.enabled);
let outcome = learning.apply_policy( let outcome = learning.apply_policy(
Instant::now(), Instant::now(),
generation.id, generation.id,
+15 -1
View File
@@ -69,7 +69,10 @@ impl WebProcessRuntime {
let Some(entry) = state.bootstraps.get(&bootstrap_hash) else { let Some(entry) = state.bootstraps.get(&bootstrap_hash) else {
return Err(ManagerError::Authentication); 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); return Err(ManagerError::Authentication);
} }
if entry.used { if entry.used {
@@ -301,6 +304,15 @@ impl WebProcessRuntime {
return Err(ManagerError::Protocol); return Err(ManagerError::Protocol);
}; };
let _operator_admission = self.try_operator_admission()?; 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) { if !admit_initial(self, &mut state, now, client_ip, profile_key, &profile) {
return Err(ManagerError::Limit); return Err(ManagerError::Limit);
} }
@@ -338,6 +350,7 @@ impl WebProcessRuntime {
recovery, recovery,
self.limits.clone(), self.limits.clone(),
issued_timeouts.clone(), issued_timeouts.clone(),
Some(user_registration),
); );
state.sessions.insert(session_hash, Arc::clone(&session)); state.sessions.insert(session_hash, Arc::clone(&session));
*state.sessions_per_ip.entry(client_ip).or_insert(0) += 1; *state.sessions_per_ip.entry(client_ip).or_insert(0) += 1;
@@ -402,6 +415,7 @@ impl WebProcessRuntime {
user_agent_id, user_agent_id,
}, },
); );
user_publication.commit();
drop(state); drop(state);
self.telemetry self.telemetry
.record_carrier_selection(carrier, learning_disposition); .record_carrier_selection(carrier, learning_disposition);
@@ -43,14 +43,24 @@ impl WebProcessRuntime {
if !valid if !valid
|| state.closed || state.closed
|| !state.issuance_enabled || !state.issuance_enabled
|| !generation
.proxy_shared
.is_user_enabled(&replacement.profile.user)
{ {
drop(state); drop(state);
self.cancel_replacement(bootstrap_hash, &replacement.old_session); self.cancel_replacement(bootstrap_hash, &replacement.old_session);
return Err(ManagerError::Closed); 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 { let Some((session_token, session_hash)) = new_unique_token(&generation, &state) else {
self.record_limit_hit(); self.record_limit_hit();
self.telemetry self.telemetry
@@ -86,6 +96,7 @@ impl WebProcessRuntime {
replacement.recovery, replacement.recovery,
self.limits.clone(), self.limits.clone(),
replacement.old_session.timeouts().clone(), replacement.old_session.timeouts().clone(),
Some(user_registration),
); );
let Some(supersede) = replacement.old_session.prepare_carrier_supersede() else { let Some(supersede) = replacement.old_session.prepare_carrier_supersede() else {
drop(state); drop(state);
@@ -146,6 +157,7 @@ impl WebProcessRuntime {
index.bootstrap_hash = bootstrap_hash; index.bootstrap_hash = bootstrap_hash;
index.attempt = replacement.attempt; index.attempt = replacement.attempt;
} }
user_publication.commit();
let identity = session.trace_identity(); let identity = session.trace_identity();
let old_identity = replacement.old_session.trace_identity(); let old_identity = replacement.old_session.trace_identity();
drop(state); drop(state);
+4
View File
@@ -10,6 +10,7 @@ use zeroize::Zeroizing;
use super::{CarrierRequest, ProfileKey, TOKEN_BYTES, TokenHash}; use super::{CarrierRequest, ProfileKey, TOKEN_BYTES, TokenHash};
use crate::config::{WebCarrier, WebRuntimeConfig, WebRuntimeProfile, WebTimeoutsConfig}; use crate::config::{WebCarrier, WebRuntimeConfig, WebRuntimeProfile, WebTimeoutsConfig};
use crate::maestro::generation::RuntimeGeneration; use crate::maestro::generation::RuntimeGeneration;
use crate::proxy::user_admission::UserSessionRegistration;
use crate::web::session::WebSession; use crate::web::session::WebSession;
use crate::web::telemetry::WebCarrierSelectionDisposition; use crate::web::telemetry::WebCarrierSelectionDisposition;
@@ -34,6 +35,8 @@ impl CarrierChainPhase {
/// One issued bootstrap and optional idempotent session-creation replay state. /// One issued bootstrap and optional idempotent session-creation replay state.
pub(super) struct Bootstrap { pub(super) struct Bootstrap {
/// User authority ownership retained for the credential lifetime.
pub(super) user_registration: UserSessionRegistration,
/// Credential and replay-state expiry deadline. /// Credential and replay-state expiry deadline.
pub(super) expires_at: Instant, pub(super) expires_at: Instant,
/// Stable ordering point used for bounded eviction. /// Stable ordering point used for bounded eviction.
@@ -273,6 +276,7 @@ pub(super) fn matching_profile(
profile.host == expected.host profile.host == expected.host
&& profile.public_addr == expected.public_addr && profile.public_addr == expected.public_addr
&& profile.user == expected.user && profile.user == expected.user
&& profile.credential_id == expected.credential_id
&& profile.secret_mode == expected.secret_mode && profile.secret_mode == expected.secret_mode
&& profile.carrier == expected.carrier && profile.carrier == expected.carrier
&& profile.carrier_negotiation_enabled == expected.carrier_negotiation_enabled && profile.carrier_negotiation_enabled == expected.carrier_negotiation_enabled
+9 -1
View File
@@ -17,6 +17,7 @@ use crate::web::frame::{self, FrameType};
use crate::web::manager::{ use crate::web::manager::{
CarrierClientClass, CarrierLearningContext, ProfileKey, TokenHash, WebProcessRuntime, CarrierClientClass, CarrierLearningContext, ProfileKey, TokenHash, WebProcessRuntime,
}; };
use crate::proxy::user_admission::UserSessionRegistration;
// Backend tasks own generation admission and authenticated MTProxy relay lifetimes. // Backend tasks own generation admission and authenticated MTProxy relay lifetimes.
mod backend; mod backend;
@@ -213,6 +214,7 @@ pub(crate) struct WebSession {
created_at: Instant, created_at: Instant,
limits: WebLimitsConfig, limits: WebLimitsConfig,
timeouts: WebTimeoutsConfig, timeouts: WebTimeoutsConfig,
_user_registration: Option<UserSessionRegistration>,
state: Mutex<SessionState>, state: Mutex<SessionState>,
carrier_health_publication: AtomicU8, carrier_health_publication: AtomicU8,
close_complete: AtomicBool, close_complete: AtomicBool,
@@ -257,8 +259,13 @@ impl WebSession {
recovery: bool, recovery: bool,
limits: WebLimitsConfig, limits: WebLimitsConfig,
timeouts: WebTimeoutsConfig, timeouts: WebTimeoutsConfig,
user_registration: Option<UserSessionRegistration>,
) -> Arc<Self> { ) -> Arc<Self> {
let created_at = Instant::now(); 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 carrier_lanes = HashMap::new();
let mut next_lane_instance = 1; let mut next_lane_instance = 1;
if selected_carrier == WebCarrier::HttpsLanes { if selected_carrier == WebCarrier::HttpsLanes {
@@ -283,6 +290,7 @@ impl WebSession {
created_at, created_at,
limits, limits,
timeouts, timeouts,
_user_registration: user_registration,
state: Mutex::new(SessionState { state: Mutex::new(SessionState {
streams: HashMap::new(), streams: HashMap::new(),
closing_streams: HashMap::new(), closing_streams: HashMap::new(),
@@ -328,7 +336,7 @@ impl WebSession {
close_notify: Notify::new(), close_notify: Notify::new(),
down_notify: Arc::new(Notify::new()), down_notify: Arc::new(Notify::new()),
lane_open_notify: Arc::new(Notify::new()), lane_open_notify: Arc::new(Notify::new()),
cancel: CancellationToken::new(), cancel,
tasks_live: AtomicUsize::new(0), tasks_live: AtomicUsize::new(0),
tasks_done: Arc::new(Notify::new()), tasks_done: Arc::new(Notify::new()),
resident: Arc::new(resident::ResidentCounters::default()), resident: Arc::new(resident::ResidentCounters::default()),
+2
View File
@@ -75,6 +75,7 @@ fn test_runtime_with_dc(
carriers: Arc::from([carrier]), carriers: Arc::from([carrier]),
carrier_negotiation_deadlines_secs: [3, 5, 8, 12], carrier_negotiation_deadlines_secs: [3, 5, 8, 12],
capability: [7; 32], capability: [7; 32],
credential_id: [0; 16],
key_fingerprint: "0000000000000000".to_string(), key_fingerprint: "0000000000000000".to_string(),
max_sessions: 4, max_sessions: 4,
max_streams: 16, max_streams: 16,
@@ -134,6 +135,7 @@ fn test_runtime_with_dc(
false, false,
limits, limits,
timeouts, timeouts,
None,
); );
TestRuntime { TestRuntime {
session, session,
+2
View File
@@ -24,6 +24,7 @@ fn session() -> (Arc<WebSession>, Arc<WebProcessRuntime>) {
carriers: Arc::from([WebCarrier::Https]), carriers: Arc::from([WebCarrier::Https]),
carrier_negotiation_deadlines_secs: [3, 5, 8, 12], carrier_negotiation_deadlines_secs: [3, 5, 8, 12],
capability: [0; 32], capability: [0; 32],
credential_id: [0; 16],
key_fingerprint: "0000000000000000".to_string(), key_fingerprint: "0000000000000000".to_string(),
max_sessions: 1, max_sessions: 1,
max_streams: 1, max_streams: 1,
@@ -50,6 +51,7 @@ fn session() -> (Arc<WebSession>, Arc<WebProcessRuntime>) {
false, false,
WebLimitsConfig::default(), WebLimitsConfig::default(),
timeouts, timeouts,
None,
); );
(session, manager) (session, manager)
} }
+2
View File
@@ -39,6 +39,7 @@ fn new_session_with_automatic(
carriers: Arc::from([WebCarrier::HttpsLanes]), carriers: Arc::from([WebCarrier::HttpsLanes]),
carrier_negotiation_deadlines_secs: [3, 5, 8, 12], carrier_negotiation_deadlines_secs: [3, 5, 8, 12],
capability: [0; 32], capability: [0; 32],
credential_id: [0; 16],
key_fingerprint: "0000000000000000".to_string(), key_fingerprint: "0000000000000000".to_string(),
max_sessions: 1, max_sessions: 1,
max_streams: 2, max_streams: 2,
@@ -65,6 +66,7 @@ fn new_session_with_automatic(
false, false,
limits, limits,
WebTimeoutsConfig::default(), WebTimeoutsConfig::default(),
None,
) )
} }
+19 -1
View File
@@ -25,6 +25,8 @@ pub(crate) enum SessionCloseReason {
WebSocketEnded, WebSocketEnded,
/// An authenticated control-plane request selected this session. /// An authenticated control-plane request selected this session.
ApiClose, ApiClose,
/// Process user authority revoked or replaced the authenticated credential.
UserDisabled,
/// A graceful operator drain reached its force-close deadline. /// A graceful operator drain reached its force-close deadline.
OperatorForce, OperatorForce,
/// Terminal process shutdown closed all remaining sessions. /// Terminal process shutdown closed all remaining sessions.
@@ -33,7 +35,7 @@ pub(crate) enum SessionCloseReason {
impl SessionCloseReason { impl SessionCloseReason {
/// Complete fixed reason set in stable API and metric order. /// 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::ClientDelete,
Self::BridgeRecovery, Self::BridgeRecovery,
Self::PeerIdle, Self::PeerIdle,
@@ -43,6 +45,7 @@ impl SessionCloseReason {
Self::Backpressure, Self::Backpressure,
Self::WebSocketEnded, Self::WebSocketEnded,
Self::ApiClose, Self::ApiClose,
Self::UserDisabled,
Self::OperatorForce, Self::OperatorForce,
Self::RuntimeShutdown, Self::RuntimeShutdown,
]; ];
@@ -59,6 +62,7 @@ impl SessionCloseReason {
Self::Backpressure => "backpressure", Self::Backpressure => "backpressure",
Self::WebSocketEnded => "websocket_ended", Self::WebSocketEnded => "websocket_ended",
Self::ApiClose => "api_close", Self::ApiClose => "api_close",
Self::UserDisabled => "user_disabled",
Self::OperatorForce => "operator_force", Self::OperatorForce => "operator_force",
Self::RuntimeShutdown => "runtime_shutdown", Self::RuntimeShutdown => "runtime_shutdown",
} }
@@ -217,6 +221,20 @@ impl WebSession {
/// Atomically closes a session only when reconnect grace is still due. /// Atomically closes a session only when reconnect grace is still due.
pub(crate) fn close_if_due(&self, now: Instant) -> bool { pub(crate) fn close_if_due(&self, now: Instant) -> bool {
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 healthy = {
let mut state = self.state.lock(); let mut state = self.state.lock();
self.carrier_health_ready_locked(&mut state, now) self.carrier_health_ready_locked(&mut state, now)
+2
View File
@@ -320,6 +320,7 @@ mod tests {
carriers: Arc::from([carrier]), carriers: Arc::from([carrier]),
carrier_negotiation_deadlines_secs: [3, 5, 8, 12], carrier_negotiation_deadlines_secs: [3, 5, 8, 12],
capability: [0; 32], capability: [0; 32],
credential_id: [0; 16],
key_fingerprint: "0000000000000000".to_string(), key_fingerprint: "0000000000000000".to_string(),
max_sessions: 1, max_sessions: 1,
max_streams: 1, max_streams: 1,
@@ -342,6 +343,7 @@ mod tests {
false, false,
WebLimitsConfig::default(), WebLimitsConfig::default(),
WebTimeoutsConfig::default(), WebTimeoutsConfig::default(),
None,
) )
} }
+2
View File
@@ -22,6 +22,7 @@ fn session_with_automatic(automatic: bool) -> Arc<WebSession> {
carriers: Arc::from([WebCarrier::Https]), carriers: Arc::from([WebCarrier::Https]),
carrier_negotiation_deadlines_secs: [3, 5, 8, 12], carrier_negotiation_deadlines_secs: [3, 5, 8, 12],
capability: [0; 32], capability: [0; 32],
credential_id: [0; 16],
key_fingerprint: "0000000000000000".to_string(), key_fingerprint: "0000000000000000".to_string(),
max_sessions: 1, max_sessions: 1,
max_streams: 1, max_streams: 1,
@@ -48,6 +49,7 @@ fn session_with_automatic(automatic: bool) -> Arc<WebSession> {
false, false,
WebLimitsConfig::default(), WebLimitsConfig::default(),
WebTimeoutsConfig::default(), WebTimeoutsConfig::default(),
None,
) )
} }
+2
View File
@@ -39,6 +39,7 @@ fn runtime(admission: bool) -> TestRuntime {
carriers: Arc::from([WebCarrier::WebsocketLanes]), carriers: Arc::from([WebCarrier::WebsocketLanes]),
carrier_negotiation_deadlines_secs: [3, 5, 8, 12], carrier_negotiation_deadlines_secs: [3, 5, 8, 12],
capability: [7; 32], capability: [7; 32],
credential_id: [0; 16],
key_fingerprint: "0000000000000000".to_string(), key_fingerprint: "0000000000000000".to_string(),
max_sessions: 2, max_sessions: 2,
max_streams: 1, max_streams: 1,
@@ -75,6 +76,7 @@ fn runtime(admission: bool) -> TestRuntime {
false, false,
limits, limits,
timeouts, timeouts,
None,
); );
TestRuntime { TestRuntime {
session, session,