This commit is contained in:
Alexey
2026-09-25 19:46:31 +03:00
parent 2d63fcf376
commit 326c0ecdb9
128 changed files with 914 additions and 1248 deletions
+2 -2
View File
@@ -6,14 +6,14 @@ use serde_json::Value as Json;
use toml::Value as Toml; use toml::Value as Toml;
use super::ApiShared; use super::ApiShared;
#[cfg(test)]
use super::config_store::write_atomic;
use super::config_store::{ use super::config_store::{
EDITABLE_SECTIONS, EDITABLE_SERVER_FIELDS, compute_snapshot_revision, is_editable_section, EDITABLE_SECTIONS, EDITABLE_SERVER_FIELDS, compute_snapshot_revision, is_editable_section,
load_candidate_snapshot, load_config_snapshot, render_server_listeners, load_candidate_snapshot, load_config_snapshot, render_server_listeners,
render_top_level_section, resolve_single_source_owner, upsert_toml_table, render_top_level_section, resolve_single_source_owner, upsert_toml_table,
write_atomic_if_unchanged, write_atomic_if_unchanged,
}; };
#[cfg(test)]
use super::config_store::write_atomic;
use super::model::ApiFailure; use super::model::ApiFailure;
use crate::config::ProxyConfig; use crate::config::ProxyConfig;
use crate::config::hot_reload::classify_config_changes; use crate::config::hot_reload::classify_config_changes;
+1 -1
View File
@@ -13,13 +13,13 @@ mod persistence;
// Compare-and-replace file persistence and metadata preservation. // Compare-and-replace file persistence and metadata preservation.
mod atomic; mod atomic;
pub(in crate::api) use atomic::{write_atomic, write_atomic_if_unchanged};
#[cfg(test)] #[cfg(test)]
use persistence::{find_toml_table_bounds, render_access_section, save_sections_to_disk}; use persistence::{find_toml_table_bounds, render_access_section, save_sections_to_disk};
pub(in crate::api) use persistence::{ pub(in crate::api) use persistence::{
render_server_listeners, render_top_level_section, save_access_sections_to_disk, render_server_listeners, render_top_level_section, save_access_sections_to_disk,
save_access_sections_to_disk_if_revision, upsert_toml_table, save_access_sections_to_disk_if_revision, upsert_toml_table,
}; };
pub(in crate::api) use atomic::{write_atomic, write_atomic_if_unchanged};
#[derive(Clone, Copy, Debug, PartialEq, Eq)] #[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) enum AccessSection { pub(super) enum AccessSection {
+5 -12
View File
@@ -294,11 +294,7 @@ fn write_atomic_sync(
let descriptor = openat( let descriptor = openat(
anchored.parent(), anchored.parent(),
temp_name.as_str(), temp_name.as_str(),
OFlag::O_WRONLY OFlag::O_WRONLY | OFlag::O_CREAT | OFlag::O_EXCL | OFlag::O_NOFOLLOW | OFlag::O_CLOEXEC,
| OFlag::O_CREAT
| OFlag::O_EXCL
| OFlag::O_NOFOLLOW
| OFlag::O_CLOEXEC,
Mode::from_bits_truncate(0o600), Mode::from_bits_truncate(0o600),
) )
.map_err(errno_to_io)?; .map_err(errno_to_io)?;
@@ -408,9 +404,9 @@ fn validate_expected_contents(
existing: Option<&ExistingTarget>, existing: Option<&ExistingTarget>,
expected_contents: Option<&str>, expected_contents: Option<&str>,
) -> std::io::Result<()> { ) -> std::io::Result<()> {
if expected_contents.is_some_and(|expected| { if expected_contents
existing.is_none_or(|target| target.contents != expected) .is_some_and(|expected| existing.is_none_or(|target| target.contents != expected))
}) { {
return Err(std::io::Error::new( return Err(std::io::Error::new(
std::io::ErrorKind::AlreadyExists, std::io::ErrorKind::AlreadyExists,
"config source changed before persistence", "config source changed before persistence",
@@ -419,10 +415,7 @@ fn validate_expected_contents(
Ok(()) Ok(())
} }
fn target_unchanged( fn target_unchanged(existing: Option<&ExistingTarget>, current: Option<&ExistingTarget>) -> bool {
existing: Option<&ExistingTarget>,
current: Option<&ExistingTarget>,
) -> bool {
match (existing, current) { match (existing, current) {
(Some(expected), Some(current)) => { (Some(expected), Some(current)) => {
same_target(&expected.metadata, &current.metadata) same_target(&expected.metadata, &current.metadata)
+3 -3
View File
@@ -6,15 +6,15 @@ use serde::Serialize;
use crate::config::{ProxyConfig, RateLimitBps}; use crate::config::{ProxyConfig, RateLimitBps};
#[cfg(test)]
use super::atomic::write_atomic;
use super::atomic::write_atomic_if_unchanged;
#[cfg(test)] #[cfg(test)]
use super::compute_revision; use super::compute_revision;
use super::{ use super::{
AccessSection, compute_snapshot_revision, load_candidate_snapshot, load_config_snapshot, AccessSection, compute_snapshot_revision, load_candidate_snapshot, load_config_snapshot,
resolve_single_source_owner, toml_path_exists, resolve_single_source_owner, toml_path_exists,
}; };
use super::atomic::write_atomic_if_unchanged;
#[cfg(test)]
use super::atomic::write_atomic;
use crate::api::model::ApiFailure; use crate::api::model::ApiFailure;
/// Re-render the given top-level tables from `cfg` and upsert each into the /// Re-render the given top-level tables from `cfg` and upsert each into the
+5 -3
View File
@@ -266,8 +266,7 @@ async fn access_mutation_rejects_source_graph_change_after_snapshot() {
let root = dir.path().join("config.toml"); let root = dir.path().join("config.toml");
let included = dir.path().join("users.toml"); let included = dir.path().join("users.toml");
let root_body = "include = \"users.toml\"\n[censorship]\ntls_domain = \"one.example\"\n"; let root_body = "include = \"users.toml\"\n[censorship]\ntls_domain = \"one.example\"\n";
let external_root = let external_root = "include = \"users.toml\"\n[censorship]\ntls_domain = \"two.example\"\n";
"include = \"users.toml\"\n[censorship]\ntls_domain = \"two.example\"\n";
let included_body = "[access.users]\nalice = \"00000000000000000000000000000000\"\n"; let included_body = "[access.users]\nalice = \"00000000000000000000000000000000\"\n";
tokio::fs::write(&root, root_body).await.unwrap(); tokio::fs::write(&root, root_body).await.unwrap();
tokio::fs::write(&included, included_body).await.unwrap(); tokio::fs::write(&included, included_body).await.unwrap();
@@ -288,7 +287,10 @@ async fn access_mutation_rejects_source_graph_change_after_snapshot() {
.unwrap_err(); .unwrap_err();
assert_eq!(error.code, "revision_conflict"); assert_eq!(error.code, "revision_conflict");
assert_eq!(tokio::fs::read_to_string(&root).await.unwrap(), external_root); assert_eq!(
tokio::fs::read_to_string(&root).await.unwrap(),
external_root
);
assert_eq!( assert_eq!(
tokio::fs::read_to_string(&included).await.unwrap(), tokio::fs::read_to_string(&included).await.unwrap(),
included_body included_body
+6 -8
View File
@@ -103,10 +103,9 @@ pub(super) async fn handle(
}; };
let runtime_cfg = config_rx.borrow().clone(); let runtime_cfg = config_rx.borrow().clone();
data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); data.in_runtime = runtime_cfg.access.users.contains_key(&data.username);
shared.runtime_events.record( shared
"api.user.disable.ok", .runtime_events
format!("username={}", base_user), .record("api.user.disable.ok", format!("username={}", base_user));
);
let status = if data.in_runtime { let status = if data.in_runtime {
StatusCode::OK StatusCode::OK
} else { } else {
@@ -328,10 +327,9 @@ pub(super) async fn handle(
return Err(error); return Err(error);
} }
}; };
shared.runtime_events.record( shared
"api.user.delete.ok", .runtime_events
format!("username={}", deleted_user), .record("api.user.delete.ok", format!("username={}", deleted_user));
);
let runtime_cfg = config_rx.borrow().clone(); let runtime_cfg = config_rx.borrow().clone();
let in_runtime = runtime_cfg.access.users.contains_key(&deleted_user); let in_runtime = runtime_cfg.access.users.contains_key(&deleted_user);
let response = DeleteUserResponse { let response = DeleteUserResponse {
+1 -3
View File
@@ -314,9 +314,7 @@ async fn recompute_connections_payload(
let mut active_users = 0usize; let mut active_users = 0usize;
for entry in shared.stats.iter_user_stats() { for entry in shared.stats.iter_user_stats() {
let user_stats = entry.value(); let user_stats = entry.value();
let current_connections = shared let current_connections = shared.stats.get_process_user_curr_connects(entry.key());
.stats
.get_process_user_curr_connects(entry.key());
let total_octets = user_stats let total_octets = user_stats
.octets_from_client .octets_from_client
.load(std::sync::atomic::Ordering::Relaxed) .load(std::sync::atomic::Ordering::Relaxed)
+2 -4
View File
@@ -4,7 +4,7 @@ use std::collections::BTreeSet;
use serde::Serialize; use serde::Serialize;
use super::{now_epoch_secs, ApiShared, SOURCE_UNAVAILABLE_REASON}; use super::{ApiShared, SOURCE_UNAVAILABLE_REASON, now_epoch_secs};
#[derive(Serialize)] #[derive(Serialize)]
struct RuntimeMePoolStateGenerationData { struct RuntimeMePoolStateGenerationData {
@@ -90,9 +90,7 @@ struct RuntimeMePoolStateData {
} }
/// Builds the bounded runtime ME pool response projection. /// Builds the bounded runtime ME pool response projection.
pub(in crate::api) async fn build_runtime_me_pool_state_data( pub(in crate::api) async fn build_runtime_me_pool_state_data(shared: &ApiShared) -> impl Serialize {
shared: &ApiShared,
) -> impl Serialize {
let now_epoch_secs = now_epoch_secs(); let now_epoch_secs = now_epoch_secs();
let Some(pool) = shared.me_pool.read().await.clone() else { let Some(pool) = shared.me_pool.read().await.clone() else {
return RuntimeMePoolStateData { return RuntimeMePoolStateData {
+2 -4
View File
@@ -63,12 +63,10 @@ pub(super) fn build_zero_all_data(stats: &Stats, configured_users: usize) -> Zer
conntrack_rule_apply_ok: stats.get_conntrack_rule_apply_ok(), conntrack_rule_apply_ok: stats.get_conntrack_rule_apply_ok(),
conntrack_rule_reconcile_success_total: stats conntrack_rule_reconcile_success_total: stats
.get_conntrack_rule_reconcile_success_total(), .get_conntrack_rule_reconcile_success_total(),
conntrack_rule_reconcile_error_total: stats conntrack_rule_reconcile_error_total: stats.get_conntrack_rule_reconcile_error_total(),
.get_conntrack_rule_reconcile_error_total(),
conntrack_rule_rollback_success_total: stats conntrack_rule_rollback_success_total: stats
.get_conntrack_rule_rollback_success_total(), .get_conntrack_rule_rollback_success_total(),
conntrack_rule_rollback_error_total: stats conntrack_rule_rollback_error_total: stats.get_conntrack_rule_rollback_error_total(),
.get_conntrack_rule_rollback_error_total(),
conntrack_delete_attempt_total: stats.get_conntrack_delete_attempt_total(), conntrack_delete_attempt_total: stats.get_conntrack_delete_attempt_total(),
conntrack_delete_success_total: stats.get_conntrack_delete_success_total(), conntrack_delete_success_total: stats.get_conntrack_delete_success_total(),
conntrack_delete_not_found_total: stats.get_conntrack_delete_not_found_total(), conntrack_delete_not_found_total: stats.get_conntrack_delete_not_found_total(),
+1 -3
View File
@@ -145,9 +145,7 @@ async fn create_user_to_completion(
Some(&base_revision), Some(&base_revision),
) )
.await?; .await?;
shared shared.proxy_shared.stage_user_credential(
.proxy_shared
.stage_user_credential(
&body.username, &body.username,
credential_id, credential_id,
cfg.access.is_user_enabled(&body.username), cfg.access.is_user_enabled(&body.username),
+5 -3
View File
@@ -54,9 +54,11 @@ async fn rotate_secret_to_completion(
Some(&base_revision), Some(&base_revision),
) )
.await?; .await?;
shared shared.proxy_shared.stage_user_credential(
.proxy_shared user,
.stage_user_credential(user, credential_id, cfg.access.is_user_enabled(user)); credential_id,
cfg.access.is_user_enabled(user),
);
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();
+10 -8
View File
@@ -154,16 +154,16 @@ async fn patch_user_to_completion(
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 staged_credential =
if touches_users || touches_user_enabled {
let secret = cfg let secret = cfg
.access .access
.users .users
.get(user) .get(user)
.ok_or_else(|| ApiFailure::internal("updated user secret is missing"))?; .ok_or_else(|| ApiFailure::internal("updated user secret is missing"))?;
Some( Some(credential_id_from_hex(secret).ok_or_else(|| {
credential_id_from_hex(secret) ApiFailure::internal("validated user secret could not be decoded")
.ok_or_else(|| ApiFailure::internal("validated user secret could not be decoded"))?, })?)
)
} else { } else {
None None
}; };
@@ -206,9 +206,11 @@ async fn patch_user_to_completion(
.await? .await?
}; };
if let Some(credential_id) = staged_credential { if let Some(credential_id) = staged_credential {
shared shared.proxy_shared.stage_user_credential(
.proxy_shared user,
.stage_user_credential(user, credential_id, cfg.access.is_user_enabled(user)); credential_id,
cfg.access.is_user_enabled(user),
);
} }
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,
+1 -1
View File
@@ -14,7 +14,7 @@ use std::os::unix::fs::MetadataExt;
#[cfg(unix)] #[cfg(unix)]
use nix::dir::Dir; use nix::dir::Dir;
#[cfg(unix)] #[cfg(unix)]
use nix::fcntl::{openat, OFlag}; use nix::fcntl::{OFlag, openat};
#[cfg(unix)] #[cfg(unix)]
use nix::sys::stat::Mode; use nix::sys::stat::Mode;
@@ -124,7 +124,15 @@ fn load_static_directory(
let relative = path.strip_prefix(root).map_err(|_| { let relative = path.strip_prefix(root).map_err(|_| {
ProxyError::Config("WEB static path escaped its configured root".to_string()) ProxyError::Config("WEB static path escaped its configured root".to_string())
})?; })?;
load_static_file(file, &metadata, relative, &path, assets, total_bytes, limits)?; load_static_file(
file,
&metadata,
relative,
&path,
assets,
total_bytes,
limits,
)?;
} }
Ok(()) Ok(())
} }
+1 -1
View File
@@ -13,10 +13,10 @@ use crate::stats::Stats;
// Privileged netfilter rule and conntrack helper execution. // Privileged netfilter rule and conntrack helper execution.
mod firewall; mod firewall;
pub(crate) use firewall::FirewallAuthority;
use firewall::{ use firewall::{
DeleteOutcome, delete_conntrack_entry, effective_conntrack_enabled, probe_runtime_support, DeleteOutcome, delete_conntrack_entry, effective_conntrack_enabled, probe_runtime_support,
}; };
pub(crate) use firewall::FirewallAuthority;
const CONNTRACK_EVENT_QUEUE_CAPACITY: usize = 32_768; const CONNTRACK_EVENT_QUEUE_CAPACITY: usize = 32_768;
const PRESSURE_RELEASE_TICKS: u8 = 3; const PRESSURE_RELEASE_TICKS: u8 = 3;
+15 -19
View File
@@ -236,21 +236,15 @@ where
} }
let desired = current.as_ref().expect("desired state is present").clone(); let desired = current.as_ref().expect("desired state is present").clone();
let interruptible = InterruptibleRunner::new( let interruptible =
&self.runner, InterruptibleRunner::new(&self.runner, &self.terminal, &process_cancellation);
&self.terminal, let result =
&process_cancellation, reconcile_once(&interruptible, &interruptible, &mut self.applied, &desired).await;
);
let result = reconcile_once(
&interruptible,
&interruptible,
&mut self.applied,
&desired,
)
.await;
match result { match result {
Ok(()) => { Ok(()) => {
desired.stats.increment_conntrack_rule_reconcile_success_total(); desired
.stats
.increment_conntrack_rule_reconcile_success_total();
desired.stats.set_conntrack_rule_apply_ok(true); desired.stats.set_conntrack_rule_apply_ok(true);
self.status_tx.send_replace(Some(ReconcileStatus { self.status_tx.send_replace(Some(ReconcileStatus {
generation: desired.generation, generation: desired.generation,
@@ -265,7 +259,9 @@ where
} }
Err(failure) if failure.cancelled => break, Err(failure) if failure.cancelled => break,
Err(failure) => { Err(failure) => {
desired.stats.increment_conntrack_rule_reconcile_error_total(); desired
.stats
.increment_conntrack_rule_reconcile_error_total();
desired.stats.set_conntrack_rule_apply_ok(false); desired.stats.set_conntrack_rule_apply_ok(false);
if let Some(rollback_succeeded) = failure.rollback_succeeded { if let Some(rollback_succeeded) = failure.rollback_succeeded {
if rollback_succeeded { if rollback_succeeded {
@@ -326,12 +322,12 @@ where
if let Some(stats) = &self.last_stats { if let Some(stats) = &self.last_stats {
stats.set_conntrack_rule_apply_ok(false); stats.set_conntrack_rule_apply_ok(false);
} }
if let Err(error) = tokio::time::timeout( if let Err(error) =
SHUTDOWN_CLEANUP_TIMEOUT, tokio::time::timeout(SHUTDOWN_CLEANUP_TIMEOUT, recover_to_empty(&self.runner))
recover_to_empty(&self.runner),
)
.await .await
.unwrap_or_else(|_| Err(CommandError::failed("firewall shutdown cleanup timed out"))) .unwrap_or_else(|_| {
Err(CommandError::failed("firewall shutdown cleanup timed out"))
})
{ {
warn!(error = %error, "Failed to clear conntrack firewall policy during shutdown"); warn!(error = %error, "Failed to clear conntrack firewall policy during shutdown");
} else { } else {
+14 -29
View File
@@ -1,6 +1,4 @@
use super::command::{ use super::command::{CommandError, CommandErrorKind, CommandSpec, FirewallCommandRunner};
CommandError, CommandErrorKind, CommandSpec, FirewallCommandRunner,
};
use super::model::{NotrackTarget, ShadowSlot}; use super::model::{NotrackTarget, ShadowSlot};
const DISPATCH_CHAIN: &str = "TELEMT_NOTRACK"; const DISPATCH_CHAIN: &str = "TELEMT_NOTRACK";
@@ -30,10 +28,7 @@ impl IpFamily {
} }
} }
pub(super) fn family_available<R: FirewallCommandRunner>( pub(super) fn family_available<R: FirewallCommandRunner>(runner: &R, family: IpFamily) -> bool {
runner: &R,
family: IpFamily,
) -> bool {
runner.available(family.command_binary()) && runner.available(family.restore_binary()) runner.available(family.command_binary()) && runner.available(family.restore_binary())
} }
@@ -83,9 +78,7 @@ pub(super) async fn activate_family<R: FirewallCommandRunner>(
.await .await
} }
pub(super) async fn cleanup_all<R: FirewallCommandRunner>( pub(super) async fn cleanup_all<R: FirewallCommandRunner>(runner: &R) -> Result<(), CommandError> {
runner: &R,
) -> Result<(), CommandError> {
let mut errors = Vec::new(); let mut errors = Vec::new();
for family in [IpFamily::V4, IpFamily::V6] { for family in [IpFamily::V4, IpFamily::V6] {
if !runner.available(family.command_binary()) { if !runner.available(family.command_binary()) {
@@ -118,7 +111,10 @@ async fn cleanup_family<R: FirewallCommandRunner>(
match result { match result {
Ok(()) => {} Ok(()) => {}
Err(error) Err(error)
if matches!(error.kind, CommandErrorKind::NotFound | CommandErrorKind::Missing) => if matches!(
error.kind,
CommandErrorKind::NotFound | CommandErrorKind::Missing
) =>
{ {
break; break;
} }
@@ -131,13 +127,13 @@ async fn cleanup_family<R: FirewallCommandRunner>(
for chain in [DISPATCH_CHAIN, SHADOW_CHAIN_A, SHADOW_CHAIN_B] { for chain in [DISPATCH_CHAIN, SHADOW_CHAIN_A, SHADOW_CHAIN_B] {
for operation in ["-F", "-X"] { for operation in ["-F", "-X"] {
let result = runner let result = runner
.run(CommandSpec::new( .run(CommandSpec::new(binary, ["-t", "raw", operation, chain]))
binary,
["-t", "raw", operation, chain],
))
.await; .await;
if let Err(error) = result if let Err(error) = result
&& !matches!(error.kind, CommandErrorKind::NotFound | CommandErrorKind::Missing) && !matches!(
error.kind,
CommandErrorKind::NotFound | CommandErrorKind::Missing
)
{ {
errors.push(error.message); errors.push(error.message);
} }
@@ -167,15 +163,7 @@ async fn ensure_prerouting_jump<R: FirewallCommandRunner>(
runner runner
.run(CommandSpec::new( .run(CommandSpec::new(
binary, binary,
[ ["-t", "raw", "-I", "PREROUTING", "1", "-j", DISPATCH_CHAIN],
"-t",
"raw",
"-I",
"PREROUTING",
"1",
"-j",
DISPATCH_CHAIN,
],
)) ))
.await .await
} }
@@ -220,10 +208,7 @@ fn require_family<R: FirewallCommandRunner>(
Ok(()) Ok(())
} }
pub(super) fn render_stage_script( pub(super) fn render_stage_script(slot: ShadowSlot, targets: &[NotrackTarget]) -> String {
slot: ShadowSlot,
targets: &[NotrackTarget],
) -> String {
let chain = shadow_chain(slot); let chain = shadow_chain(slot);
let mut script = format!("*raw\n-F {chain}\n"); let mut script = format!("*raw\n-F {chain}\n");
for target in targets { for target in targets {
-1
View File
@@ -18,7 +18,6 @@ impl ShadowSlot {
Self::B => Self::A, Self::B => Self::A,
} }
} }
} }
#[derive(Clone, Debug, Eq, Ord, PartialEq, PartialOrd)] #[derive(Clone, Debug, Eq, Ord, PartialEq, PartialOrd)]
+6 -7
View File
@@ -54,9 +54,7 @@ pub(super) async fn deactivate<R: FirewallCommandRunner>(
delete_table_if_present(runner, table(slot)).await delete_table_if_present(runner, table(slot)).await
} }
pub(super) async fn cleanup_all<R: FirewallCommandRunner>( pub(super) async fn cleanup_all<R: FirewallCommandRunner>(runner: &R) -> Result<(), CommandError> {
runner: &R,
) -> Result<(), CommandError> {
if !runner.available("nft") { if !runner.available("nft") {
return Ok(()); return Ok(());
} }
@@ -86,7 +84,10 @@ async fn delete_table_if_present<R: FirewallCommandRunner>(
{ {
Ok(()) => Ok(()), Ok(()) => Ok(()),
Err(error) Err(error)
if matches!(error.kind, CommandErrorKind::NotFound | CommandErrorKind::Missing) => if matches!(
error.kind,
CommandErrorKind::NotFound | CommandErrorKind::Missing
) =>
{ {
Ok(()) Ok(())
} }
@@ -111,9 +112,7 @@ pub(super) fn render_stage_script(
v6: &[NotrackTarget], v6: &[NotrackTarget],
) -> String { ) -> String {
let table = table(slot); let table = table(slot);
let mut script = format!( let mut script = format!("add table inet {table}\nadd chain inet {table} rules\n");
"add table inet {table}\nadd chain inet {table} rules\n"
);
for target in v4 { for target in v4 {
script.push_str("add rule inet "); script.push_str("add rule inet ");
script.push_str(table); script.push_str(table);
+3 -8
View File
@@ -1,7 +1,7 @@
use std::collections::{BTreeMap, BTreeSet}; use std::collections::{BTreeMap, BTreeSet};
use std::future::pending; use std::future::pending;
use std::sync::{Arc, Mutex};
use std::sync::atomic::Ordering; use std::sync::atomic::Ordering;
use std::sync::{Arc, Mutex};
use std::time::Duration; use std::time::Duration;
use tokio::sync::{Notify, watch}; use tokio::sync::{Notify, watch};
@@ -102,8 +102,7 @@ impl FirewallCommandRunner for FakeRunner {
let operation = spec.args.get(2).map(String::as_str); let operation = spec.args.get(2).map(String::as_str);
if (matches!(spec.binary, "iptables" | "ip6tables") if (matches!(spec.binary, "iptables" | "ip6tables")
&& matches!(operation, Some("-C" | "-D" | "-F" | "-X"))) && matches!(operation, Some("-C" | "-D" | "-F" | "-X")))
|| (spec.binary == "nft" || (spec.binary == "nft" && spec.args.first().map(String::as_str) == Some("delete"))
&& spec.args.first().map(String::as_str) == Some("delete"))
{ {
return Err(CommandError { return Err(CommandError {
kind: CommandErrorKind::NotFound, kind: CommandErrorKind::NotFound,
@@ -344,11 +343,7 @@ async fn transaction_cancellation_does_not_claim_a_new_applied_plan() {
}; };
let terminal = CancellationToken::new(); let terminal = CancellationToken::new();
let process_cancellation = CancellationToken::new(); let process_cancellation = CancellationToken::new();
let interruptible = InterruptibleRunner::new( let interruptible = InterruptibleRunner::new(&runner, &terminal, &process_cancellation);
&runner,
&terminal,
&process_cancellation,
);
let mut applied = AppliedState::Known(AppliedPlan::Empty); let mut applied = AppliedState::Known(AppliedPlan::Empty);
let desired = desired(1, dual_stack_policy(443)); let desired = desired(1, dual_stack_policy(443));
let failure = { let failure = {
@@ -73,14 +73,9 @@ fn hybrid_policy_is_a_sorted_deduplicated_address_port_product() {
#[test] #[test]
fn restore_renderers_keep_staging_detached_from_activation() { fn restore_renderers_keep_staging_detached_from_activation() {
let stage = iptables::render_stage_script( let stage = iptables::render_stage_script(ShadowSlot::B, &[target(Some("192.0.2.20"), 443)]);
ShadowSlot::B,
&[target(Some("192.0.2.20"), 443)],
);
assert!(stage.contains("-F TELEMT_NT_B\n")); assert!(stage.contains("-F TELEMT_NT_B\n"));
assert!(stage.contains( assert!(stage.contains("-A TELEMT_NT_B -p tcp --dport 443 -d 192.0.2.20 -j CT --notrack\n"));
"-A TELEMT_NT_B -p tcp --dport 443 -d 192.0.2.20 -j CT --notrack\n"
));
assert!(!stage.contains("-A TELEMT_NOTRACK -j TELEMT_NT_B")); assert!(!stage.contains("-A TELEMT_NOTRACK -j TELEMT_NT_B"));
assert!(!stage.contains(":TELEMT_")); assert!(!stage.contains(":TELEMT_"));
@@ -4,9 +4,7 @@ use tokio_util::sync::CancellationToken;
use crate::config::ConntrackBackend; use crate::config::ConntrackBackend;
use super::command::{ use super::command::{CommandError, CommandErrorKind, CommandSpec, FirewallCommandRunner};
CommandError, CommandErrorKind, CommandSpec, FirewallCommandRunner,
};
use super::iptables::{self, IpFamily}; use super::iptables::{self, IpFamily};
use super::model::{AppliedPlan, AppliedState, DesiredPolicy, DesiredState, ShadowSlot}; use super::model::{AppliedPlan, AppliedState, DesiredPolicy, DesiredState, ShadowSlot};
use super::nftables; use super::nftables;
@@ -251,9 +249,13 @@ pub(super) async fn transition_plan<R: FirewallCommandRunner>(
iptables::activate_family(runner, IpFamily::V6, None).await?; iptables::activate_family(runner, IpFamily::V6, None).await?;
} }
} }
(AppliedPlan::Nftables { slot: previous_slot, .. }, AppliedPlan::Nftables { slot, .. }) (
if previous_slot != slot => AppliedPlan::Nftables {
{ slot: previous_slot,
..
},
AppliedPlan::Nftables { slot, .. },
) if previous_slot != slot => {
nftables::deactivate(runner, *previous_slot).await?; nftables::deactivate(runner, *previous_slot).await?;
} }
(_, _) if previous != target => clear_plan(runner, previous).await?, (_, _) if previous != target => clear_plan(runner, previous).await?,
+18 -32
View File
@@ -1,9 +1,9 @@
use std::ffi::OsStr; use std::ffi::OsStr;
use std::fs::{self, File}; use std::fs::{self, File};
use std::io::{self, ErrorKind, Read, Write}; use std::io::{self, ErrorKind, Read, Write};
use std::os::unix::fs::{MetadataExt, PermissionsExt};
#[cfg(target_os = "linux")] #[cfg(target_os = "linux")]
use std::os::fd::{FromRawFd, OwnedFd}; use std::os::fd::{FromRawFd, OwnedFd};
use std::os::unix::fs::{MetadataExt, PermissionsExt};
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use nix::fcntl::{Flock, FlockArg, OFlag, openat}; use nix::fcntl::{Flock, FlockArg, OFlag, openat};
@@ -66,15 +66,14 @@ impl PidFile {
/// ///
/// Fails if another owner holds the lock or the existing PID names a running process. /// Fails if another owner holds the lock or the existing PID names a running process.
pub fn acquire(&mut self) -> Result<(), DaemonError> { pub fn acquire(&mut self) -> Result<(), DaemonError> {
let anchor = AnchoredPath::open_trusted_parent_or_create(&self.path, 0o755).map_err( let anchor =
|error| { AnchoredPath::open_trusted_parent_or_create(&self.path, 0o755).map_err(|error| {
DaemonError::PidFile(format!( DaemonError::PidFile(format!(
"cannot open trusted parent for {}: {}", "cannot open trusted parent for {}: {}",
self.path.display(), self.path.display(),
error error
)) ))
}, })?;
)?;
let lock_name = self.lock_path.file_name().ok_or_else(|| { let lock_name = self.lock_path.file_name().ok_or_else(|| {
DaemonError::PidFile(format!( DaemonError::PidFile(format!(
"lock path {} has no file name", "lock path {} has no file name",
@@ -132,7 +131,11 @@ impl PidFile {
// Validate the opened inode before modifying it so a hard-link substitution // Validate the opened inode before modifying it so a hard-link substitution
// cannot turn PID publication into truncation of an unrelated file. // cannot turn PID publication into truncation of an unrelated file.
pid_file.set_len(0).map_err(|error| { pid_file.set_len(0).map_err(|error| {
DaemonError::PidFile(format!("cannot truncate {}: {}", self.path.display(), error)) DaemonError::PidFile(format!(
"cannot truncate {}: {}",
self.path.display(),
error
))
})?; })?;
let pid = getpid(); let pid = getpid();
writeln!(pid_file, "{}", pid).map_err(|error| { writeln!(pid_file, "{}", pid).map_err(|error| {
@@ -236,18 +239,8 @@ fn normalize_pid_path(path: &Path) -> PathBuf {
} }
} }
fn open_file_at( fn open_file_at(anchor: &AnchoredPath, name: &OsStr, flags: OFlag, mode: u32) -> io::Result<File> {
anchor: &AnchoredPath, let descriptor = openat(anchor.parent(), name, flags, Mode::from_bits_truncate(mode))
name: &OsStr,
flags: OFlag,
mode: u32,
) -> io::Result<File> {
let descriptor = openat(
anchor.parent(),
name,
flags,
Mode::from_bits_truncate(mode),
)
.map_err(|error| io::Error::from_raw_os_error(error as i32))?; .map_err(|error| io::Error::from_raw_os_error(error as i32))?;
Ok(File::from(descriptor)) Ok(File::from(descriptor))
} }
@@ -337,12 +330,7 @@ fn remove_owned_pid_file(
))); )));
} }
drop(file); drop(file);
unlinkat( unlinkat(anchor.parent(), anchor.name(), UnlinkatFlags::NoRemoveDir).map_err(|error| {
anchor.parent(),
anchor.name(),
UnlinkatFlags::NoRemoveDir,
)
.map_err(|error| {
DaemonError::PidFile(format!( DaemonError::PidFile(format!(
"cannot remove {}: {}", "cannot remove {}: {}",
path.display(), path.display(),
@@ -351,10 +339,7 @@ fn remove_owned_pid_file(
}) })
} }
fn validate_regular_single_link( fn validate_regular_single_link(file: &File, path: &Path) -> Result<fs::Metadata, DaemonError> {
file: &File,
path: &Path,
) -> Result<fs::Metadata, DaemonError> {
let metadata = file.metadata().map_err(|error| { let metadata = file.metadata().map_err(|error| {
DaemonError::PidFile(format!("cannot inspect {}: {}", path.display(), error)) DaemonError::PidFile(format!("cannot inspect {}: {}", path.display(), error))
})?; })?;
@@ -419,9 +404,7 @@ pub enum DaemonStatus {
pub fn check_status<P: AsRef<Path>>(path: P) -> DaemonStatus { pub fn check_status<P: AsRef<Path>>(path: P) -> DaemonStatus {
let path = normalize_pid_path(path.as_ref()); let path = normalize_pid_path(path.as_ref());
match read_pid_file_if_exists(&path) { match read_pid_file_if_exists(&path) {
Ok(Some(pid)) Ok(Some(pid)) if daemon_lock_is_held(&path).unwrap_or(false) && is_process_running(pid) => {
if daemon_lock_is_held(&path).unwrap_or(false) && is_process_running(pid) =>
{
DaemonStatus::Running(pid) DaemonStatus::Running(pid)
} }
Ok(Some(pid)) => DaemonStatus::Stale(pid), Ok(Some(pid)) => DaemonStatus::Stale(pid),
@@ -443,7 +426,10 @@ fn daemon_lock_is_held(path: &Path) -> Result<bool, DaemonError> {
} }
}; };
let lock_name = lock_path.file_name().ok_or_else(|| { let lock_name = lock_path.file_name().ok_or_else(|| {
DaemonError::PidFile(format!("lock path {} has no file name", lock_path.display())) DaemonError::PidFile(format!(
"lock path {} has no file name",
lock_path.display()
))
})?; })?;
let file = match open_file_at( let file = match open_file_at(
&anchor, &anchor,
+5 -1
View File
@@ -238,7 +238,11 @@ fn release_does_not_remove_replacement_path() {
let error = pid_file.release().unwrap_err(); let error = pid_file.release().unwrap_err();
assert!(error.to_string().contains("refusing to remove replaced PID file")); assert!(
error
.to_string()
.contains("refusing to remove replaced PID file")
);
assert_eq!(fs::read(&pid_path).unwrap(), b"replacement\n"); assert_eq!(fs::read(&pid_path).unwrap(), b"replacement\n");
} }
+7 -3
View File
@@ -41,8 +41,7 @@ struct CleanupShard {
queue: Mutex<CleanupQueue>, queue: Mutex<CleanupQueue>,
} }
type CleanupQueue = type CleanupQueue = HashMap<String, HashMap<UserIncarnation, HashMap<IpAddr, usize>>>;
HashMap<String, HashMap<UserIncarnation, HashMap<IpAddr, usize>>>;
type CleanupBatch = HashMap<(String, UserIncarnation, IpAddr), usize>; type CleanupBatch = HashMap<(String, UserIncarnation, IpAddr), usize>;
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
@@ -208,7 +207,12 @@ impl UserIpTracker {
) -> Option<(String, UserIncarnation, IpAddr, usize)> { ) -> Option<(String, UserIncarnation, IpAddr, usize)> {
let user = queue.keys().next().cloned()?; let user = queue.keys().next().cloned()?;
let incarnation = queue.get(&user)?.keys().next().copied()?; let incarnation = queue.get(&user)?.keys().next().copied()?;
let ip = queue.get(&user)?.get(&incarnation)?.keys().next().copied()?; let ip = queue
.get(&user)?
.get(&incarnation)?
.keys()
.next()
.copied()?;
let incarnations = queue.get_mut(&user)?; let incarnations = queue.get_mut(&user)?;
let ips = incarnations.get_mut(&incarnation)?; let ips = incarnations.get_mut(&incarnation)?;
let count = ips.remove(&ip)?; let count = ips.remove(&ip)?;
+2 -10
View File
@@ -56,10 +56,7 @@ impl Drop for DetachedCleanupBatch<'_> {
); );
} }
} }
UserIpTracker::decrement_counter( UserIpTracker::decrement_counter(&self.tracker.cleanup_queue_len, duplicate_entries);
&self.tracker.cleanup_queue_len,
duplicate_entries,
);
} }
} }
@@ -187,12 +184,7 @@ impl UserIpTracker {
continue; continue;
} }
removed_active_entries = removed_active_entries.saturating_add( removed_active_entries = removed_active_entries.saturating_add(
Self::apply_active_cleanup( Self::apply_active_cleanup(&mut shard.active_ips, queued_user, *ip, *pending_count),
&mut shard.active_ips,
queued_user,
*ip,
*pending_count,
),
); );
} }
Self::decrement_counter(&self.active_entry_count, removed_active_entries); Self::decrement_counter(&self.active_entry_count, removed_active_entries);
+4 -24
View File
@@ -203,9 +203,7 @@ async fn stale_incarnation_cleanup_cannot_release_recreated_user_ip() {
.check_and_add_for_incarnation("test_user", 1, old_ip) .check_and_add_for_incarnation("test_user", 1, old_ip)
.await .await
.unwrap(); .unwrap();
tracker tracker.clear_user_ips_if_not_newer("test_user", 2).await;
.clear_user_ips_if_not_newer("test_user", 2)
.await;
tracker tracker
.check_and_add_for_incarnation("test_user", 3, current_ip) .check_and_add_for_incarnation("test_user", 3, current_ip)
.await .await
@@ -275,13 +273,7 @@ async fn stale_runtime_cannot_overwrite_newer_ip_policy() {
newer.insert("alice".to_string(), 5); newer.insert("alice".to_string(), 5);
assert!( assert!(
tracker tracker
.apply_policy_from_source( .apply_policy_from_source(2, 7, &newer, UserMaxUniqueIpsMode::Combined, 90,)
2,
7,
&newer,
UserMaxUniqueIpsMode::Combined,
90,
)
.await .await
); );
@@ -289,13 +281,7 @@ async fn stale_runtime_cannot_overwrite_newer_ip_policy() {
stale.insert("alice".to_string(), 1); stale.insert("alice".to_string(), 1);
assert!( assert!(
!tracker !tracker
.apply_policy_from_source( .apply_policy_from_source(1, 1, &stale, UserMaxUniqueIpsMode::ActiveWindow, 1,)
1,
1,
&stale,
UserMaxUniqueIpsMode::ActiveWindow,
1,
)
.await .await
); );
@@ -326,13 +312,7 @@ async fn active_runtime_can_publish_coherent_same_generation_ip_policy() {
assert!( assert!(
tracker tracker
.apply_policy_from_source( .apply_policy_from_source(3, 6, &limits, UserMaxUniqueIpsMode::TimeWindow, 30,)
3,
6,
&limits,
UserMaxUniqueIpsMode::TimeWindow,
30,
)
.await .await
); );
+9 -7
View File
@@ -77,7 +77,10 @@ fn clear_all_serializes_queue_reset_with_concurrent_enqueue() {
Err(std::sync::TryLockError::WouldBlock) => break, Err(std::sync::TryLockError::WouldBlock) => break,
Err(std::sync::TryLockError::Poisoned(_)) => panic!("cleanup queue lock poisoned"), Err(std::sync::TryLockError::Poisoned(_)) => panic!("cleanup queue lock poisoned"),
} }
assert!(Instant::now() < wait_deadline, "clear_all did not reach queue reset"); assert!(
Instant::now() < wait_deadline,
"clear_all did not reach queue reset"
);
std::thread::yield_now(); std::thread::yield_now();
} }
@@ -86,16 +89,15 @@ fn clear_all_serializes_queue_reset_with_concurrent_enqueue() {
let enqueue_tracker = Arc::clone(&tracker); let enqueue_tracker = Arc::clone(&tracker);
let enqueue = std::thread::spawn(move || { let enqueue = std::thread::spawn(move || {
started_tx.send(()).unwrap(); started_tx.send(()).unwrap();
enqueue_tracker.enqueue_cleanup( enqueue_tracker.enqueue_cleanup(first_shard_user, test_ipv4(10, 2, 2, 1));
first_shard_user,
test_ipv4(10, 2, 2, 1),
);
completed_tx.send(()).unwrap(); completed_tx.send(()).unwrap();
}); });
started_rx.recv().unwrap(); started_rx.recv().unwrap();
assert!(completed_rx assert!(
completed_rx
.recv_timeout(Duration::from_millis(50)) .recv_timeout(Duration::from_millis(50))
.is_err()); .is_err()
);
drop(last_queue_guard); drop(last_queue_guard);
clear.join().unwrap(); clear.join().unwrap();
+2 -2
View File
@@ -142,8 +142,8 @@ pub fn init_logging(
} }
LogDestination::File { options } => { LogDestination::File { options } => {
let file_appender = file::BoundedFileAppender::new(options.clone()) let file_appender =
.expect("Failed to open log file"); file::BoundedFileAppender::new(options.clone()).expect("Failed to open log file");
let (non_blocking, guard) = tracing_appender::non_blocking(file_appender); let (non_blocking, guard) = tracing_appender::non_blocking(file_appender);
let fmt_layer = fmt::Layer::default() let fmt_layer = fmt::Layer::default()
+9 -15
View File
@@ -1,6 +1,6 @@
use std::fs::{self, File};
#[cfg(not(unix))] #[cfg(not(unix))]
use std::fs::OpenOptions; use std::fs::OpenOptions;
use std::fs::{self, File};
use std::io::{self, Write}; use std::io::{self, Write};
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use std::time::{Duration, SystemTime, UNIX_EPOCH}; use std::time::{Duration, SystemTime, UNIX_EPOCH};
@@ -223,12 +223,7 @@ impl BoundedFileAppender {
let archive_name = archive_path.file_name().ok_or_else(|| { let archive_name = archive_path.file_name().ok_or_else(|| {
io::Error::new(io::ErrorKind::InvalidInput, "archive path has no file name") io::Error::new(io::ErrorKind::InvalidInput, "archive path has no file name")
})?; })?;
match renameat( match renameat(&self.dir_fd, current_name, &self.dir_fd, archive_name) {
&self.dir_fd,
current_name,
&self.dir_fd,
archive_name,
) {
Ok(()) => Ok(()), Ok(()) => Ok(()),
Err(nix::errno::Errno::ENOENT) => Ok(()), Err(nix::errno::Errno::ENOENT) => Ok(()),
Err(error) => Err(io::Error::from_raw_os_error(error as i32)), Err(error) => Err(io::Error::from_raw_os_error(error as i32)),
@@ -247,11 +242,10 @@ impl BoundedFileAppender {
fn collect_candidates(&self) -> io::Result<Vec<LogFileCandidate>> { fn collect_candidates(&self) -> io::Result<Vec<LogFileCandidate>> {
use std::os::unix::fs::MetadataExt; use std::os::unix::fs::MetadataExt;
let descriptor = dup(&self.dir_fd).map_err(|error| { let descriptor =
io::Error::from_raw_os_error(error as i32) dup(&self.dir_fd).map_err(|error| io::Error::from_raw_os_error(error as i32))?;
})?; let mut directory =
let mut directory = Dir::from_fd(descriptor) Dir::from_fd(descriptor).map_err(|error| io::Error::from_raw_os_error(error as i32))?;
.map_err(|error| io::Error::from_raw_os_error(error as i32))?;
let mut candidates = Vec::new(); let mut candidates = Vec::new();
let prefix = format!("{}.", self.base_name); let prefix = format!("{}.", self.base_name);
for entry in directory.iter().flatten() { for entry in directory.iter().flatten() {
@@ -375,9 +369,9 @@ struct LogFileCandidate {
#[cfg(unix)] #[cfg(unix)]
fn open_append_file(dir_fd: &OwnedFd, path: &Path) -> io::Result<(File, u64)> { fn open_append_file(dir_fd: &OwnedFd, path: &Path) -> io::Result<(File, u64)> {
let name = path.file_name().ok_or_else(|| { let name = path
io::Error::new(io::ErrorKind::InvalidInput, "log path has no file name") .file_name()
})?; .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "log path has no file name"))?;
let file = crate::util::secure_fs::open_append_regular_at(dir_fd, name, 0o640)?; let file = crate::util::secure_fs::open_append_regular_at(dir_fd, name, 0o640)?;
let current_size = file.metadata()?.len(); let current_size = file.metadata()?.len();
Ok((file, current_size)) Ok((file, current_size))
+1 -4
View File
@@ -137,10 +137,7 @@ fn appender_rejects_group_writable_log_directory() {
fs::set_permissions(dir.path(), fs::Permissions::from_mode(0o770)).unwrap(); fs::set_permissions(dir.path(), fs::Permissions::from_mode(0o770)).unwrap();
assert!( assert!(
BoundedFileAppender::with_now( BoundedFileAppender::with_now(options(dir.path().join("telemt.log")), Box::new(fixed_now),)
options(dir.path().join("telemt.log")),
Box::new(fixed_now),
)
.is_err() .is_err()
); );
} }
+4 -3
View File
@@ -137,10 +137,11 @@ mod tests {
let drain_generation = Arc::clone(&generation); let drain_generation = Arc::clone(&generation);
let drain = tokio::spawn(async move { let drain = tokio::spawn(async move {
drain_generation.drain_sessions(Duration::from_secs(60)).await drain_generation
.drain_sessions(Duration::from_secs(60))
.await
}); });
while generation.session_admission.state.load(Ordering::Acquire) while generation.session_admission.state.load(Ordering::Acquire) & SESSION_ADMISSION_CLOSED
& SESSION_ADMISSION_CLOSED
== 0 == 0
{ {
tokio::task::yield_now().await; tokio::task::yield_now().await;
+56 -66
View File
@@ -1,26 +1,26 @@
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use super::{ use super::{
expected_handshake_close_description, format_maestro_line, is_expected_handshake_eof, expected_handshake_close_description, format_maestro_line, is_expected_handshake_eof,
peer_close_description, resolve_runtime_base_dir, resolve_runtime_config_path, peer_close_description, resolve_runtime_base_dir, resolve_runtime_config_path,
}; };
use crate::error::{ProxyError, StreamError}; use crate::error::{ProxyError, StreamError};
#[test] #[test]
fn maestro_line_formatter_respects_disabled_colors() { fn maestro_line_formatter_respects_disabled_colors() {
let plain = format_maestro_line("boot", false); let plain = format_maestro_line("boot", false);
assert_eq!(plain, "MAESTRO: boot"); assert_eq!(plain, "MAESTRO: boot");
assert!(!plain.contains('\x1b')); assert!(!plain.contains('\x1b'));
} }
#[test] #[test]
fn maestro_line_formatter_keeps_color_when_enabled() { fn maestro_line_formatter_keeps_color_when_enabled() {
let colored = format_maestro_line("boot", true); let colored = format_maestro_line("boot", true);
assert!(colored.contains("\x1b[92mMAESTRO\x1b[0m")); assert!(colored.contains("\x1b[92mMAESTRO\x1b[0m"));
} }
#[test] #[test]
fn resolve_runtime_config_path_anchors_relative_to_startup_cwd() { fn resolve_runtime_config_path_anchors_relative_to_startup_cwd() {
let nonce = std::time::SystemTime::now() let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH) .duration_since(std::time::UNIX_EPOCH)
.unwrap() .unwrap()
@@ -35,10 +35,10 @@
let _ = std::fs::remove_file(&target); let _ = std::fs::remove_file(&target);
let _ = std::fs::remove_dir(&startup_cwd); let _ = std::fs::remove_dir(&startup_cwd);
} }
#[test] #[test]
fn resolve_runtime_config_path_keeps_absolute_for_missing_file() { fn resolve_runtime_config_path_keeps_absolute_for_missing_file() {
let nonce = std::time::SystemTime::now() let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH) .duration_since(std::time::UNIX_EPOCH)
.unwrap() .unwrap()
@@ -50,11 +50,11 @@
assert_eq!(resolved, startup_cwd.join("missing.toml")); assert_eq!(resolved, startup_cwd.join("missing.toml"));
let _ = std::fs::remove_dir(&startup_cwd); let _ = std::fs::remove_dir(&startup_cwd);
} }
#[cfg(unix)] #[cfg(unix)]
#[test] #[test]
fn runtime_paths_preserve_symlinks_for_descriptor_validation() { fn runtime_paths_preserve_symlinks_for_descriptor_validation() {
use std::os::unix::fs::symlink; use std::os::unix::fs::symlink;
let dir = tempfile::tempdir().unwrap(); let dir = tempfile::tempdir().unwrap();
@@ -65,30 +65,20 @@
symlink(&real_dir, &linked_dir).unwrap(); symlink(&real_dir, &linked_dir).unwrap();
let linked_config = linked_dir.join("config.toml"); let linked_config = linked_dir.join("config.toml");
let config = resolve_runtime_config_path( let config = resolve_runtime_config_path(linked_config.to_str().unwrap(), dir.path(), true);
linked_config.to_str().unwrap(), let runtime = resolve_runtime_base_dir(&linked_config, dir.path(), true, Some(&linked_dir));
dir.path(),
true,
);
let runtime = resolve_runtime_base_dir(
&linked_config,
dir.path(),
true,
Some(&linked_dir),
);
assert_eq!(config, linked_config); assert_eq!(config, linked_config);
assert_eq!(runtime, linked_dir); assert_eq!(runtime, linked_dir);
} }
#[test] #[test]
fn resolve_runtime_config_path_uses_startup_candidates_when_not_explicit() { fn resolve_runtime_config_path_uses_startup_candidates_when_not_explicit() {
let nonce = std::time::SystemTime::now() let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH) .duration_since(std::time::UNIX_EPOCH)
.unwrap() .unwrap()
.as_nanos(); .as_nanos();
let startup_cwd = let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_startup_candidates_{nonce}"));
std::env::temp_dir().join(format!("telemt_cfg_startup_candidates_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap(); std::fs::create_dir_all(&startup_cwd).unwrap();
let telemt = startup_cwd.join("telemt.toml"); let telemt = startup_cwd.join("telemt.toml");
std::fs::write(&telemt, " ").unwrap(); std::fs::write(&telemt, " ").unwrap();
@@ -98,10 +88,10 @@
let _ = std::fs::remove_file(&telemt); let _ = std::fs::remove_file(&telemt);
let _ = std::fs::remove_dir(&startup_cwd); let _ = std::fs::remove_dir(&startup_cwd);
} }
#[test] #[test]
fn resolve_runtime_config_path_defaults_to_startup_config_when_none_found() { fn resolve_runtime_config_path_defaults_to_startup_config_when_none_found() {
let nonce = std::time::SystemTime::now() let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH) .duration_since(std::time::UNIX_EPOCH)
.unwrap() .unwrap()
@@ -113,10 +103,10 @@
assert_eq!(resolved, startup_cwd.join("config.toml")); assert_eq!(resolved, startup_cwd.join("config.toml"));
let _ = std::fs::remove_dir(&startup_cwd); let _ = std::fs::remove_dir(&startup_cwd);
} }
#[test] #[test]
fn resolve_runtime_base_dir_prefers_cli_data_path() { fn resolve_runtime_base_dir_prefers_cli_data_path() {
let nonce = std::time::SystemTime::now() let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH) .duration_since(std::time::UNIX_EPOCH)
.unwrap() .unwrap()
@@ -136,10 +126,10 @@
let _ = std::fs::remove_dir(&data_path); let _ = std::fs::remove_dir(&data_path);
let _ = std::fs::remove_dir(&startup_cwd); let _ = std::fs::remove_dir(&startup_cwd);
} }
#[test] #[test]
fn resolve_runtime_base_dir_uses_working_directory_before_explicit_config_parent() { fn resolve_runtime_base_dir_uses_working_directory_before_explicit_config_parent() {
let nonce = std::time::SystemTime::now() let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH) .duration_since(std::time::UNIX_EPOCH)
.unwrap() .unwrap()
@@ -155,10 +145,10 @@
let _ = std::fs::remove_dir(&config_dir); let _ = std::fs::remove_dir(&config_dir);
let _ = std::fs::remove_dir(&startup_cwd); let _ = std::fs::remove_dir(&startup_cwd);
} }
#[test] #[test]
fn resolve_runtime_base_dir_uses_explicit_config_parent_from_root() { fn resolve_runtime_base_dir_uses_explicit_config_parent_from_root() {
let nonce = std::time::SystemTime::now() let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH) .duration_since(std::time::UNIX_EPOCH)
.unwrap() .unwrap()
@@ -171,10 +161,10 @@
assert_eq!(resolved, config_dir.canonicalize().unwrap()); assert_eq!(resolved, config_dir.canonicalize().unwrap());
let _ = std::fs::remove_dir(&config_dir); let _ = std::fs::remove_dir(&config_dir);
} }
#[test] #[test]
fn resolve_runtime_base_dir_uses_systemd_working_directory_before_etc() { fn resolve_runtime_base_dir_uses_systemd_working_directory_before_etc() {
let nonce = std::time::SystemTime::now() let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH) .duration_since(std::time::UNIX_EPOCH)
.unwrap() .unwrap()
@@ -187,10 +177,10 @@
assert_eq!(resolved, startup_cwd.canonicalize().unwrap()); assert_eq!(resolved, startup_cwd.canonicalize().unwrap());
let _ = std::fs::remove_dir(&startup_cwd); let _ = std::fs::remove_dir(&startup_cwd);
} }
#[test] #[test]
fn resolve_runtime_base_dir_falls_back_to_etc_from_root() { fn resolve_runtime_base_dir_falls_back_to_etc_from_root() {
let resolved = resolve_runtime_base_dir( let resolved = resolve_runtime_base_dir(
Path::new("/etc/telemt/config.toml"), Path::new("/etc/telemt/config.toml"),
Path::new("/"), Path::new("/"),
@@ -198,24 +188,24 @@
None, None,
); );
assert_eq!(resolved, PathBuf::from("/etc/telemt")); assert_eq!(resolved, PathBuf::from("/etc/telemt"));
} }
#[test] #[test]
fn expected_handshake_eof_matches_connection_reset() { fn expected_handshake_eof_matches_connection_reset() {
let err = ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionReset)); let err = ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionReset));
assert!(is_expected_handshake_eof(&err)); assert!(is_expected_handshake_eof(&err));
} }
#[test] #[test]
fn expected_handshake_eof_matches_stream_io_unexpected_eof() { fn expected_handshake_eof_matches_stream_io_unexpected_eof() {
let err = ProxyError::Stream(StreamError::Io(std::io::Error::from( let err = ProxyError::Stream(StreamError::Io(std::io::Error::from(
std::io::ErrorKind::UnexpectedEof, std::io::ErrorKind::UnexpectedEof,
))); )));
assert!(is_expected_handshake_eof(&err)); assert!(is_expected_handshake_eof(&err));
} }
#[test] #[test]
fn peer_close_description_is_human_readable_for_all_peer_close_kinds() { fn peer_close_description_is_human_readable_for_all_peer_close_kinds() {
let cases = [ let cases = [
( (
std::io::ErrorKind::ConnectionReset, std::io::ErrorKind::ConnectionReset,
@@ -239,10 +229,10 @@
let err = ProxyError::Io(std::io::Error::from(kind)); let err = ProxyError::Io(std::io::Error::from(kind));
assert_eq!(peer_close_description(&err), Some(expected)); assert_eq!(peer_close_description(&err), Some(expected));
} }
} }
#[test] #[test]
fn handshake_close_description_is_human_readable_for_all_expected_kinds() { fn handshake_close_description_is_human_readable_for_all_expected_kinds() {
let cases = [ let cases = [
( (
ProxyError::Io(std::io::Error::from(std::io::ErrorKind::UnexpectedEof)), ProxyError::Io(std::io::Error::from(std::io::ErrorKind::UnexpectedEof)),
@@ -273,4 +263,4 @@
for (err, expected) in cases { for (err, expected) in cases {
assert_eq!(expected_handshake_close_description(&err), Some(expected)); assert_eq!(expected_handshake_close_description(&err), Some(expected));
} }
} }
+10 -2
View File
@@ -234,7 +234,10 @@ fn remove_stale_unix_socket(path: &Path) -> std::io::Result<()> {
{ {
return Err(IoError::new( return Err(IoError::new(
ErrorKind::AlreadyExists, ErrorKind::AlreadyExists,
format!("Unix listener path {} changed during cleanup", path.display()), format!(
"Unix listener path {} changed during cleanup",
path.display()
),
)); ));
} }
std::fs::remove_file(path) std::fs::remove_file(path)
@@ -373,7 +376,12 @@ mod tests {
assert!(remove_stale_unix_socket(&regular).is_err()); assert!(remove_stale_unix_socket(&regular).is_err());
assert!(remove_stale_unix_socket(&link).is_err()); assert!(remove_stale_unix_socket(&link).is_err());
assert_eq!(std::fs::read(&regular).unwrap(), b"preserve"); assert_eq!(std::fs::read(&regular).unwrap(), b"preserve");
assert!(std::fs::symlink_metadata(&link).unwrap().file_type().is_symlink()); assert!(
std::fs::symlink_metadata(&link)
.unwrap()
.file_type()
.is_symlink()
);
} }
#[test] #[test]
+1 -3
View File
@@ -49,9 +49,7 @@ pub(super) async fn run_telemt_core(
} = bootstrap::bootstrap(privilege_drop_requested).await?; } = bootstrap::bootstrap(privilege_drop_requested).await?;
if privilege_drop_requested && config.server.conntrack_control.inline_conntrack_control { if privilege_drop_requested && config.server.conntrack_control.inline_conntrack_control {
warn!( warn!("Inline conntrack control is disabled when process privileges are dropped");
"Inline conntrack control is disabled when process privileges are dropped"
);
config.server.conntrack_control.inline_conntrack_control = false; config.server.conntrack_control.inline_conntrack_control = false;
} }
+2 -2
View File
@@ -366,8 +366,8 @@ impl ReloadSupervisor {
self.runtime_watch_tx self.runtime_watch_tx
.send_replace(Some(new_runtime.watch_state())); .send_replace(Some(new_runtime.watch_state()));
if !conntrack_firewall_published { if !conntrack_firewall_published {
let warning = "conntrack firewall reconciler is unavailable after runtime activation" let warning =
.to_string(); "conntrack firewall reconciler is unavailable after runtime activation".to_string();
warn!(reload_id = command.reload_id, warning = %warning); warn!(reload_id = command.reload_id, warning = %warning);
self.control.add_warning(command.reload_id, warning).await; self.control.add_warning(command.reload_id, warning).await;
} }
+2 -6
View File
@@ -11,9 +11,7 @@ use crate::config::{
use crate::crypto::SecureRandom; use crate::crypto::SecureRandom;
use crate::ip_tracker::UserIpTracker; use crate::ip_tracker::UserIpTracker;
use crate::network::probe::{decide_network_capabilities, run_probe}; use crate::network::probe::{decide_network_capabilities, run_probe};
use crate::proxy::direct_buffer_budget::{ use crate::proxy::direct_buffer_budget::{DirectBufferBudget, run_direct_buffer_budget_controller};
DirectBufferBudget, run_direct_buffer_budget_controller,
};
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::traffic_limiter::TrafficLimiter; use crate::proxy::traffic_limiter::TrafficLimiter;
@@ -29,9 +27,7 @@ use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool; use crate::transport::middle_proxy::MePool;
use super::admission; use super::admission;
use super::generation::{ use super::generation::{RuntimeGeneration, RuntimeTaskScope, RuntimeTaskScopePreparationGuard};
RuntimeGeneration, RuntimeTaskScope, RuntimeTaskScopePreparationGuard,
};
use super::listeners::listener_rebind_supported; use super::listeners::listener_rebind_supported;
use super::runtime_tasks::RuntimeLogFilter; use super::runtime_tasks::RuntimeLogFilter;
use super::{me_startup, runtime_tasks, tls_bootstrap}; use super::{me_startup, runtime_tasks, tls_bootstrap};
+3 -8
View File
@@ -135,14 +135,9 @@ fn conntrack_control_policy_is_restart_only_as_one_process_owned_unit() {
!old.server.conntrack_control.inline_conntrack_control; !old.server.conntrack_control.inline_conntrack_control;
desired.server.conntrack_control.mode = crate::config::ConntrackMode::Notrack; desired.server.conntrack_control.mode = crate::config::ConntrackMode::Notrack;
desired.server.conntrack_control.backend = crate::config::ConntrackBackend::Iptables; desired.server.conntrack_control.backend = crate::config::ConntrackBackend::Iptables;
desired.server.conntrack_control.profile = desired.server.conntrack_control.profile = crate::config::ConntrackPressureProfile::Aggressive;
crate::config::ConntrackPressureProfile::Aggressive; desired.server.conntrack_control.hybrid_listener_ips = vec!["192.0.2.10".parse().unwrap()];
desired.server.conntrack_control.hybrid_listener_ips = desired.server.conntrack_control.pressure_high_watermark_pct = 90;
vec!["192.0.2.10".parse().unwrap()];
desired
.server
.conntrack_control
.pressure_high_watermark_pct = 90;
desired.server.conntrack_control.pressure_low_watermark_pct = 40; desired.server.conntrack_control.pressure_low_watermark_pct = 40;
desired.server.conntrack_control.delete_budget_per_sec = old desired.server.conntrack_control.delete_budget_per_sec = old
.server .server
+3 -11
View File
@@ -3,11 +3,7 @@ use std::fmt::Write;
use crate::transport::middle_proxy::MeApiHardswapSnapshot; use crate::transport::middle_proxy::MeApiHardswapSnapshot;
/// Renders fixed-cardinality hardswap and writer-replacement gauges. /// Renders fixed-cardinality hardswap and writer-replacement gauges.
pub(super) fn render( pub(super) fn render(out: &mut String, snapshot: Option<&MeApiHardswapSnapshot>, enabled: bool) {
out: &mut String,
snapshot: Option<&MeApiHardswapSnapshot>,
enabled: bool,
) {
let snapshot = enabled.then_some(snapshot).flatten(); let snapshot = enabled.then_some(snapshot).flatten();
let pending = snapshot.is_some_and(|value| value.pending); let pending = snapshot.is_some_and(|value| value.pending);
let pending_age_secs = snapshot let pending_age_secs = snapshot
@@ -147,11 +143,7 @@ mod tests {
assert!(out.contains("telemt_me_hardswap_pending 1")); assert!(out.contains("telemt_me_hardswap_pending 1"));
assert!(out.contains("telemt_me_hardswap_pending_age_seconds 42")); assert!(out.contains("telemt_me_hardswap_pending_age_seconds 42"));
assert!(out.contains("telemt_me_hardswap_pending_writer_deficit 4")); assert!(out.contains("telemt_me_hardswap_pending_writer_deficit 4"));
assert!(out.contains( assert!(out.contains("telemt_me_writer_replacement_current{state=\"preparing\"} 5"));
"telemt_me_writer_replacement_current{state=\"preparing\"} 5" assert!(out.contains("telemt_me_writer_replacement_current{state=\"retiring\"} 6"));
));
assert!(out.contains(
"telemt_me_writer_replacement_current{state=\"retiring\"} 6"
));
} }
} }
+16 -5
View File
@@ -122,22 +122,33 @@ pub(super) fn render(
out, out,
"# HELP telemt_rate_limiter_cas_retry_exhausted_total Traffic limiter operations that exhausted their bounded CAS attempt budget" "# HELP telemt_rate_limiter_cas_retry_exhausted_total Traffic limiter operations that exhausted their bounded CAS attempt budget"
); );
let _ = writeln!(out, "# TYPE telemt_rate_limiter_cas_retry_exhausted_total counter"); let _ = writeln!(
out,
"# TYPE telemt_rate_limiter_cas_retry_exhausted_total counter"
);
for (scope, direction, reserve, refund) in [ for (scope, direction, reserve, refund) in [
( (
"user", "up", limiter_metrics.user_reserve_cas_retry_exhausted_up_total, "user",
"up",
limiter_metrics.user_reserve_cas_retry_exhausted_up_total,
limiter_metrics.user_refund_cas_retry_exhausted_up_total, limiter_metrics.user_refund_cas_retry_exhausted_up_total,
), ),
( (
"user", "down", limiter_metrics.user_reserve_cas_retry_exhausted_down_total, "user",
"down",
limiter_metrics.user_reserve_cas_retry_exhausted_down_total,
limiter_metrics.user_refund_cas_retry_exhausted_down_total, limiter_metrics.user_refund_cas_retry_exhausted_down_total,
), ),
( (
"cidr", "up", limiter_metrics.cidr_reserve_cas_retry_exhausted_up_total, "cidr",
"up",
limiter_metrics.cidr_reserve_cas_retry_exhausted_up_total,
limiter_metrics.cidr_refund_cas_retry_exhausted_up_total, limiter_metrics.cidr_refund_cas_retry_exhausted_up_total,
), ),
( (
"cidr", "down", limiter_metrics.cidr_reserve_cas_retry_exhausted_down_total, "cidr",
"down",
limiter_metrics.cidr_reserve_cas_retry_exhausted_down_total,
limiter_metrics.cidr_refund_cas_retry_exhausted_down_total, limiter_metrics.cidr_refund_cas_retry_exhausted_down_total,
), ),
] { ] {
+1 -3
View File
@@ -127,9 +127,7 @@ async fn test_render_metrics_format() {
); );
assert!(output.contains("telemt_handshake_timeouts_total 1")); assert!(output.contains("telemt_handshake_timeouts_total 1"));
assert!(output.contains("telemt_handshake_failures_by_class_total{class=\"timeout\"} 1")); assert!(output.contains("telemt_handshake_failures_by_class_total{class=\"timeout\"} 1"));
assert!( assert!(output.contains("telemt_conntrack_rule_reconcile_total{result=\"success\"} 1"));
output.contains("telemt_conntrack_rule_reconcile_total{result=\"success\"} 1")
);
assert!(output.contains("telemt_conntrack_rule_reconcile_total{result=\"error\"} 1")); assert!(output.contains("telemt_conntrack_rule_reconcile_total{result=\"error\"} 1"));
assert!(output.contains("telemt_conntrack_rule_rollback_total{result=\"success\"} 1")); assert!(output.contains("telemt_conntrack_rule_rollback_total{result=\"success\"} 1"));
assert!(output.contains("telemt_conntrack_rule_rollback_total{result=\"error\"} 1")); assert!(output.contains("telemt_conntrack_rule_rollback_total{result=\"error\"} 1"));
+5 -18
View File
@@ -290,11 +290,8 @@ impl Drop for UserIpPermit {
let Some(owner) = self.owner.take() else { let Some(owner) = self.owner.take() else {
return; return;
}; };
self.tracker.enqueue_cleanup_for_incarnation( self.tracker
owner.user, .enqueue_cleanup_for_incarnation(owner.user, owner.incarnation, owner.ip);
owner.incarnation,
owner.ip,
);
} }
} }
@@ -338,9 +335,7 @@ impl UserConnectionReservation {
stats_observation: Option<UserConnectionObservation>, stats_observation: Option<UserConnectionObservation>,
tracks_ip: bool, tracks_ip: bool,
) -> Self { ) -> Self {
let ip_permit = tracks_ip.then(|| { let ip_permit = tracks_ip.then(|| UserIpPermit::new(ip_tracker, user, incarnation, ip));
UserIpPermit::new(ip_tracker, user, incarnation, ip)
});
Self { Self {
stats, stats,
quota_handle, quota_handle,
@@ -388,12 +383,7 @@ pub(crate) async fn acquire_user_connection_reservation(
ip_tracker: Arc<UserIpTracker>, ip_tracker: Arc<UserIpTracker>,
) -> Result<UserConnectionReservation> { ) -> Result<UserConnectionReservation> {
acquire_user_connection_reservation_for_incarnation( acquire_user_connection_reservation_for_incarnation(
user, user, 0, config, stats, peer_addr, ip_tracker,
0,
config,
stats,
peer_addr,
ip_tracker,
) )
.await .await
} }
@@ -435,10 +425,7 @@ async fn acquire_user_connection_reservation_for_incarnation(
.or((config.access.user_max_tcp_conns_global_each > 0) .or((config.access.user_max_tcp_conns_global_each > 0)
.then_some(config.access.user_max_tcp_conns_global_each)) .then_some(config.access.user_max_tcp_conns_global_each))
.map(|value| value as u64); .map(|value| value as u64);
let Some(connection_permit) = stats let Some(connection_permit) = stats.connection_authority().try_acquire(user, limit) else {
.connection_authority()
.try_acquire(user, limit)
else {
return Err(ProxyError::ConnectionLimitExceeded { return Err(ProxyError::ConnectionLimitExceeded {
user: user.to_string(), user: user.to_string(),
}); });
+1 -4
View File
@@ -155,10 +155,7 @@ impl RunningClientHandler {
.or((config.access.user_max_tcp_conns_global_each > 0) .or((config.access.user_max_tcp_conns_global_each > 0)
.then_some(config.access.user_max_tcp_conns_global_each)) .then_some(config.access.user_max_tcp_conns_global_each))
.map(|v| v as u64); .map(|v| v as u64);
let Some(_connection_permit) = stats let Some(_connection_permit) = stats.connection_authority().try_acquire(user, limit) else {
.connection_authority()
.try_acquire(user, limit)
else {
return Err(ProxyError::ConnectionLimitExceeded { return Err(ProxyError::ConnectionLimitExceeded {
user: user.to_string(), user: user.to_string(),
}); });
+3 -6
View File
@@ -7,11 +7,11 @@ use tokio::sync::watch;
// Process controller and system-memory sampling remain outside data-plane accounting. // Process controller and system-memory sampling remain outside data-plane accounting.
mod controller; mod controller;
#[cfg(test)]
use controller::connection_fill_pct;
pub(crate) use controller::{ pub(crate) use controller::{
resolve_direct_buffer_hard_limit, run_direct_buffer_budget_controller, resolve_direct_buffer_hard_limit, run_direct_buffer_budget_controller,
}; };
#[cfg(test)]
use controller::connection_fill_pct;
/// Accounting granularity for process-wide Direct copy-buffer reservations. /// Accounting granularity for process-wide Direct copy-buffer reservations.
pub(crate) const DIRECT_BUFFER_UNIT_BYTES: usize = 4 * 1024; pub(crate) const DIRECT_BUFFER_UNIT_BYTES: usize = 4 * 1024;
@@ -135,10 +135,7 @@ impl DirectBufferBudget {
.fetch_max(generation, Ordering::AcqRel); .fetch_max(generation, Ordering::AcqRel);
} }
fn begin_controller_update( fn begin_controller_update(&self, generation: u64) -> Option<ParkingMutexGuard<'_, ()>> {
&self,
generation: u64,
) -> Option<ParkingMutexGuard<'_, ()>> {
let controller_update = self.controller_update.lock(); let controller_update = self.controller_update.lock();
(self.active_controller_generation.load(Ordering::Acquire) == generation) (self.active_controller_generation.load(Ordering::Acquire) == generation)
.then_some(controller_update) .then_some(controller_update)
+2 -5
View File
@@ -145,11 +145,8 @@ pub(super) fn connection_fill_pct(
return None; return None;
} }
let max_connections = max_connections as usize; let max_connections = max_connections as usize;
let active = max_connections.saturating_sub( let active =
connection_slots max_connections.saturating_sub(connection_slots.available_permits().min(max_connections));
.available_permits()
.min(max_connections),
);
Some((active.saturating_mul(100) / max_connections).min(100) as u8) Some((active.saturating_mul(100) / max_connections).min(100) as u8)
} }
+2 -2
View File
@@ -74,8 +74,8 @@ pub(crate) use self::auth_probe::{
auth_probe_saturation_is_throttled_at_for_testing_in_shared, auth_probe_saturation_is_throttled_at_for_testing_in_shared,
auth_probe_saturation_is_throttled_for_testing_in_shared, auth_probe_saturation_is_throttled_for_testing_in_shared,
auth_probe_saturation_state_for_testing_in_shared, auth_probe_saturation_state_for_testing_in_shared,
auth_probe_saturation_state_lock_for_testing_in_shared, auth_probe_state_for_testing_in_shared, auth_probe_saturation_state_lock_for_testing_in_shared, auth_probe_slots_for_testing_in_shared,
auth_probe_slots_for_testing_in_shared, clear_auth_probe_state_for_testing_in_shared, auth_probe_state_for_testing_in_shared, clear_auth_probe_state_for_testing_in_shared,
clear_unknown_sni_warn_state_for_testing_in_shared, clear_warned_secrets_for_testing_in_shared, clear_unknown_sni_warn_state_for_testing_in_shared, clear_warned_secrets_for_testing_in_shared,
insert_auth_probe_state_for_testing_in_shared, insert_auth_probe_state_for_testing_in_shared,
should_emit_unknown_sni_warn_for_testing_in_shared, warned_secrets_for_testing_in_shared, should_emit_unknown_sni_warn_for_testing_in_shared, warned_secrets_for_testing_in_shared,
+8 -2
View File
@@ -403,7 +403,10 @@ mod bounded_registry_tests {
for index in (worker..ATTEMPTS).step_by(16) { for index in (worker..ATTEMPTS).step_by(16) {
let octets = (index as u32).to_be_bytes(); let octets = (index as u32).to_be_bytes();
let peer_ip = IpAddr::V4(std::net::Ipv4Addr::new( let peer_ip = IpAddr::V4(std::net::Ipv4Addr::new(
octets[1], octets[2], octets[3], worker as u8, octets[1],
octets[2],
octets[3],
worker as u8,
)); ));
sticky_hint_record_success_in( sticky_hint_record_success_in(
shared.as_ref(), shared.as_ref(),
@@ -416,7 +419,10 @@ mod bounded_registry_tests {
} }
}); });
assert_eq!(shared.handshake.sticky_user_by_ip.len(), STICKY_HINT_MAX_ENTRIES); assert_eq!(
shared.handshake.sticky_user_by_ip.len(),
STICKY_HINT_MAX_ENTRIES
);
assert_eq!( assert_eq!(
shared.handshake.sticky_user_by_ip_prefix.len(), shared.handshake.sticky_user_by_ip_prefix.len(),
STICKY_HINT_MAX_ENTRIES STICKY_HINT_MAX_ENTRIES
+1 -2
View File
@@ -397,8 +397,7 @@ fn auth_probe_record_failure_with_state_and_budget_in(
}; };
if state if state
.remove_if(&evict_key, |_, current| { .remove_if(&evict_key, |_, current| {
current.fail_streak == evict_fail_streak current.fail_streak == evict_fail_streak && current.last_seen == evict_last_seen
&& current.last_seen == evict_last_seen
}) })
.is_some() .is_some()
&& let Some(slots) = slots && let Some(slots) = slots
+8 -2
View File
@@ -160,7 +160,10 @@ fn parallel_distinct_failures_respect_exact_auth_probe_capacity() {
for index in (worker..ATTEMPTS).step_by(16) { for index in (worker..ATTEMPTS).step_by(16) {
let octets = (index as u32).to_be_bytes(); let octets = (index as u32).to_be_bytes();
let peer_ip = IpAddr::V4(std::net::Ipv4Addr::new( let peer_ip = IpAddr::V4(std::net::Ipv4Addr::new(
octets[1], octets[2], octets[3], worker as u8, octets[1],
octets[2],
octets[3],
worker as u8,
)); ));
auth_probe_record_failure_in(shared.as_ref(), peer_ip, Instant::now()); auth_probe_record_failure_in(shared.as_ref(), peer_ip, Instant::now());
} }
@@ -168,7 +171,10 @@ fn parallel_distinct_failures_respect_exact_auth_probe_capacity() {
} }
}); });
assert_eq!(shared.handshake.auth_probe.len(), AUTH_PROBE_TRACK_MAX_ENTRIES); assert_eq!(
shared.handshake.auth_probe.len(),
AUTH_PROBE_TRACK_MAX_ENTRIES
);
assert_eq!( assert_eq!(
auth_probe_slots_for_testing_in_shared(shared.as_ref()), auth_probe_slots_for_testing_in_shared(shared.as_ref()),
AUTH_PROBE_TRACK_MAX_ENTRIES AUTH_PROBE_TRACK_MAX_ENTRIES
+4 -4
View File
@@ -151,10 +151,10 @@ where
if let Some(snapshot) = config.runtime_user_auth() { if let Some(snapshot) = config.runtime_user_auth() {
let sticky_ip_hint = sticky_hint_get_by_ip(shared, peer.ip()); let sticky_ip_hint = sticky_hint_get_by_ip(shared, peer.ip());
let sticky_prefix_hint = sticky_hint_get_by_ip_prefix(shared, peer.ip()); let sticky_prefix_hint = sticky_hint_get_by_ip_prefix(shared, peer.ip());
let sticky_ip_candidates = sticky_ip_hint let sticky_ip_candidates =
.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key)); sticky_ip_hint.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key));
let sticky_prefix_candidates = sticky_prefix_hint let sticky_prefix_candidates =
.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key)); sticky_prefix_hint.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key));
let preferred_user_id = preferred_user.and_then(|user| snapshot.user_id_by_name(user)); let preferred_user_id = preferred_user.and_then(|user| snapshot.user_id_by_name(user));
let exact_user_id = exact_user.and_then(|user| snapshot.user_id_by_name(user)); let exact_user_id = exact_user.and_then(|user| snapshot.user_id_by_name(user));
let has_hint = sticky_ip_candidates.is_some_and(|ids| !ids.is_empty()) let has_hint = sticky_ip_candidates.is_some_and(|ids| !ids.is_empty())
+1 -6
View File
@@ -400,12 +400,7 @@ where
.runtime_user_auth() .runtime_user_auth()
.and_then(|snapshot| snapshot.entry_by_id(user_id)) .and_then(|snapshot| snapshot.entry_by_id(user_id))
{ {
sticky_hint_record_success_in( sticky_hint_record_success_in(shared, peer.ip(), entry.hint_key, client_sni.as_deref());
shared,
peer.ip(),
entry.hint_key,
client_sni.as_deref(),
);
record_recent_user_success_in(shared, entry.hint_key); record_recent_user_success_in(shared, entry.hint_key);
} }
} }
+6 -6
View File
@@ -40,17 +40,17 @@ pub(super) async fn validate_tls_client(
}; };
let sticky_ip_hint = sticky_hint_get_by_ip(shared, peer.ip()); let sticky_ip_hint = sticky_hint_get_by_ip(shared, peer.ip());
let sticky_ip_candidates = sticky_ip_hint let sticky_ip_candidates =
.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key)); sticky_ip_hint.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key));
let preferred_user_id = preferred_user_hint.and_then(|user| snapshot.user_id_by_name(user)); let preferred_user_id = preferred_user_hint.and_then(|user| snapshot.user_id_by_name(user));
let sticky_sni_hint = client_sni let sticky_sni_hint = client_sni
.as_deref() .as_deref()
.and_then(|sni| sticky_hint_get_by_sni(shared, sni)); .and_then(|sni| sticky_hint_get_by_sni(shared, sni));
let sticky_sni_candidates = sticky_sni_hint let sticky_sni_candidates =
.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key)); sticky_sni_hint.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key));
let sticky_prefix_hint = sticky_hint_get_by_ip_prefix(shared, peer.ip()); let sticky_prefix_hint = sticky_hint_get_by_ip_prefix(shared, peer.ip());
let sticky_prefix_candidates = sticky_prefix_hint let sticky_prefix_candidates =
.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key)); sticky_prefix_hint.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key));
let sni_candidates = client_sni let sni_candidates = client_sni
.as_deref() .as_deref()
.and_then(|sni| snapshot.sni_candidates(sni)); .and_then(|sni| snapshot.sni_candidates(sni));
+6 -1
View File
@@ -191,7 +191,12 @@ where
if let (Some(limit), Some(quota_handle)) = (quota_limit, quota_handle) { 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(
quota_handle, data_len, soft_limit, stats, cancel, None, quota_handle,
data_len,
soft_limit,
stats,
cancel,
None,
) )
.await .await
{ {
+3 -3
View File
@@ -44,9 +44,9 @@ mod tests {
c2me_sender: AbortOnDropHandle::new(tokio::spawn(pending_child(DropSignal( c2me_sender: AbortOnDropHandle::new(tokio::spawn(pending_child(DropSignal(
Arc::clone(&dropped), Arc::clone(&dropped),
)))), )))),
me_writer: AbortOnDropHandle::new(tokio::spawn(pending_child(DropSignal( me_writer: AbortOnDropHandle::new(tokio::spawn(pending_child(DropSignal(Arc::clone(
Arc::clone(&dropped), &dropped,
)))), ))))),
flow_cancel: flow_cancel.clone(), flow_cancel: flow_cancel.clone(),
stop_tx: Some(stop_tx), stop_tx: Some(stop_tx),
}; };
+1 -7
View File
@@ -414,13 +414,7 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
if quota_reservation.is_none() { if quota_reservation.is_none() {
this.stats.increment_quota_contention_timeout_total(); this.stats.increment_quota_contention_timeout_total();
Self::arm_wait(&mut this.quota_wait, false, false); Self::arm_wait(&mut this.quota_wait, false, false);
if Self::poll_wait( if Self::poll_wait(&mut this.quota_wait, cx, None, RateDirection::Up).is_ready()
&mut this.quota_wait,
cx,
None,
RateDirection::Up,
)
.is_ready()
{ {
cx.waker().wake_by_ref(); cx.waker().wake_by_ref();
} }
+7 -4
View File
@@ -219,8 +219,12 @@ impl ProxySharedState {
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.activate_config_source(
.activate_config_source(source_generation, expected_epoch, users, user_enabled) source_generation,
expected_epoch,
users,
user_enabled,
)
} }
/// Applies an update only from the active runtime generation. /// Applies an update only from the active runtime generation.
@@ -276,8 +280,7 @@ impl ProxySharedState {
user: &str, user: &str,
credential_id: UserCredentialId, credential_id: UserCredentialId,
) -> Option<UserAdmissionPublication<'_>> { ) -> Option<UserAdmissionPublication<'_>> {
self.user_admission self.user_admission.claim_authenticated(user, credential_id)
.claim_authenticated(user, credential_id)
} }
pub(crate) fn register_user_session( pub(crate) fn register_user_session(
+2 -8
View File
@@ -310,10 +310,7 @@ async fn cancelled_ip_admission_releases_process_connection_permit() {
let user = "cancelled-admission-user"; let user = "cancelled-admission-user";
let peer_addr: SocketAddr = "198.51.100.210:50000".parse().unwrap(); let peer_addr: SocketAddr = "198.51.100.210:50000".parse().unwrap();
let mut config = ProxyConfig::default(); let mut config = ProxyConfig::default();
config config.access.user_max_tcp_conns.insert(user.to_string(), 1);
.access
.user_max_tcp_conns
.insert(user.to_string(), 1);
let (entered_tx, entered_rx) = tokio::sync::oneshot::channel(); let (entered_tx, entered_rx) = tokio::sync::oneshot::channel();
let (release_tx, release_rx) = tokio::sync::oneshot::channel(); let (release_tx, release_rx) = tokio::sync::oneshot::channel();
@@ -370,10 +367,7 @@ async fn cancelled_async_release_preserves_ip_cleanup_ownership() {
let user = "cancelled-release-user"; let user = "cancelled-release-user";
let peer_addr: SocketAddr = "198.51.100.211:50001".parse().unwrap(); let peer_addr: SocketAddr = "198.51.100.211:50001".parse().unwrap();
let mut config = ProxyConfig::default(); let mut config = ProxyConfig::default();
config config.access.user_max_tcp_conns.insert(user.to_string(), 1);
.access
.user_max_tcp_conns
.insert(user.to_string(), 1);
let reservation = acquire_user_connection_reservation( let reservation = acquire_user_connection_reservation(
user, user,
@@ -77,9 +77,11 @@ fn controller_handoff_waits_for_inflight_update_and_fences_old_generation() {
activated_tx.send(()).unwrap(); activated_tx.send(()).unwrap();
}); });
assert!(activated_rx assert!(
activated_rx
.recv_timeout(Duration::from_millis(50)) .recv_timeout(Duration::from_millis(50))
.is_err()); .is_err()
);
drop(update); drop(update);
activated_rx.recv_timeout(Duration::from_secs(1)).unwrap(); activated_rx.recv_timeout(Duration::from_secs(1)).unwrap();
activation.join().unwrap(); activation.join().unwrap();
@@ -81,10 +81,8 @@ fn adversarial_intermediate_parent_swap_is_blocked_by_component_walk() {
let parent = directory.path().join("parent"); let parent = directory.path().join("parent");
let moved = directory.path().join("moved"); let moved = directory.path().join("moved");
let outside = directory.path().join("outside"); let outside = directory.path().join("outside");
fs::create_dir_all(parent.join("nested")) fs::create_dir_all(parent.join("nested")).expect("original nested directory must be creatable");
.expect("original nested directory must be creatable"); fs::create_dir_all(outside.join("nested")).expect("outside nested directory must be creatable");
fs::create_dir_all(outside.join("nested"))
.expect("outside nested directory must be creatable");
let candidate = parent.join("nested/unknown-dc.log"); let candidate = parent.join("nested/unknown-dc.log");
let sanitized = sanitize_unknown_dc_log_path( let sanitized = sanitize_unknown_dc_log_path(
+1 -4
View File
@@ -154,10 +154,7 @@ enum BucketReserveError {
impl BucketReserveError { impl BucketReserveError {
fn exhausted_reserve_budget(self) -> bool { fn exhausted_reserve_budget(self) -> bool {
matches!( matches!(self, Self::Contended | Self::ReserveAndRefundContended)
self,
Self::Contended | Self::ReserveAndRefundContended
)
} }
} }
+11 -24
View File
@@ -67,12 +67,8 @@ impl DirectionBucket {
if self.should_force_reserve_failure() { if self.should_force_reserve_failure() {
return Err(current); return Err(current);
} }
self.state.compare_exchange( self.state
current, .compare_exchange(current, next, Ordering::Relaxed, Ordering::Relaxed)
next,
Ordering::Relaxed,
Ordering::Relaxed,
)
} }
#[inline(always)] #[inline(always)]
@@ -81,12 +77,8 @@ impl DirectionBucket {
if self.should_force_refund_failure() { if self.should_force_refund_failure() {
return Err(current); return Err(current);
} }
self.state.compare_exchange( self.state
current, .compare_exchange(current, next, Ordering::Relaxed, Ordering::Relaxed)
next,
Ordering::Relaxed,
Ordering::Relaxed,
)
} }
fn unpack(state: u64) -> (u64, u64) { fn unpack(state: u64) -> (u64, u64) {
@@ -355,8 +347,8 @@ impl CidrDirectionBucket {
}); });
}; };
let user_granted = user_debit.granted(); let user_granted = user_debit.granted();
let Some(aggregate_debit) = self let Some(aggregate_debit) =
.used self.used
.try_reserve_at(epoch, cap_epoch, user_granted, budget)? .try_reserve_at(epoch, cap_epoch, user_granted, budget)?
else { else {
return Ok(CidrReservation { return Ok(CidrReservation {
@@ -410,12 +402,8 @@ impl CidrUserDirectionState {
if observed_epoch > epoch { if observed_epoch > epoch {
return Err(BucketReserveError::StaleEpoch); return Err(BucketReserveError::StaleEpoch);
} }
let Some(mut active_debit) = active_users.try_reserve_at( let Some(mut active_debit) =
epoch, active_users.try_reserve_at(epoch, PACKED_USAGE_MASK, 1, budget)?
PACKED_USAGE_MASK,
1,
budget,
)?
else { else {
return Ok(false); return Ok(false);
}; };
@@ -532,10 +520,9 @@ impl CidrBucket {
} }
let cap_epoch = bytes_per_epoch(cap_bps); let cap_epoch = bytes_per_epoch(cap_bps);
match direction { match direction {
RateDirection::Up => { RateDirection::Up => self
self.up .up
.try_reserve(&share.up, epoch, cap_epoch, requested, budget) .try_reserve(&share.up, epoch, cap_epoch, requested, budget),
}
RateDirection::Down => { RateDirection::Down => {
self.down self.down
.try_reserve(&share.down, epoch, cap_epoch, requested, budget) .try_reserve(&share.down, epoch, cap_epoch, requested, budget)
+6 -9
View File
@@ -61,12 +61,8 @@ 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) = binding.user_bucket.as_ref() { if let Some(user_bucket) = binding.user_bucket.as_ref() {
let user_reservation = match user_bucket.try_reserve( let user_reservation =
direction, match user_bucket.try_reserve(direction, epoch, granted, &mut budget) {
epoch,
granted,
&mut budget,
) {
Ok(reservation) => reservation, Ok(reservation) => reservation,
Err(error) => { Err(error) => {
if error.exhausted_reserve_budget() { if error.exhausted_reserve_budget() {
@@ -107,9 +103,10 @@ 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)) = (
(binding.cidr_bucket.as_ref(), binding.cidr_user_share.as_ref()) binding.cidr_bucket.as_ref(),
{ binding.cidr_user_share.as_ref(),
) {
let cidr_reservation = match cidr_bucket.try_reserve_for_user( let cidr_reservation = match cidr_bucket.try_reserve_for_user(
direction, direction,
cidr_user_share, cidr_user_share,
+32 -8
View File
@@ -330,14 +330,38 @@ impl TrafficLimiter {
cidr_refund_down, cidr_refund_down,
] = values; ] = values;
for (counter, value) in [ for (counter, value) in [
(&self.user_scope.contention_up.reserve_exhausted_total, user_reserve_up), (
(&self.user_scope.contention_down.reserve_exhausted_total, user_reserve_down), &self.user_scope.contention_up.reserve_exhausted_total,
(&self.user_scope.contention_up.refund_exhausted_total, user_refund_up), user_reserve_up,
(&self.user_scope.contention_down.refund_exhausted_total, user_refund_down), ),
(&self.cidr_scope.contention_up.reserve_exhausted_total, cidr_reserve_up), (
(&self.cidr_scope.contention_down.reserve_exhausted_total, cidr_reserve_down), &self.user_scope.contention_down.reserve_exhausted_total,
(&self.cidr_scope.contention_up.refund_exhausted_total, cidr_refund_up), user_reserve_down,
(&self.cidr_scope.contention_down.refund_exhausted_total, cidr_refund_down), ),
(
&self.user_scope.contention_up.refund_exhausted_total,
user_refund_up,
),
(
&self.user_scope.contention_down.refund_exhausted_total,
user_refund_down,
),
(
&self.cidr_scope.contention_up.reserve_exhausted_total,
cidr_reserve_up,
),
(
&self.cidr_scope.contention_down.reserve_exhausted_total,
cidr_reserve_down,
),
(
&self.cidr_scope.contention_up.refund_exhausted_total,
cidr_refund_up,
),
(
&self.cidr_scope.contention_down.refund_exhausted_total,
cidr_refund_down,
),
] { ] {
counter.store(value, Ordering::Relaxed); counter.store(value, Ordering::Relaxed);
} }
+2 -11
View File
@@ -65,12 +65,7 @@ fn reserve_at(
cap: u64, cap: u64,
requested: u64, requested: u64,
) -> Result<Option<DirectionDebit>, BucketReserveError> { ) -> Result<Option<DirectionDebit>, BucketReserveError> {
bucket.try_reserve_at( bucket.try_reserve_at(epoch, cap, requested, &mut ReserveCasBudget::new())
epoch,
cap,
requested,
&mut ReserveCasBudget::new(),
)
} }
#[test] #[test]
@@ -346,11 +341,7 @@ fn concurrent_first_use_counts_one_active_cidr_user() {
let barrier = Arc::clone(&barrier); let barrier = Arc::clone(&barrier);
threads.push(std::thread::spawn(move || { threads.push(std::thread::spawn(move || {
barrier.wait(); barrier.wait();
user.ensure_active( user.ensure_active(13, &bucket.active_users, &mut ReserveCasBudget::new())
13,
&bucket.active_users,
&mut ReserveCasBudget::new(),
)
})); }));
} }
let results: Vec<_> = threads let results: Vec<_> = threads
@@ -9,10 +9,7 @@ fn reserve_stops_after_the_attempt_limit() {
let reservation = bucket.try_reserve_at(1, 100, 1, &mut budget); let reservation = bucket.try_reserve_at(1, 100, 1, &mut budget);
assert!(matches!( assert!(matches!(reservation, Err(BucketReserveError::Contended)));
reservation,
Err(BucketReserveError::Contended)
));
assert_eq!( assert_eq!(
bucket.reserve_cas_attempts(), bucket.reserve_cas_attempts(),
RESERVE_CAS_ATTEMPT_LIMIT as u64 RESERVE_CAS_ATTEMPT_LIMIT as u64
@@ -32,7 +29,10 @@ fn reserve_succeeds_on_the_last_allowed_attempt() {
.unwrap(); .unwrap();
assert_eq!(debit.commit_all(), 80); assert_eq!(debit.commit_all(), 80);
assert_eq!(bucket.reserve_cas_attempts(), RESERVE_CAS_ATTEMPT_LIMIT as u64); assert_eq!(
bucket.reserve_cas_attempts(),
RESERVE_CAS_ATTEMPT_LIMIT as u64
);
assert!(budget.is_exhausted()); assert!(budget.is_exhausted());
assert_eq!(bucket.used_at(1), Some(80)); assert_eq!(bucket.used_at(1), Some(80));
} }
@@ -121,15 +121,7 @@ fn lease_contention_is_not_reported_as_throttling() {
let lease = limiter let lease = limiter
.acquire_lease("alice", "203.0.113.7".parse().unwrap()) .acquire_lease("alice", "203.0.113.7".parse().unwrap())
.unwrap(); .unwrap();
let bucket = Arc::clone( let bucket = Arc::clone(&lease.binding.load_full().user_bucket.as_ref().unwrap().down);
&lease
.binding
.load_full()
.user_bucket
.as_ref()
.unwrap()
.down,
);
bucket.force_reserve_failures(RESERVE_CAS_ATTEMPT_LIMIT); bucket.force_reserve_failures(RESERVE_CAS_ATTEMPT_LIMIT);
let result = lease.try_consume(RateDirection::Down, 1); let result = lease.try_consume(RateDirection::Down, 1);
@@ -187,12 +179,7 @@ fn cidr_contention_rolls_back_provisional_user_debits() {
assert!(!result.blocked_user); assert!(!result.blocked_user);
assert!(!result.blocked_cidr); assert!(!result.blocked_cidr);
assert_eq!( assert_eq!(
binding binding.user_bucket.as_ref().unwrap().down.used_at(epoch),
.user_bucket
.as_ref()
.unwrap()
.down
.used_at(epoch),
Some(0) Some(0)
); );
assert_eq!(cidr_bucket.down.used.used_at(epoch), None); assert_eq!(cidr_bucket.down.used.used_at(epoch), None);
@@ -275,8 +262,7 @@ fn contention_snapshot_preserves_scope_direction_and_operation() {
fn cidr_activation_consumes_one_shared_attempt_budget() { fn cidr_activation_consumes_one_shared_attempt_budget() {
let bucket = CidrDirectionBucket::default(); let bucket = CidrDirectionBucket::default();
let user = CidrUserDirectionState::default(); let user = CidrUserDirectionState::default();
user.used user.used.force_reserve_failures(RESERVE_CAS_ATTEMPT_LIMIT);
.force_reserve_failures(RESERVE_CAS_ATTEMPT_LIMIT);
let mut budget = ReserveCasBudget::new(); let mut budget = ReserveCasBudget::new();
let activation = user.ensure_active(13, &bucket.active_users, &mut budget); let activation = user.ensure_active(13, &bucket.active_users, &mut budget);
@@ -360,19 +346,11 @@ fn cidr_first_grants_preserve_the_current_soft_fair_share() {
let first = CidrUserDirectionState::default(); let first = CidrUserDirectionState::default();
let second = CidrUserDirectionState::default(); let second = CidrUserDirectionState::default();
assert_eq!( assert_eq!(
first.ensure_active( first.ensure_active(17, &bucket.active_users, &mut ReserveCasBudget::new(),),
17,
&bucket.active_users,
&mut ReserveCasBudget::new(),
),
Ok(true) Ok(true)
); );
assert_eq!( assert_eq!(
second.ensure_active( second.ensure_active(17, &bucket.active_users, &mut ReserveCasBudget::new(),),
17,
&bucket.active_users,
&mut ReserveCasBudget::new(),
),
Ok(true) Ok(true)
); );
@@ -391,21 +369,13 @@ fn cidr_first_grants_preserve_the_current_soft_fair_share() {
.as_mut() .as_mut()
.unwrap() .unwrap()
.commit_all(); .commit_all();
first_reservation first_reservation.user_debit.as_mut().unwrap().commit_all();
.user_debit
.as_mut()
.unwrap()
.commit_all();
second_reservation second_reservation
.aggregate_debit .aggregate_debit
.as_mut() .as_mut()
.unwrap() .unwrap()
.commit_all(); .commit_all();
second_reservation second_reservation.user_debit.as_mut().unwrap().commit_all();
.user_debit
.as_mut()
.unwrap()
.commit_all();
assert_eq!(bucket.used.used_at(17), Some(100)); assert_eq!(bucket.used.used_at(17), Some(100));
assert_eq!(first.used.used_at(17), Some(50)); assert_eq!(first.used.used_at(17), Some(50));
assert_eq!(second.used.used_at(17), Some(50)); assert_eq!(second.used.used_at(17), Some(50));
+5 -10
View File
@@ -335,8 +335,7 @@ impl UserAdmissionAuthority {
record.incarnation = incarnation; record.incarnation = incarnation;
if identity_changed { if identity_changed {
if previous.is_some() { if previous.is_some() {
self.quota_store self.quota_store.advance_preserving_usage(user, incarnation);
.advance_preserving_usage(user, incarnation);
} else { } else {
self.quota_store.activate_fresh(user, incarnation); self.quota_store.activate_fresh(user, incarnation);
} }
@@ -420,7 +419,8 @@ impl UserAdmissionAuthority {
} }
let record = state.users.get(user)?; let record = state.users.get(user)?;
let effective = record.effective()?; let effective = record.effective()?;
(effective.enabled && effective.credential_id == credential_id).then_some(record.incarnation) (effective.enabled && effective.credential_id == credential_id)
.then_some(record.incarnation)
} }
/// Starts a short publication critical section for one authenticated owner. /// Starts a short publication critical section for one authenticated owner.
@@ -456,10 +456,7 @@ impl UserAdmissionAuthority {
} }
/// Registers a legacy owner when no credential snapshot is available. /// Registers a legacy owner when no credential snapshot is available.
pub(crate) fn register_legacy( pub(crate) fn register_legacy(self: &Arc<Self>, user: &str) -> Option<UserSessionRegistration> {
self: &Arc<Self>,
user: &str,
) -> Option<UserSessionRegistration> {
let credential_id = { let credential_id = {
let state = self.state.lock(); let state = self.state.lock();
if !state.initialized { if !state.initialized {
@@ -520,9 +517,7 @@ pub(crate) fn credential_id_from_hex(secret: &str) -> Option<UserCredentialId> {
Some(credential_id(&secret)) Some(credential_id(&secret))
} }
fn cancel_owners( fn cancel_owners(cancellations: Vec<(String, Vec<CancellationToken>)>) -> Vec<(String, usize)> {
cancellations: Vec<(String, Vec<CancellationToken>)>,
) -> Vec<(String, usize)> {
cancellations cancellations
.into_iter() .into_iter()
.map(|(user, tokens)| { .map(|(user, tokens)| {
+3 -12
View File
@@ -51,12 +51,7 @@ fn stale_candidate_cannot_overwrite_newer_mutation() {
assert!( assert!(
authority authority
.activate_config_source( .activate_config_source(2, Some(candidate_epoch), &users(secret), &HashMap::new(),)
2,
Some(candidate_epoch),
&users(secret),
&HashMap::new(),
)
.is_none() .is_none()
); );
assert!(!authority.is_user_enabled("alice")); assert!(!authority.is_user_enabled("alice"));
@@ -105,9 +100,7 @@ fn registration_dropped_before_publication_cannot_leave_an_owner() {
let secret = "00112233445566778899aabbccddeeff"; let secret = "00112233445566778899aabbccddeeff";
authority.apply_config(&users(secret), &HashMap::new()); authority.apply_config(&users(secret), &HashMap::new());
let credential = credential_id_from_hex(secret).unwrap(); let credential = credential_id_from_hex(secret).unwrap();
let mut publication = authority let mut publication = authority.claim_authenticated("alice", credential).unwrap();
.claim_authenticated("alice", credential)
.unwrap();
let registration = publication.take_registration().unwrap(); let registration = publication.take_registration().unwrap();
drop(registration); drop(registration);
@@ -126,9 +119,7 @@ fn quota_identity_follows_credential_rotation_and_recreation() {
let old_incarnation = authority let old_incarnation = authority
.authenticated_incarnation("alice", credential_id_from_hex(old_secret).unwrap()) .authenticated_incarnation("alice", credential_id_from_hex(old_secret).unwrap())
.unwrap(); .unwrap();
let old_quota = quota_store let old_quota = quota_store.handle_exact("alice", old_incarnation).unwrap();
.handle_exact("alice", old_incarnation)
.unwrap();
old_quota.charge(40); old_quota.charge(40);
let rotated = authority.stage_user("alice", new_secret, true).unwrap(); let rotated = authority.stage_user("alice", new_secret, true).unwrap();
+4 -1
View File
@@ -54,7 +54,10 @@ impl SlotBudget {
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |current| { .fetch_update(Ordering::AcqRel, Ordering::Acquire, |current| {
current.checked_sub(amount) current.checked_sub(amount)
}); });
debug_assert!(released.is_ok(), "slot budget release must match acquisitions"); debug_assert!(
released.is_ok(),
"slot budget release must match acquisitions"
);
} }
/// Returns the exact number of currently committed or reserved slots. /// Returns the exact number of currently committed or reserved slots.
+3 -6
View File
@@ -22,13 +22,13 @@ 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, UserQuotaHandle}; pub(crate) use self::quota_store::{QuotaReservation, QuotaStore, UserQuotaHandle};
pub(crate) use self::users::UserConnectionObservation;
#[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;
use crate::proxy::user_connection_authority::UserConnectionAuthority;
pub use self::tls_fingerprints::TlsFingerprintSnapshotRow; pub use self::tls_fingerprints::TlsFingerprintSnapshotRow;
pub(crate) use self::users::UserConnectionObservation;
use crate::config::MeWriterPickMode; use crate::config::MeWriterPickMode;
use crate::proxy::user_connection_authority::UserConnectionAuthority;
const ME_HANDSHAKE_ERROR_CODE_MAX: usize = 64; const ME_HANDSHAKE_ERROR_CODE_MAX: usize = 64;
@@ -432,10 +432,7 @@ impl Stats {
#[cfg(test)] #[cfg(test)]
pub(crate) fn with_quota_store(quota_store: Arc<QuotaStore>) -> Self { pub(crate) fn with_quota_store(quota_store: Arc<QuotaStore>) -> Self {
Self::with_process_authorities( Self::with_process_authorities(quota_store, Arc::new(UserConnectionAuthority::default()))
quota_store,
Arc::new(UserConnectionAuthority::default()),
)
} }
/// Creates generation telemetry around process-owned enforcement authorities. /// Creates generation telemetry around process-owned enforcement authorities.
+1 -5
View File
@@ -146,11 +146,7 @@ impl QuotaStore {
} }
/// Advances a credential incarnation while preserving usage captured at the transition. /// Advances a credential incarnation while preserving usage captured at the transition.
pub(crate) fn advance_preserving_usage( pub(crate) fn advance_preserving_usage(&self, user: &str, incarnation: UserIncarnation) {
&self,
user: &str,
incarnation: UserIncarnation,
) {
let slot = self.slot(user); let slot = self.slot(user);
let mut state = slot.state.lock(); let mut state = slot.state.lock();
if incarnation <= state.high_water { if incarnation <= state.high_water {
+1 -1
View File
@@ -1,6 +1,6 @@
use std::borrow::Borrow; use std::borrow::Borrow;
use std::collections::{HashMap, VecDeque};
use std::collections::hash_map::DefaultHasher; use std::collections::hash_map::DefaultHasher;
use std::collections::{HashMap, VecDeque};
use std::hash::{Hash, Hasher}; use std::hash::{Hash, Hasher};
use std::num::NonZeroUsize; use std::num::NonZeroUsize;
use std::sync::Arc; use std::sync::Arc;
+9 -8
View File
@@ -17,26 +17,27 @@ fn test_stats_shared_counters() {
fn runtime_stats_share_process_connection_admission_authority() { fn runtime_stats_share_process_connection_admission_authority() {
let quota_store = Arc::new(QuotaStore::default()); let quota_store = Arc::new(QuotaStore::default());
let authority = Arc::new(UserConnectionAuthority::default()); let authority = Arc::new(UserConnectionAuthority::default());
let first = Stats::with_process_authorities( let first = Stats::with_process_authorities(Arc::clone(&quota_store), Arc::clone(&authority));
Arc::clone(&quota_store),
Arc::clone(&authority),
);
let second = Stats::with_process_authorities(quota_store, authority); let second = Stats::with_process_authorities(quota_store, authority);
let permit = first let permit = first
.connection_authority() .connection_authority()
.try_acquire("alice", Some(1)) .try_acquire("alice", Some(1))
.unwrap(); .unwrap();
assert!(second assert!(
second
.connection_authority() .connection_authority()
.try_acquire("alice", Some(1)) .try_acquire("alice", Some(1))
.is_none()); .is_none()
);
drop(permit); drop(permit);
assert!(second assert!(
second
.connection_authority() .connection_authority()
.try_acquire("alice", Some(1)) .try_acquire("alice", Some(1))
.is_some()); .is_some()
);
} }
#[test] #[test]
+1 -2
View File
@@ -350,8 +350,7 @@ impl TlsFingerprintCollector {
let mut removed = 0usize; let mut removed = 0usize;
self.entries.retain(|_, entry| { self.entries.retain(|_, entry| {
let last_seen = entry.last_seen_epoch_secs.load(Ordering::Relaxed); let last_seen = entry.last_seen_epoch_secs.load(Ordering::Relaxed);
let retained = let retained = ttl_secs != 0 && now_epoch_secs.saturating_sub(last_seen) <= ttl_secs;
ttl_secs != 0 && now_epoch_secs.saturating_sub(last_seen) <= ttl_secs;
if !retained { if !retained {
removed += 1; removed += 1;
} }
+2 -4
View File
@@ -76,10 +76,8 @@ pub fn parse_proxy_config_text(text: &str, http_status: u16) -> ProxyConfigData
pub async fn load_proxy_config_cache(path: &str) -> Result<ProxyConfigData> { pub async fn load_proxy_config_cache(path: &str) -> Result<ProxyConfigData> {
#[cfg(unix)] #[cfg(unix)]
let bytes = read_regular_limited_async( let bytes =
Path::new(path).to_path_buf(), read_regular_limited_async(Path::new(path).to_path_buf(), HTTPS_RESPONSE_BODY_MAX_BYTES)
HTTPS_RESPONSE_BODY_MAX_BYTES,
)
.await; .await;
#[cfg(not(unix))] #[cfg(not(unix))]
let bytes = tokio::fs::read(path).await; let bytes = tokio::fs::read(path).await;
@@ -234,7 +234,8 @@ pub(super) async fn maybe_refresh_idle_writer_for_dc(
purpose, purpose,
&mut reservation, &mut reservation,
); );
let rotate_ok = match tokio::time::timeout(pool.reconnect_runtime.me_one_timeout, replace).await { let rotate_ok = match tokio::time::timeout(pool.reconnect_runtime.me_one_timeout, replace).await
{
Ok(Ok(())) => true, Ok(Ok(())) => true,
Ok(Err(error)) => { Ok(Err(error)) => {
debug!( debug!(
+2 -4
View File
@@ -270,10 +270,8 @@ async fn under_floor_idle_writer_still_enters_transactional_refresh() {
let writer = insert_active_writer_at(&pool, writer_id, 2, endpoint).await; let writer = insert_active_writer_at(&pool, writer_id, 2, endpoint).await;
let key = (2, IpFamily::V4); let key = (2, IpFamily::V4);
let live_writer_ids_by_addr = HashMap::from([((2, endpoint), vec![writer_id])]); let live_writer_ids_by_addr = HashMap::from([((2, endpoint), vec![writer_id])]);
let writer_idle_since = HashMap::from([( let writer_idle_since =
writer_id, HashMap::from([(writer_id, MePool::now_epoch_secs().saturating_sub(60))]);
MePool::now_epoch_secs().saturating_sub(60),
)]);
let bound_clients_by_writer = HashMap::from([(writer_id, 0)]); let bound_clients_by_writer = HashMap::from([(writer_id, 0)]);
let mut next_attempt = HashMap::new(); let mut next_attempt = HashMap::new();
let rng = Arc::new(SecureRandom::new()); let rng = Arc::new(SecureRandom::new());
+1 -1
View File
@@ -64,9 +64,9 @@ pub use ping::{
MePingFamily, MePingReport, MePingSample, format_me_route, format_sample_line, run_me_ping, MePingFamily, MePingReport, MePingSample, format_me_route, format_sample_line, run_me_ping,
}; };
pub use pool::MePool; pub use pool::MePool;
pub(crate) use pool_status::MeApiHardswapSnapshot;
#[allow(unused_imports)] #[allow(unused_imports)]
pub use pool_nat::{detect_public_ip, stun_probe}; pub use pool_nat::{detect_public_ip, stun_probe};
pub(crate) use pool_status::MeApiHardswapSnapshot;
pub(crate) use registry::ConnLease; pub(crate) use registry::ConnLease;
pub use registry::ConnRegistry; pub use registry::ConnRegistry;
pub use rotation::{MeReinitTrigger, me_reinit_scheduler, me_rotation_task}; pub use rotation::{MeReinitTrigger, me_reinit_scheduler, me_rotation_task};
@@ -120,12 +120,8 @@ impl MePool {
me_route_inline_recovery_wait_ms: u64, me_route_inline_recovery_wait_ms: u64,
me_connection_cleanup_capacity: usize, me_connection_cleanup_capacity: usize,
) -> Arc<Self> { ) -> Arc<Self> {
let endpoint_snapshot = Self::build_endpoint_snapshot( let endpoint_snapshot =
&decision, Self::build_endpoint_snapshot(&decision, proxy_map_v4, proxy_map_v6, 1);
proxy_map_v4,
proxy_map_v6,
1,
);
let registry = Arc::new(ConnRegistry::with_route_and_cleanup_capacity( let registry = Arc::new(ConnRegistry::with_route_and_cleanup_capacity(
me_route_channel_capacity, me_route_channel_capacity,
me_connection_cleanup_capacity, me_connection_cleanup_capacity,
+1 -5
View File
@@ -198,11 +198,7 @@ impl MePool {
} }
fn mirror_negative_dcs(map: &mut HashMap<i32, Vec<(IpAddr, u16)>>) { fn mirror_negative_dcs(map: &mut HashMap<i32, Vec<(IpAddr, u16)>>) {
let positive_dcs = map let positive_dcs = map.keys().copied().filter(|dc| *dc > 0).collect::<Vec<_>>();
.keys()
.copied()
.filter(|dc| *dc > 0)
.collect::<Vec<_>>();
for dc in positive_dcs { for dc in positive_dcs {
if !map.contains_key(&-dc) if !map.contains_key(&-dc)
&& let Some(endpoints) = map.get(&dc).cloned() && let Some(endpoints) = map.get(&dc).cloned()
@@ -203,7 +203,9 @@ impl MePool {
.max(1) .max(1)
.min(WRITER_REPLACEMENT_OPEN_LIMIT_MAX); .min(WRITER_REPLACEMENT_OPEN_LIMIT_MAX);
loop { loop {
let reserved = self.writer_replacement_open_reserved.load(Ordering::Acquire); let reserved = self
.writer_replacement_open_reserved
.load(Ordering::Acquire);
if reserved >= replacement_limit { if reserved >= replacement_limit {
return None; return None;
} }
+5 -8
View File
@@ -252,12 +252,7 @@ impl MePool {
let addr = candidates[idx]; let addr = candidates[idx];
match self match self
.connect_one_with_generation_contour_for_dc_with_intent( .connect_one_with_generation_contour_for_dc_with_intent(
addr, addr, rng, generation, contour, dc, intent,
rng,
generation,
contour,
dc,
intent,
) )
.await .await
{ {
@@ -296,8 +291,10 @@ impl MePool {
let status = self.reinit.status.load(); let status = self.reinit.status.load();
let role_is_authoritative = match target.contour { let role_is_authoritative = match target.contour {
WriterContour::Active => target.generation == status.active_generation, WriterContour::Active => target.generation == status.active_generation,
WriterContour::Warm => status.pending_hardswap_generation != 0 WriterContour::Warm => {
&& target.generation == status.pending_hardswap_generation, status.pending_hardswap_generation != 0
&& target.generation == status.pending_hardswap_generation
}
WriterContour::Draining => false, WriterContour::Draining => false,
}; };
let pending_revision_matches = target.contour != WriterContour::Warm let pending_revision_matches = target.contour != WriterContour::Warm
+2 -3
View File
@@ -10,8 +10,8 @@ use rand::seq::SliceRandom;
use std::collections::hash_map::DefaultHasher; use std::collections::hash_map::DefaultHasher;
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
use crate::crypto::SecureRandom;
use crate::config::MeBindStaleMode; use crate::config::MeBindStaleMode;
use crate::crypto::SecureRandom;
use crate::network::IpFamily; use crate::network::IpFamily;
use super::pool::{ use super::pool::{
@@ -117,8 +117,7 @@ fn publish_reinit_state(reinit: &ReinitCore, state: &ReinitCoordinatorState) {
pending_hardswap_started_at_epoch_secs: pending pending_hardswap_started_at_epoch_secs: pending
.map_or(0, |value| value.started_at_epoch_secs), .map_or(0, |value| value.started_at_epoch_secs),
pending_hardswap_map_hash: pending.map_or(0, |value| value.map_hash), pending_hardswap_map_hash: pending.map_or(0, |value| value.map_hash),
pending_hardswap_endpoint_revision: pending pending_hardswap_endpoint_revision: pending.map_or(0, |value| value.endpoint_revision),
.map_or(0, |value| value.endpoint_revision),
inflight: state.attempts.len(), inflight: state.attempts.len(),
}; };
reinit reinit
@@ -318,15 +318,13 @@ impl MePool {
if alive >= required { if alive >= required {
covered = covered.saturating_add(1); covered = covered.saturating_add(1);
} else { } else {
writer_deficit = writer_deficit = writer_deficit.saturating_add(required.saturating_sub(alive));
writer_deficit.saturating_add(required.saturating_sub(alive));
missing_groups.push(DcFamilyGroup { dc: *dc, family }); missing_groups.push(DcFamilyGroup { dc: *dc, family });
} }
} }
} }
missing_groups.sort_unstable_by_key(|group| { missing_groups
(group.dc, matches!(group.family, IpFamily::V6)) .sort_unstable_by_key(|group| (group.dc, matches!(group.family, IpFamily::V6)));
});
HardswapCoverage { HardswapCoverage {
ratio: if total == 0 { ratio: if total == 0 {
1.0 1.0
@@ -453,8 +451,8 @@ impl MePool {
let authoritative_warm = contour == WriterContour::Warm let authoritative_warm = contour == WriterContour::Warm
&& pending_generation == Some(writer.generation) && pending_generation == Some(writer.generation)
&& endpoint_is_current; && endpoint_is_current;
let stale_active = contour == WriterContour::Active let stale_active =
&& writer.generation != active_generation; contour == WriterContour::Active && writer.generation != active_generation;
if authoritative_warm || (contour == WriterContour::Active && !stale_active) { if authoritative_warm || (contour == WriterContour::Active && !stale_active) {
continue; continue;
} }
@@ -535,8 +533,7 @@ impl MePool {
.filter(|w| !w.draining.load(Ordering::Relaxed)) .filter(|w| !w.draining.load(Ordering::Relaxed))
.filter(|w| w.generation == generation) .filter(|w| w.generation == generation)
.filter(|w| { .filter(|w| {
WriterContour::from_u8(w.contour.load(Ordering::Acquire)) WriterContour::from_u8(w.contour.load(Ordering::Acquire)) == WriterContour::Active
== WriterContour::Active
}) })
.filter(|w| w.writer_dc == dc) .filter(|w| w.writer_dc == dc)
.filter(|w| endpoints.contains(&w.addr)) .filter(|w| endpoints.contains(&w.addr))
@@ -34,11 +34,7 @@ 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 !self.hardswap_warmup_is_authoritative( if !self.hardswap_warmup_is_authoritative(generation, map_hash, endpoint_revision) {
generation,
map_hash,
endpoint_revision,
) {
return; return;
} }
for family in [IpFamily::V4, IpFamily::V6] { for family in [IpFamily::V4, IpFamily::V6] {
@@ -119,11 +115,7 @@ impl MePool {
} }
last_fresh_count = self last_fresh_count = self
.fresh_writer_count_for_dc_endpoints( .fresh_writer_count_for_dc_endpoints(generation, *dc, &family_endpoints)
generation,
*dc,
&family_endpoints,
)
.await; .await;
if last_fresh_count >= required { if last_fresh_count >= required {
completed = true; completed = true;
@@ -324,8 +316,7 @@ impl MePool {
Err(ReinitCommitFailure::Superseded) => { Err(ReinitCommitFailure::Superseded) => {
debug!( debug!(
previous_generation, previous_generation,
generation, generation, "ME reinit result discarded after a newer desired-map attempt"
"ME reinit result discarded after a newer desired-map attempt"
); );
return false; return false;
} }
@@ -228,7 +228,10 @@ async fn partial_hardswap_is_rejected_when_stale_binding_is_disabled() {
.commit_reinit_attempt(&reservation.attempt, &desired_by_dc, 0.5) .commit_reinit_attempt(&reservation.attempt, &desired_by_dc, 0.5)
.await; .await;
assert!(matches!(result, Err(ReinitCommitFailure::Redundancy { .. }))); assert!(matches!(
result,
Err(ReinitCommitFailure::Redundancy { .. })
));
assert_eq!(pool.current_generation(), active_generation); assert_eq!(pool.current_generation(), active_generation);
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));
@@ -303,24 +306,8 @@ async fn partial_hardswap_preserves_fallback_only_for_underfloor_family() {
let v6 = addr_v6(1, 2001); let v6 = addr_v6(1, 2001);
let desired_by_dc = HashMap::from([(1, HashSet::from([v4, v6]))]); let desired_by_dc = HashMap::from([(1, HashSet::from([v4, v6]))]);
let active_generation = pool.current_generation(); let active_generation = pool.current_generation();
let old_v4 = insert_writer( let old_v4 = insert_writer(&pool, 451, 1, v4, active_generation, WriterContour::Active).await;
&pool, let old_v6 = insert_writer(&pool, 452, 1, v6, active_generation, WriterContour::Active).await;
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 map_hash = MePool::desired_map_hash(&desired_by_dc);
let endpoint_revision = pool.endpoint_snapshot.load().revision; let endpoint_revision = pool.endpoint_snapshot.load().revision;
let reservation = pool let reservation = pool
@@ -81,8 +81,7 @@ impl MePool {
let (replacement_preparing_current, replacement_retiring_current) = let (replacement_preparing_current, replacement_retiring_current) =
self.registry.writer_replacement_counts(); self.registry.writer_replacement_counts();
let pending_age_secs = pending.then(|| { let pending_age_secs = pending.then(|| {
Self::now_epoch_secs() Self::now_epoch_secs().saturating_sub(reinit.pending_hardswap_started_at_epoch_secs)
.saturating_sub(reinit.pending_hardswap_started_at_epoch_secs)
}); });
MeApiHardswapSnapshot { MeApiHardswapSnapshot {
@@ -83,12 +83,8 @@ impl MePool {
if endpoint_count == 0 { if endpoint_count == 0 {
continue; continue;
} }
let required = let required = self.required_writers_for_dc_with_floor_mode(endpoint_count, false);
self.required_writers_for_dc_with_floor_mode(endpoint_count, false); let alive = live_writers_by_group.get(&(dc, ipv4)).copied().unwrap_or(0);
let alive = live_writers_by_group
.get(&(dc, ipv4))
.copied()
.unwrap_or(0);
if alive < required { if alive < required {
return false; return false;
} }
@@ -133,9 +129,7 @@ impl MePool {
.map(|endpoints| { .map(|endpoints| {
endpoint_family_counts(endpoints) endpoint_family_counts(endpoints)
.into_iter() .into_iter()
.map(|(_, count)| { .map(|(_, count)| self.required_writers_for_dc_with_floor_mode(count, false))
self.required_writers_for_dc_with_floor_mode(count, false)
})
.sum::<usize>() .sum::<usize>()
}) })
.sum(); .sum();
@@ -254,9 +248,7 @@ impl MePool {
let dc_required_writers = family_counts let dc_required_writers = family_counts
.iter() .iter()
.filter(|(_, count)| *count > 0) .filter(|(_, count)| *count > 0)
.map(|(_, count)| { .map(|(_, count)| self.required_writers_for_dc_with_floor_mode(*count, false))
self.required_writers_for_dc_with_floor_mode(*count, false)
})
.sum::<usize>(); .sum::<usize>();
let floor_min = family_counts let floor_min = family_counts
.iter() .iter()
@@ -289,8 +281,7 @@ impl MePool {
.me_adaptive_floor_max_extra_writers_multi_per_core .me_adaptive_floor_max_extra_writers_multi_per_core
.load(Ordering::Relaxed) as usize .load(Ordering::Relaxed) as usize
}; };
family_base family_base.saturating_add(adaptive_cpu_cores.saturating_mul(extra_per_core))
.saturating_add(adaptive_cpu_cores.saturating_mul(extra_per_core))
}) })
.sum::<usize>(); .sum::<usize>();
let floor_capped = let floor_capped =
@@ -357,6 +348,9 @@ impl MePool {
} }
fn endpoint_family_counts(endpoints: &BTreeSet<SocketAddr>) -> [(bool, usize); 2] { fn endpoint_family_counts(endpoints: &BTreeSet<SocketAddr>) -> [(bool, usize); 2] {
let ipv4 = endpoints.iter().filter(|endpoint| endpoint.is_ipv4()).count(); let ipv4 = endpoints
.iter()
.filter(|endpoint| endpoint.is_ipv4())
.count();
[(true, ipv4), (false, endpoints.len().saturating_sub(ipv4))] [(true, ipv4), (false, endpoints.len().saturating_sub(ipv4))]
} }
@@ -12,12 +12,7 @@ use crate::transport::middle_proxy::codec::WriterCommand;
use crate::transport::middle_proxy::pool::{MePool, MeWriter, WriterContour}; use crate::transport::middle_proxy::pool::{MePool, MeWriter, WriterContour};
use crate::transport::middle_proxy::pool_writer_security_tests::make_pool_with_decision; use crate::transport::middle_proxy::pool_writer_security_tests::make_pool_with_decision;
fn writer( fn writer(pool: &Arc<MePool>, id: u64, dc: i32, addr: SocketAddr) -> MeWriter {
pool: &Arc<MePool>,
id: u64,
dc: i32,
addr: SocketAddr,
) -> MeWriter {
let (tx, _rx) = mpsc::channel::<WriterCommand>(8); let (tx, _rx) = mpsc::channel::<WriterCommand>(8);
MeWriter { MeWriter {
id, id,
@@ -60,12 +55,7 @@ async fn dual_family_status_reports_each_family_floor() {
let mut writers = pool.writers.write().await; let mut writers = pool.writers.write().await;
for (group, dc) in [2, -2].into_iter().enumerate() { for (group, dc) in [2, -2].into_iter().enumerate() {
for offset in 0..required_per_family { for offset in 0..required_per_family {
writers.push(writer( writers.push(writer(&pool, (group as u64 * 100) + offset as u64, dc, v4));
&pool,
(group as u64 * 100) + offset as u64,
dc,
v4,
));
} }
} }
drop(writers); drop(writers);
@@ -122,11 +122,9 @@ impl MePool {
!candidate.draining.load(Ordering::Acquire) !candidate.draining.load(Ordering::Acquire)
&& candidate.writer_dc == writer.writer_dc && candidate.writer_dc == writer.writer_dc
&& 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
&& candidate.addr.is_ipv4() == writer.addr.is_ipv4() && candidate.addr.is_ipv4() == writer.addr.is_ipv4()
&& endpoint_snapshot && endpoint_snapshot.contains_dc_endpoint(candidate.writer_dc, candidate.addr)
.contains_dc_endpoint(candidate.writer_dc, candidate.addr)
}) })
.count(); .count();
if current >= required { if current >= required {
@@ -138,11 +138,7 @@ impl MePool {
writers.push(writer); writers.push(writer);
self.conn_count.fetch_add(1, Ordering::Relaxed); self.conn_count.fetch_add(1, Ordering::Relaxed);
writers.publish_current(); writers.publish_current();
self.apply_writer_draining_state( self.apply_writer_draining_state(&writers[victim_pos], self.force_close_timeout(), false);
&writers[victim_pos],
self.force_close_timeout(),
false,
);
self.lifecycle self.lifecycle
.spawn_registered_writer(task_registration, writer_task); .spawn_registered_writer(task_registration, writer_task);
reservation.mark_committed(); reservation.mark_committed();
@@ -264,10 +260,7 @@ mod tests {
async fn replacement_commit_publishes_successor_before_draining_victim() { async fn replacement_commit_publishes_successor_before_draining_victim() {
let pool = make_pool().await; let pool = make_pool().await;
let addr = endpoint(1); let addr = endpoint(1);
pool.update_proxy_maps( pool.update_proxy_maps(HashMap::from([(2, vec![(addr.ip(), addr.port())])]), None)
HashMap::from([(2, vec![(addr.ip(), addr.port())])]),
None,
)
.await; .await;
let victim = install_writer(&pool, 1001, 2, addr).await; let victim = install_writer(&pool, 1001, 2, addr).await;
let expected_role = WriterRole::from_writer(&victim); let expected_role = WriterRole::from_writer(&victim);
@@ -307,10 +300,7 @@ mod tests {
async fn cancelled_replacement_waiting_for_publication_restores_all_reservations() { async fn cancelled_replacement_waiting_for_publication_restores_all_reservations() {
let pool = make_pool().await; let pool = make_pool().await;
let addr = endpoint(2); let addr = endpoint(2);
pool.update_proxy_maps( pool.update_proxy_maps(HashMap::from([(2, vec![(addr.ip(), addr.port())])]), None)
HashMap::from([(2, vec![(addr.ip(), addr.port())])]),
None,
)
.await; .await;
let victim = install_writer(&pool, 2001, 2, addr).await; let victim = install_writer(&pool, 2001, 2, addr).await;
let expected_role = WriterRole::from_writer(&victim); let expected_role = WriterRole::from_writer(&victim);
@@ -336,10 +326,21 @@ mod tests {
assert!(result.is_err()); assert!(result.is_err());
drop(writers_guard); drop(writers_guard);
drop(reservation); drop(reservation);
assert_eq!(pool.writer_replacement_open_reserved.load(Ordering::Acquire), 0); assert_eq!(
pool.writer_replacement_open_reserved
.load(Ordering::Acquire),
0
);
assert_eq!(pool.registry.writer_replacement_counts(), (0, 0)); assert_eq!(pool.registry.writer_replacement_counts(), (0, 0));
assert!(!victim.draining.load(Ordering::Acquire)); assert!(!victim.draining.load(Ordering::Acquire));
assert!(!pool.writers.read().await.iter().any(|writer| writer.id == 2002)); assert!(
!pool
.writers
.read()
.await
.iter()
.any(|writer| writer.id == 2002)
);
} }
#[tokio::test] #[tokio::test]
@@ -379,7 +380,14 @@ mod tests {
assert!(result.is_err()); assert!(result.is_err());
drop(reservation); drop(reservation);
assert!(!victim.draining.load(Ordering::Acquire)); assert!(!victim.draining.load(Ordering::Acquire));
assert!(!pool.writers.read().await.iter().any(|writer| writer.id == 3002)); assert!(
!pool
.writers
.read()
.await
.iter()
.any(|writer| writer.id == 3002)
);
assert_eq!(pool.registry.writer_replacement_counts(), (0, 0)); assert_eq!(pool.registry.writer_replacement_counts(), (0, 0));
} }
} }
@@ -92,12 +92,7 @@ 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( .reserve_writer_open(contour, intent, writer_dc, addr)
contour,
intent,
writer_dc,
addr,
)
.await .await
else { else {
return Err(ProxyError::Proxy(format!( return Err(ProxyError::Proxy(format!(
@@ -8,11 +8,11 @@ use tokio::sync::mpsc::error::TrySendError;
use super::super::codec::WriterCommand; use super::super::codec::WriterCommand;
use super::super::{MeResponse, RouteBytePermit}; use super::super::{MeResponse, RouteBytePermit};
use super::replacement::WriterBindOutcome;
use super::{ use super::{
BoundConn, ConnMeta, ConnRegistry, ConnWriter, HotConnBinding, RouteResult, BoundConn, ConnMeta, ConnRegistry, ConnWriter, HotConnBinding, RouteResult,
WriterActivitySnapshot, WriterActivitySnapshot,
}; };
use super::replacement::WriterBindOutcome;
impl ConnRegistry { impl ConnRegistry {
fn set_writer_bound_count(&self, writer_id: u64, count: usize) { fn set_writer_bound_count(&self, writer_id: u64, count: usize) {
+1 -3
View File
@@ -169,8 +169,7 @@ impl MePool {
0..self.route_runtime.me_route_inline_recovery_attempts.max(1) 0..self.route_runtime.me_route_inline_recovery_attempts.max(1)
{ {
let endpoint_snapshot = self.endpoint_snapshot.load_full(); let endpoint_snapshot = self.endpoint_snapshot.load_full();
for (dc, addrs) in for (dc, addrs) in &endpoint_snapshot.preferred_endpoints_by_dc
&endpoint_snapshot.preferred_endpoints_by_dc
{ {
for addr in addrs { for addr in addrs {
let _ = self let _ = self
@@ -544,5 +543,4 @@ impl MePool {
return Ok(()); return Ok(());
} }
} }
} }
+2 -1
View File
@@ -31,7 +31,8 @@ impl MePool {
writer_reserved_bytes: usize, writer_reserved_bytes: usize,
payload_permit: Option<OwnedSemaphorePermit>, payload_permit: Option<OwnedSemaphorePermit>,
) -> Result<BoundWriterSendOutcome> { ) -> Result<BoundWriterSendOutcome> {
let Some((current, current_meta)) = self.registry.get_writer_with_meta(conn_id).await else { let Some((current, current_meta)) = self.registry.get_writer_with_meta(conn_id).await
else {
return Ok(BoundWriterSendOutcome::Retry(payload_permit)); return Ok(BoundWriterSendOutcome::Retry(payload_permit));
}; };
let deadline = writer_send_deadline(self.route_runtime.me_route_blocking_send_timeout); let deadline = writer_send_deadline(self.route_runtime.me_route_blocking_send_timeout);
+2 -2
View File
@@ -9,8 +9,8 @@ use super::super::MePool;
use super::super::codec::{ProxyReqCommand, WriterCommand}; use super::super::codec::{ProxyReqCommand, WriterCommand};
use super::reservation::{ use super::reservation::{
WriterByteReserveError, WriterCommandReserveError, proxy_req_payload_from_command, WriterByteReserveError, WriterCommandReserveError, proxy_req_payload_from_command,
proxy_req_resident_permits, proxy_tag_array, reserve_writer_bytes, proxy_req_resident_permits, proxy_tag_array, reserve_writer_bytes, reserve_writer_command_slot,
reserve_writer_command_slot, writer_send_deadline, writer_send_deadline,
}; };
use crate::error::{ProxyError, Result}; use crate::error::{ProxyError, Result};
use crate::stream::PooledBuffer; use crate::stream::PooledBuffer;
@@ -34,9 +34,7 @@ pub(super) fn proxy_req_payload_from_command(
} }
} }
pub(super) fn payload_permit_from_data_command( pub(super) fn payload_permit_from_data_command(cmd: WriterCommand) -> Option<OwnedSemaphorePermit> {
cmd: WriterCommand,
) -> Option<OwnedSemaphorePermit> {
match cmd { match cmd {
WriterCommand::Data { _permit, .. } => _permit, WriterCommand::Data { _permit, .. } => _permit,
_ => None, _ => None,
+1 -5
View File
@@ -127,11 +127,7 @@ impl MePool {
WriterContour::Draining => PICK_PENALTY_DRAINING, WriterContour::Draining => PICK_PENALTY_DRAINING,
}; };
let stale = (writer.generation < current_generation) as usize; let stale = (writer.generation < current_generation) as usize;
let stale_penalty = if stale != 0 { let stale_penalty = if stale != 0 { PICK_PENALTY_STALE } else { 0 };
PICK_PENALTY_STALE
} else {
0
};
let degraded = writer.degraded.load(Ordering::Relaxed) as usize; let degraded = writer.degraded.load(Ordering::Relaxed) as usize;
let degraded_penalty = if degraded != 0 { let degraded_penalty = if degraded != 0 {
PICK_PENALTY_DEGRADED PICK_PENALTY_DEGRADED
@@ -10,9 +10,7 @@ use crate::network::IpFamily;
use crate::network::probe::NetworkDecision; use crate::network::probe::NetworkDecision;
use crate::stats::Stats; use crate::stats::Stats;
use super::pool::{ use super::pool::{MePool, ReinitStatusSnapshot, WriterContour, WriterRole};
MePool, ReinitStatusSnapshot, WriterContour, WriterRole,
};
use super::pool_writer_security_tests::make_pool_with_decision; use super::pool_writer_security_tests::make_pool_with_decision;
async fn make_pool() -> Arc<MePool> { async fn make_pool() -> Arc<MePool> {
@@ -195,10 +193,7 @@ async fn refill_accepts_enabled_nonpreferred_family_without_multipath() {
let v6_addr = SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 443); let v6_addr = SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 443);
pool.update_proxy_maps( pool.update_proxy_maps(
HashMap::from([(2, vec![(v4_addr.ip(), v4_addr.port())])]), HashMap::from([(2, vec![(v4_addr.ip(), v4_addr.port())])]),
Some(HashMap::from([( Some(HashMap::from([(2, vec![(v6_addr.ip(), v6_addr.port())])])),
2,
vec![(v6_addr.ip(), v6_addr.port())],
)])),
) )
.await; .await;
@@ -98,13 +98,7 @@ async fn normal_active_publication_cannot_race_past_the_family_floor() {
let generation = pool.current_generation(); let generation = pool.current_generation();
let writers = (1..=3) let writers = (1..=3)
.map(|writer_id| { .map(|writer_id| {
unregistered_writer( unregistered_writer(&pool, writer_id, addr, generation, WriterContour::Active)
&pool,
writer_id,
addr,
generation,
WriterContour::Active,
)
}) })
.collect::<Vec<_>>(); .collect::<Vec<_>>();
let candidate = unregistered_writer(&pool, 4, addr, generation, WriterContour::Active); let candidate = unregistered_writer(&pool, 4, addr, generation, WriterContour::Active);
@@ -135,10 +129,7 @@ async fn stale_same_family_writers_do_not_satisfy_current_endpoint_coverage() {
let current_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); let current_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443);
let stale_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 2)), 443); let stale_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 2)), 443);
pool.update_proxy_maps( pool.update_proxy_maps(
std::collections::HashMap::from([( std::collections::HashMap::from([(2, vec![(current_addr.ip(), current_addr.port())])]),
2,
vec![(current_addr.ip(), current_addr.port())],
)]),
None, None,
) )
.await; .await;
@@ -244,13 +235,7 @@ async fn normal_active_publication_allows_adaptive_growth_above_family_floor() {
let generation = pool.current_generation(); let generation = pool.current_generation();
let writers = (1..=3) let writers = (1..=3)
.map(|writer_id| { .map(|writer_id| {
unregistered_writer( unregistered_writer(&pool, writer_id, addr, generation, WriterContour::Active)
&pool,
writer_id,
addr,
generation,
WriterContour::Active,
)
}) })
.collect::<Vec<_>>(); .collect::<Vec<_>>();
let candidate = unregistered_writer(&pool, 4, addr, generation, WriterContour::Active); let candidate = unregistered_writer(&pool, 4, addr, generation, WriterContour::Active);
@@ -287,13 +272,7 @@ async fn normal_active_publication_rejects_configured_contour_cap() {
let generation = pool.current_generation(); let generation = pool.current_generation();
let writers = (1..=3) let writers = (1..=3)
.map(|writer_id| { .map(|writer_id| {
unregistered_writer( unregistered_writer(&pool, writer_id, addr, generation, WriterContour::Active)
&pool,
writer_id,
addr,
generation,
WriterContour::Active,
)
}) })
.collect::<Vec<_>>(); .collect::<Vec<_>>();
let candidate = unregistered_writer(&pool, 4, addr, generation, WriterContour::Active); let candidate = unregistered_writer(&pool, 4, addr, generation, WriterContour::Active);
@@ -238,10 +238,7 @@ async fn send_proxy_req_uses_live_same_dc_writer_while_preferred_endpoint_refill
.await; .await;
assert!(pool.admission_ready_conditional_cast().await); assert!(pool.admission_ready_conditional_cast().await);
assert_eq!( assert_eq!(pool.preferred_endpoints_for_dc(2).await, vec![new_addr]);
pool.preferred_endpoints_for_dc(2).await,
vec![new_addr]
);
let (conn_id, _rx) = pool.registry.register().await; let (conn_id, _rx) = pool.registry.register().await;
let result = pool let result = pool

Some files were not shown because too many files have changed in this diff Show More