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 super::ApiShared;
#[cfg(test)]
use super::config_store::write_atomic;
use super::config_store::{
EDITABLE_SECTIONS, EDITABLE_SERVER_FIELDS, compute_snapshot_revision, is_editable_section,
load_candidate_snapshot, load_config_snapshot, render_server_listeners,
render_top_level_section, resolve_single_source_owner, upsert_toml_table,
write_atomic_if_unchanged,
};
#[cfg(test)]
use super::config_store::write_atomic;
use super::model::ApiFailure;
use crate::config::ProxyConfig;
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.
mod atomic;
pub(in crate::api) use atomic::{write_atomic, write_atomic_if_unchanged};
#[cfg(test)]
use persistence::{find_toml_table_bounds, render_access_section, save_sections_to_disk};
pub(in crate::api) use persistence::{
render_server_listeners, render_top_level_section, save_access_sections_to_disk,
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)]
pub(super) enum AccessSection {
+8 -15
View File
@@ -86,9 +86,9 @@ pub(in crate::api) async fn write_atomic(
let _lock = ConfigWriteLock::acquire(&path)?;
write_atomic_sync(&path, None, &contents, None).map(|_| ())
})
.await
.map_err(|error| ApiFailure::internal(format!("failed to join writer: {error}")))?
.map_err(|error| ApiFailure::internal(format!("failed to write config: {error}")))
.await
.map_err(|error| ApiFailure::internal(format!("failed to join writer: {error}")))?
.map_err(|error| ApiFailure::internal(format!("failed to write config: {error}")))
}
/// Replaces one source only if both its graph revision and owner contents are unchanged.
@@ -294,11 +294,7 @@ fn write_atomic_sync(
let descriptor = openat(
anchored.parent(),
temp_name.as_str(),
OFlag::O_WRONLY
| OFlag::O_CREAT
| OFlag::O_EXCL
| OFlag::O_NOFOLLOW
| OFlag::O_CLOEXEC,
OFlag::O_WRONLY | OFlag::O_CREAT | OFlag::O_EXCL | OFlag::O_NOFOLLOW | OFlag::O_CLOEXEC,
Mode::from_bits_truncate(0o600),
)
.map_err(errno_to_io)?;
@@ -408,9 +404,9 @@ fn validate_expected_contents(
existing: Option<&ExistingTarget>,
expected_contents: Option<&str>,
) -> std::io::Result<()> {
if expected_contents.is_some_and(|expected| {
existing.is_none_or(|target| target.contents != expected)
}) {
if expected_contents
.is_some_and(|expected| existing.is_none_or(|target| target.contents != expected))
{
return Err(std::io::Error::new(
std::io::ErrorKind::AlreadyExists,
"config source changed before persistence",
@@ -419,10 +415,7 @@ fn validate_expected_contents(
Ok(())
}
fn target_unchanged(
existing: Option<&ExistingTarget>,
current: Option<&ExistingTarget>,
) -> bool {
fn target_unchanged(existing: Option<&ExistingTarget>, current: Option<&ExistingTarget>) -> bool {
match (existing, current) {
(Some(expected), Some(current)) => {
same_target(&expected.metadata, &current.metadata)
+3 -3
View File
@@ -6,15 +6,15 @@ use serde::Serialize;
use crate::config::{ProxyConfig, RateLimitBps};
#[cfg(test)]
use super::atomic::write_atomic;
use super::atomic::write_atomic_if_unchanged;
#[cfg(test)]
use super::compute_revision;
use super::{
AccessSection, compute_snapshot_revision, load_candidate_snapshot, load_config_snapshot,
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;
/// 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 included = dir.path().join("users.toml");
let root_body = "include = \"users.toml\"\n[censorship]\ntls_domain = \"one.example\"\n";
let external_root =
"include = \"users.toml\"\n[censorship]\ntls_domain = \"two.example\"\n";
let external_root = "include = \"users.toml\"\n[censorship]\ntls_domain = \"two.example\"\n";
let included_body = "[access.users]\nalice = \"00000000000000000000000000000000\"\n";
tokio::fs::write(&root, root_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();
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!(
tokio::fs::read_to_string(&included).await.unwrap(),
included_body
+6 -8
View File
@@ -103,10 +103,9 @@ pub(super) async fn handle(
};
let runtime_cfg = config_rx.borrow().clone();
data.in_runtime = runtime_cfg.access.users.contains_key(&data.username);
shared.runtime_events.record(
"api.user.disable.ok",
format!("username={}", base_user),
);
shared
.runtime_events
.record("api.user.disable.ok", format!("username={}", base_user));
let status = if data.in_runtime {
StatusCode::OK
} else {
@@ -328,10 +327,9 @@ pub(super) async fn handle(
return Err(error);
}
};
shared.runtime_events.record(
"api.user.delete.ok",
format!("username={}", deleted_user),
);
shared
.runtime_events
.record("api.user.delete.ok", format!("username={}", deleted_user));
let runtime_cfg = config_rx.borrow().clone();
let in_runtime = runtime_cfg.access.users.contains_key(&deleted_user);
let response = DeleteUserResponse {
+1 -3
View File
@@ -314,9 +314,7 @@ async fn recompute_connections_payload(
let mut active_users = 0usize;
for entry in shared.stats.iter_user_stats() {
let user_stats = entry.value();
let current_connections = shared
.stats
.get_process_user_curr_connects(entry.key());
let current_connections = shared.stats.get_process_user_curr_connects(entry.key());
let total_octets = user_stats
.octets_from_client
.load(std::sync::atomic::Ordering::Relaxed)
+2 -4
View File
@@ -4,7 +4,7 @@ use std::collections::BTreeSet;
use serde::Serialize;
use super::{now_epoch_secs, ApiShared, SOURCE_UNAVAILABLE_REASON};
use super::{ApiShared, SOURCE_UNAVAILABLE_REASON, now_epoch_secs};
#[derive(Serialize)]
struct RuntimeMePoolStateGenerationData {
@@ -90,9 +90,7 @@ struct RuntimeMePoolStateData {
}
/// Builds the bounded runtime ME pool response projection.
pub(in crate::api) async fn build_runtime_me_pool_state_data(
shared: &ApiShared,
) -> impl Serialize {
pub(in crate::api) async fn build_runtime_me_pool_state_data(shared: &ApiShared) -> impl Serialize {
let now_epoch_secs = now_epoch_secs();
let Some(pool) = shared.me_pool.read().await.clone() else {
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_reconcile_success_total: stats
.get_conntrack_rule_reconcile_success_total(),
conntrack_rule_reconcile_error_total: stats
.get_conntrack_rule_reconcile_error_total(),
conntrack_rule_reconcile_error_total: stats.get_conntrack_rule_reconcile_error_total(),
conntrack_rule_rollback_success_total: stats
.get_conntrack_rule_rollback_success_total(),
conntrack_rule_rollback_error_total: stats
.get_conntrack_rule_rollback_error_total(),
conntrack_rule_rollback_error_total: stats.get_conntrack_rule_rollback_error_total(),
conntrack_delete_attempt_total: stats.get_conntrack_delete_attempt_total(),
conntrack_delete_success_total: stats.get_conntrack_delete_success_total(),
conntrack_delete_not_found_total: stats.get_conntrack_delete_not_found_total(),
+5 -7
View File
@@ -145,13 +145,11 @@ async fn create_user_to_completion(
Some(&base_revision),
)
.await?;
shared
.proxy_shared
.stage_user_credential(
&body.username,
credential_id,
cfg.access.is_user_enabled(&body.username),
);
shared.proxy_shared.stage_user_credential(
&body.username,
credential_id,
cfg.access.is_user_enabled(&body.username),
);
if let Some(limit) = updated_limit {
shared
+5 -3
View File
@@ -54,9 +54,11 @@ async fn rotate_secret_to_completion(
Some(&base_revision),
)
.await?;
shared
.proxy_shared
.stage_user_credential(user, credential_id, cfg.access.is_user_enabled(user));
shared.proxy_shared.stage_user_credential(
user,
credential_id,
cfg.access.is_user_enabled(user),
);
drop(_guard);
let (detected_ip_v4, detected_ip_v6) = shared.detected_link_ips();
+18 -16
View File
@@ -154,19 +154,19 @@ async fn patch_user_to_completion(
cfg.validate()
.map_err(|e| ApiFailure::bad_request(format!("config validation failed: {}", e)))?;
let staged_credential = if touches_users || touches_user_enabled {
let secret = cfg
.access
.users
.get(user)
.ok_or_else(|| ApiFailure::internal("updated user secret is missing"))?;
Some(
credential_id_from_hex(secret)
.ok_or_else(|| ApiFailure::internal("validated user secret could not be decoded"))?,
)
} else {
None
};
let staged_credential =
if touches_users || touches_user_enabled {
let secret = cfg
.access
.users
.get(user)
.ok_or_else(|| ApiFailure::internal("updated user secret is missing"))?;
Some(credential_id_from_hex(secret).ok_or_else(|| {
ApiFailure::internal("validated user secret could not be decoded")
})?)
} else {
None
};
let mut touched_sections = Vec::new();
if touches_users {
@@ -206,9 +206,11 @@ async fn patch_user_to_completion(
.await?
};
if let Some(credential_id) = staged_credential {
shared
.proxy_shared
.stage_user_credential(user, credential_id, cfg.access.is_user_enabled(user));
shared.proxy_shared.stage_user_credential(
user,
credential_id,
cfg.access.is_user_enabled(user),
);
}
match max_unique_ips_change {
Some(Some(limit)) => shared.ip_tracker.set_user_limit(user, limit).await,
+1 -1
View File
@@ -14,7 +14,7 @@ use std::os::unix::fs::MetadataExt;
#[cfg(unix)]
use nix::dir::Dir;
#[cfg(unix)]
use nix::fcntl::{openat, OFlag};
use nix::fcntl::{OFlag, openat};
#[cfg(unix)]
use nix::sys::stat::Mode;
@@ -124,7 +124,15 @@ fn load_static_directory(
let relative = path.strip_prefix(root).map_err(|_| {
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(())
}
+1 -1
View File
@@ -13,10 +13,10 @@ use crate::stats::Stats;
// Privileged netfilter rule and conntrack helper execution.
mod firewall;
pub(crate) use firewall::FirewallAuthority;
use firewall::{
DeleteOutcome, delete_conntrack_entry, effective_conntrack_enabled, probe_runtime_support,
};
pub(crate) use firewall::FirewallAuthority;
const CONNTRACK_EVENT_QUEUE_CAPACITY: usize = 32_768;
const PRESSURE_RELEASE_TICKS: u8 = 3;
+16 -20
View File
@@ -236,21 +236,15 @@ where
}
let desired = current.as_ref().expect("desired state is present").clone();
let interruptible = InterruptibleRunner::new(
&self.runner,
&self.terminal,
&process_cancellation,
);
let result = reconcile_once(
&interruptible,
&interruptible,
&mut self.applied,
&desired,
)
.await;
let interruptible =
InterruptibleRunner::new(&self.runner, &self.terminal, &process_cancellation);
let result =
reconcile_once(&interruptible, &interruptible, &mut self.applied, &desired).await;
match result {
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);
self.status_tx.send_replace(Some(ReconcileStatus {
generation: desired.generation,
@@ -265,7 +259,9 @@ where
}
Err(failure) if failure.cancelled => break,
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);
if let Some(rollback_succeeded) = failure.rollback_succeeded {
if rollback_succeeded {
@@ -326,12 +322,12 @@ where
if let Some(stats) = &self.last_stats {
stats.set_conntrack_rule_apply_ok(false);
}
if let Err(error) = tokio::time::timeout(
SHUTDOWN_CLEANUP_TIMEOUT,
recover_to_empty(&self.runner),
)
.await
.unwrap_or_else(|_| Err(CommandError::failed("firewall shutdown cleanup timed out")))
if let Err(error) =
tokio::time::timeout(SHUTDOWN_CLEANUP_TIMEOUT, recover_to_empty(&self.runner))
.await
.unwrap_or_else(|_| {
Err(CommandError::failed("firewall shutdown cleanup timed out"))
})
{
warn!(error = %error, "Failed to clear conntrack firewall policy during shutdown");
} else {
+14 -29
View File
@@ -1,6 +1,4 @@
use super::command::{
CommandError, CommandErrorKind, CommandSpec, FirewallCommandRunner,
};
use super::command::{CommandError, CommandErrorKind, CommandSpec, FirewallCommandRunner};
use super::model::{NotrackTarget, ShadowSlot};
const DISPATCH_CHAIN: &str = "TELEMT_NOTRACK";
@@ -30,10 +28,7 @@ impl IpFamily {
}
}
pub(super) fn family_available<R: FirewallCommandRunner>(
runner: &R,
family: IpFamily,
) -> bool {
pub(super) fn family_available<R: FirewallCommandRunner>(runner: &R, family: IpFamily) -> bool {
runner.available(family.command_binary()) && runner.available(family.restore_binary())
}
@@ -83,9 +78,7 @@ pub(super) async fn activate_family<R: FirewallCommandRunner>(
.await
}
pub(super) async fn cleanup_all<R: FirewallCommandRunner>(
runner: &R,
) -> Result<(), CommandError> {
pub(super) async fn cleanup_all<R: FirewallCommandRunner>(runner: &R) -> Result<(), CommandError> {
let mut errors = Vec::new();
for family in [IpFamily::V4, IpFamily::V6] {
if !runner.available(family.command_binary()) {
@@ -118,7 +111,10 @@ async fn cleanup_family<R: FirewallCommandRunner>(
match result {
Ok(()) => {}
Err(error)
if matches!(error.kind, CommandErrorKind::NotFound | CommandErrorKind::Missing) =>
if matches!(
error.kind,
CommandErrorKind::NotFound | CommandErrorKind::Missing
) =>
{
break;
}
@@ -131,13 +127,13 @@ async fn cleanup_family<R: FirewallCommandRunner>(
for chain in [DISPATCH_CHAIN, SHADOW_CHAIN_A, SHADOW_CHAIN_B] {
for operation in ["-F", "-X"] {
let result = runner
.run(CommandSpec::new(
binary,
["-t", "raw", operation, chain],
))
.run(CommandSpec::new(binary, ["-t", "raw", operation, chain]))
.await;
if let Err(error) = result
&& !matches!(error.kind, CommandErrorKind::NotFound | CommandErrorKind::Missing)
&& !matches!(
error.kind,
CommandErrorKind::NotFound | CommandErrorKind::Missing
)
{
errors.push(error.message);
}
@@ -167,15 +163,7 @@ async fn ensure_prerouting_jump<R: FirewallCommandRunner>(
runner
.run(CommandSpec::new(
binary,
[
"-t",
"raw",
"-I",
"PREROUTING",
"1",
"-j",
DISPATCH_CHAIN,
],
["-t", "raw", "-I", "PREROUTING", "1", "-j", DISPATCH_CHAIN],
))
.await
}
@@ -220,10 +208,7 @@ fn require_family<R: FirewallCommandRunner>(
Ok(())
}
pub(super) fn render_stage_script(
slot: ShadowSlot,
targets: &[NotrackTarget],
) -> String {
pub(super) fn render_stage_script(slot: ShadowSlot, targets: &[NotrackTarget]) -> String {
let chain = shadow_chain(slot);
let mut script = format!("*raw\n-F {chain}\n");
for target in targets {
-1
View File
@@ -18,7 +18,6 @@ impl ShadowSlot {
Self::B => Self::A,
}
}
}
#[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
}
pub(super) async fn cleanup_all<R: FirewallCommandRunner>(
runner: &R,
) -> Result<(), CommandError> {
pub(super) async fn cleanup_all<R: FirewallCommandRunner>(runner: &R) -> Result<(), CommandError> {
if !runner.available("nft") {
return Ok(());
}
@@ -86,7 +84,10 @@ async fn delete_table_if_present<R: FirewallCommandRunner>(
{
Ok(()) => Ok(()),
Err(error)
if matches!(error.kind, CommandErrorKind::NotFound | CommandErrorKind::Missing) =>
if matches!(
error.kind,
CommandErrorKind::NotFound | CommandErrorKind::Missing
) =>
{
Ok(())
}
@@ -111,9 +112,7 @@ pub(super) fn render_stage_script(
v6: &[NotrackTarget],
) -> String {
let table = table(slot);
let mut script = format!(
"add table inet {table}\nadd chain inet {table} rules\n"
);
let mut script = format!("add table inet {table}\nadd chain inet {table} rules\n");
for target in v4 {
script.push_str("add rule inet ");
script.push_str(table);
+3 -8
View File
@@ -1,7 +1,7 @@
use std::collections::{BTreeMap, BTreeSet};
use std::future::pending;
use std::sync::{Arc, Mutex};
use std::sync::atomic::Ordering;
use std::sync::{Arc, Mutex};
use std::time::Duration;
use tokio::sync::{Notify, watch};
@@ -102,8 +102,7 @@ impl FirewallCommandRunner for FakeRunner {
let operation = spec.args.get(2).map(String::as_str);
if (matches!(spec.binary, "iptables" | "ip6tables")
&& matches!(operation, Some("-C" | "-D" | "-F" | "-X")))
|| (spec.binary == "nft"
&& spec.args.first().map(String::as_str) == Some("delete"))
|| (spec.binary == "nft" && spec.args.first().map(String::as_str) == Some("delete"))
{
return Err(CommandError {
kind: CommandErrorKind::NotFound,
@@ -344,11 +343,7 @@ async fn transaction_cancellation_does_not_claim_a_new_applied_plan() {
};
let terminal = CancellationToken::new();
let process_cancellation = CancellationToken::new();
let interruptible = InterruptibleRunner::new(
&runner,
&terminal,
&process_cancellation,
);
let interruptible = InterruptibleRunner::new(&runner, &terminal, &process_cancellation);
let mut applied = AppliedState::Known(AppliedPlan::Empty);
let desired = desired(1, dual_stack_policy(443));
let failure = {
@@ -73,14 +73,9 @@ fn hybrid_policy_is_a_sorted_deduplicated_address_port_product() {
#[test]
fn restore_renderers_keep_staging_detached_from_activation() {
let stage = iptables::render_stage_script(
ShadowSlot::B,
&[target(Some("192.0.2.20"), 443)],
);
let stage = iptables::render_stage_script(ShadowSlot::B, &[target(Some("192.0.2.20"), 443)]);
assert!(stage.contains("-F TELEMT_NT_B\n"));
assert!(stage.contains(
"-A TELEMT_NT_B -p tcp --dport 443 -d 192.0.2.20 -j CT --notrack\n"
));
assert!(stage.contains("-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(":TELEMT_"));
@@ -4,9 +4,7 @@ use tokio_util::sync::CancellationToken;
use crate::config::ConntrackBackend;
use super::command::{
CommandError, CommandErrorKind, CommandSpec, FirewallCommandRunner,
};
use super::command::{CommandError, CommandErrorKind, CommandSpec, FirewallCommandRunner};
use super::iptables::{self, IpFamily};
use super::model::{AppliedPlan, AppliedState, DesiredPolicy, DesiredState, ShadowSlot};
use super::nftables;
@@ -251,9 +249,13 @@ pub(super) async fn transition_plan<R: FirewallCommandRunner>(
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?;
}
(_, _) if previous != target => clear_plan(runner, previous).await?,
+19 -33
View File
@@ -1,9 +1,9 @@
use std::ffi::OsStr;
use std::fs::{self, File};
use std::io::{self, ErrorKind, Read, Write};
use std::os::unix::fs::{MetadataExt, PermissionsExt};
#[cfg(target_os = "linux")]
use std::os::fd::{FromRawFd, OwnedFd};
use std::os::unix::fs::{MetadataExt, PermissionsExt};
use std::path::{Path, PathBuf};
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.
pub fn acquire(&mut self) -> Result<(), DaemonError> {
let anchor = AnchoredPath::open_trusted_parent_or_create(&self.path, 0o755).map_err(
|error| {
let anchor =
AnchoredPath::open_trusted_parent_or_create(&self.path, 0o755).map_err(|error| {
DaemonError::PidFile(format!(
"cannot open trusted parent for {}: {}",
self.path.display(),
error
))
},
)?;
})?;
let lock_name = self.lock_path.file_name().ok_or_else(|| {
DaemonError::PidFile(format!(
"lock path {} has no file name",
@@ -132,7 +131,11 @@ impl PidFile {
// Validate the opened inode before modifying it so a hard-link substitution
// cannot turn PID publication into truncation of an unrelated file.
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();
writeln!(pid_file, "{}", pid).map_err(|error| {
@@ -236,19 +239,9 @@ fn normalize_pid_path(path: &Path) -> PathBuf {
}
}
fn open_file_at(
anchor: &AnchoredPath,
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))?;
fn open_file_at(anchor: &AnchoredPath, 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))?;
Ok(File::from(descriptor))
}
@@ -337,12 +330,7 @@ fn remove_owned_pid_file(
)));
}
drop(file);
unlinkat(
anchor.parent(),
anchor.name(),
UnlinkatFlags::NoRemoveDir,
)
.map_err(|error| {
unlinkat(anchor.parent(), anchor.name(), UnlinkatFlags::NoRemoveDir).map_err(|error| {
DaemonError::PidFile(format!(
"cannot remove {}: {}",
path.display(),
@@ -351,10 +339,7 @@ fn remove_owned_pid_file(
})
}
fn validate_regular_single_link(
file: &File,
path: &Path,
) -> Result<fs::Metadata, DaemonError> {
fn validate_regular_single_link(file: &File, path: &Path) -> Result<fs::Metadata, DaemonError> {
let metadata = file.metadata().map_err(|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 {
let path = normalize_pid_path(path.as_ref());
match read_pid_file_if_exists(&path) {
Ok(Some(pid))
if daemon_lock_is_held(&path).unwrap_or(false) && is_process_running(pid) =>
{
Ok(Some(pid)) if daemon_lock_is_held(&path).unwrap_or(false) && is_process_running(pid) => {
DaemonStatus::Running(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(|| {
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(
&anchor,
+5 -1
View File
@@ -238,7 +238,11 @@ fn release_does_not_remove_replacement_path() {
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");
}
+7 -3
View File
@@ -41,8 +41,7 @@ struct CleanupShard {
queue: Mutex<CleanupQueue>,
}
type CleanupQueue =
HashMap<String, HashMap<UserIncarnation, HashMap<IpAddr, usize>>>;
type CleanupQueue = HashMap<String, HashMap<UserIncarnation, HashMap<IpAddr, usize>>>;
type CleanupBatch = HashMap<(String, UserIncarnation, IpAddr), usize>;
#[derive(Debug, Clone)]
@@ -208,7 +207,12 @@ impl UserIpTracker {
) -> Option<(String, UserIncarnation, IpAddr, usize)> {
let user = queue.keys().next().cloned()?;
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 ips = incarnations.get_mut(&incarnation)?;
let count = ips.remove(&ip)?;
+2 -10
View File
@@ -56,10 +56,7 @@ impl Drop for DetachedCleanupBatch<'_> {
);
}
}
UserIpTracker::decrement_counter(
&self.tracker.cleanup_queue_len,
duplicate_entries,
);
UserIpTracker::decrement_counter(&self.tracker.cleanup_queue_len, duplicate_entries);
}
}
@@ -187,12 +184,7 @@ impl UserIpTracker {
continue;
}
removed_active_entries = removed_active_entries.saturating_add(
Self::apply_active_cleanup(
&mut shard.active_ips,
queued_user,
*ip,
*pending_count,
),
Self::apply_active_cleanup(&mut shard.active_ips, queued_user, *ip, *pending_count),
);
}
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)
.await
.unwrap();
tracker
.clear_user_ips_if_not_newer("test_user", 2)
.await;
tracker.clear_user_ips_if_not_newer("test_user", 2).await;
tracker
.check_and_add_for_incarnation("test_user", 3, current_ip)
.await
@@ -275,13 +273,7 @@ async fn stale_runtime_cannot_overwrite_newer_ip_policy() {
newer.insert("alice".to_string(), 5);
assert!(
tracker
.apply_policy_from_source(
2,
7,
&newer,
UserMaxUniqueIpsMode::Combined,
90,
)
.apply_policy_from_source(2, 7, &newer, UserMaxUniqueIpsMode::Combined, 90,)
.await
);
@@ -289,13 +281,7 @@ async fn stale_runtime_cannot_overwrite_newer_ip_policy() {
stale.insert("alice".to_string(), 1);
assert!(
!tracker
.apply_policy_from_source(
1,
1,
&stale,
UserMaxUniqueIpsMode::ActiveWindow,
1,
)
.apply_policy_from_source(1, 1, &stale, UserMaxUniqueIpsMode::ActiveWindow, 1,)
.await
);
@@ -326,13 +312,7 @@ async fn active_runtime_can_publish_coherent_same_generation_ip_policy() {
assert!(
tracker
.apply_policy_from_source(
3,
6,
&limits,
UserMaxUniqueIpsMode::TimeWindow,
30,
)
.apply_policy_from_source(3, 6, &limits, UserMaxUniqueIpsMode::TimeWindow, 30,)
.await
);
+10 -8
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::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();
}
@@ -86,16 +89,15 @@ fn clear_all_serializes_queue_reset_with_concurrent_enqueue() {
let enqueue_tracker = Arc::clone(&tracker);
let enqueue = std::thread::spawn(move || {
started_tx.send(()).unwrap();
enqueue_tracker.enqueue_cleanup(
first_shard_user,
test_ipv4(10, 2, 2, 1),
);
enqueue_tracker.enqueue_cleanup(first_shard_user, test_ipv4(10, 2, 2, 1));
completed_tx.send(()).unwrap();
});
started_rx.recv().unwrap();
assert!(completed_rx
.recv_timeout(Duration::from_millis(50))
.is_err());
assert!(
completed_rx
.recv_timeout(Duration::from_millis(50))
.is_err()
);
drop(last_queue_guard);
clear.join().unwrap();
+2 -2
View File
@@ -142,8 +142,8 @@ pub fn init_logging(
}
LogDestination::File { options } => {
let file_appender = file::BoundedFileAppender::new(options.clone())
.expect("Failed to open log file");
let file_appender =
file::BoundedFileAppender::new(options.clone()).expect("Failed to open log file");
let (non_blocking, guard) = tracing_appender::non_blocking(file_appender);
let fmt_layer = fmt::Layer::default()
+9 -15
View File
@@ -1,6 +1,6 @@
use std::fs::{self, File};
#[cfg(not(unix))]
use std::fs::OpenOptions;
use std::fs::{self, File};
use std::io::{self, Write};
use std::path::{Path, PathBuf};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
@@ -223,12 +223,7 @@ impl BoundedFileAppender {
let archive_name = archive_path.file_name().ok_or_else(|| {
io::Error::new(io::ErrorKind::InvalidInput, "archive path has no file name")
})?;
match renameat(
&self.dir_fd,
current_name,
&self.dir_fd,
archive_name,
) {
match renameat(&self.dir_fd, current_name, &self.dir_fd, archive_name) {
Ok(()) => Ok(()),
Err(nix::errno::Errno::ENOENT) => Ok(()),
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>> {
use std::os::unix::fs::MetadataExt;
let descriptor = dup(&self.dir_fd).map_err(|error| {
io::Error::from_raw_os_error(error as i32)
})?;
let mut directory = Dir::from_fd(descriptor)
.map_err(|error| io::Error::from_raw_os_error(error as i32))?;
let descriptor =
dup(&self.dir_fd).map_err(|error| io::Error::from_raw_os_error(error as i32))?;
let mut directory =
Dir::from_fd(descriptor).map_err(|error| io::Error::from_raw_os_error(error as i32))?;
let mut candidates = Vec::new();
let prefix = format!("{}.", self.base_name);
for entry in directory.iter().flatten() {
@@ -375,9 +369,9 @@ struct LogFileCandidate {
#[cfg(unix)]
fn open_append_file(dir_fd: &OwnedFd, path: &Path) -> io::Result<(File, u64)> {
let name = path.file_name().ok_or_else(|| {
io::Error::new(io::ErrorKind::InvalidInput, "log path has no file name")
})?;
let name = path
.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 current_size = file.metadata()?.len();
Ok((file, current_size))
+2 -5
View File
@@ -137,10 +137,7 @@ fn appender_rejects_group_writable_log_directory() {
fs::set_permissions(dir.path(), fs::Permissions::from_mode(0o770)).unwrap();
assert!(
BoundedFileAppender::with_now(
options(dir.path().join("telemt.log")),
Box::new(fixed_now),
)
.is_err()
BoundedFileAppender::with_now(options(dir.path().join("telemt.log")), Box::new(fixed_now),)
.is_err()
);
}
+4 -3
View File
@@ -137,10 +137,11 @@ mod tests {
let drain_generation = Arc::clone(&generation);
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)
& SESSION_ADMISSION_CLOSED
while generation.session_admission.state.load(Ordering::Acquire) & SESSION_ADMISSION_CLOSED
== 0
{
tokio::task::yield_now().await;
+261 -271
View File
@@ -1,276 +1,266 @@
use std::path::{Path, PathBuf};
use std::path::{Path, PathBuf};
use super::{
expected_handshake_close_description, format_maestro_line, is_expected_handshake_eof,
peer_close_description, resolve_runtime_base_dir, resolve_runtime_config_path,
};
use crate::error::{ProxyError, StreamError};
use super::{
expected_handshake_close_description, format_maestro_line, is_expected_handshake_eof,
peer_close_description, resolve_runtime_base_dir, resolve_runtime_config_path,
};
use crate::error::{ProxyError, StreamError};
#[test]
fn maestro_line_formatter_respects_disabled_colors() {
let plain = format_maestro_line("boot", false);
assert_eq!(plain, "MAESTRO: boot");
assert!(!plain.contains('\x1b'));
#[test]
fn maestro_line_formatter_respects_disabled_colors() {
let plain = format_maestro_line("boot", false);
assert_eq!(plain, "MAESTRO: boot");
assert!(!plain.contains('\x1b'));
}
#[test]
fn maestro_line_formatter_keeps_color_when_enabled() {
let colored = format_maestro_line("boot", true);
assert!(colored.contains("\x1b[92mMAESTRO\x1b[0m"));
}
#[test]
fn resolve_runtime_config_path_anchors_relative_to_startup_cwd() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_path_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
let target = startup_cwd.join("config.toml");
std::fs::write(&target, " ").unwrap();
let resolved = resolve_runtime_config_path("config.toml", &startup_cwd, true);
assert_eq!(resolved, target.canonicalize().unwrap());
let _ = std::fs::remove_file(&target);
let _ = std::fs::remove_dir(&startup_cwd);
}
#[test]
fn resolve_runtime_config_path_keeps_absolute_for_missing_file() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_path_missing_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
let resolved = resolve_runtime_config_path("missing.toml", &startup_cwd, true);
assert_eq!(resolved, startup_cwd.join("missing.toml"));
let _ = std::fs::remove_dir(&startup_cwd);
}
#[cfg(unix)]
#[test]
fn runtime_paths_preserve_symlinks_for_descriptor_validation() {
use std::os::unix::fs::symlink;
let dir = tempfile::tempdir().unwrap();
let real_dir = dir.path().join("real");
let linked_dir = dir.path().join("linked");
std::fs::create_dir(&real_dir).unwrap();
std::fs::write(real_dir.join("config.toml"), " ").unwrap();
symlink(&real_dir, &linked_dir).unwrap();
let linked_config = linked_dir.join("config.toml");
let config = resolve_runtime_config_path(linked_config.to_str().unwrap(), dir.path(), true);
let runtime = resolve_runtime_base_dir(&linked_config, dir.path(), true, Some(&linked_dir));
assert_eq!(config, linked_config);
assert_eq!(runtime, linked_dir);
}
#[test]
fn resolve_runtime_config_path_uses_startup_candidates_when_not_explicit() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_startup_candidates_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
let telemt = startup_cwd.join("telemt.toml");
std::fs::write(&telemt, " ").unwrap();
let resolved = resolve_runtime_config_path("config.toml", &startup_cwd, false);
assert_eq!(resolved, telemt.canonicalize().unwrap());
let _ = std::fs::remove_file(&telemt);
let _ = std::fs::remove_dir(&startup_cwd);
}
#[test]
fn resolve_runtime_config_path_defaults_to_startup_config_when_none_found() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_startup_default_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
let resolved = resolve_runtime_config_path("config.toml", &startup_cwd, false);
assert_eq!(resolved, startup_cwd.join("config.toml"));
let _ = std::fs::remove_dir(&startup_cwd);
}
#[test]
fn resolve_runtime_base_dir_prefers_cli_data_path() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd = std::env::temp_dir().join(format!("telemt_runtime_base_cwd_{nonce}"));
let data_path = std::env::temp_dir().join(format!("telemt_runtime_base_data_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
std::fs::create_dir_all(&data_path).unwrap();
let resolved = resolve_runtime_base_dir(
&startup_cwd.join("config.toml"),
&startup_cwd,
true,
Some(&data_path),
);
assert_eq!(resolved, data_path.canonicalize().unwrap());
let _ = std::fs::remove_dir(&data_path);
let _ = std::fs::remove_dir(&startup_cwd);
}
#[test]
fn resolve_runtime_base_dir_uses_working_directory_before_explicit_config_parent() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd = std::env::temp_dir().join(format!("telemt_runtime_base_start_{nonce}"));
let config_dir = std::env::temp_dir().join(format!("telemt_runtime_base_cfg_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
std::fs::create_dir_all(&config_dir).unwrap();
let resolved =
resolve_runtime_base_dir(&config_dir.join("telemt.toml"), &startup_cwd, true, None);
assert_eq!(resolved, startup_cwd.canonicalize().unwrap());
let _ = std::fs::remove_dir(&config_dir);
let _ = std::fs::remove_dir(&startup_cwd);
}
#[test]
fn resolve_runtime_base_dir_uses_explicit_config_parent_from_root() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let config_dir = std::env::temp_dir().join(format!("telemt_runtime_base_root_cfg_{nonce}"));
std::fs::create_dir_all(&config_dir).unwrap();
let resolved =
resolve_runtime_base_dir(&config_dir.join("telemt.toml"), Path::new("/"), true, None);
assert_eq!(resolved, config_dir.canonicalize().unwrap());
let _ = std::fs::remove_dir(&config_dir);
}
#[test]
fn resolve_runtime_base_dir_uses_systemd_working_directory_before_etc() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd = std::env::temp_dir().join(format!("telemt_runtime_base_systemd_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
let resolved =
resolve_runtime_base_dir(&startup_cwd.join("config.toml"), &startup_cwd, false, None);
assert_eq!(resolved, startup_cwd.canonicalize().unwrap());
let _ = std::fs::remove_dir(&startup_cwd);
}
#[test]
fn resolve_runtime_base_dir_falls_back_to_etc_from_root() {
let resolved = resolve_runtime_base_dir(
Path::new("/etc/telemt/config.toml"),
Path::new("/"),
false,
None,
);
assert_eq!(resolved, PathBuf::from("/etc/telemt"));
}
#[test]
fn expected_handshake_eof_matches_connection_reset() {
let err = ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionReset));
assert!(is_expected_handshake_eof(&err));
}
#[test]
fn expected_handshake_eof_matches_stream_io_unexpected_eof() {
let err = ProxyError::Stream(StreamError::Io(std::io::Error::from(
std::io::ErrorKind::UnexpectedEof,
)));
assert!(is_expected_handshake_eof(&err));
}
#[test]
fn peer_close_description_is_human_readable_for_all_peer_close_kinds() {
let cases = [
(
std::io::ErrorKind::ConnectionReset,
"Peer reset TCP connection (RST)",
),
(
std::io::ErrorKind::ConnectionAborted,
"Peer aborted TCP connection during transport",
),
(
std::io::ErrorKind::BrokenPipe,
"Peer closed write side (broken pipe)",
),
(
std::io::ErrorKind::NotConnected,
"Socket was already closed by peer",
),
];
for (kind, expected) in cases {
let err = ProxyError::Io(std::io::Error::from(kind));
assert_eq!(peer_close_description(&err), Some(expected));
}
}
#[test]
fn maestro_line_formatter_keeps_color_when_enabled() {
let colored = format_maestro_line("boot", true);
assert!(colored.contains("\x1b[92mMAESTRO\x1b[0m"));
}
#[test]
fn resolve_runtime_config_path_anchors_relative_to_startup_cwd() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_path_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
let target = startup_cwd.join("config.toml");
std::fs::write(&target, " ").unwrap();
let resolved = resolve_runtime_config_path("config.toml", &startup_cwd, true);
assert_eq!(resolved, target.canonicalize().unwrap());
let _ = std::fs::remove_file(&target);
let _ = std::fs::remove_dir(&startup_cwd);
}
#[test]
fn resolve_runtime_config_path_keeps_absolute_for_missing_file() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_path_missing_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
let resolved = resolve_runtime_config_path("missing.toml", &startup_cwd, true);
assert_eq!(resolved, startup_cwd.join("missing.toml"));
let _ = std::fs::remove_dir(&startup_cwd);
}
#[cfg(unix)]
#[test]
fn runtime_paths_preserve_symlinks_for_descriptor_validation() {
use std::os::unix::fs::symlink;
let dir = tempfile::tempdir().unwrap();
let real_dir = dir.path().join("real");
let linked_dir = dir.path().join("linked");
std::fs::create_dir(&real_dir).unwrap();
std::fs::write(real_dir.join("config.toml"), " ").unwrap();
symlink(&real_dir, &linked_dir).unwrap();
let linked_config = linked_dir.join("config.toml");
let config = resolve_runtime_config_path(
linked_config.to_str().unwrap(),
dir.path(),
true,
);
let runtime = resolve_runtime_base_dir(
&linked_config,
dir.path(),
true,
Some(&linked_dir),
);
assert_eq!(config, linked_config);
assert_eq!(runtime, linked_dir);
}
#[test]
fn resolve_runtime_config_path_uses_startup_candidates_when_not_explicit() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd =
std::env::temp_dir().join(format!("telemt_cfg_startup_candidates_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
let telemt = startup_cwd.join("telemt.toml");
std::fs::write(&telemt, " ").unwrap();
let resolved = resolve_runtime_config_path("config.toml", &startup_cwd, false);
assert_eq!(resolved, telemt.canonicalize().unwrap());
let _ = std::fs::remove_file(&telemt);
let _ = std::fs::remove_dir(&startup_cwd);
}
#[test]
fn resolve_runtime_config_path_defaults_to_startup_config_when_none_found() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd = std::env::temp_dir().join(format!("telemt_cfg_startup_default_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
let resolved = resolve_runtime_config_path("config.toml", &startup_cwd, false);
assert_eq!(resolved, startup_cwd.join("config.toml"));
let _ = std::fs::remove_dir(&startup_cwd);
}
#[test]
fn resolve_runtime_base_dir_prefers_cli_data_path() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd = std::env::temp_dir().join(format!("telemt_runtime_base_cwd_{nonce}"));
let data_path = std::env::temp_dir().join(format!("telemt_runtime_base_data_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
std::fs::create_dir_all(&data_path).unwrap();
let resolved = resolve_runtime_base_dir(
&startup_cwd.join("config.toml"),
&startup_cwd,
true,
Some(&data_path),
);
assert_eq!(resolved, data_path.canonicalize().unwrap());
let _ = std::fs::remove_dir(&data_path);
let _ = std::fs::remove_dir(&startup_cwd);
}
#[test]
fn resolve_runtime_base_dir_uses_working_directory_before_explicit_config_parent() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd = std::env::temp_dir().join(format!("telemt_runtime_base_start_{nonce}"));
let config_dir = std::env::temp_dir().join(format!("telemt_runtime_base_cfg_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
std::fs::create_dir_all(&config_dir).unwrap();
let resolved =
resolve_runtime_base_dir(&config_dir.join("telemt.toml"), &startup_cwd, true, None);
assert_eq!(resolved, startup_cwd.canonicalize().unwrap());
let _ = std::fs::remove_dir(&config_dir);
let _ = std::fs::remove_dir(&startup_cwd);
}
#[test]
fn resolve_runtime_base_dir_uses_explicit_config_parent_from_root() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let config_dir = std::env::temp_dir().join(format!("telemt_runtime_base_root_cfg_{nonce}"));
std::fs::create_dir_all(&config_dir).unwrap();
let resolved =
resolve_runtime_base_dir(&config_dir.join("telemt.toml"), Path::new("/"), true, None);
assert_eq!(resolved, config_dir.canonicalize().unwrap());
let _ = std::fs::remove_dir(&config_dir);
}
#[test]
fn resolve_runtime_base_dir_uses_systemd_working_directory_before_etc() {
let nonce = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let startup_cwd = std::env::temp_dir().join(format!("telemt_runtime_base_systemd_{nonce}"));
std::fs::create_dir_all(&startup_cwd).unwrap();
let resolved =
resolve_runtime_base_dir(&startup_cwd.join("config.toml"), &startup_cwd, false, None);
assert_eq!(resolved, startup_cwd.canonicalize().unwrap());
let _ = std::fs::remove_dir(&startup_cwd);
}
#[test]
fn resolve_runtime_base_dir_falls_back_to_etc_from_root() {
let resolved = resolve_runtime_base_dir(
Path::new("/etc/telemt/config.toml"),
Path::new("/"),
false,
None,
);
assert_eq!(resolved, PathBuf::from("/etc/telemt"));
}
#[test]
fn expected_handshake_eof_matches_connection_reset() {
let err = ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionReset));
assert!(is_expected_handshake_eof(&err));
}
#[test]
fn expected_handshake_eof_matches_stream_io_unexpected_eof() {
let err = ProxyError::Stream(StreamError::Io(std::io::Error::from(
std::io::ErrorKind::UnexpectedEof,
)));
assert!(is_expected_handshake_eof(&err));
}
#[test]
fn peer_close_description_is_human_readable_for_all_peer_close_kinds() {
let cases = [
(
std::io::ErrorKind::ConnectionReset,
"Peer reset TCP connection (RST)",
),
(
std::io::ErrorKind::ConnectionAborted,
"Peer aborted TCP connection during transport",
),
(
std::io::ErrorKind::BrokenPipe,
"Peer closed write side (broken pipe)",
),
(
std::io::ErrorKind::NotConnected,
"Socket was already closed by peer",
),
];
for (kind, expected) in cases {
let err = ProxyError::Io(std::io::Error::from(kind));
assert_eq!(peer_close_description(&err), Some(expected));
}
}
#[test]
fn handshake_close_description_is_human_readable_for_all_expected_kinds() {
let cases = [
(
ProxyError::Io(std::io::Error::from(std::io::ErrorKind::UnexpectedEof)),
"Peer closed before sending full 64-byte MTProto handshake",
),
(
ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionReset)),
"Peer reset TCP connection during initial MTProto handshake",
),
(
ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionAborted)),
"Peer aborted TCP connection during initial MTProto handshake",
),
(
ProxyError::Io(std::io::Error::from(std::io::ErrorKind::BrokenPipe)),
"Peer closed write side before MTProto handshake completed",
),
(
ProxyError::Io(std::io::Error::from(std::io::ErrorKind::NotConnected)),
"Handshake socket was already closed by peer",
),
(
ProxyError::Stream(StreamError::UnexpectedEof),
"Peer closed before sending full 64-byte MTProto handshake",
),
];
for (err, expected) in cases {
assert_eq!(expected_handshake_close_description(&err), Some(expected));
}
#[test]
fn handshake_close_description_is_human_readable_for_all_expected_kinds() {
let cases = [
(
ProxyError::Io(std::io::Error::from(std::io::ErrorKind::UnexpectedEof)),
"Peer closed before sending full 64-byte MTProto handshake",
),
(
ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionReset)),
"Peer reset TCP connection during initial MTProto handshake",
),
(
ProxyError::Io(std::io::Error::from(std::io::ErrorKind::ConnectionAborted)),
"Peer aborted TCP connection during initial MTProto handshake",
),
(
ProxyError::Io(std::io::Error::from(std::io::ErrorKind::BrokenPipe)),
"Peer closed write side before MTProto handshake completed",
),
(
ProxyError::Io(std::io::Error::from(std::io::ErrorKind::NotConnected)),
"Handshake socket was already closed by peer",
),
(
ProxyError::Stream(StreamError::UnexpectedEof),
"Peer closed before sending full 64-byte MTProto handshake",
),
];
for (err, expected) in cases {
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(
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)
@@ -373,7 +376,12 @@ mod tests {
assert!(remove_stale_unix_socket(&regular).is_err());
assert!(remove_stale_unix_socket(&link).is_err());
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]
+1 -3
View File
@@ -49,9 +49,7 @@ pub(super) async fn run_telemt_core(
} = bootstrap::bootstrap(privilege_drop_requested).await?;
if privilege_drop_requested && config.server.conntrack_control.inline_conntrack_control {
warn!(
"Inline conntrack control is disabled when process privileges are dropped"
);
warn!("Inline conntrack control is disabled when process privileges are dropped");
config.server.conntrack_control.inline_conntrack_control = false;
}
+2 -2
View File
@@ -366,8 +366,8 @@ impl ReloadSupervisor {
self.runtime_watch_tx
.send_replace(Some(new_runtime.watch_state()));
if !conntrack_firewall_published {
let warning = "conntrack firewall reconciler is unavailable after runtime activation"
.to_string();
let warning =
"conntrack firewall reconciler is unavailable after runtime activation".to_string();
warn!(reload_id = command.reload_id, warning = %warning);
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::ip_tracker::UserIpTracker;
use crate::network::probe::{decide_network_capabilities, run_probe};
use crate::proxy::direct_buffer_budget::{
DirectBufferBudget, run_direct_buffer_budget_controller,
};
use crate::proxy::direct_buffer_budget::{DirectBufferBudget, run_direct_buffer_budget_controller};
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::shared_state::ProxySharedState;
use crate::proxy::traffic_limiter::TrafficLimiter;
@@ -29,9 +27,7 @@ use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool;
use super::admission;
use super::generation::{
RuntimeGeneration, RuntimeTaskScope, RuntimeTaskScopePreparationGuard,
};
use super::generation::{RuntimeGeneration, RuntimeTaskScope, RuntimeTaskScopePreparationGuard};
use super::listeners::listener_rebind_supported;
use super::runtime_tasks::RuntimeLogFilter;
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;
desired.server.conntrack_control.mode = crate::config::ConntrackMode::Notrack;
desired.server.conntrack_control.backend = crate::config::ConntrackBackend::Iptables;
desired.server.conntrack_control.profile =
crate::config::ConntrackPressureProfile::Aggressive;
desired.server.conntrack_control.hybrid_listener_ips =
vec!["192.0.2.10".parse().unwrap()];
desired
.server
.conntrack_control
.pressure_high_watermark_pct = 90;
desired.server.conntrack_control.profile = crate::config::ConntrackPressureProfile::Aggressive;
desired.server.conntrack_control.hybrid_listener_ips = 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.delete_budget_per_sec = old
.server
+3 -11
View File
@@ -3,11 +3,7 @@ use std::fmt::Write;
use crate::transport::middle_proxy::MeApiHardswapSnapshot;
/// Renders fixed-cardinality hardswap and writer-replacement gauges.
pub(super) fn render(
out: &mut String,
snapshot: Option<&MeApiHardswapSnapshot>,
enabled: bool,
) {
pub(super) fn render(out: &mut String, snapshot: Option<&MeApiHardswapSnapshot>, enabled: bool) {
let snapshot = enabled.then_some(snapshot).flatten();
let pending = snapshot.is_some_and(|value| value.pending);
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_age_seconds 42"));
assert!(out.contains("telemt_me_hardswap_pending_writer_deficit 4"));
assert!(out.contains(
"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=\"preparing\"} 5"));
assert!(out.contains("telemt_me_writer_replacement_current{state=\"retiring\"} 6"));
}
}
+16 -5
View File
@@ -122,22 +122,33 @@ pub(super) fn render(
out,
"# 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 [
(
"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,
),
(
"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,
),
(
"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,
),
(
"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,
),
] {
+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_failures_by_class_total{class=\"timeout\"} 1"));
assert!(
output.contains("telemt_conntrack_rule_reconcile_total{result=\"success\"} 1")
);
assert!(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_rollback_total{result=\"success\"} 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 {
return;
};
self.tracker.enqueue_cleanup_for_incarnation(
owner.user,
owner.incarnation,
owner.ip,
);
self.tracker
.enqueue_cleanup_for_incarnation(owner.user, owner.incarnation, owner.ip);
}
}
@@ -338,9 +335,7 @@ impl UserConnectionReservation {
stats_observation: Option<UserConnectionObservation>,
tracks_ip: bool,
) -> Self {
let ip_permit = tracks_ip.then(|| {
UserIpPermit::new(ip_tracker, user, incarnation, ip)
});
let ip_permit = tracks_ip.then(|| UserIpPermit::new(ip_tracker, user, incarnation, ip));
Self {
stats,
quota_handle,
@@ -388,12 +383,7 @@ pub(crate) async fn acquire_user_connection_reservation(
ip_tracker: Arc<UserIpTracker>,
) -> Result<UserConnectionReservation> {
acquire_user_connection_reservation_for_incarnation(
user,
0,
config,
stats,
peer_addr,
ip_tracker,
user, 0, config, stats, peer_addr, ip_tracker,
)
.await
}
@@ -435,10 +425,7 @@ async fn acquire_user_connection_reservation_for_incarnation(
.or((config.access.user_max_tcp_conns_global_each > 0)
.then_some(config.access.user_max_tcp_conns_global_each))
.map(|value| value as u64);
let Some(connection_permit) = stats
.connection_authority()
.try_acquire(user, limit)
else {
let Some(connection_permit) = stats.connection_authority().try_acquire(user, limit) else {
return Err(ProxyError::ConnectionLimitExceeded {
user: user.to_string(),
});
+1 -4
View File
@@ -155,10 +155,7 @@ impl RunningClientHandler {
.or((config.access.user_max_tcp_conns_global_each > 0)
.then_some(config.access.user_max_tcp_conns_global_each))
.map(|v| v as u64);
let Some(_connection_permit) = stats
.connection_authority()
.try_acquire(user, limit)
else {
let Some(_connection_permit) = stats.connection_authority().try_acquire(user, limit) else {
return Err(ProxyError::ConnectionLimitExceeded {
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.
mod controller;
#[cfg(test)]
use controller::connection_fill_pct;
pub(crate) use 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.
pub(crate) const DIRECT_BUFFER_UNIT_BYTES: usize = 4 * 1024;
@@ -135,10 +135,7 @@ impl DirectBufferBudget {
.fetch_max(generation, Ordering::AcqRel);
}
fn begin_controller_update(
&self,
generation: u64,
) -> Option<ParkingMutexGuard<'_, ()>> {
fn begin_controller_update(&self, generation: u64) -> Option<ParkingMutexGuard<'_, ()>> {
let controller_update = self.controller_update.lock();
(self.active_controller_generation.load(Ordering::Acquire) == generation)
.then_some(controller_update)
+2 -5
View File
@@ -145,11 +145,8 @@ pub(super) fn connection_fill_pct(
return None;
}
let max_connections = max_connections as usize;
let active = max_connections.saturating_sub(
connection_slots
.available_permits()
.min(max_connections),
);
let active =
max_connections.saturating_sub(connection_slots.available_permits().min(max_connections));
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_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_slots_for_testing_in_shared, clear_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_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,
insert_auth_probe_state_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) {
let octets = (index as u32).to_be_bytes();
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(
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!(
shared.handshake.sticky_user_by_ip_prefix.len(),
STICKY_HINT_MAX_ENTRIES
+1 -2
View File
@@ -397,8 +397,7 @@ fn auth_probe_record_failure_with_state_and_budget_in(
};
if state
.remove_if(&evict_key, |_, current| {
current.fail_streak == evict_fail_streak
&& current.last_seen == evict_last_seen
current.fail_streak == evict_fail_streak && current.last_seen == evict_last_seen
})
.is_some()
&& 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) {
let octets = (index as u32).to_be_bytes();
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());
}
@@ -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!(
auth_probe_slots_for_testing_in_shared(shared.as_ref()),
AUTH_PROBE_TRACK_MAX_ENTRIES
+4 -4
View File
@@ -151,10 +151,10 @@ where
if let Some(snapshot) = config.runtime_user_auth() {
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_ip_candidates = sticky_ip_hint
.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key));
let sticky_prefix_candidates = sticky_prefix_hint
.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key));
let sticky_ip_candidates =
sticky_ip_hint.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key));
let sticky_prefix_candidates =
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 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())
+1 -6
View File
@@ -400,12 +400,7 @@ where
.runtime_user_auth()
.and_then(|snapshot| snapshot.entry_by_id(user_id))
{
sticky_hint_record_success_in(
shared,
peer.ip(),
entry.hint_key,
client_sni.as_deref(),
);
sticky_hint_record_success_in(shared, peer.ip(), entry.hint_key, client_sni.as_deref());
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_candidates = sticky_ip_hint
.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key));
let sticky_ip_candidates =
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 sticky_sni_hint = client_sni
.as_deref()
.and_then(|sni| sticky_hint_get_by_sni(shared, sni));
let sticky_sni_candidates = sticky_sni_hint
.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key));
let sticky_sni_candidates =
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_candidates = sticky_prefix_hint
.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key));
let sticky_prefix_candidates =
sticky_prefix_hint.and_then(|hint_key| snapshot.candidate_ids_by_hint_key(hint_key));
let sni_candidates = client_sni
.as_deref()
.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) {
let soft_limit = quota_soft_cap(limit, quota_soft_overshoot_bytes);
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
{
+3 -3
View File
@@ -44,9 +44,9 @@ mod tests {
c2me_sender: AbortOnDropHandle::new(tokio::spawn(pending_child(DropSignal(
Arc::clone(&dropped),
)))),
me_writer: AbortOnDropHandle::new(tokio::spawn(pending_child(DropSignal(
Arc::clone(&dropped),
)))),
me_writer: AbortOnDropHandle::new(tokio::spawn(pending_child(DropSignal(Arc::clone(
&dropped,
))))),
flow_cancel: flow_cancel.clone(),
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() {
this.stats.increment_quota_contention_timeout_total();
Self::arm_wait(&mut this.quota_wait, false, false);
if Self::poll_wait(
&mut this.quota_wait,
cx,
None,
RateDirection::Up,
)
.is_ready()
if Self::poll_wait(&mut this.quota_wait, cx, None, RateDirection::Up).is_ready()
{
cx.waker().wake_by_ref();
}
+7 -4
View File
@@ -219,8 +219,12 @@ impl ProxySharedState {
users: &HashMap<String, String>,
user_enabled: &HashMap<String, bool>,
) -> Option<Vec<(String, usize)>> {
self.user_admission
.activate_config_source(source_generation, expected_epoch, users, user_enabled)
self.user_admission.activate_config_source(
source_generation,
expected_epoch,
users,
user_enabled,
)
}
/// Applies an update only from the active runtime generation.
@@ -276,8 +280,7 @@ impl ProxySharedState {
user: &str,
credential_id: UserCredentialId,
) -> Option<UserAdmissionPublication<'_>> {
self.user_admission
.claim_authenticated(user, credential_id)
self.user_admission.claim_authenticated(user, credential_id)
}
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 peer_addr: SocketAddr = "198.51.100.210:50000".parse().unwrap();
let mut config = ProxyConfig::default();
config
.access
.user_max_tcp_conns
.insert(user.to_string(), 1);
config.access.user_max_tcp_conns.insert(user.to_string(), 1);
let (entered_tx, entered_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 peer_addr: SocketAddr = "198.51.100.211:50001".parse().unwrap();
let mut config = ProxyConfig::default();
config
.access
.user_max_tcp_conns
.insert(user.to_string(), 1);
config.access.user_max_tcp_conns.insert(user.to_string(), 1);
let reservation = acquire_user_connection_reservation(
user,
@@ -77,9 +77,11 @@ fn controller_handoff_waits_for_inflight_update_and_fences_old_generation() {
activated_tx.send(()).unwrap();
});
assert!(activated_rx
.recv_timeout(Duration::from_millis(50))
.is_err());
assert!(
activated_rx
.recv_timeout(Duration::from_millis(50))
.is_err()
);
drop(update);
activated_rx.recv_timeout(Duration::from_secs(1)).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 moved = directory.path().join("moved");
let outside = directory.path().join("outside");
fs::create_dir_all(parent.join("nested"))
.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(parent.join("nested")).expect("original 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 sanitized = sanitize_unknown_dc_log_path(
+1 -4
View File
@@ -154,10 +154,7 @@ enum BucketReserveError {
impl BucketReserveError {
fn exhausted_reserve_budget(self) -> bool {
matches!(
self,
Self::Contended | Self::ReserveAndRefundContended
)
matches!(self, Self::Contended | Self::ReserveAndRefundContended)
}
}
+12 -25
View File
@@ -67,12 +67,8 @@ impl DirectionBucket {
if self.should_force_reserve_failure() {
return Err(current);
}
self.state.compare_exchange(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
)
self.state
.compare_exchange(current, next, Ordering::Relaxed, Ordering::Relaxed)
}
#[inline(always)]
@@ -81,12 +77,8 @@ impl DirectionBucket {
if self.should_force_refund_failure() {
return Err(current);
}
self.state.compare_exchange(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
)
self.state
.compare_exchange(current, next, Ordering::Relaxed, Ordering::Relaxed)
}
fn unpack(state: u64) -> (u64, u64) {
@@ -355,9 +347,9 @@ impl CidrDirectionBucket {
});
};
let user_granted = user_debit.granted();
let Some(aggregate_debit) = self
.used
.try_reserve_at(epoch, cap_epoch, user_granted, budget)?
let Some(aggregate_debit) =
self.used
.try_reserve_at(epoch, cap_epoch, user_granted, budget)?
else {
return Ok(CidrReservation {
granted: 0,
@@ -410,12 +402,8 @@ impl CidrUserDirectionState {
if observed_epoch > epoch {
return Err(BucketReserveError::StaleEpoch);
}
let Some(mut active_debit) = active_users.try_reserve_at(
epoch,
PACKED_USAGE_MASK,
1,
budget,
)?
let Some(mut active_debit) =
active_users.try_reserve_at(epoch, PACKED_USAGE_MASK, 1, budget)?
else {
return Ok(false);
};
@@ -532,10 +520,9 @@ impl CidrBucket {
}
let cap_epoch = bytes_per_epoch(cap_bps);
match direction {
RateDirection::Up => {
self.up
.try_reserve(&share.up, epoch, cap_epoch, requested, budget)
}
RateDirection::Up => self
.up
.try_reserve(&share.up, epoch, cap_epoch, requested, budget),
RateDirection::Down => {
self.down
.try_reserve(&share.down, epoch, cap_epoch, requested, budget)
+25 -28
View File
@@ -61,32 +61,28 @@ impl TrafficLease {
let mut granted = requested;
let mut user_debit = None;
if let Some(user_bucket) = binding.user_bucket.as_ref() {
let user_reservation = match user_bucket.try_reserve(
direction,
epoch,
granted,
&mut budget,
) {
Ok(reservation) => reservation,
Err(error) => {
if error.exhausted_reserve_budget() {
self.limiter
.user_scope
.reserve_cas_retry_exhausted(direction);
let user_reservation =
match user_bucket.try_reserve(direction, epoch, granted, &mut budget) {
Ok(reservation) => reservation,
Err(error) => {
if error.exhausted_reserve_budget() {
self.limiter
.user_scope
.reserve_cas_retry_exhausted(direction);
}
return TrafficReservation {
result: TrafficConsumeResult {
granted: 0,
blocked_user: false,
blocked_cidr: false,
},
_binding: binding,
user: None,
cidr: None,
cidr_user: None,
};
}
return TrafficReservation {
result: TrafficConsumeResult {
granted: 0,
blocked_user: false,
blocked_cidr: false,
},
_binding: binding,
user: None,
cidr: None,
cidr_user: None,
};
}
};
};
user_debit = user_reservation.debit;
if user_reservation.granted == 0 {
self.limiter.observe_throttle(direction, true, false);
@@ -107,9 +103,10 @@ impl TrafficLease {
let mut cidr_debit = None;
let mut cidr_user_debit = None;
if let (Some(cidr_bucket), Some(cidr_user_share)) =
(binding.cidr_bucket.as_ref(), binding.cidr_user_share.as_ref())
{
if let (Some(cidr_bucket), Some(cidr_user_share)) = (
binding.cidr_bucket.as_ref(),
binding.cidr_user_share.as_ref(),
) {
let cidr_reservation = match cidr_bucket.try_reserve_for_user(
direction,
cidr_user_share,
+32 -8
View File
@@ -330,14 +330,38 @@ impl TrafficLimiter {
cidr_refund_down,
] = values;
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.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),
(
&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.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);
}
+2 -11
View File
@@ -65,12 +65,7 @@ fn reserve_at(
cap: u64,
requested: u64,
) -> Result<Option<DirectionDebit>, BucketReserveError> {
bucket.try_reserve_at(
epoch,
cap,
requested,
&mut ReserveCasBudget::new(),
)
bucket.try_reserve_at(epoch, cap, requested, &mut ReserveCasBudget::new())
}
#[test]
@@ -346,11 +341,7 @@ fn concurrent_first_use_counts_one_active_cidr_user() {
let barrier = Arc::clone(&barrier);
threads.push(std::thread::spawn(move || {
barrier.wait();
user.ensure_active(
13,
&bucket.active_users,
&mut ReserveCasBudget::new(),
)
user.ensure_active(13, &bucket.active_users, &mut ReserveCasBudget::new())
}));
}
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);
assert!(matches!(
reservation,
Err(BucketReserveError::Contended)
));
assert!(matches!(reservation, Err(BucketReserveError::Contended)));
assert_eq!(
bucket.reserve_cas_attempts(),
RESERVE_CAS_ATTEMPT_LIMIT as u64
@@ -32,7 +29,10 @@ fn reserve_succeeds_on_the_last_allowed_attempt() {
.unwrap();
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_eq!(bucket.used_at(1), Some(80));
}
@@ -121,15 +121,7 @@ fn lease_contention_is_not_reported_as_throttling() {
let lease = limiter
.acquire_lease("alice", "203.0.113.7".parse().unwrap())
.unwrap();
let bucket = Arc::clone(
&lease
.binding
.load_full()
.user_bucket
.as_ref()
.unwrap()
.down,
);
let bucket = Arc::clone(&lease.binding.load_full().user_bucket.as_ref().unwrap().down);
bucket.force_reserve_failures(RESERVE_CAS_ATTEMPT_LIMIT);
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_cidr);
assert_eq!(
binding
.user_bucket
.as_ref()
.unwrap()
.down
.used_at(epoch),
binding.user_bucket.as_ref().unwrap().down.used_at(epoch),
Some(0)
);
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() {
let bucket = CidrDirectionBucket::default();
let user = CidrUserDirectionState::default();
user.used
.force_reserve_failures(RESERVE_CAS_ATTEMPT_LIMIT);
user.used.force_reserve_failures(RESERVE_CAS_ATTEMPT_LIMIT);
let mut budget = ReserveCasBudget::new();
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 second = CidrUserDirectionState::default();
assert_eq!(
first.ensure_active(
17,
&bucket.active_users,
&mut ReserveCasBudget::new(),
),
first.ensure_active(17, &bucket.active_users, &mut ReserveCasBudget::new(),),
Ok(true)
);
assert_eq!(
second.ensure_active(
17,
&bucket.active_users,
&mut ReserveCasBudget::new(),
),
second.ensure_active(17, &bucket.active_users, &mut ReserveCasBudget::new(),),
Ok(true)
);
@@ -391,21 +369,13 @@ fn cidr_first_grants_preserve_the_current_soft_fair_share() {
.as_mut()
.unwrap()
.commit_all();
first_reservation
.user_debit
.as_mut()
.unwrap()
.commit_all();
first_reservation.user_debit.as_mut().unwrap().commit_all();
second_reservation
.aggregate_debit
.as_mut()
.unwrap()
.commit_all();
second_reservation
.user_debit
.as_mut()
.unwrap()
.commit_all();
second_reservation.user_debit.as_mut().unwrap().commit_all();
assert_eq!(bucket.used.used_at(17), Some(100));
assert_eq!(first.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;
if identity_changed {
if previous.is_some() {
self.quota_store
.advance_preserving_usage(user, incarnation);
self.quota_store.advance_preserving_usage(user, incarnation);
} else {
self.quota_store.activate_fresh(user, incarnation);
}
@@ -420,7 +419,8 @@ impl UserAdmissionAuthority {
}
let record = state.users.get(user)?;
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.
@@ -456,10 +456,7 @@ impl UserAdmissionAuthority {
}
/// Registers a legacy owner when no credential snapshot is available.
pub(crate) fn register_legacy(
self: &Arc<Self>,
user: &str,
) -> Option<UserSessionRegistration> {
pub(crate) fn register_legacy(self: &Arc<Self>, user: &str) -> Option<UserSessionRegistration> {
let credential_id = {
let state = self.state.lock();
if !state.initialized {
@@ -520,9 +517,7 @@ pub(crate) fn credential_id_from_hex(secret: &str) -> Option<UserCredentialId> {
Some(credential_id(&secret))
}
fn cancel_owners(
cancellations: Vec<(String, Vec<CancellationToken>)>,
) -> Vec<(String, usize)> {
fn cancel_owners(cancellations: Vec<(String, Vec<CancellationToken>)>) -> Vec<(String, usize)> {
cancellations
.into_iter()
.map(|(user, tokens)| {
+3 -12
View File
@@ -51,12 +51,7 @@ fn stale_candidate_cannot_overwrite_newer_mutation() {
assert!(
authority
.activate_config_source(
2,
Some(candidate_epoch),
&users(secret),
&HashMap::new(),
)
.activate_config_source(2, Some(candidate_epoch), &users(secret), &HashMap::new(),)
.is_none()
);
assert!(!authority.is_user_enabled("alice"));
@@ -105,9 +100,7 @@ fn registration_dropped_before_publication_cannot_leave_an_owner() {
let secret = "00112233445566778899aabbccddeeff";
authority.apply_config(&users(secret), &HashMap::new());
let credential = credential_id_from_hex(secret).unwrap();
let mut publication = authority
.claim_authenticated("alice", credential)
.unwrap();
let mut publication = authority.claim_authenticated("alice", credential).unwrap();
let registration = publication.take_registration().unwrap();
drop(registration);
@@ -126,9 +119,7 @@ fn quota_identity_follows_credential_rotation_and_recreation() {
let old_incarnation = authority
.authenticated_incarnation("alice", credential_id_from_hex(old_secret).unwrap())
.unwrap();
let old_quota = quota_store
.handle_exact("alice", old_incarnation)
.unwrap();
let old_quota = quota_store.handle_exact("alice", old_incarnation).unwrap();
old_quota.charge(40);
let rotated = authority.stage_user("alice", new_secret, true).unwrap();
+64 -64
View File
@@ -180,27 +180,27 @@ async fn read_state_file(path: &Path) -> std::io::Result<Option<QuotaStateFile>>
};
#[cfg(not(unix))]
let payload = {
let file = match tokio::fs::File::open(path).await {
Ok(file) => file,
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
Err(error) => return Err(error),
};
if file.metadata().await?.len() > QUOTA_STATE_MAX_BYTES {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"quota state file exceeds the 16 MiB limit",
));
}
let mut payload = Vec::new();
file.take(QUOTA_STATE_MAX_BYTES.saturating_add(1))
.read_to_end(&mut payload)
.await?;
if payload.len() as u64 > QUOTA_STATE_MAX_BYTES {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"quota state file grew beyond the 16 MiB limit while reading",
));
}
let file = match tokio::fs::File::open(path).await {
Ok(file) => file,
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
Err(error) => return Err(error),
};
if file.metadata().await?.len() > QUOTA_STATE_MAX_BYTES {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"quota state file exceeds the 16 MiB limit",
));
}
let mut payload = Vec::new();
file.take(QUOTA_STATE_MAX_BYTES.saturating_add(1))
.read_to_end(&mut payload)
.await?;
if payload.len() as u64 > QUOTA_STATE_MAX_BYTES {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"quota state file grew beyond the 16 MiB limit while reading",
));
}
payload
};
let state = serde_json::from_slice(&payload).map_err(|error| {
@@ -241,53 +241,53 @@ fn write_state_file_blocking(path: &Path, state: &QuotaStateFile) -> std::io::Re
}
#[cfg(not(unix))]
{
use std::io::Write;
use std::io::Write;
let parent = path
.parent()
.filter(|parent| !parent.as_os_str().is_empty())
.unwrap_or_else(|| Path::new("."));
std::fs::create_dir_all(parent)?;
let parent = path
.parent()
.filter(|parent| !parent.as_os_str().is_empty())
.unwrap_or_else(|| Path::new("."));
std::fs::create_dir_all(parent)?;
let mut last_collision = None;
for _ in 0..8 {
let tmp_path = path.with_extension(format!(
"tmp.{}.{}",
std::process::id(),
rand::random::<u64>()
));
let mut file = match std::fs::OpenOptions::new()
.write(true)
.create_new(true)
.open(&tmp_path)
{
Ok(file) => file,
Err(error) if error.kind() == std::io::ErrorKind::AlreadyExists => {
last_collision = Some(error);
continue;
let mut last_collision = None;
for _ in 0..8 {
let tmp_path = path.with_extension(format!(
"tmp.{}.{}",
std::process::id(),
rand::random::<u64>()
));
let mut file = match std::fs::OpenOptions::new()
.write(true)
.create_new(true)
.open(&tmp_path)
{
Ok(file) => file,
Err(error) if error.kind() == std::io::ErrorKind::AlreadyExists => {
last_collision = Some(error);
continue;
}
Err(error) => return Err(error),
};
let result = (|| {
file.write_all(&payload)?;
file.sync_all()?;
drop(file);
std::fs::rename(&tmp_path, path)?;
#[cfg(unix)]
std::fs::File::open(parent)?.sync_all()?;
Ok(())
})();
if result.is_err() {
let _ = std::fs::remove_file(&tmp_path);
}
Err(error) => return Err(error),
};
let result = (|| {
file.write_all(&payload)?;
file.sync_all()?;
drop(file);
std::fs::rename(&tmp_path, path)?;
#[cfg(unix)]
std::fs::File::open(parent)?.sync_all()?;
Ok(())
})();
if result.is_err() {
let _ = std::fs::remove_file(&tmp_path);
return result;
}
return result;
}
Err(last_collision.unwrap_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::AlreadyExists,
"failed to allocate a unique quota checkpoint temporary file",
)
}))
Err(last_collision.unwrap_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::AlreadyExists,
"failed to allocate a unique quota checkpoint temporary file",
)
}))
}
}
+4 -1
View File
@@ -54,7 +54,10 @@ impl SlotBudget {
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |current| {
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.
+3 -6
View File
@@ -22,13 +22,13 @@ use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, AtomicUsize, Ordering};
use std::time::Instant;
pub(crate) use self::quota_store::{QuotaReservation, QuotaStore, UserQuotaHandle};
pub(crate) use self::users::UserConnectionObservation;
#[allow(unused_imports)]
pub use self::replay::{ReplayChecker, ReplayStats};
use self::telemetry::TelemetryPolicy;
use crate::proxy::user_connection_authority::UserConnectionAuthority;
pub use self::tls_fingerprints::TlsFingerprintSnapshotRow;
pub(crate) use self::users::UserConnectionObservation;
use crate::config::MeWriterPickMode;
use crate::proxy::user_connection_authority::UserConnectionAuthority;
const ME_HANDSHAKE_ERROR_CODE_MAX: usize = 64;
@@ -432,10 +432,7 @@ impl Stats {
#[cfg(test)]
pub(crate) fn with_quota_store(quota_store: Arc<QuotaStore>) -> Self {
Self::with_process_authorities(
quota_store,
Arc::new(UserConnectionAuthority::default()),
)
Self::with_process_authorities(quota_store, Arc::new(UserConnectionAuthority::default()))
}
/// 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.
pub(crate) fn advance_preserving_usage(
&self,
user: &str,
incarnation: UserIncarnation,
) {
pub(crate) fn advance_preserving_usage(&self, user: &str, incarnation: UserIncarnation) {
let slot = self.slot(user);
let mut state = slot.state.lock();
if incarnation <= state.high_water {
+1 -1
View File
@@ -1,6 +1,6 @@
use std::borrow::Borrow;
use std::collections::{HashMap, VecDeque};
use std::collections::hash_map::DefaultHasher;
use std::collections::{HashMap, VecDeque};
use std::hash::{Hash, Hasher};
use std::num::NonZeroUsize;
use std::sync::Arc;
+13 -12
View File
@@ -17,26 +17,27 @@ fn test_stats_shared_counters() {
fn runtime_stats_share_process_connection_admission_authority() {
let quota_store = Arc::new(QuotaStore::default());
let authority = Arc::new(UserConnectionAuthority::default());
let first = Stats::with_process_authorities(
Arc::clone(&quota_store),
Arc::clone(&authority),
);
let first = Stats::with_process_authorities(Arc::clone(&quota_store), Arc::clone(&authority));
let second = Stats::with_process_authorities(quota_store, authority);
let permit = first
.connection_authority()
.try_acquire("alice", Some(1))
.unwrap();
assert!(second
.connection_authority()
.try_acquire("alice", Some(1))
.is_none());
assert!(
second
.connection_authority()
.try_acquire("alice", Some(1))
.is_none()
);
drop(permit);
assert!(second
.connection_authority()
.try_acquire("alice", Some(1))
.is_some());
assert!(
second
.connection_authority()
.try_acquire("alice", Some(1))
.is_some()
);
}
#[test]
+1 -2
View File
@@ -350,8 +350,7 @@ impl TlsFingerprintCollector {
let mut removed = 0usize;
self.entries.retain(|_, entry| {
let last_seen = entry.last_seen_epoch_secs.load(Ordering::Relaxed);
let retained =
ttl_secs != 0 && now_epoch_secs.saturating_sub(last_seen) <= ttl_secs;
let retained = ttl_secs != 0 && now_epoch_secs.saturating_sub(last_seen) <= ttl_secs;
if !retained {
removed += 1;
}
+1 -1
View File
@@ -39,7 +39,7 @@ pub(super) async fn run_command(
.map_err(|e| format!("wait {binary} failed: {e}"))
})
.await
.map_err(|_| format!("{binary} timed out after {}s", COMMAND_TIMEOUT.as_secs()))??;
.map_err(|_| format!("{binary} timed out after {}s", COMMAND_TIMEOUT.as_secs()))??;
if output.status.success() {
return Ok(());
}
+3 -5
View File
@@ -76,11 +76,9 @@ pub fn parse_proxy_config_text(text: &str, http_status: u16) -> ProxyConfigData
pub async fn load_proxy_config_cache(path: &str) -> Result<ProxyConfigData> {
#[cfg(unix)]
let bytes = read_regular_limited_async(
Path::new(path).to_path_buf(),
HTTPS_RESPONSE_BODY_MAX_BYTES,
)
.await;
let bytes =
read_regular_limited_async(Path::new(path).to_path_buf(), HTTPS_RESPONSE_BODY_MAX_BYTES)
.await;
#[cfg(not(unix))]
let bytes = tokio::fs::read(path).await;
let bytes = bytes.map_err(|e| {
@@ -234,7 +234,8 @@ pub(super) async fn maybe_refresh_idle_writer_for_dc(
purpose,
&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(Err(error)) => {
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 key = (2, IpFamily::V4);
let live_writer_ids_by_addr = HashMap::from([((2, endpoint), vec![writer_id])]);
let writer_idle_since = HashMap::from([(
writer_id,
MePool::now_epoch_secs().saturating_sub(60),
)]);
let writer_idle_since =
HashMap::from([(writer_id, MePool::now_epoch_secs().saturating_sub(60))]);
let bound_clients_by_writer = HashMap::from([(writer_id, 0)]);
let mut next_attempt = HashMap::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,
};
pub use pool::MePool;
pub(crate) use pool_status::MeApiHardswapSnapshot;
#[allow(unused_imports)]
pub use pool_nat::{detect_public_ip, stun_probe};
pub(crate) use pool_status::MeApiHardswapSnapshot;
pub(crate) use registry::ConnLease;
pub use registry::ConnRegistry;
pub use rotation::{MeReinitTrigger, me_reinit_scheduler, me_rotation_task};
@@ -120,12 +120,8 @@ impl MePool {
me_route_inline_recovery_wait_ms: u64,
me_connection_cleanup_capacity: usize,
) -> Arc<Self> {
let endpoint_snapshot = Self::build_endpoint_snapshot(
&decision,
proxy_map_v4,
proxy_map_v6,
1,
);
let endpoint_snapshot =
Self::build_endpoint_snapshot(&decision, proxy_map_v4, proxy_map_v6, 1);
let registry = Arc::new(ConnRegistry::with_route_and_cleanup_capacity(
me_route_channel_capacity,
me_connection_cleanup_capacity,
+1 -5
View File
@@ -198,11 +198,7 @@ impl MePool {
}
fn mirror_negative_dcs(map: &mut HashMap<i32, Vec<(IpAddr, u16)>>) {
let positive_dcs = map
.keys()
.copied()
.filter(|dc| *dc > 0)
.collect::<Vec<_>>();
let positive_dcs = map.keys().copied().filter(|dc| *dc > 0).collect::<Vec<_>>();
for dc in positive_dcs {
if !map.contains_key(&-dc)
&& let Some(endpoints) = map.get(&dc).cloned()
@@ -203,7 +203,9 @@ impl MePool {
.max(1)
.min(WRITER_REPLACEMENT_OPEN_LIMIT_MAX);
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 {
return None;
}
+5 -8
View File
@@ -252,12 +252,7 @@ impl MePool {
let addr = candidates[idx];
match self
.connect_one_with_generation_contour_for_dc_with_intent(
addr,
rng,
generation,
contour,
dc,
intent,
addr, rng, generation, contour, dc, intent,
)
.await
{
@@ -296,8 +291,10 @@ impl MePool {
let status = self.reinit.status.load();
let role_is_authoritative = match target.contour {
WriterContour::Active => target.generation == status.active_generation,
WriterContour::Warm => status.pending_hardswap_generation != 0
&& target.generation == status.pending_hardswap_generation,
WriterContour::Warm => {
status.pending_hardswap_generation != 0
&& target.generation == status.pending_hardswap_generation
}
WriterContour::Draining => false,
};
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 tracing::{debug, info, warn};
use crate::crypto::SecureRandom;
use crate::config::MeBindStaleMode;
use crate::crypto::SecureRandom;
use crate::network::IpFamily;
use super::pool::{
@@ -117,8 +117,7 @@ fn publish_reinit_state(reinit: &ReinitCore, state: &ReinitCoordinatorState) {
pending_hardswap_started_at_epoch_secs: pending
.map_or(0, |value| value.started_at_epoch_secs),
pending_hardswap_map_hash: pending.map_or(0, |value| value.map_hash),
pending_hardswap_endpoint_revision: pending
.map_or(0, |value| value.endpoint_revision),
pending_hardswap_endpoint_revision: pending.map_or(0, |value| value.endpoint_revision),
inflight: state.attempts.len(),
};
reinit
@@ -318,15 +318,13 @@ impl MePool {
if alive >= required {
covered = covered.saturating_add(1);
} else {
writer_deficit =
writer_deficit.saturating_add(required.saturating_sub(alive));
writer_deficit = writer_deficit.saturating_add(required.saturating_sub(alive));
missing_groups.push(DcFamilyGroup { dc: *dc, family });
}
}
}
missing_groups.sort_unstable_by_key(|group| {
(group.dc, matches!(group.family, IpFamily::V6))
});
missing_groups
.sort_unstable_by_key(|group| (group.dc, matches!(group.family, IpFamily::V6)));
HardswapCoverage {
ratio: if total == 0 {
1.0
@@ -453,8 +451,8 @@ impl MePool {
let authoritative_warm = contour == WriterContour::Warm
&& pending_generation == Some(writer.generation)
&& endpoint_is_current;
let stale_active = contour == WriterContour::Active
&& writer.generation != active_generation;
let stale_active =
contour == WriterContour::Active && writer.generation != active_generation;
if authoritative_warm || (contour == WriterContour::Active && !stale_active) {
continue;
}
@@ -535,8 +533,7 @@ impl MePool {
.filter(|w| !w.draining.load(Ordering::Relaxed))
.filter(|w| w.generation == generation)
.filter(|w| {
WriterContour::from_u8(w.contour.load(Ordering::Acquire))
== WriterContour::Active
WriterContour::from_u8(w.contour.load(Ordering::Acquire)) == WriterContour::Active
})
.filter(|w| w.writer_dc == dc)
.filter(|w| endpoints.contains(&w.addr))
@@ -34,11 +34,7 @@ impl MePool {
let total_passes = 1 + extra_passes;
for (dc, endpoints) in desired_by_dc {
if !self.hardswap_warmup_is_authoritative(
generation,
map_hash,
endpoint_revision,
) {
if !self.hardswap_warmup_is_authoritative(generation, map_hash, endpoint_revision) {
return;
}
for family in [IpFamily::V4, IpFamily::V6] {
@@ -119,11 +115,7 @@ impl MePool {
}
last_fresh_count = self
.fresh_writer_count_for_dc_endpoints(
generation,
*dc,
&family_endpoints,
)
.fresh_writer_count_for_dc_endpoints(generation, *dc, &family_endpoints)
.await;
if last_fresh_count >= required {
completed = true;
@@ -324,8 +316,7 @@ impl MePool {
Err(ReinitCommitFailure::Superseded) => {
debug!(
previous_generation,
generation,
"ME reinit result discarded after a newer desired-map attempt"
generation, "ME reinit result discarded after a newer desired-map attempt"
);
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)
.await;
assert!(matches!(result, Err(ReinitCommitFailure::Redundancy { .. })));
assert!(matches!(
result,
Err(ReinitCommitFailure::Redundancy { .. })
));
assert_eq!(pool.current_generation(), active_generation);
assert!(!old_dc1.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 desired_by_dc = HashMap::from([(1, HashSet::from([v4, v6]))]);
let active_generation = pool.current_generation();
let old_v4 = insert_writer(
&pool,
451,
1,
v4,
active_generation,
WriterContour::Active,
)
.await;
let old_v6 = insert_writer(
&pool,
452,
1,
v6,
active_generation,
WriterContour::Active,
)
.await;
let old_v4 = insert_writer(&pool, 451, 1, v4, active_generation, WriterContour::Active).await;
let old_v6 = insert_writer(&pool, 452, 1, v6, active_generation, WriterContour::Active).await;
let map_hash = MePool::desired_map_hash(&desired_by_dc);
let endpoint_revision = pool.endpoint_snapshot.load().revision;
let reservation = pool
@@ -81,8 +81,7 @@ impl MePool {
let (replacement_preparing_current, replacement_retiring_current) =
self.registry.writer_replacement_counts();
let pending_age_secs = pending.then(|| {
Self::now_epoch_secs()
.saturating_sub(reinit.pending_hardswap_started_at_epoch_secs)
Self::now_epoch_secs().saturating_sub(reinit.pending_hardswap_started_at_epoch_secs)
});
MeApiHardswapSnapshot {
@@ -83,12 +83,8 @@ impl MePool {
if endpoint_count == 0 {
continue;
}
let required =
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 required = self.required_writers_for_dc_with_floor_mode(endpoint_count, false);
let alive = live_writers_by_group.get(&(dc, ipv4)).copied().unwrap_or(0);
if alive < required {
return false;
}
@@ -133,9 +129,7 @@ impl MePool {
.map(|endpoints| {
endpoint_family_counts(endpoints)
.into_iter()
.map(|(_, count)| {
self.required_writers_for_dc_with_floor_mode(count, false)
})
.map(|(_, count)| self.required_writers_for_dc_with_floor_mode(count, false))
.sum::<usize>()
})
.sum();
@@ -254,9 +248,7 @@ impl MePool {
let dc_required_writers = family_counts
.iter()
.filter(|(_, count)| *count > 0)
.map(|(_, count)| {
self.required_writers_for_dc_with_floor_mode(*count, false)
})
.map(|(_, count)| self.required_writers_for_dc_with_floor_mode(*count, false))
.sum::<usize>();
let floor_min = family_counts
.iter()
@@ -289,8 +281,7 @@ impl MePool {
.me_adaptive_floor_max_extra_writers_multi_per_core
.load(Ordering::Relaxed) as usize
};
family_base
.saturating_add(adaptive_cpu_cores.saturating_mul(extra_per_core))
family_base.saturating_add(adaptive_cpu_cores.saturating_mul(extra_per_core))
})
.sum::<usize>();
let floor_capped =
@@ -357,6 +348,9 @@ impl MePool {
}
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))]
}
@@ -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_writer_security_tests::make_pool_with_decision;
fn writer(
pool: &Arc<MePool>,
id: u64,
dc: i32,
addr: SocketAddr,
) -> MeWriter {
fn writer(pool: &Arc<MePool>, id: u64, dc: i32, addr: SocketAddr) -> MeWriter {
let (tx, _rx) = mpsc::channel::<WriterCommand>(8);
MeWriter {
id,
@@ -60,12 +55,7 @@ async fn dual_family_status_reports_each_family_floor() {
let mut writers = pool.writers.write().await;
for (group, dc) in [2, -2].into_iter().enumerate() {
for offset in 0..required_per_family {
writers.push(writer(
&pool,
(group as u64 * 100) + offset as u64,
dc,
v4,
));
writers.push(writer(&pool, (group as u64 * 100) + offset as u64, dc, v4));
}
}
drop(writers);
@@ -122,11 +122,9 @@ impl MePool {
!candidate.draining.load(Ordering::Acquire)
&& candidate.writer_dc == writer.writer_dc
&& candidate.generation == writer.generation
&& WriterContour::from_u8(candidate.contour.load(Ordering::Acquire))
== contour
&& WriterContour::from_u8(candidate.contour.load(Ordering::Acquire)) == contour
&& candidate.addr.is_ipv4() == writer.addr.is_ipv4()
&& endpoint_snapshot
.contains_dc_endpoint(candidate.writer_dc, candidate.addr)
&& endpoint_snapshot.contains_dc_endpoint(candidate.writer_dc, candidate.addr)
})
.count();
if current >= required {
@@ -138,11 +138,7 @@ impl MePool {
writers.push(writer);
self.conn_count.fetch_add(1, Ordering::Relaxed);
writers.publish_current();
self.apply_writer_draining_state(
&writers[victim_pos],
self.force_close_timeout(),
false,
);
self.apply_writer_draining_state(&writers[victim_pos], self.force_close_timeout(), false);
self.lifecycle
.spawn_registered_writer(task_registration, writer_task);
reservation.mark_committed();
@@ -264,11 +260,8 @@ mod tests {
async fn replacement_commit_publishes_successor_before_draining_victim() {
let pool = make_pool().await;
let addr = endpoint(1);
pool.update_proxy_maps(
HashMap::from([(2, vec![(addr.ip(), addr.port())])]),
None,
)
.await;
pool.update_proxy_maps(HashMap::from([(2, vec![(addr.ip(), addr.port())])]), None)
.await;
let victim = install_writer(&pool, 1001, 2, addr).await;
let expected_role = WriterRole::from_writer(&victim);
let mut reservation = pool
@@ -307,11 +300,8 @@ mod tests {
async fn cancelled_replacement_waiting_for_publication_restores_all_reservations() {
let pool = make_pool().await;
let addr = endpoint(2);
pool.update_proxy_maps(
HashMap::from([(2, vec![(addr.ip(), addr.port())])]),
None,
)
.await;
pool.update_proxy_maps(HashMap::from([(2, vec![(addr.ip(), addr.port())])]), None)
.await;
let victim = install_writer(&pool, 2001, 2, addr).await;
let expected_role = WriterRole::from_writer(&victim);
let mut reservation = pool
@@ -336,10 +326,21 @@ mod tests {
assert!(result.is_err());
drop(writers_guard);
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!(!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]
@@ -379,7 +380,14 @@ mod tests {
assert!(result.is_err());
drop(reservation);
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));
}
}
@@ -92,12 +92,7 @@ impl MePool {
intent: WriterOpenIntent,
) -> Result<PreparedWriter<'a>> {
let Some(writer_open_reservation) = self
.reserve_writer_open(
contour,
intent,
writer_dc,
addr,
)
.reserve_writer_open(contour, intent, writer_dc, addr)
.await
else {
return Err(ProxyError::Proxy(format!(
@@ -76,10 +76,10 @@ impl WriterRegistrationGuard<'_> {
|| !Arc::ptr_eq(&route_state, reservation.state())
|| reservation.requires_idle()
&& self
.binding
.conns_for_writer
.get(&reservation.writer_id())
.is_none_or(|conn_ids| !conn_ids.is_empty())
.binding
.conns_for_writer
.get(&reservation.writer_id())
.is_none_or(|conn_ids| !conn_ids.is_empty())
{
return false;
}
@@ -98,9 +98,9 @@ impl ConnRegistry {
.map(|route| Arc::clone(&route.replacement_state))?;
if require_idle
&& binding
.conns_for_writer
.get(&writer_id)
.is_none_or(|conn_ids| !conn_ids.is_empty())
.conns_for_writer
.get(&writer_id)
.is_none_or(|conn_ids| !conn_ids.is_empty())
{
return None;
}
@@ -8,11 +8,11 @@ use tokio::sync::mpsc::error::TrySendError;
use super::super::codec::WriterCommand;
use super::super::{MeResponse, RouteBytePermit};
use super::replacement::WriterBindOutcome;
use super::{
BoundConn, ConnMeta, ConnRegistry, ConnWriter, HotConnBinding, RouteResult,
WriterActivitySnapshot,
};
use super::replacement::WriterBindOutcome;
impl ConnRegistry {
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)
{
let endpoint_snapshot = self.endpoint_snapshot.load_full();
for (dc, addrs) in
&endpoint_snapshot.preferred_endpoints_by_dc
for (dc, addrs) in &endpoint_snapshot.preferred_endpoints_by_dc
{
for addr in addrs {
let _ = self
@@ -544,5 +543,4 @@ impl MePool {
return Ok(());
}
}
}
+2 -1
View File
@@ -31,7 +31,8 @@ impl MePool {
writer_reserved_bytes: usize,
payload_permit: Option<OwnedSemaphorePermit>,
) -> 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));
};
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::reservation::{
WriterByteReserveError, WriterCommandReserveError, proxy_req_payload_from_command,
proxy_req_resident_permits, proxy_tag_array, reserve_writer_bytes,
reserve_writer_command_slot, writer_send_deadline,
proxy_req_resident_permits, proxy_tag_array, reserve_writer_bytes, reserve_writer_command_slot,
writer_send_deadline,
};
use crate::error::{ProxyError, Result};
use crate::stream::PooledBuffer;
@@ -34,9 +34,7 @@ pub(super) fn proxy_req_payload_from_command(
}
}
pub(super) fn payload_permit_from_data_command(
cmd: WriterCommand,
) -> Option<OwnedSemaphorePermit> {
pub(super) fn payload_permit_from_data_command(cmd: WriterCommand) -> Option<OwnedSemaphorePermit> {
match cmd {
WriterCommand::Data { _permit, .. } => _permit,
_ => None,

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