Hardswap Invariants in tests + Quota fixes

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