Races in admission + accounting + publication,+ PID fixed

This commit is contained in:
Alexey
2026-09-09 22:45:02 +03:00
parent 021ad1fe68
commit 844e41ea34
34 changed files with 2237 additions and 1207 deletions
+1 -4
View File
@@ -35,13 +35,10 @@ pub(super) async fn create_user_route(
let runtime_cfg = config_rx.borrow().clone();
data.user.in_runtime = runtime_cfg.access.users.contains_key(&data.user.username);
if let Some(enabled) = requested_enabled {
shared
let (_, cancelled) = shared
.proxy_shared
.set_user_enabled(&data.user.username, enabled);
if !enabled {
let cancelled = shared
.proxy_shared
.cancel_user_sessions(&data.user.username);
if cancelled > 0 {
shared.runtime_events.record(
"api.user.disable.runtime",
+2 -4
View File
@@ -104,8 +104,7 @@ pub(super) async fn handle(
};
let runtime_cfg = config_rx.borrow().clone();
data.in_runtime = runtime_cfg.access.users.contains_key(&data.username);
let newly_disabled = shared.proxy_shared.set_user_enabled(base_user, false);
let cancelled = shared.proxy_shared.cancel_user_sessions(base_user);
let (newly_disabled, cancelled) = shared.proxy_shared.set_user_enabled(base_user, false);
shared.runtime_events.record(
"api.user.disable.ok",
format!(
@@ -290,11 +289,10 @@ pub(super) async fn handle(
let runtime_cfg = config_rx.borrow().clone();
data.in_runtime = runtime_cfg.access.users.contains_key(&data.username);
if let Some(enabled) = enabled_update {
shared
let (_, cancelled) = shared
.proxy_shared
.set_user_enabled(&data.username, enabled);
if !enabled {
let cancelled = shared.proxy_shared.cancel_user_sessions(&data.username);
shared.runtime_events.record(
"api.user.disable.runtime",
format!(
+23 -496
View File
@@ -8,15 +8,22 @@
//! - `run [OPTIONS] [config.toml]` - Run in foreground (default behavior)
//! - `healthcheck [OPTIONS] [config.toml]` - Run control-plane health probe
use rand::RngExt;
use std::fs;
use std::path::{Path, PathBuf};
use std::process::Command;
use std::path::PathBuf;
use crate::healthcheck::{self, HealthcheckMode};
#[cfg(unix)]
use crate::daemon::{self, DEFAULT_PID_FILE, DaemonOptions};
use crate::daemon::{DEFAULT_PID_FILE, DaemonOptions};
// Unix daemon control and argument parsing.
#[cfg(unix)]
mod daemon_commands;
// Fire-and-forget installation workflow.
mod init;
#[cfg(unix)]
pub use daemon_commands::parse_daemon_args;
pub use init::{InitOptions, parse_init_args, run_init};
/// CLI subcommand to execute.
#[derive(Debug, Clone, PartialEq, Eq)]
@@ -40,13 +47,20 @@ pub enum Subcommand {
/// Parsed subcommand with its options.
#[derive(Debug)]
pub struct ParsedCommand {
/// Selected command mode.
pub subcommand: Subcommand,
/// PID file used by daemon-control commands.
pub pid_file: PathBuf,
/// Configuration file passed to runtime or healthcheck.
pub config_path: String,
/// Requested healthcheck mode.
pub healthcheck_mode: HealthcheckMode,
/// Invalid healthcheck mode retained for command diagnostics.
pub healthcheck_mode_invalid: Option<String>,
#[cfg(unix)]
/// Unix daemon lifecycle options.
pub daemon_opts: DaemonOptions,
/// Fire-and-forget initialization options.
pub init_opts: Option<InitOptions>,
}
@@ -79,7 +93,6 @@ pub fn parse_command(args: &[String]) -> ParsedCommand {
return cmd;
}
// Check for subcommand as first argument
if let Some(first) = args.first() {
match first.as_str() {
"start" => {
@@ -120,11 +133,9 @@ pub fn parse_command(args: &[String]) -> ParsedCommand {
}
}
// Parse remaining options
let mut i = 0;
while i < args.len() {
match args[i].as_str() {
// Skip subcommand names
"start" | "stop" | "reload" | "status" | "run" | "healthcheck" => {}
"--mode" => {
i += 1;
@@ -154,7 +165,6 @@ pub fn parse_command(args: &[String]) -> ParsedCommand {
}
}
}
// PID file option (for stop/reload/status)
"--pid-file" => {
i += 1;
if i < args.len() {
@@ -189,9 +199,9 @@ pub fn parse_command(args: &[String]) -> ParsedCommand {
#[cfg(unix)]
pub fn execute_subcommand(cmd: &ParsedCommand) -> Option<i32> {
match cmd.subcommand {
Subcommand::Stop => Some(cmd_stop(&cmd.pid_file)),
Subcommand::Reload => Some(cmd_reload(&cmd.pid_file)),
Subcommand::Status => Some(cmd_status(&cmd.pid_file)),
Subcommand::Stop => Some(daemon_commands::stop(&cmd.pid_file)),
Subcommand::Reload => Some(daemon_commands::reload(&cmd.pid_file)),
Subcommand::Status => Some(daemon_commands::status(&cmd.pid_file)),
Subcommand::Healthcheck => {
if let Some(invalid_mode) = cmd.healthcheck_mode_invalid.as_ref() {
if invalid_mode.is_empty() {
@@ -224,6 +234,7 @@ pub fn execute_subcommand(cmd: &ParsedCommand) -> Option<i32> {
}
}
/// Executes a non-server subcommand on platforms without daemon support.
#[cfg(not(unix))]
pub fn execute_subcommand(cmd: &ParsedCommand) -> Option<i32> {
match cmd.subcommand {
@@ -261,487 +272,3 @@ pub fn execute_subcommand(cmd: &ParsedCommand) -> Option<i32> {
Subcommand::Run | Subcommand::Start => None,
}
}
/// Stop command: send SIGTERM to the running daemon.
#[cfg(unix)]
fn cmd_stop(pid_file: &Path) -> i32 {
use nix::sys::signal::Signal;
println!("Stopping telemt daemon...");
match daemon::signal_pid_file(pid_file, Signal::SIGTERM) {
Ok(()) => {
println!("Stop signal sent successfully");
// Wait for process to exit (up to 10 seconds)
for _ in 0..20 {
std::thread::sleep(std::time::Duration::from_millis(500));
if let daemon::DaemonStatus::NotRunning = daemon::check_status(pid_file) {
println!("Daemon stopped");
return 0;
}
}
println!("Daemon may still be shutting down");
0
}
Err(e) => {
eprintln!("Failed to stop daemon: {}", e);
1
}
}
}
/// Reload command: send SIGHUP to trigger config reload.
#[cfg(unix)]
fn cmd_reload(pid_file: &Path) -> i32 {
use nix::sys::signal::Signal;
println!("Reloading telemt configuration...");
match daemon::signal_pid_file(pid_file, Signal::SIGHUP) {
Ok(()) => {
println!("Reload signal sent successfully");
0
}
Err(e) => {
eprintln!("Failed to reload daemon: {}", e);
1
}
}
}
/// Status command: check if daemon is running.
#[cfg(unix)]
fn cmd_status(pid_file: &Path) -> i32 {
match daemon::check_status(pid_file) {
daemon::DaemonStatus::Running(pid) => {
println!("telemt is running (pid {})", pid);
0
}
daemon::DaemonStatus::Stale(pid) => {
println!("telemt is not running (stale pid file, was pid {})", pid);
// Clean up stale PID file
let _ = std::fs::remove_file(pid_file);
1
}
daemon::DaemonStatus::NotRunning => {
println!("telemt is not running");
1
}
}
}
/// Options for the init command
#[derive(Debug, Clone)]
pub struct InitOptions {
pub port: u16,
pub domain: String,
pub secret: Option<String>,
pub username: String,
pub config_dir: PathBuf,
pub no_start: bool,
}
/// Parse daemon-related options from CLI args.
#[cfg(unix)]
pub fn parse_daemon_args(args: &[String]) -> DaemonOptions {
let mut opts = DaemonOptions::default();
let mut i = 0;
while i < args.len() {
match args[i].as_str() {
"--daemon" | "-d" => {
opts.daemonize = true;
}
"--foreground" | "-f" => {
opts.foreground = true;
}
"--pid-file" => {
i += 1;
if i < args.len() {
opts.pid_file = Some(PathBuf::from(&args[i]));
}
}
s if s.starts_with("--pid-file=") => {
opts.pid_file = Some(PathBuf::from(s.trim_start_matches("--pid-file=")));
}
"--run-as-user" => {
i += 1;
if i < args.len() {
opts.user = Some(args[i].clone());
}
}
s if s.starts_with("--run-as-user=") => {
opts.user = Some(s.trim_start_matches("--run-as-user=").to_string());
}
"--run-as-group" => {
i += 1;
if i < args.len() {
opts.group = Some(args[i].clone());
}
}
s if s.starts_with("--run-as-group=") => {
opts.group = Some(s.trim_start_matches("--run-as-group=").to_string());
}
"--working-dir" => {
i += 1;
if i < args.len() {
opts.working_dir = Some(PathBuf::from(&args[i]));
}
}
s if s.starts_with("--working-dir=") => {
opts.working_dir = Some(PathBuf::from(s.trim_start_matches("--working-dir=")));
}
_ => {}
}
i += 1;
}
opts
}
impl Default for InitOptions {
fn default() -> Self {
Self {
port: 443,
domain: "www.google.com".to_string(),
secret: None,
username: "user".to_string(),
config_dir: PathBuf::from("/etc/telemt"),
no_start: false,
}
}
}
/// Parse --init subcommand options from CLI args.
///
/// Returns `Some(InitOptions)` if `--init` was found, `None` otherwise.
pub fn parse_init_args(args: &[String]) -> Option<InitOptions> {
if !args.iter().any(|a| a == "--init") {
return None;
}
let mut opts = InitOptions::default();
let mut i = 0;
while i < args.len() {
match args[i].as_str() {
"--port" => {
i += 1;
if i < args.len() {
opts.port = args[i].parse().unwrap_or(443);
}
}
"--domain" => {
i += 1;
if i < args.len() {
opts.domain = args[i].clone();
}
}
"--secret" => {
i += 1;
if i < args.len() {
opts.secret = Some(args[i].clone());
}
}
"--user" => {
i += 1;
if i < args.len() {
opts.username = args[i].clone();
}
}
"--config-dir" => {
i += 1;
if i < args.len() {
opts.config_dir = PathBuf::from(&args[i]);
}
}
"--no-start" => {
opts.no_start = true;
}
_ => {}
}
i += 1;
}
Some(opts)
}
/// Run the fire-and-forget setup.
pub fn run_init(opts: InitOptions) -> Result<(), Box<dyn std::error::Error>> {
use crate::service::{self, InitSystem, ServiceOptions};
eprintln!("[telemt] Fire-and-forget setup");
eprintln!();
// 1. Detect init system
let init_system = service::detect_init_system();
eprintln!("[+] Detected init system: {}", init_system);
// 2. Generate or validate secret
let secret = match opts.secret {
Some(s) => {
if s.len() != 32 || !s.chars().all(|c| c.is_ascii_hexdigit()) {
eprintln!("[error] Secret must be exactly 32 hex characters");
std::process::exit(1);
}
s
}
None => generate_secret(),
};
eprintln!("[+] Secret: {}", secret);
eprintln!("[+] User: {}", opts.username);
eprintln!("[+] Port: {}", opts.port);
eprintln!("[+] Domain: {}", opts.domain);
// 3. Create config directory
fs::create_dir_all(&opts.config_dir)?;
let config_path = opts.config_dir.join("config.toml");
// 4. Write config
let config_content = generate_config(&opts.username, &secret, opts.port, &opts.domain);
fs::write(&config_path, &config_content)?;
eprintln!("[+] Config written to {}", config_path.display());
// 5. Generate and write service file
let exe_path =
std::env::current_exe().unwrap_or_else(|_| PathBuf::from("/usr/local/bin/telemt"));
let service_opts = ServiceOptions {
exe_path: &exe_path,
config_path: &config_path,
user: None, // Let systemd/init handle user
group: None,
pid_file: "/var/run/telemt.pid",
working_dir: Some("/var/lib/telemt"),
description: "Telemt MTProxy - Telegram MTProto Proxy",
};
let service_path = service::service_file_path(init_system);
let service_content = service::generate_service_file(init_system, &service_opts);
// Ensure parent directory exists
if let Some(parent) = Path::new(service_path).parent() {
let _ = fs::create_dir_all(parent);
}
match fs::write(service_path, &service_content) {
Ok(()) => {
eprintln!("[+] Service file written to {}", service_path);
// Make script executable for OpenRC/FreeBSD
#[cfg(unix)]
if init_system == InitSystem::OpenRC || init_system == InitSystem::FreeBSDRc {
use std::os::unix::fs::PermissionsExt;
let mut perms = fs::metadata(service_path)?.permissions();
perms.set_mode(0o755);
fs::set_permissions(service_path, perms)?;
}
}
Err(e) => {
eprintln!("[!] Cannot write service file (run as root?): {}", e);
eprintln!("[!] Manual service file content:");
eprintln!("{}", service_content);
// Still print links and installation instructions
eprintln!();
eprintln!("{}", service::installation_instructions(init_system));
print_links(&opts.username, &secret, opts.port, &opts.domain);
return Ok(());
}
}
// 6. Install and enable service based on init system
match init_system {
InitSystem::Systemd => {
run_cmd("systemctl", &["daemon-reload"]);
run_cmd("systemctl", &["enable", "telemt.service"]);
eprintln!("[+] Service enabled");
if !opts.no_start {
run_cmd("systemctl", &["start", "telemt.service"]);
eprintln!("[+] Service started");
std::thread::sleep(std::time::Duration::from_secs(1));
let status = Command::new("systemctl")
.args(["is-active", "telemt.service"])
.output();
match status {
Ok(out) if out.status.success() => {
eprintln!("[+] Service is running");
}
_ => {
eprintln!("[!] Service may not have started correctly");
eprintln!("[!] Check: journalctl -u telemt.service -n 20");
}
}
} else {
eprintln!("[+] Service not started (--no-start)");
eprintln!("[+] Start manually: systemctl start telemt.service");
}
}
InitSystem::OpenRC => {
run_cmd("rc-update", &["add", "telemt", "default"]);
eprintln!("[+] Service enabled");
if !opts.no_start {
run_cmd("rc-service", &["telemt", "start"]);
eprintln!("[+] Service started");
} else {
eprintln!("[+] Service not started (--no-start)");
eprintln!("[+] Start manually: rc-service telemt start");
}
}
InitSystem::FreeBSDRc => {
run_cmd("sysrc", &["telemt_enable=YES"]);
eprintln!("[+] Service enabled");
if !opts.no_start {
run_cmd("service", &["telemt", "start"]);
eprintln!("[+] Service started");
} else {
eprintln!("[+] Service not started (--no-start)");
eprintln!("[+] Start manually: service telemt start");
}
}
InitSystem::Unknown => {
eprintln!("[!] Unknown init system - service file written but not installed");
eprintln!("[!] You may need to install it manually");
}
}
eprintln!();
// 7. Print links
print_links(&opts.username, &secret, opts.port, &opts.domain);
Ok(())
}
fn generate_secret() -> String {
let mut rng = rand::rng();
let bytes: Vec<u8> = (0..16).map(|_| rng.random::<u8>()).collect();
hex::encode(bytes)
}
fn generate_config(username: &str, secret: &str, port: u16, domain: &str) -> String {
format!(
r#"# Telemt MTProxy — auto-generated config
# Re-run `telemt --init` to regenerate
show_link = ["{username}"]
[general]
# prefer_ipv6 is deprecated; use [network].prefer
prefer_ipv6 = false
fast_mode = true
use_middle_proxy = false
log_level = "normal"
desync_all_full = false
update_every = 43200
hardswap = false
me_pool_drain_ttl_secs = 90
me_instadrain = false
me_pool_drain_threshold = 32
me_pool_drain_soft_evict_grace_secs = 10
me_pool_drain_soft_evict_per_writer = 2
me_pool_drain_soft_evict_budget_per_core = 16
me_pool_drain_soft_evict_cooldown_ms = 1000
me_bind_stale_mode = "never"
me_pool_min_fresh_ratio = 0.8
me_reinit_drain_timeout_secs = 90
tg_connect = 10
[network]
ipv4 = true
ipv6 = true
prefer = 4
multipath = false
[general.modes]
classic = false
secure = false
tls = true
[server]
listen_addr_ipv4 = "0.0.0.0"
listen_addr_ipv6 = "::"
[[server.listeners]]
ip = "0.0.0.0"
port = {port}
# reuse_allow = false # Set true only when intentionally running multiple telemt instances on same port
[[server.listeners]]
ip = "::"
port = {port}
[timeouts]
client_first_byte_idle_secs = 300
client_handshake = 60
client_keepalive = 60
client_ack = 300
[censorship]
tls_domain = "{domain}"
mask = true
mask_port = 443
fake_cert_len = 2048
serverhello_compact = false
tls_full_cert_ttl_secs = 90
[access]
user_max_tcp_conns_global_each = 0
replay_check_len = 65536
replay_window_secs = 120
ignore_time_skew = false
[access.users]
{username} = "{secret}"
[[upstreams]]
type = "direct"
enabled = true
weight = 10
# Optional per-upstream DC family policy:
# ipv6 = true
# prefer = 6
"#,
username = username,
secret = secret,
port = port,
domain = domain,
)
}
fn run_cmd(cmd: &str, args: &[&str]) {
match Command::new(cmd).args(args).output() {
Ok(output) => {
if !output.status.success() {
let stderr = String::from_utf8_lossy(&output.stderr);
eprintln!("[!] {} {} failed: {}", cmd, args.join(" "), stderr.trim());
}
}
Err(e) => {
eprintln!("[!] Failed to run {} {}: {}", cmd, args.join(" "), e);
}
}
}
fn print_links(username: &str, secret: &str, port: u16, domain: &str) {
let domain_hex = hex::encode(domain);
println!("=== Proxy Links ===");
println!("[{}]", username);
println!(
" EE-TLS: tg://proxy?server=YOUR_SERVER_IP&port={}&secret=ee{}{}",
port, secret, domain_hex
);
println!();
println!("Replace YOUR_SERVER_IP with your server's public IP.");
println!("The proxy will auto-detect and display the correct link on startup.");
println!("Check: journalctl -u telemt.service | head -30");
println!("===================");
}
+141
View File
@@ -0,0 +1,141 @@
use std::path::{Path, PathBuf};
use crate::daemon::{self, DaemonOptions};
/// Parses daemon-related options from CLI arguments.
pub fn parse_daemon_args(args: &[String]) -> DaemonOptions {
let mut opts = DaemonOptions::default();
let mut i = 0;
while i < args.len() {
match args[i].as_str() {
"--daemon" | "-d" => {
opts.daemonize = true;
}
"--foreground" | "-f" => {
opts.foreground = true;
}
"--pid-file" => {
i += 1;
if i < args.len() {
opts.pid_file = Some(PathBuf::from(&args[i]));
}
}
s if s.starts_with("--pid-file=") => {
opts.pid_file = Some(PathBuf::from(s.trim_start_matches("--pid-file=")));
}
"--run-as-user" => {
i += 1;
if i < args.len() {
opts.user = Some(args[i].clone());
}
}
s if s.starts_with("--run-as-user=") => {
opts.user = Some(s.trim_start_matches("--run-as-user=").to_string());
}
"--run-as-group" => {
i += 1;
if i < args.len() {
opts.group = Some(args[i].clone());
}
}
s if s.starts_with("--run-as-group=") => {
opts.group = Some(s.trim_start_matches("--run-as-group=").to_string());
}
"--working-dir" => {
i += 1;
if i < args.len() {
opts.working_dir = Some(PathBuf::from(&args[i]));
}
}
s if s.starts_with("--working-dir=") => {
opts.working_dir = Some(PathBuf::from(s.trim_start_matches("--working-dir=")));
}
_ => {}
}
i += 1;
}
opts
}
/// Sends SIGTERM and waits briefly for graceful PID-file cleanup.
pub(super) fn stop(pid_file: &Path) -> i32 {
use nix::sys::signal::Signal;
println!("Stopping telemt daemon...");
match daemon::signal_pid_file(pid_file, Signal::SIGTERM) {
Ok(()) => {
println!("Stop signal sent successfully");
// Wait for process to exit for up to ten seconds.
for _ in 0..20 {
std::thread::sleep(std::time::Duration::from_millis(500));
if let daemon::DaemonStatus::NotRunning = daemon::check_status(pid_file) {
println!("Daemon stopped");
return 0;
}
}
println!("Daemon may still be shutting down");
0
}
Err(e) => {
eprintln!("Failed to stop daemon: {}", e);
1
}
}
}
/// Sends SIGHUP to trigger configuration reload.
pub(super) fn reload(pid_file: &Path) -> i32 {
use nix::sys::signal::Signal;
println!("Reloading telemt configuration...");
match daemon::signal_pid_file(pid_file, Signal::SIGHUP) {
Ok(()) => {
println!("Reload signal sent successfully");
0
}
Err(e) => {
eprintln!("Failed to reload daemon: {}", e);
1
}
}
}
/// Reports daemon status without mutating PID lifecycle state.
pub(super) fn status(pid_file: &Path) -> i32 {
match daemon::check_status(pid_file) {
daemon::DaemonStatus::Running(pid) => {
println!("telemt is running (pid {})", pid);
0
}
daemon::DaemonStatus::Stale(pid) => {
println!("telemt is not running (stale pid file, was pid {})", pid);
1
}
daemon::DaemonStatus::NotRunning => {
println!("telemt is not running");
1
}
}
}
#[cfg(test)]
mod tests {
use std::fs;
use super::*;
#[test]
fn status_does_not_remove_stale_pid_file() {
let directory = tempfile::tempdir().unwrap();
let pid_file = directory.path().join("telemt.pid");
fs::write(&pid_file, b"2000000000\n").unwrap();
assert_eq!(status(&pid_file), 1);
assert!(pid_file.exists());
}
}
+353
View File
@@ -0,0 +1,353 @@
use std::fs;
use std::path::{Path, PathBuf};
use std::process::Command;
use rand::RngExt;
/// Options for the fire-and-forget init command.
#[derive(Debug, Clone)]
pub struct InitOptions {
/// Public listener port.
pub port: u16,
/// TLS camouflage domain.
pub domain: String,
/// Optional pre-generated proxy secret.
pub secret: Option<String>,
/// Initial access username.
pub username: String,
/// Destination directory for generated configuration.
pub config_dir: PathBuf,
/// Generate service files without starting the service.
pub no_start: bool,
}
impl Default for InitOptions {
fn default() -> Self {
Self {
port: 443,
domain: "www.google.com".to_string(),
secret: None,
username: "user".to_string(),
config_dir: PathBuf::from("/etc/telemt"),
no_start: false,
}
}
}
/// Parse --init subcommand options from CLI args.
///
/// Returns `Some(InitOptions)` if `--init` was found, `None` otherwise.
pub fn parse_init_args(args: &[String]) -> Option<InitOptions> {
if !args.iter().any(|a| a == "--init") {
return None;
}
let mut opts = InitOptions::default();
let mut i = 0;
while i < args.len() {
match args[i].as_str() {
"--port" => {
i += 1;
if i < args.len() {
opts.port = args[i].parse().unwrap_or(443);
}
}
"--domain" => {
i += 1;
if i < args.len() {
opts.domain = args[i].clone();
}
}
"--secret" => {
i += 1;
if i < args.len() {
opts.secret = Some(args[i].clone());
}
}
"--user" => {
i += 1;
if i < args.len() {
opts.username = args[i].clone();
}
}
"--config-dir" => {
i += 1;
if i < args.len() {
opts.config_dir = PathBuf::from(&args[i]);
}
}
"--no-start" => {
opts.no_start = true;
}
_ => {}
}
i += 1;
}
Some(opts)
}
/// Run the fire-and-forget setup.
pub fn run_init(opts: InitOptions) -> Result<(), Box<dyn std::error::Error>> {
use crate::service::{self, InitSystem, ServiceOptions};
eprintln!("[telemt] Fire-and-forget setup");
eprintln!();
let init_system = service::detect_init_system();
eprintln!("[+] Detected init system: {}", init_system);
let secret = match opts.secret {
Some(s) => {
if s.len() != 32 || !s.chars().all(|c| c.is_ascii_hexdigit()) {
eprintln!("[error] Secret must be exactly 32 hex characters");
std::process::exit(1);
}
s
}
None => generate_secret(),
};
eprintln!("[+] Secret: {}", secret);
eprintln!("[+] User: {}", opts.username);
eprintln!("[+] Port: {}", opts.port);
eprintln!("[+] Domain: {}", opts.domain);
fs::create_dir_all(&opts.config_dir)?;
let config_path = opts.config_dir.join("config.toml");
let config_content = generate_config(&opts.username, &secret, opts.port, &opts.domain);
fs::write(&config_path, &config_content)?;
eprintln!("[+] Config written to {}", config_path.display());
let exe_path =
std::env::current_exe().unwrap_or_else(|_| PathBuf::from("/usr/local/bin/telemt"));
let service_opts = ServiceOptions {
exe_path: &exe_path,
config_path: &config_path,
// Let the selected init system manage process identity.
user: None,
group: None,
pid_file: "/var/run/telemt.pid",
working_dir: Some("/var/lib/telemt"),
description: "Telemt MTProxy - Telegram MTProto Proxy",
};
let service_path = service::service_file_path(init_system);
let service_content = service::generate_service_file(init_system, &service_opts);
if let Some(parent) = Path::new(service_path).parent() {
let _ = fs::create_dir_all(parent);
}
match fs::write(service_path, &service_content) {
Ok(()) => {
eprintln!("[+] Service file written to {}", service_path);
// OpenRC and FreeBSD service scripts must be executable.
#[cfg(unix)]
if init_system == InitSystem::OpenRC || init_system == InitSystem::FreeBSDRc {
use std::os::unix::fs::PermissionsExt;
let mut perms = fs::metadata(service_path)?.permissions();
perms.set_mode(0o755);
fs::set_permissions(service_path, perms)?;
}
}
Err(e) => {
eprintln!("[!] Cannot write service file (run as root?): {}", e);
eprintln!("[!] Manual service file content:");
eprintln!("{}", service_content);
eprintln!();
eprintln!("{}", service::installation_instructions(init_system));
print_links(&opts.username, &secret, opts.port, &opts.domain);
return Ok(());
}
}
match init_system {
InitSystem::Systemd => {
run_cmd("systemctl", &["daemon-reload"]);
run_cmd("systemctl", &["enable", "telemt.service"]);
eprintln!("[+] Service enabled");
if !opts.no_start {
run_cmd("systemctl", &["start", "telemt.service"]);
eprintln!("[+] Service started");
std::thread::sleep(std::time::Duration::from_secs(1));
let status = Command::new("systemctl")
.args(["is-active", "telemt.service"])
.output();
match status {
Ok(out) if out.status.success() => {
eprintln!("[+] Service is running");
}
_ => {
eprintln!("[!] Service may not have started correctly");
eprintln!("[!] Check: journalctl -u telemt.service -n 20");
}
}
} else {
eprintln!("[+] Service not started (--no-start)");
eprintln!("[+] Start manually: systemctl start telemt.service");
}
}
InitSystem::OpenRC => {
run_cmd("rc-update", &["add", "telemt", "default"]);
eprintln!("[+] Service enabled");
if !opts.no_start {
run_cmd("rc-service", &["telemt", "start"]);
eprintln!("[+] Service started");
} else {
eprintln!("[+] Service not started (--no-start)");
eprintln!("[+] Start manually: rc-service telemt start");
}
}
InitSystem::FreeBSDRc => {
run_cmd("sysrc", &["telemt_enable=YES"]);
eprintln!("[+] Service enabled");
if !opts.no_start {
run_cmd("service", &["telemt", "start"]);
eprintln!("[+] Service started");
} else {
eprintln!("[+] Service not started (--no-start)");
eprintln!("[+] Start manually: service telemt start");
}
}
InitSystem::Unknown => {
eprintln!("[!] Unknown init system - service file written but not installed");
eprintln!("[!] You may need to install it manually");
}
}
eprintln!();
print_links(&opts.username, &secret, opts.port, &opts.domain);
Ok(())
}
fn generate_secret() -> String {
let mut rng = rand::rng();
let bytes: Vec<u8> = (0..16).map(|_| rng.random::<u8>()).collect();
hex::encode(bytes)
}
fn generate_config(username: &str, secret: &str, port: u16, domain: &str) -> String {
format!(
r#"# Telemt MTProxy — auto-generated config
# Re-run `telemt --init` to regenerate
show_link = ["{username}"]
[general]
# prefer_ipv6 is deprecated; use [network].prefer
prefer_ipv6 = false
fast_mode = true
use_middle_proxy = false
log_level = "normal"
desync_all_full = false
update_every = 43200
hardswap = false
me_pool_drain_ttl_secs = 90
me_instadrain = false
me_pool_drain_threshold = 32
me_pool_drain_soft_evict_grace_secs = 10
me_pool_drain_soft_evict_per_writer = 2
me_pool_drain_soft_evict_budget_per_core = 16
me_pool_drain_soft_evict_cooldown_ms = 1000
me_bind_stale_mode = "never"
me_pool_min_fresh_ratio = 0.8
me_reinit_drain_timeout_secs = 90
tg_connect = 10
[network]
ipv4 = true
ipv6 = true
prefer = 4
multipath = false
[general.modes]
classic = false
secure = false
tls = true
[server]
listen_addr_ipv4 = "0.0.0.0"
listen_addr_ipv6 = "::"
[[server.listeners]]
ip = "0.0.0.0"
port = {port}
# reuse_allow = false # Set true only when intentionally running multiple telemt instances on same port
[[server.listeners]]
ip = "::"
port = {port}
[timeouts]
client_first_byte_idle_secs = 300
client_handshake = 60
client_keepalive = 60
client_ack = 300
[censorship]
tls_domain = "{domain}"
mask = true
mask_port = 443
fake_cert_len = 2048
serverhello_compact = false
tls_full_cert_ttl_secs = 90
[access]
user_max_tcp_conns_global_each = 0
replay_check_len = 65536
replay_window_secs = 120
ignore_time_skew = false
[access.users]
{username} = "{secret}"
[[upstreams]]
type = "direct"
enabled = true
weight = 10
# Optional per-upstream DC family policy:
# ipv6 = true
# prefer = 6
"#,
username = username,
secret = secret,
port = port,
domain = domain,
)
}
fn run_cmd(cmd: &str, args: &[&str]) {
match Command::new(cmd).args(args).output() {
Ok(output) => {
if !output.status.success() {
let stderr = String::from_utf8_lossy(&output.stderr);
eprintln!("[!] {} {} failed: {}", cmd, args.join(" "), stderr.trim());
}
}
Err(e) => {
eprintln!("[!] Failed to run {} {}: {}", cmd, args.join(" "), e);
}
}
}
fn print_links(username: &str, secret: &str, port: u16, domain: &str) {
let domain_hex = hex::encode(domain);
println!("=== Proxy Links ===");
println!("[{}]", username);
println!(
" EE-TLS: tg://proxy?server=YOUR_SERVER_IP&port={}&secret=ee{}{}",
port, secret, domain_hex
);
println!();
println!("Replace YOUR_SERVER_IP with your server's public IP.");
println!("The proxy will auto-detect and display the correct link on startup.");
println!("Check: journalctl -u telemt.service | head -30");
println!("===================");
}
+14
View File
@@ -196,6 +196,13 @@ pub(super) fn validate(config: &mut ProxyConfig) -> Result<()> {
"access.user_rate_limits.{user} must set at least one non-zero direction"
)));
}
for (direction, value) in [("up_bps", limit.up_bps), ("down_bps", limit.down_bps)] {
if value > MAX_RATE_LIMIT_BPS {
return Err(ProxyError::Config(format!(
"access.user_rate_limits.{user}.{direction} must be within [0, {MAX_RATE_LIMIT_BPS}]"
)));
}
}
}
for (cidr, limit) in &config.access.cidr_rate_limits {
@@ -204,6 +211,13 @@ pub(super) fn validate(config: &mut ProxyConfig) -> Result<()> {
"access.cidr_rate_limits.{cidr} must set at least one non-zero direction"
)));
}
for (direction, value) in [("up_bps", limit.up_bps), ("down_bps", limit.down_bps)] {
if value > MAX_RATE_LIMIT_BPS {
return Err(ProxyError::Config(format!(
"access.cidr_rate_limits.{cidr}.{direction} must be within [0, {MAX_RATE_LIMIT_BPS}]"
)));
}
}
}
let mut cidr_auto_templates = HashSet::new();
for cidr in config.access.cidr_rate_limits.keys() {
@@ -288,6 +288,68 @@ fn cidr_rate_limits_reject_duplicate_normalized_auto_templates() {
assert!(error.contains("duplicates normalized auto-template *6/128"));
}
#[test]
fn rate_limits_accept_the_packed_counter_maximum() {
let cfg = load_config_from_temp_toml(
r#"
[censorship]
tls_domain = "example.com"
[access.users]
user = "00000000000000000000000000000000"
[access.user_rate_limits]
user = { up_bps = 100000000000, down_bps = 0 }
[access.cidr_rate_limits]
"203.0.113.0/24" = { up_bps = 0, down_bps = 100000000000 }
"#,
);
assert_eq!(cfg.access.user_rate_limits["user"].up_bps, 100_000_000_000);
assert_eq!(
cfg.access.cidr_rate_limits[&CidrRateLimitKey::Network("203.0.113.0/24".parse().unwrap())]
.down_bps,
100_000_000_000
);
}
#[test]
fn user_rate_limits_reject_values_above_the_packed_counter_maximum() {
let error = load_config_error_from_temp_toml(
r#"
[censorship]
tls_domain = "example.com"
[access.users]
user = "00000000000000000000000000000000"
[access.user_rate_limits]
user = { up_bps = 100000000001, down_bps = 0 }
"#,
);
assert!(error.contains("access.user_rate_limits.user.up_bps must be within"));
}
#[test]
fn cidr_rate_limits_reject_values_above_the_packed_counter_maximum() {
let error = load_config_error_from_temp_toml(
r#"
[censorship]
tls_domain = "example.com"
[access.users]
user = "00000000000000000000000000000000"
[access.cidr_rate_limits]
"203.0.113.0/24" = { up_bps = 0, down_bps = 100000000001 }
"#,
);
assert!(error.contains("access.cidr_rate_limits.203.0.113.0/24.down_bps must be within"));
}
#[test]
fn file_logging_requires_path() {
let error = load_config_error_from_temp_toml(
+1 -1
View File
@@ -31,7 +31,7 @@ mod web_debug;
pub use access::{AccessConfig, CidrRateLimitKey, RateLimitBps};
#[allow(unused_imports)]
pub(crate) use access::{CidrAutoTemplate, CidrAutoTemplateFamily};
pub(crate) use access::{CidrAutoTemplate, CidrAutoTemplateFamily, MAX_RATE_LIMIT_BPS};
pub use api::{ApiConfig, ApiGrayAction};
pub use censorship::{
AntiCensorshipConfig, ExclusiveMaskTarget, TlsFetchConfig, TlsFetchProfile, UnknownSniAction,
+5 -2
View File
@@ -1,5 +1,8 @@
use super::*;
/// Highest rate that fits one packed 20 ms shaping epoch.
pub(crate) const MAX_RATE_LIMIT_BPS: u64 = 100_000_000_000;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct AccessConfig {
#[serde(default = "default_access_users")]
@@ -260,10 +263,10 @@ fn parse_cidr_auto_prefix(
/// Transport rate limit in bits-per-second.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct RateLimitBps {
/// Upload direction limit in bits-per-second; `0` means unlimited.
/// Upload limit in bits-per-second within `0..=100_000_000_000`; `0` means unlimited.
#[serde(default)]
pub up_bps: u64,
/// Download direction limit in bits-per-second; `0` means unlimited.
/// Download limit in bits-per-second within `0..=100_000_000_000`; `0` means unlimited.
#[serde(default)]
pub down_bps: u64,
}
+28 -279
View File
@@ -4,14 +4,20 @@
//! and privilege dropping for running telemt as a background service.
use std::fs::{self, File, OpenOptions};
use std::io::{self, Read, Write};
use std::io;
use std::os::unix::fs::OpenOptionsExt;
use std::path::{Path, PathBuf};
use nix::errno::Errno;
use nix::fcntl::{Flock, FlockArg};
use nix::unistd::{self, ForkResult, Gid, Pid, Uid, chdir, close, fork, getpid, setsid};
use tracing::{debug, info, warn};
use nix::unistd::{self, ForkResult, Gid, Uid, chdir, close, fork, getpid, setsid};
use tracing::info;
// PID file ownership and process-control helpers.
mod pid_file;
pub use pid_file::DaemonStatus;
#[allow(unused_imports)]
pub use pid_file::{PidFile, check_status, read_pid_file, signal_pid_file};
/// Default PID file location.
pub const DEFAULT_PID_FILE: &str = "/var/run/telemt.pid";
@@ -51,36 +57,47 @@ impl DaemonOptions {
/// Error types for daemon operations.
#[derive(Debug, thiserror::Error)]
pub enum DaemonError {
/// A daemonization fork failed.
#[error("fork failed: {0}")]
ForkFailed(#[source] nix::Error),
/// Creation of the detached session failed.
#[error("setsid failed: {0}")]
SetsidFailed(#[source] nix::Error),
/// Switching to the configured working directory failed.
#[error("chdir failed: {0}")]
ChdirFailed(#[source] nix::Error),
/// Opening `/dev/null` for standard-stream redirection failed.
#[error("failed to open /dev/null: {0}")]
DevNullFailed(#[source] io::Error),
/// Redirecting a standard file descriptor failed.
#[error("failed to redirect stdio: {0}")]
RedirectFailed(#[source] nix::Error),
/// A PID lifecycle operation failed.
#[error("PID file error: {0}")]
PidFile(String),
/// Another process owns the daemon PID lifecycle.
#[error("another instance is already running (pid {0})")]
AlreadyRunning(i32),
/// The configured runtime user does not exist.
#[error("user '{0}' not found")]
UserNotFound(String),
/// The configured runtime group does not exist.
#[error("group '{0}' not found")]
GroupNotFound(String),
/// Applying the configured runtime identity failed.
#[error("failed to set uid/gid: {0}")]
PrivilegeDrop(#[source] nix::Error),
/// An underlying filesystem operation failed.
#[error("io error: {0}")]
Io(#[from] io::Error),
}
@@ -106,38 +123,28 @@ pub enum DaemonizeResult {
/// Returns `DaemonizeResult::Parent` in the original parent (which should exit),
/// or `DaemonizeResult::Child` in the final daemon child.
pub fn daemonize(working_dir: Option<&Path>) -> Result<DaemonizeResult, DaemonError> {
// First fork
match unsafe { fork() } {
Ok(ForkResult::Parent { .. }) => {
// Parent exits
return Ok(DaemonizeResult::Parent);
}
Ok(ForkResult::Child) => {
// Child continues
}
Ok(ForkResult::Child) => {}
Err(e) => return Err(DaemonError::ForkFailed(e)),
}
// Create new session, become session leader
setsid().map_err(DaemonError::SetsidFailed)?;
// Second fork to ensure we can never acquire a controlling terminal
match unsafe { fork() } {
Ok(ForkResult::Parent { .. }) => {
// Intermediate parent exits
std::process::exit(0);
}
Ok(ForkResult::Child) => {
// Final daemon child continues
}
Ok(ForkResult::Child) => {}
Err(e) => return Err(DaemonError::ForkFailed(e)),
}
// Change working directory
let target_dir = working_dir.unwrap_or(Path::new("/"));
chdir(target_dir).map_err(DaemonError::ChdirFailed)?;
// Redirect stdin, stdout, stderr to /dev/null
redirect_stdio_to_devnull()?;
Ok(DaemonizeResult::Child)
@@ -156,21 +163,17 @@ fn redirect_stdio_to_devnull() -> Result<(), DaemonError> {
// Use libc::dup2 directly for redirecting standard file descriptors
// nix 0.31's dup2 requires OwnedFd which doesn't work well with stdio fds
unsafe {
// Redirect stdin (fd 0)
if libc::dup2(devnull_fd, 0) < 0 {
return Err(DaemonError::RedirectFailed(Errno::last()));
}
// Redirect stdout (fd 1)
if libc::dup2(devnull_fd, 1) < 0 {
return Err(DaemonError::RedirectFailed(Errno::last()));
}
// Redirect stderr (fd 2)
if libc::dup2(devnull_fd, 2) < 0 {
return Err(DaemonError::RedirectFailed(Errno::last()));
}
}
// Close original devnull fd if it's not one of the standard fds
if devnull_fd > 2 {
let _ = close(devnull_fd);
}
@@ -178,166 +181,6 @@ fn redirect_stdio_to_devnull() -> Result<(), DaemonError> {
Ok(())
}
/// PID file manager with flock-based locking.
pub struct PidFile {
path: PathBuf,
file: Option<File>,
locked: bool,
}
impl PidFile {
/// Creates a new PID file manager for the given path.
pub fn new<P: AsRef<Path>>(path: P) -> Self {
Self {
path: path.as_ref().to_path_buf(),
file: None,
locked: false,
}
}
/// Checks if another instance is already running.
///
/// Returns the PID of the running instance if one exists.
pub fn check_running(&self) -> Result<Option<i32>, DaemonError> {
if !self.path.exists() {
return Ok(None);
}
// Try to read existing PID
let mut contents = String::new();
File::open(&self.path)
.and_then(|mut f| f.read_to_string(&mut contents))
.map_err(|e| {
DaemonError::PidFile(format!("cannot read {}: {}", self.path.display(), e))
})?;
let pid: i32 = contents
.trim()
.parse()
.map_err(|_| DaemonError::PidFile(format!("invalid PID in {}", self.path.display())))?;
// Check if process is still running
if is_process_running(pid) {
Ok(Some(pid))
} else {
// Stale PID file
debug!(pid, path = %self.path.display(), "Removing stale PID file");
let _ = fs::remove_file(&self.path);
Ok(None)
}
}
/// Acquires the PID file lock and writes the current PID.
///
/// Fails if another instance is already running.
pub fn acquire(&mut self) -> Result<(), DaemonError> {
// Check for running instance first
if let Some(pid) = self.check_running()? {
return Err(DaemonError::AlreadyRunning(pid));
}
// Ensure parent directory exists
if let Some(parent) = self.path.parent() {
if !parent.exists() {
fs::create_dir_all(parent).map_err(|e| {
DaemonError::PidFile(format!(
"cannot create directory {}: {}",
parent.display(),
e
))
})?;
}
}
// Open/create PID file with exclusive lock
let file = OpenOptions::new()
.write(true)
.create(true)
.truncate(true)
.mode(0o644)
.open(&self.path)
.map_err(|e| {
DaemonError::PidFile(format!("cannot open {}: {}", self.path.display(), e))
})?;
// Try to acquire exclusive lock (non-blocking)
let flock = Flock::lock(file, FlockArg::LockExclusiveNonblock).map_err(|(_, errno)| {
// Check if another instance grabbed the lock
if let Some(pid) = self.check_running().ok().flatten() {
DaemonError::AlreadyRunning(pid)
} else {
DaemonError::PidFile(format!("cannot lock {}: {}", self.path.display(), errno))
}
})?;
// Write our PID
let pid = getpid();
let mut file = flock
.unlock()
.map_err(|(_, errno)| DaemonError::PidFile(format!("unlock failed: {}", errno)))?;
writeln!(file, "{}", pid).map_err(|e| {
DaemonError::PidFile(format!(
"cannot write PID to {}: {}",
self.path.display(),
e
))
})?;
// Re-acquire lock and keep it
let flock = Flock::lock(file, FlockArg::LockExclusiveNonblock).map_err(|(_, errno)| {
DaemonError::PidFile(format!("cannot re-lock {}: {}", self.path.display(), errno))
})?;
self.file = Some(flock.unlock().map_err(|(_, errno)| {
DaemonError::PidFile(format!("unlock for storage failed: {}", errno))
})?);
self.locked = true;
info!(pid = pid.as_raw(), path = %self.path.display(), "PID file created");
Ok(())
}
/// Releases the PID file lock and removes the file.
pub fn release(&mut self) -> Result<(), DaemonError> {
if let Some(file) = self.file.take() {
drop(file);
}
self.locked = false;
if self.path.exists() {
fs::remove_file(&self.path).map_err(|e| {
DaemonError::PidFile(format!("cannot remove {}: {}", self.path.display(), e))
})?;
debug!(path = %self.path.display(), "PID file removed");
}
Ok(())
}
/// Returns the path to this PID file.
#[allow(dead_code)]
pub fn path(&self) -> &Path {
&self.path
}
}
impl Drop for PidFile {
fn drop(&mut self) {
if self.locked {
if let Err(e) = self.release() {
warn!(error = %e, "Failed to clean up PID file on drop");
}
}
}
}
/// Checks if a process with the given PID is running.
fn is_process_running(pid: i32) -> bool {
// kill(pid, 0) checks if process exists without sending a signal
nix::sys::signal::kill(Pid::from_raw(pid), None).is_ok()
}
// macOS gates nix::unistd::setgroups differently in the current dependency set,
// so call libc directly there while preserving the original nix path elsewhere.
fn set_supplementary_groups(gid: Gid) -> Result<(), nix::Error> {
@@ -383,9 +226,11 @@ pub fn drop_privileges(
};
if (target_uid.is_some() || target_gid.is_some())
&& let Some(file) = pid_file.and_then(|pid| pid.file.as_ref())
&& let Some(pid_file) = pid_file
{
unistd::fchown(file, target_uid, target_gid).map_err(DaemonError::PrivilegeDrop)?;
for file in pid_file.ownership_file_handles().into_iter().flatten() {
unistd::fchown(file, target_uid, target_gid).map_err(DaemonError::PrivilegeDrop)?;
}
}
if let Some(gid) = target_gid {
@@ -401,7 +246,7 @@ pub fn drop_privileges(
if uid.as_raw() != 0
&& let Some(pid) = pid_file
{
let parent = pid.path.parent().unwrap_or(Path::new("."));
let parent = pid.path().parent().unwrap_or(Path::new("."));
let probe_path = parent.join(format!(
".telemt_pid_probe_{}_{}",
std::process::id(),
@@ -436,7 +281,6 @@ pub fn drop_privileges(
/// Looks up a user by name and returns their UID.
fn lookup_user(name: &str) -> Result<Uid, DaemonError> {
// Use libc getpwnam
let c_name =
std::ffi::CString::new(name).map_err(|_| DaemonError::UserNotFound(name.to_string()))?;
@@ -480,76 +324,6 @@ fn lookup_group(name: &str) -> Result<Gid, DaemonError> {
}
}
/// Reads PID from a PID file.
#[allow(dead_code)]
pub fn read_pid_file<P: AsRef<Path>>(path: P) -> Result<i32, DaemonError> {
let path = path.as_ref();
let mut contents = String::new();
File::open(path)
.and_then(|mut f| f.read_to_string(&mut contents))
.map_err(|e| DaemonError::PidFile(format!("cannot read {}: {}", path.display(), e)))?;
contents
.trim()
.parse()
.map_err(|_| DaemonError::PidFile(format!("invalid PID in {}", path.display())))
}
/// Sends a signal to the process specified in a PID file.
#[allow(dead_code)]
pub fn signal_pid_file<P: AsRef<Path>>(
path: P,
signal: nix::sys::signal::Signal,
) -> Result<(), DaemonError> {
let pid = read_pid_file(&path)?;
if !is_process_running(pid) {
return Err(DaemonError::PidFile(format!(
"process {} from {} is not running",
pid,
path.as_ref().display()
)));
}
nix::sys::signal::kill(Pid::from_raw(pid), signal)
.map_err(|e| DaemonError::PidFile(format!("cannot signal process {}: {}", pid, e)))?;
Ok(())
}
/// Returns the status of the daemon based on PID file.
#[allow(dead_code)]
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum DaemonStatus {
/// Daemon is running with the given PID.
Running(i32),
/// PID file exists but process is not running.
Stale(i32),
/// No PID file exists.
NotRunning,
}
/// Checks the daemon status from a PID file.
#[allow(dead_code)]
pub fn check_status<P: AsRef<Path>>(path: P) -> DaemonStatus {
let path = path.as_ref();
if !path.exists() {
return DaemonStatus::NotRunning;
}
match read_pid_file(path) {
Ok(pid) => {
if is_process_running(pid) {
DaemonStatus::Running(pid)
} else {
DaemonStatus::Stale(pid)
}
}
Err(_) => DaemonStatus::NotRunning,
}
}
#[cfg(test)]
mod tests {
use super::*;
@@ -571,29 +345,4 @@ mod tests {
};
assert!(!opts.should_daemonize());
}
#[test]
fn test_check_status_not_running() {
let path = "/tmp/telemt_test_nonexistent.pid";
assert_eq!(check_status(path), DaemonStatus::NotRunning);
}
#[test]
fn test_pid_file_basic() {
let path = "/tmp/telemt_test_pidfile.pid";
let _ = fs::remove_file(path);
let mut pf = PidFile::new(path);
assert!(pf.check_running().unwrap().is_none());
pf.acquire().unwrap();
assert!(Path::new(path).exists());
// Read it back
let pid = read_pid_file(path).unwrap();
assert_eq!(pid, std::process::id() as i32);
pf.release().unwrap();
assert!(!Path::new(path).exists());
}
}
+394
View File
@@ -0,0 +1,394 @@
use std::fs::{self, File, OpenOptions};
use std::io::{ErrorKind, Read, Write};
use std::os::unix::fs::OpenOptionsExt;
use std::path::{Path, PathBuf};
use nix::fcntl::{Flock, FlockArg};
use nix::unistd::{Pid, getpid};
use tracing::{debug, info, warn};
use super::DaemonError;
/// PID file manager backed by a persistent sibling lock file.
pub struct PidFile {
path: PathBuf,
lock_path: PathBuf,
pid_file: Option<File>,
lock_file: Option<Flock<File>>,
}
impl PidFile {
/// Creates a new PID file manager for the given path.
pub fn new<P: AsRef<Path>>(path: P) -> Self {
let path = path.as_ref().to_path_buf();
let lock_path = sibling_lock_path(&path);
Self {
path,
lock_path,
pid_file: None,
lock_file: None,
}
}
/// Checks whether the PID file names a running process without modifying either file.
pub fn check_running(&self) -> Result<Option<i32>, DaemonError> {
let Some(pid) = read_pid_file_if_exists(&self.path)? else {
return Ok(None);
};
Ok(is_process_running(pid).then_some(pid))
}
/// Acquires the persistent sibling lock and writes the current PID.
///
/// Fails if another owner holds the lock or the existing PID names a running process.
pub fn acquire(&mut self) -> Result<(), DaemonError> {
if let Some(parent) = self.path.parent()
&& !parent.exists()
{
fs::create_dir_all(parent).map_err(|error| {
DaemonError::PidFile(format!(
"cannot create directory {}: {}",
parent.display(),
error
))
})?;
}
let lock_file = OpenOptions::new()
.read(true)
.write(true)
.create(true)
.truncate(false)
.mode(0o644)
.open(&self.lock_path)
.map_err(|error| {
DaemonError::PidFile(format!(
"cannot open lock file {}: {}",
self.lock_path.display(),
error
))
})?;
let lock_file =
Flock::lock(lock_file, FlockArg::LockExclusiveNonblock).map_err(|(_, errno)| {
if let Some(pid) = self.check_running().ok().flatten() {
DaemonError::AlreadyRunning(pid)
} else {
DaemonError::PidFile(format!(
"cannot lock {}: {}",
self.lock_path.display(),
errno
))
}
})?;
if let Some(pid) = self.check_running()? {
return Err(DaemonError::AlreadyRunning(pid));
}
let mut pid_file = OpenOptions::new()
.read(true)
.write(true)
.create(true)
.truncate(true)
.mode(0o644)
.open(&self.path)
.map_err(|error| {
DaemonError::PidFile(format!("cannot open {}: {}", self.path.display(), error))
})?;
let pid = getpid();
writeln!(pid_file, "{}", pid).map_err(|error| {
DaemonError::PidFile(format!(
"cannot write PID to {}: {}",
self.path.display(),
error
))
})?;
self.pid_file = Some(pid_file);
self.lock_file = Some(lock_file);
info!(pid = pid.as_raw(), path = %self.path.display(), "PID file created");
Ok(())
}
/// Removes the PID file while retaining exclusive lock ownership until cleanup completes.
pub fn release(&mut self) -> Result<(), DaemonError> {
if self.lock_file.is_none() {
self.pid_file = None;
return Ok(());
}
let removal = match fs::remove_file(&self.path) {
Ok(()) => Ok(()),
Err(error) if error.kind() == ErrorKind::NotFound => Ok(()),
Err(error) => Err(DaemonError::PidFile(format!(
"cannot remove {}: {}",
self.path.display(),
error
))),
};
self.pid_file = None;
self.lock_file = None;
removal?;
debug!(path = %self.path.display(), "PID file removed");
Ok(())
}
/// Returns the path to this PID file.
pub fn path(&self) -> &Path {
&self.path
}
/// Returns open files whose ownership must follow the target runtime identity.
pub(super) fn ownership_file_handles(&self) -> [Option<&File>; 2] {
[self.pid_file.as_ref(), self.lock_file.as_deref()]
}
}
impl Drop for PidFile {
fn drop(&mut self) {
if self.lock_file.is_some()
&& let Err(error) = self.release()
{
warn!(error = %error, "Failed to clean up PID file on drop");
}
}
}
fn sibling_lock_path(path: &Path) -> PathBuf {
let mut lock_path = path.as_os_str().to_os_string();
lock_path.push(".lock");
lock_path.into()
}
fn read_pid_file_if_exists(path: &Path) -> Result<Option<i32>, DaemonError> {
let mut file = match File::open(path) {
Ok(file) => file,
Err(error) if error.kind() == ErrorKind::NotFound => return Ok(None),
Err(error) => {
return Err(DaemonError::PidFile(format!(
"cannot read {}: {}",
path.display(),
error
)));
}
};
let mut contents = String::new();
file.read_to_string(&mut contents).map_err(|error| {
DaemonError::PidFile(format!("cannot read {}: {}", path.display(), error))
})?;
let pid = contents
.trim()
.parse()
.map_err(|_| DaemonError::PidFile(format!("invalid PID in {}", path.display())))?;
Ok(Some(pid))
}
/// Reads a PID from a PID file.
#[allow(dead_code)]
pub fn read_pid_file<P: AsRef<Path>>(path: P) -> Result<i32, DaemonError> {
let path = path.as_ref();
read_pid_file_if_exists(path)?.ok_or_else(|| {
DaemonError::PidFile(format!(
"cannot read {}: file does not exist",
path.display()
))
})
}
/// Sends a signal to the process specified in a PID file.
#[allow(dead_code)]
pub fn signal_pid_file<P: AsRef<Path>>(
path: P,
signal: nix::sys::signal::Signal,
) -> Result<(), DaemonError> {
let pid = read_pid_file(&path)?;
if !is_process_running(pid) {
return Err(DaemonError::PidFile(format!(
"process {} from {} is not running",
pid,
path.as_ref().display()
)));
}
nix::sys::signal::kill(Pid::from_raw(pid), signal)
.map_err(|error| DaemonError::PidFile(format!("cannot signal process {}: {}", pid, error)))
}
/// Daemon state derived from the PID file.
#[allow(dead_code)]
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum DaemonStatus {
/// Daemon is running with the given PID.
Running(i32),
/// PID file exists but the named process is not running.
Stale(i32),
/// No readable PID file exists.
NotRunning,
}
/// Checks daemon status without modifying the PID or lock file.
#[allow(dead_code)]
pub fn check_status<P: AsRef<Path>>(path: P) -> DaemonStatus {
let path = path.as_ref();
match read_pid_file_if_exists(path) {
Ok(Some(pid)) if is_process_running(pid) => DaemonStatus::Running(pid),
Ok(Some(pid)) => DaemonStatus::Stale(pid),
Ok(None) | Err(_) => DaemonStatus::NotRunning,
}
}
fn is_process_running(pid: i32) -> bool {
nix::sys::signal::kill(Pid::from_raw(pid), None).is_ok()
}
#[cfg(test)]
mod tests {
use std::os::unix::fs::MetadataExt;
use std::process::{Child, Command, Stdio};
use std::thread;
use std::time::{Duration, Instant};
use super::*;
const HELPER_PID_PATH: &str = "TELEMT_PID_LOCK_HELPER_PATH";
const HELPER_READY_PATH: &str = "TELEMT_PID_LOCK_HELPER_READY";
const HELPER_STOP_PATH: &str = "TELEMT_PID_LOCK_HELPER_STOP";
fn wait_for_path(path: &Path, timeout: Duration) -> bool {
let deadline = Instant::now() + timeout;
while Instant::now() < deadline {
if path.exists() {
return true;
}
thread::sleep(Duration::from_millis(10));
}
false
}
fn wait_for_child(child: &mut Child, timeout: Duration) -> Option<std::process::ExitStatus> {
let deadline = Instant::now() + timeout;
while Instant::now() < deadline {
if let Some(status) = child.try_wait().unwrap() {
return Some(status);
}
thread::sleep(Duration::from_millis(10));
}
None
}
#[test]
fn pid_file_remains_send_and_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<PidFile>();
}
#[test]
fn lock_holder_subprocess() {
let Some(pid_path) = std::env::var_os(HELPER_PID_PATH) else {
return;
};
let ready_path = PathBuf::from(std::env::var_os(HELPER_READY_PATH).unwrap());
let stop_path = PathBuf::from(std::env::var_os(HELPER_STOP_PATH).unwrap());
let mut pid_file = PidFile::new(PathBuf::from(pid_path));
pid_file.acquire().unwrap();
fs::write(&ready_path, b"ready").unwrap();
assert!(wait_for_path(&stop_path, Duration::from_secs(10)));
pid_file.release().unwrap();
}
#[test]
fn persistent_sibling_lock_serializes_processes_after_pid_unlink() {
let directory = tempfile::tempdir().unwrap();
let pid_path = directory.path().join("telemt.pid");
let lock_path = sibling_lock_path(&pid_path);
let ready_path = directory.path().join("ready");
let stop_path = directory.path().join("stop");
let mut child = Command::new(std::env::current_exe().unwrap())
.args([
"--exact",
"daemon::pid_file::tests::lock_holder_subprocess",
"--nocapture",
])
.env(HELPER_PID_PATH, &pid_path)
.env(HELPER_READY_PATH, &ready_path)
.env(HELPER_STOP_PATH, &stop_path)
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.unwrap();
if !wait_for_path(&ready_path, Duration::from_secs(5)) {
let _ = child.kill();
let _ = child.wait();
panic!("PID lock holder did not become ready");
}
let lock_inode = fs::metadata(&lock_path).unwrap().ino();
fs::remove_file(&pid_path).unwrap();
let mut contender = PidFile::new(&pid_path);
assert!(contender.acquire().is_err());
fs::write(&stop_path, b"stop").unwrap();
let status = wait_for_child(&mut child, Duration::from_secs(5)).unwrap_or_else(|| {
let _ = child.kill();
child.wait().unwrap()
});
assert!(status.success());
assert_eq!(fs::metadata(&lock_path).unwrap().ino(), lock_inode);
contender.acquire().unwrap();
assert_eq!(fs::metadata(&lock_path).unwrap().ino(), lock_inode);
contender.release().unwrap();
assert!(!pid_path.exists());
assert!(lock_path.exists());
}
#[test]
fn stale_pid_checks_are_read_only() {
let directory = tempfile::tempdir().unwrap();
let pid_path = directory.path().join("telemt.pid");
fs::write(&pid_path, b"2000000000\n").unwrap();
let pid_file = PidFile::new(&pid_path);
assert_eq!(pid_file.check_running().unwrap(), None);
assert_eq!(check_status(&pid_path), DaemonStatus::Stale(2_000_000_000));
assert!(pid_path.exists());
}
#[test]
fn unowned_release_does_not_remove_pid_file() {
let directory = tempfile::tempdir().unwrap();
let pid_path = directory.path().join("telemt.pid");
fs::write(&pid_path, b"2000000000\n").unwrap();
let mut pid_file = PidFile::new(&pid_path);
pid_file.release().unwrap();
assert!(pid_path.exists());
}
#[test]
fn pid_file_release_keeps_lock_inode() {
let directory = tempfile::tempdir().unwrap();
let pid_path = directory.path().join("telemt.pid");
let lock_path = sibling_lock_path(&pid_path);
let mut pid_file = PidFile::new(&pid_path);
pid_file.acquire().unwrap();
assert!(
pid_file
.ownership_file_handles()
.into_iter()
.all(|file| file.is_some())
);
assert_eq!(read_pid_file(&pid_path).unwrap(), std::process::id() as i32);
let lock_inode = fs::metadata(&lock_path).unwrap().ino();
pid_file.release().unwrap();
assert!(!pid_path.exists());
assert!(lock_path.exists());
pid_file.acquire().unwrap();
assert_eq!(fs::metadata(&lock_path).unwrap().ino(), lock_inode);
pid_file.release().unwrap();
}
}
+3 -2
View File
@@ -288,8 +288,9 @@ pub(crate) async fn spawn_runtime_tasks(
break;
}
let cfg = config_rx_user_enabled.borrow_and_update().clone();
for user in shared_user_enabled.apply_user_enabled_config(&cfg.access.user_enabled) {
let cancelled = shared_user_enabled.cancel_user_sessions(&user);
for (user, cancelled) in
shared_user_enabled.apply_user_enabled_config(&cfg.access.user_enabled)
{
if cancelled > 0 {
info!(
user = %user,
+17 -1
View File
@@ -79,7 +79,11 @@ where
let route_snapshot = deps.route_runtime.snapshot();
let session_id = deps.rng.u64();
let user_session = deps.shared.register_user_session(&user, session_id);
let Some(user_session) = deps.shared.register_user_session(&user, session_id) else {
user_reservation.release_deferred();
warn!(user = %user, "Disabled user rejected during final admission");
return Err(ProxyError::UserDisabled { user });
};
let session_cancel = user_session.token();
let selected_me_pool = if deps.config.general.use_middle_proxy
&& matches!(route_snapshot.mode, RelayRouteMode::Middle)
@@ -246,6 +250,18 @@ impl UserConnectionReservation {
}
self.stats.decrement_user_curr_connects(&self.user);
}
/// Defers IP cleanup when admission fails after the asynchronous reservation step.
pub(crate) fn release_deferred(mut self) {
if !self.active {
return;
}
self.active = false;
self.stats.decrement_user_curr_connects(&self.user);
if self.tracks_ip {
self.ip_tracker.enqueue_cleanup(self.user.clone(), self.ip);
}
}
}
impl Drop for UserConnectionReservation {
+40 -61
View File
@@ -16,10 +16,7 @@ mod quota;
pub(super) use self::combined::CombinedStream;
pub(super) use self::counters::SharedCounters;
pub(super) use self::quota::is_quota_io_error;
use self::quota::{
QUOTA_RESERVE_MAX_ROUNDS, QUOTA_RESERVE_SPIN_RETRIES, quota_io_error,
refund_reserved_quota_bytes,
};
use self::quota::{QUOTA_RESERVE_MAX_ROUNDS, QUOTA_RESERVE_SPIN_RETRIES, quota_io_error};
pub(super) use self::quota::{quota_adaptive_interval_bytes, should_immediate_quota_check};
/// Transparent I/O wrapper that tracks per-user statistics and activity.
@@ -213,7 +210,7 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
}
let mut remaining_before = None;
let mut reserved_read_bytes = 0u64;
let mut quota_reservation = None;
let mut read_limit = buf.remaining();
if let Some(limit) = this.quota_limit {
let used_before = this.user_stats.quota_used();
@@ -231,11 +228,11 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
let desired = read_limit as u64;
let mut reserve_rounds = 0usize;
while reserved_read_bytes == 0 {
while quota_reservation.is_none() {
for _ in 0..QUOTA_RESERVE_SPIN_RETRIES {
match this.user_stats.quota_try_reserve(desired, limit) {
Ok(_) => {
reserved_read_bytes = desired;
match this.user_stats.quota_reserve(desired, limit) {
Ok(reservation) => {
quota_reservation = Some(reservation);
break;
}
Err(crate::stats::QuotaReserveError::LimitExceeded) => {
@@ -248,7 +245,7 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
}
}
if reserved_read_bytes == 0 {
if quota_reservation.is_none() {
reserve_rounds = reserve_rounds.saturating_add(1);
if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS {
this.stats.increment_quota_contention_timeout_total();
@@ -287,9 +284,9 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
match read_result {
Poll::Ready(Ok(n)) => {
if reserved_read_bytes > n as u64 {
let refund_bytes = reserved_read_bytes - n as u64;
refund_reserved_quota_bytes(this.user_stats.as_ref(), refund_bytes);
if let Some(reservation) = quota_reservation.take() {
let refund_bytes = reservation.reserved_bytes().saturating_sub(n as u64);
reservation.settle(n as u64);
this.stats.add_quota_refund_bytes_total(refund_bytes);
}
if n > 0 {
@@ -333,16 +330,16 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
Poll::Ready(Ok(()))
}
Poll::Pending => {
if reserved_read_bytes > 0 {
refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_read_bytes);
this.stats.add_quota_refund_bytes_total(reserved_read_bytes);
if let Some(reservation) = quota_reservation.take() {
this.stats
.add_quota_refund_bytes_total(reservation.reserved_bytes());
}
Poll::Pending
}
Poll::Ready(Err(err)) => {
if reserved_read_bytes > 0 {
refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_read_bytes);
this.stats.add_quota_refund_bytes_total(reserved_read_bytes);
if let Some(reservation) = quota_reservation.take() {
this.stats
.add_quota_refund_bytes_total(reservation.reserved_bytes());
}
Poll::Ready(Err(err))
}
@@ -361,14 +358,15 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
return Poll::Ready(Err(quota_io_error()));
}
let mut shaper_reserved_bytes = 0u64;
let mut shaper_reservation = None;
let mut write_buf = buf;
if let Some(lease) = this.traffic_lease.as_ref() {
if !buf.is_empty() {
loop {
let consume = lease.try_consume(RateDirection::Down, buf.len() as u64);
let reservation = lease.try_reserve(RateDirection::Down, buf.len() as u64);
let consume = reservation.result();
if consume.granted > 0 {
shaper_reserved_bytes = consume.granted;
shaper_reservation = Some(reservation);
if consume.granted < buf.len() as u64 {
write_buf = &buf[..consume.granted as usize];
}
@@ -398,17 +396,14 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
}
let mut remaining_before = None;
let mut reserved_bytes = 0u64;
let mut quota_reservation = None;
if let Some(limit) = this.quota_limit {
if !write_buf.is_empty() {
let mut reserve_rounds = 0usize;
while reserved_bytes == 0 {
while quota_reservation.is_none() {
let used_before = this.user_stats.quota_used();
let remaining = limit.saturating_sub(used_before);
if remaining == 0 {
if let Some(lease) = this.traffic_lease.as_ref() {
lease.refund(RateDirection::Down, shaper_reserved_bytes);
}
this.quota_exceeded.store(true, Ordering::Release);
return Poll::Ready(Err(quota_io_error()));
}
@@ -417,9 +412,9 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
let desired = remaining.min(write_buf.len() as u64);
let mut saw_contention = false;
for _ in 0..QUOTA_RESERVE_SPIN_RETRIES {
match this.user_stats.quota_try_reserve(desired, limit) {
Ok(_) => {
reserved_bytes = desired;
match this.user_stats.quota_reserve(desired, limit) {
Ok(reservation) => {
quota_reservation = Some(reservation);
write_buf = &write_buf[..desired as usize];
break;
}
@@ -433,14 +428,13 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
}
}
if reserved_bytes == 0 {
if quota_reservation.is_none() {
reserve_rounds = reserve_rounds.saturating_add(1);
if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS {
this.stats.increment_quota_contention_timeout_total();
if let Some(lease) = this.traffic_lease.as_ref() {
lease.refund(RateDirection::Down, shaper_reserved_bytes);
}
let _ = this.arm_quota_wait(cx);
Self::arm_wait(&mut this.quota_wait, false, false);
let _ =
Self::poll_wait(&mut this.quota_wait, cx, None, RateDirection::Up);
return Poll::Pending;
} else if saw_contention {
std::hint::spin_loop();
@@ -451,9 +445,6 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
let used_before = this.user_stats.quota_used();
let remaining = limit.saturating_sub(used_before);
if remaining == 0 {
if let Some(lease) = this.traffic_lease.as_ref() {
lease.refund(RateDirection::Down, shaper_reserved_bytes);
}
this.quota_exceeded.store(true, Ordering::Release);
return Poll::Ready(Err(quota_io_error()));
}
@@ -463,15 +454,13 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
match Pin::new(&mut this.inner).poll_write(cx, write_buf) {
Poll::Ready(Ok(n)) => {
if reserved_bytes > n as u64 {
let refund_bytes = reserved_bytes - n as u64;
refund_reserved_quota_bytes(this.user_stats.as_ref(), refund_bytes);
if let Some(reservation) = quota_reservation.take() {
let refund_bytes = reservation.reserved_bytes().saturating_sub(n as u64);
reservation.settle(n as u64);
this.stats.add_quota_refund_bytes_total(refund_bytes);
}
if shaper_reserved_bytes > n as u64
&& let Some(lease) = this.traffic_lease.as_ref()
{
lease.refund(RateDirection::Down, shaper_reserved_bytes - n as u64);
if let Some(reservation) = shaper_reservation.take() {
reservation.settle_written(n as u64);
}
if n > 0 {
if let Some(lease) = this.traffic_lease.as_ref() {
@@ -513,26 +502,16 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
Poll::Ready(Ok(n))
}
Poll::Ready(Err(err)) => {
if reserved_bytes > 0 {
refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_bytes);
this.stats.add_quota_refund_bytes_total(reserved_bytes);
}
if shaper_reserved_bytes > 0
&& let Some(lease) = this.traffic_lease.as_ref()
{
lease.refund(RateDirection::Down, shaper_reserved_bytes);
if let Some(reservation) = quota_reservation.take() {
this.stats
.add_quota_refund_bytes_total(reservation.reserved_bytes());
}
Poll::Ready(Err(err))
}
Poll::Pending => {
if reserved_bytes > 0 {
refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_bytes);
this.stats.add_quota_refund_bytes_total(reserved_bytes);
}
if shaper_reserved_bytes > 0
&& let Some(lease) = this.traffic_lease.as_ref()
{
lease.refund(RateDirection::Down, shaper_reserved_bytes);
if let Some(reservation) = quota_reservation.take() {
this.stats
.add_quota_refund_bytes_total(reservation.reserved_bytes());
}
Poll::Pending
}
-8
View File
@@ -1,4 +1,3 @@
use crate::stats::UserStats;
use std::io;
#[derive(Debug)]
@@ -46,10 +45,3 @@ pub(in crate::proxy::relay) fn should_immediate_quota_check(
) -> bool {
remaining_before <= QUOTA_NEAR_LIMIT_BYTES || charge_bytes >= QUOTA_LARGE_CHARGE_BYTES
}
pub(super) fn refund_reserved_quota_bytes(user_stats: &UserStats, reserved_bytes: u64) {
if reserved_bytes == 0 {
return;
}
user_stats.refund_quota(reserved_bytes);
}
+150 -43
View File
@@ -6,6 +6,7 @@ use std::sync::{Arc, Mutex};
use std::time::Instant;
use dashmap::DashMap;
use parking_lot::Mutex as ParkingMutex;
use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc};
use tokio_util::sync::CancellationToken;
@@ -75,13 +76,18 @@ pub(crate) struct MiddleRelaySharedState {
pub(crate) relay_idle_mark_seq: AtomicU64,
}
#[derive(Default)]
struct UserAdmissionState {
disabled_users: HashSet<String>,
sessions_by_user: HashMap<String, HashMap<u64, CancellationToken>>,
}
pub(crate) struct ProxySharedState {
pub(crate) handshake: HandshakeSharedState,
pub(crate) middle_relay: MiddleRelaySharedState,
pub(crate) traffic_limiter: Arc<TrafficLimiter>,
pub(crate) direct_buffer_budget: Arc<DirectBufferBudget>,
disabled_users: DashMap<String, ()>,
active_user_sessions: DashMap<(String, u64), CancellationToken>,
user_admission: ParkingMutex<UserAdmissionState>,
pub(crate) conntrack_pressure_active: AtomicBool,
pub(crate) conntrack_close_tx: Mutex<Option<mpsc::Sender<ConntrackCloseEvent>>>,
masking_fallback_permits: Arc<Semaphore>,
@@ -106,7 +112,18 @@ struct UserSessionGuard {
impl Drop for UserSessionGuard {
fn drop(&mut self) {
self.shared.active_user_sessions.remove(&self.key);
let mut admission = self.shared.user_admission.lock();
let remove_user = admission
.sessions_by_user
.get_mut(&self.key.0)
.map(|sessions| {
sessions.remove(&self.key.1);
sessions.is_empty()
})
.unwrap_or(false);
if remove_user {
admission.sessions_by_user.remove(&self.key.0);
}
}
}
@@ -150,8 +167,7 @@ impl ProxySharedState {
},
traffic_limiter: TrafficLimiter::new(),
direct_buffer_budget,
disabled_users: DashMap::new(),
active_user_sessions: DashMap::new(),
user_admission: ParkingMutex::new(UserAdmissionState::default()),
conntrack_pressure_active: AtomicBool::new(false),
conntrack_close_tx: Mutex::new(None),
masking_fallback_permits: Arc::new(Semaphore::new(MASKING_FALLBACK_MAX_CONCURRENT)),
@@ -167,68 +183,102 @@ impl ProxySharedState {
}
pub(crate) fn is_user_enabled(&self, user: &str) -> bool {
!self.disabled_users.contains_key(user)
!self.user_admission.lock().disabled_users.contains(user)
}
pub(crate) fn set_user_enabled(&self, user: &str, enabled: bool) -> bool {
if enabled {
self.disabled_users.remove(user);
false
} else {
self.disabled_users.insert(user.to_string(), ()).is_none()
pub(crate) fn set_user_enabled(&self, user: &str, enabled: bool) -> (bool, usize) {
let (newly_disabled, tokens) = {
let mut admission = self.user_admission.lock();
if enabled {
admission.disabled_users.remove(user);
(false, Vec::new())
} else {
let newly_disabled = admission.disabled_users.insert(user.to_string());
let tokens = admission
.sessions_by_user
.get(user)
.map(|sessions| sessions.values().cloned().collect())
.unwrap_or_default();
(newly_disabled, tokens)
}
};
for token in &tokens {
token.cancel();
}
(newly_disabled, tokens.len())
}
pub(crate) fn apply_user_enabled_config(
&self,
user_enabled: &HashMap<String, bool>,
) -> Vec<String> {
) -> Vec<(String, usize)> {
let desired_disabled = user_enabled
.iter()
.filter_map(|(user, enabled)| (!*enabled).then_some(user.clone()))
.collect::<HashSet<_>>();
let current_disabled = self
.disabled_users
.iter()
.map(|entry| entry.key().clone())
.collect::<HashSet<_>>();
for user in current_disabled.difference(&desired_disabled) {
self.disabled_users.remove(user);
}
let newly_disabled = desired_disabled
.difference(&current_disabled)
.cloned()
.collect::<Vec<_>>();
for user in desired_disabled {
self.disabled_users.insert(user, ());
}
newly_disabled
let cancellations = {
let mut admission = self.user_admission.lock();
let newly_disabled = desired_disabled
.difference(&admission.disabled_users)
.cloned()
.collect::<Vec<_>>();
admission.disabled_users = desired_disabled;
newly_disabled
.into_iter()
.map(|user| {
let tokens = admission
.sessions_by_user
.get(&user)
.map(|sessions| sessions.values().cloned().collect())
.unwrap_or_default();
(user, tokens)
})
.collect::<Vec<(String, Vec<CancellationToken>)>>()
};
cancellations
.into_iter()
.map(|(user, tokens)| {
for token in &tokens {
token.cancel();
}
(user, tokens.len())
})
.collect()
}
pub(crate) fn register_user_session(
self: &Arc<Self>,
user: &str,
session_id: u64,
) -> UserSessionRegistration {
) -> Option<UserSessionRegistration> {
let token = CancellationToken::new();
let key = (user.to_string(), session_id);
self.active_user_sessions.insert(key.clone(), token.clone());
UserSessionRegistration {
let mut admission = self.user_admission.lock();
if admission.disabled_users.contains(user) {
return None;
}
admission
.sessions_by_user
.entry(key.0.clone())
.or_default()
.insert(session_id, token.clone());
Some(UserSessionRegistration {
token,
_guard: UserSessionGuard {
shared: Arc::clone(self),
key,
},
}
})
}
pub(crate) fn cancel_user_sessions(&self, user: &str) -> usize {
let tokens = self
.active_user_sessions
.iter()
.filter_map(|entry| (entry.key().0 == user).then(|| entry.value().clone()))
.collect::<Vec<_>>();
let tokens: Vec<CancellationToken> = self
.user_admission
.lock()
.sessions_by_user
.get(user)
.map(|sessions| sessions.values().cloned().collect())
.unwrap_or_default();
for token in &tokens {
token.cancel();
}
@@ -311,7 +361,7 @@ mod tests {
let mut newly_disabled = shared.apply_user_enabled_config(&user_enabled);
newly_disabled.sort();
assert_eq!(newly_disabled, vec!["alice".to_string()]);
assert_eq!(newly_disabled, vec![("alice".to_string(), 0)]);
assert!(!shared.is_user_enabled("alice"));
assert!(shared.is_user_enabled("bob"));
@@ -325,9 +375,9 @@ mod tests {
#[test]
fn cancel_user_sessions_cancels_only_registered_matching_user() {
let shared = ProxySharedState::new();
let alice_1 = shared.register_user_session("alice", 1);
let alice_2 = shared.register_user_session("alice", 2);
let bob = shared.register_user_session("bob", 1);
let alice_1 = shared.register_user_session("alice", 1).unwrap();
let alice_2 = shared.register_user_session("alice", 2).unwrap();
let bob = shared.register_user_session("bob", 1).unwrap();
let alice_1_token = alice_1.token();
let alice_2_token = alice_2.token();
let bob_token = bob.token();
@@ -339,4 +389,61 @@ mod tests {
assert!(alice_2_token.is_cancelled());
assert!(!bob_token.is_cancelled());
}
#[test]
fn disabled_user_cannot_register_after_the_cancellation_snapshot() {
let shared = ProxySharedState::new();
assert_eq!(shared.set_user_enabled("alice", false), (true, 0));
assert_eq!(shared.cancel_user_sessions("alice"), 0);
let late = shared.register_user_session("alice", 1);
assert!(
late.is_none(),
"a session registered after disable returned must be rejected"
);
}
#[test]
fn disabling_user_cancels_existing_sessions_before_return() {
let shared = ProxySharedState::new();
let registration = shared.register_user_session("alice", 1).unwrap();
let token = registration.token();
assert_eq!(shared.set_user_enabled("alice", false), (true, 1));
assert!(token.is_cancelled());
assert!(shared.register_user_session("alice", 2).is_none());
assert_eq!(shared.set_user_enabled("alice", true), (false, 0));
assert!(shared.register_user_session("alice", 3).is_some());
}
#[test]
fn concurrent_disable_and_registration_never_leave_a_live_session() {
const ITERATIONS: usize = 10_000;
let shared = ProxySharedState::new();
let barrier = Arc::new(std::sync::Barrier::new(2));
let register_shared = Arc::clone(&shared);
let register_barrier = Arc::clone(&barrier);
let register = std::thread::spawn(move || {
let mut registrations = Vec::with_capacity(ITERATIONS);
for session_id in 0..ITERATIONS as u64 {
let user = format!("user-{session_id}");
register_barrier.wait();
registrations.push(register_shared.register_user_session(&user, session_id));
}
registrations
});
for session_id in 0..ITERATIONS as u64 {
let user = format!("user-{session_id}");
barrier.wait();
shared.set_user_enabled(&user, false);
}
for registration in register.join().unwrap().into_iter().flatten() {
assert!(registration.token().is_cancelled());
}
}
}
+26 -7
View File
@@ -6,6 +6,7 @@ use std::sync::atomic::{AtomicU64, Ordering};
use arc_swap::ArcSwap;
use dashmap::DashMap;
use ipnetwork::IpNetwork;
use parking_lot::Mutex as ParkingMutex;
use crate::config::RateLimitBps;
@@ -32,6 +33,9 @@ const REGISTRY_SHARDS: usize = 64;
const FAIR_EPOCH_MS: u64 = 20;
const MAX_BORROW_CHUNK_BYTES: u64 = 32 * 1024;
const CLEANUP_INTERVAL_SECS: u64 = 60;
const PACKED_USAGE_BITS: u32 = 28;
const PACKED_USAGE_MASK: u64 = (1u64 << PACKED_USAGE_BITS) - 1;
const PACKED_EPOCH_MAX: u64 = (1u64 << (u64::BITS - PACKED_USAGE_BITS)) - 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RateDirection {
@@ -76,12 +80,12 @@ struct ScopeMetrics {
struct AtomicRatePair {
up_bps: AtomicU64,
down_bps: AtomicU64,
revision: ParkingMutex<u64>,
}
#[derive(Default)]
struct DirectionBucket {
epoch: AtomicU64,
used: AtomicU64,
state: AtomicU64,
}
struct UserBucket {
@@ -93,15 +97,13 @@ struct UserBucket {
#[derive(Default)]
struct CidrDirectionBucket {
epoch: AtomicU64,
used: AtomicU64,
active_users: AtomicU64,
used: DirectionBucket,
active_users: DirectionBucket,
}
#[derive(Default)]
struct CidrUserDirectionState {
epoch: AtomicU64,
used: AtomicU64,
used: DirectionBucket,
}
struct CidrUserShare {
@@ -139,6 +141,7 @@ enum CidrPolicyMatch<'a> {
#[derive(Default)]
struct PolicySnapshot {
revision: u64,
user_limits: HashMap<String, RateLimitBps>,
cidr_rules_v4: Vec<CidrRule>,
cidr_rules_v6: Vec<CidrRule>,
@@ -162,9 +165,25 @@ pub struct TrafficLease {
pub struct TrafficLimiter {
policy: ArcSwap<PolicySnapshot>,
policy_update: ParkingMutex<()>,
user_buckets: ShardedRegistry<UserBucket>,
cidr_buckets: ShardedRegistry<CidrBucket>,
user_scope: ScopeMetrics,
cidr_scope: ScopeMetrics,
last_cleanup_epoch_secs: AtomicU64,
}
struct DirectionDebit<'a> {
bucket: &'a DirectionBucket,
epoch: u64,
refundable: u64,
}
/// Refunds uncommitted shaping budget when an I/O attempt is cancelled.
#[must_use = "traffic reservations must be settled after the I/O attempt"]
pub(crate) struct TrafficReservation<'a> {
result: TrafficConsumeResult,
user: Option<DirectionDebit<'a>>,
cidr: Option<DirectionDebit<'a>>,
cidr_user: Option<DirectionDebit<'a>>,
}
+218 -164
View File
@@ -26,203 +26,277 @@ impl ScopeMetrics {
}
impl AtomicRatePair {
pub(super) fn set(&self, limits: RateLimitBps) {
self.up_bps.store(limits.up_bps, Ordering::Relaxed);
self.down_bps.store(limits.down_bps, Ordering::Relaxed);
pub(super) fn new(revision: u64, limits: RateLimitBps) -> Self {
let rates = Self::default();
rates.set(revision, limits);
rates
}
pub(super) fn set(&self, revision: u64, limits: RateLimitBps) {
let mut current_revision = self.revision.lock();
if revision < *current_revision {
return;
}
self.up_bps.store(limits.up_bps, Ordering::Release);
self.down_bps.store(limits.down_bps, Ordering::Release);
*current_revision = revision;
}
pub(super) fn get(&self, direction: RateDirection) -> u64 {
match direction {
RateDirection::Up => self.up_bps.load(Ordering::Relaxed),
RateDirection::Down => self.down_bps.load(Ordering::Relaxed),
RateDirection::Up => self.up_bps.load(Ordering::Acquire),
RateDirection::Down => self.down_bps.load(Ordering::Acquire),
}
}
}
impl DirectionBucket {
pub(super) fn sync_epoch(&self, epoch: u64) {
let current = self.epoch.load(Ordering::Relaxed);
if current == epoch {
return;
}
if current < epoch
&& self
.epoch
.compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed)
.is_ok()
{
self.used.store(0, Ordering::Relaxed);
}
fn unpack(state: u64) -> (u64, u64) {
(state >> PACKED_USAGE_BITS, state & PACKED_USAGE_MASK)
}
pub(super) fn try_consume(&self, cap_bps: u64, requested: u64) -> u64 {
if requested == 0 {
return 0;
fn pack(epoch: u64, used: u64) -> Option<u64> {
if epoch > PACKED_EPOCH_MAX || used > PACKED_USAGE_MASK {
return None;
}
if cap_bps == 0 {
return requested;
Some((epoch << PACKED_USAGE_BITS) | used)
}
pub(super) fn used_at(&self, epoch: u64) -> Option<u64> {
if epoch > PACKED_EPOCH_MAX {
return None;
}
let (current_epoch, used) = Self::unpack(self.state.load(Ordering::Relaxed));
(current_epoch == epoch).then_some(used)
}
let epoch = current_epoch();
self.sync_epoch(epoch);
let cap_epoch = bytes_per_epoch(cap_bps);
pub(super) fn try_reserve_at(
&self,
epoch: u64,
cap: u64,
requested: u64,
) -> Option<DirectionDebit<'_>> {
if requested == 0 || cap == 0 || epoch > PACKED_EPOCH_MAX {
return None;
}
let cap = cap.min(PACKED_USAGE_MASK);
let mut observed = self.state.load(Ordering::Relaxed);
loop {
let used = self.used.load(Ordering::Relaxed);
if used >= cap_epoch {
return 0;
let (observed_epoch, observed_used) = Self::unpack(observed);
if observed_epoch > epoch {
return None;
}
let remaining = cap_epoch.saturating_sub(used);
let used = if observed_epoch == epoch {
observed_used
} else {
0
};
if used >= cap {
return None;
}
let remaining = cap - used;
let grant = requested.min(remaining);
if grant == 0 {
return 0;
return None;
}
let next = used.saturating_add(grant);
if self
.used
.compare_exchange_weak(used, next, Ordering::Relaxed, Ordering::Relaxed)
.is_ok()
{
return grant;
let next = Self::pack(epoch, used + grant)?;
match self.state.compare_exchange_weak(
observed,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => {
return Some(DirectionDebit {
bucket: self,
epoch,
refundable: grant,
});
}
Err(actual) => observed = actual,
}
}
}
pub(super) fn refund(&self, bytes: u64) {
if bytes == 0 {
fn refund_at(&self, epoch: u64, bytes: u64) {
if bytes == 0 || epoch > PACKED_EPOCH_MAX {
return;
}
decrement_atomic_saturating(&self.used, bytes);
let mut observed = self.state.load(Ordering::Relaxed);
loop {
let (observed_epoch, used) = Self::unpack(observed);
if observed_epoch != epoch || used == 0 {
return;
}
let next = Self::pack(epoch, used.saturating_sub(bytes)).unwrap_or(observed);
match self.state.compare_exchange_weak(
observed,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(actual) => observed = actual,
}
}
}
}
impl DirectionDebit<'_> {
fn granted(&self) -> u64 {
self.refundable
}
pub(super) fn shrink_to(&mut self, retained: u64) {
let retained = retained.min(self.refundable);
self.bucket
.refund_at(self.epoch, self.refundable - retained);
self.refundable = retained;
}
pub(super) fn settle(&mut self, committed: u64) {
self.shrink_to(committed);
self.refundable = 0;
}
pub(super) fn commit_all(&mut self) -> u64 {
let committed = self.refundable;
self.refundable = 0;
committed
}
}
impl Drop for DirectionDebit<'_> {
fn drop(&mut self) {
self.bucket.refund_at(self.epoch, self.refundable);
}
}
impl UserBucket {
pub(super) fn new(limits: RateLimitBps) -> Self {
let rates = AtomicRatePair::default();
rates.set(limits);
pub(super) fn new(revision: u64, limits: RateLimitBps) -> Self {
Self {
rates,
rates: AtomicRatePair::new(revision, limits),
up: DirectionBucket::default(),
down: DirectionBucket::default(),
active_leases: AtomicU64::new(0),
}
}
pub(super) fn set_rates(&self, limits: RateLimitBps) {
self.rates.set(limits);
pub(super) fn set_rates(&self, revision: u64, limits: RateLimitBps) {
self.rates.set(revision, limits);
}
pub(super) fn try_consume(&self, direction: RateDirection, requested: u64) -> u64 {
pub(super) fn try_reserve(
&self,
direction: RateDirection,
requested: u64,
) -> (u64, Option<DirectionDebit<'_>>) {
let cap_bps = self.rates.get(direction);
match direction {
RateDirection::Up => self.up.try_consume(cap_bps, requested),
RateDirection::Down => self.down.try_consume(cap_bps, requested),
}
}
pub(super) fn refund(&self, direction: RateDirection, bytes: u64) {
match direction {
RateDirection::Up => self.up.refund(bytes),
RateDirection::Down => self.down.refund(bytes),
if cap_bps == 0 {
return (requested, None);
}
let cap = bytes_per_epoch(cap_bps);
let debit = match direction {
RateDirection::Up => self.up.try_reserve_at(current_epoch(), cap, requested),
RateDirection::Down => self.down.try_reserve_at(current_epoch(), cap, requested),
};
let granted = debit.as_ref().map(DirectionDebit::granted).unwrap_or(0);
(granted, debit)
}
}
impl CidrDirectionBucket {
pub(super) fn sync_epoch(&self, epoch: u64) {
let current = self.epoch.load(Ordering::Relaxed);
if current == epoch {
return;
}
if current < epoch
&& self
.epoch
.compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed)
.is_ok()
{
self.used.store(0, Ordering::Relaxed);
self.active_users.store(0, Ordering::Relaxed);
}
}
pub(super) fn try_consume(
&self,
user_state: &CidrUserDirectionState,
pub(super) fn try_reserve<'a>(
&'a self,
user_state: &'a CidrUserDirectionState,
cap_epoch: u64,
requested: u64,
) -> u64 {
) -> (u64, Option<DirectionDebit<'a>>, Option<DirectionDebit<'a>>) {
if requested == 0 || cap_epoch == 0 {
return 0;
return (0, None, None);
}
let epoch = current_epoch();
self.sync_epoch(epoch);
user_state.sync_epoch_and_mark_active(epoch, &self.active_users);
let active_users = self.active_users.load(Ordering::Relaxed).max(1);
if !user_state.ensure_active(epoch, &self.active_users) {
return (0, None, None);
}
let Some(active_users) = self.active_users.used_at(epoch) else {
return (0, None, None);
};
let active_users = active_users.max(1);
let fair_share = cap_epoch.saturating_div(active_users).max(1);
loop {
let total_used = self.used.load(Ordering::Relaxed);
if total_used >= cap_epoch {
return 0;
}
let total_remaining = cap_epoch.saturating_sub(total_used);
let user_used = user_state.used.load(Ordering::Relaxed);
let guaranteed_remaining = fair_share.saturating_sub(user_used);
let grant = if guaranteed_remaining > 0 {
requested.min(guaranteed_remaining).min(total_remaining)
} else {
requested.min(total_remaining).min(MAX_BORROW_CHUNK_BYTES)
let Some(user_used) = user_state.used.used_at(epoch) else {
return (0, None, None);
};
if grant == 0 {
return 0;
}
let next_total = total_used.saturating_add(grant);
if self
.used
.compare_exchange_weak(total_used, next_total, Ordering::Relaxed, Ordering::Relaxed)
.is_ok()
{
user_state.used.fetch_add(grant, Ordering::Relaxed);
return grant;
let guaranteed_remaining = fair_share.saturating_sub(user_used);
let (user_cap, desired) = if guaranteed_remaining > 0 {
(fair_share, requested.min(guaranteed_remaining))
} else {
(PACKED_USAGE_MASK, requested.min(MAX_BORROW_CHUNK_BYTES))
};
let Some(mut user_debit) = user_state.used.try_reserve_at(epoch, user_cap, desired)
else {
if guaranteed_remaining > 0 {
continue;
}
return (0, None, None);
};
let user_granted = user_debit.granted();
let Some(aggregate_debit) = self.used.try_reserve_at(epoch, cap_epoch, user_granted)
else {
return (0, None, None);
};
let granted = aggregate_debit.granted();
if granted < user_granted {
user_debit.shrink_to(granted);
}
return (granted, Some(aggregate_debit), Some(user_debit));
}
}
pub(super) fn refund(&self, bytes: u64) {
if bytes == 0 {
return;
}
decrement_atomic_saturating(&self.used, bytes);
}
}
impl CidrUserDirectionState {
pub(super) fn sync_epoch_and_mark_active(&self, epoch: u64, active_users: &AtomicU64) {
let current = self.epoch.load(Ordering::Relaxed);
if current == epoch {
return;
pub(super) fn ensure_active(&self, epoch: u64, active_users: &DirectionBucket) -> bool {
if epoch > PACKED_EPOCH_MAX {
return false;
}
if current < epoch
&& self
.epoch
.compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed)
.is_ok()
{
self.used.store(0, Ordering::Relaxed);
active_users.fetch_add(1, Ordering::Relaxed);
let mut observed = self.used.state.load(Ordering::Relaxed);
loop {
let (observed_epoch, _) = DirectionBucket::unpack(observed);
if observed_epoch == epoch {
return true;
}
if observed_epoch > epoch {
return false;
}
let Some(mut active_debit) = active_users.try_reserve_at(epoch, PACKED_USAGE_MASK, 1)
else {
return false;
};
let Some(next) = DirectionBucket::pack(epoch, 0) else {
return false;
};
match self.used.state.compare_exchange(
observed,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => {
active_debit.commit_all();
return true;
}
Err(actual) => {
drop(active_debit);
observed = actual;
}
}
}
}
pub(super) fn refund(&self, bytes: u64) {
if bytes == 0 {
return;
}
decrement_atomic_saturating(&self.used, bytes);
}
}
impl CidrUserShare {
@@ -236,11 +310,9 @@ impl CidrUserShare {
}
impl CidrBucket {
pub(super) fn new(limits: RateLimitBps) -> Self {
let rates = AtomicRatePair::default();
rates.set(limits);
pub(super) fn new(revision: u64, limits: RateLimitBps) -> Self {
Self {
rates,
rates: AtomicRatePair::new(revision, limits),
up: CidrDirectionBucket::default(),
down: CidrDirectionBucket::default(),
users: ShardedRegistry::new(REGISTRY_SHARDS),
@@ -248,8 +320,8 @@ impl CidrBucket {
}
}
pub(super) fn set_rates(&self, limits: RateLimitBps) {
self.rates.set(limits);
pub(super) fn set_rates(&self, revision: u64, limits: RateLimitBps) {
self.rates.set(revision, limits);
}
pub(super) fn acquire_user_share(&self, user: &str) -> Arc<CidrUserShare> {
@@ -268,38 +340,20 @@ impl CidrBucket {
});
}
pub(super) fn try_consume_for_user(
&self,
pub(super) fn try_reserve_for_user<'a>(
&'a self,
direction: RateDirection,
share: &CidrUserShare,
share: &'a CidrUserShare,
requested: u64,
) -> u64 {
) -> (u64, Option<DirectionDebit<'a>>, Option<DirectionDebit<'a>>) {
let cap_bps = self.rates.get(direction);
if cap_bps == 0 {
return requested;
return (requested, None, None);
}
let cap_epoch = bytes_per_epoch(cap_bps);
match direction {
RateDirection::Up => self.up.try_consume(&share.up, cap_epoch, requested),
RateDirection::Down => self.down.try_consume(&share.down, cap_epoch, requested),
}
}
pub(super) fn refund_for_user(
&self,
direction: RateDirection,
share: &CidrUserShare,
bytes: u64,
) {
match direction {
RateDirection::Up => {
self.up.refund(bytes);
share.up.refund(bytes);
}
RateDirection::Down => {
self.down.refund(bytes);
share.down.refund(bytes);
}
RateDirection::Up => self.up.try_reserve(&share.up, cap_epoch, requested),
RateDirection::Down => self.down.try_reserve(&share.down, cap_epoch, requested),
}
}
+1 -1
View File
@@ -56,7 +56,7 @@ pub(super) fn auto_cidr_bucket_key(ip: IpAddr, prefix_len: u8) -> Option<String>
pub(super) fn current_epoch() -> u64 {
let start = limiter_epoch_start();
let elapsed_ms = start.elapsed().as_millis() as u64;
elapsed_ms / FAIR_EPOCH_MS
elapsed_ms / FAIR_EPOCH_MS + 1
}
pub(super) fn limiter_epoch_start() -> &'static Instant {
+79 -35
View File
@@ -1,70 +1,93 @@
use super::*;
impl TrafficLease {
pub fn try_consume(&self, direction: RateDirection, requested: u64) -> TrafficConsumeResult {
/// Reserves shaping budget until the associated I/O result is settled.
pub(crate) fn try_reserve(
&self,
direction: RateDirection,
requested: u64,
) -> TrafficReservation<'_> {
if requested == 0 {
return TrafficConsumeResult {
granted: 0,
blocked_user: false,
blocked_cidr: false,
return TrafficReservation {
result: TrafficConsumeResult {
granted: 0,
blocked_user: false,
blocked_cidr: false,
},
user: None,
cidr: None,
cidr_user: None,
};
}
let mut granted = requested;
let mut user_debit = None;
if let Some(user_bucket) = self.user_bucket.as_ref() {
let user_granted = user_bucket.try_consume(direction, granted);
let (user_granted, debit) = user_bucket.try_reserve(direction, granted);
user_debit = debit;
if user_granted == 0 {
self.limiter.observe_throttle(direction, true, false);
return TrafficConsumeResult {
granted: 0,
blocked_user: true,
blocked_cidr: false,
return TrafficReservation {
result: TrafficConsumeResult {
granted: 0,
blocked_user: true,
blocked_cidr: false,
},
user: user_debit,
cidr: None,
cidr_user: None,
};
}
granted = user_granted;
}
let mut cidr_debit = None;
let mut cidr_user_debit = None;
if let (Some(cidr_bucket), Some(cidr_user_share)) =
(self.cidr_bucket.as_ref(), self.cidr_user_share.as_ref())
{
let cidr_granted =
cidr_bucket.try_consume_for_user(direction, cidr_user_share, granted);
let (cidr_granted, aggregate_debit, share_debit) =
cidr_bucket.try_reserve_for_user(direction, cidr_user_share, granted);
cidr_debit = aggregate_debit;
cidr_user_debit = share_debit;
if cidr_granted < granted
&& let Some(user_bucket) = self.user_bucket.as_ref()
&& let Some(debit) = user_debit.as_mut()
{
user_bucket.refund(direction, granted.saturating_sub(cidr_granted));
debit.shrink_to(cidr_granted);
}
if cidr_granted == 0 {
self.limiter.observe_throttle(direction, false, true);
return TrafficConsumeResult {
granted: 0,
blocked_user: false,
blocked_cidr: true,
return TrafficReservation {
result: TrafficConsumeResult {
granted: 0,
blocked_user: false,
blocked_cidr: true,
},
user: user_debit,
cidr: cidr_debit,
cidr_user: cidr_user_debit,
};
}
granted = cidr_granted;
}
TrafficConsumeResult {
granted,
blocked_user: false,
blocked_cidr: false,
TrafficReservation {
result: TrafficConsumeResult {
granted,
blocked_user: false,
blocked_cidr: false,
},
user: user_debit,
cidr: cidr_debit,
cidr_user: cidr_user_debit,
}
}
pub fn refund(&self, direction: RateDirection, bytes: u64) {
if bytes == 0 {
return;
}
if let Some(user_bucket) = self.user_bucket.as_ref() {
user_bucket.refund(direction, bytes);
}
if let (Some(cidr_bucket), Some(cidr_user_share)) =
(self.cidr_bucket.as_ref(), self.cidr_user_share.as_ref())
{
cidr_bucket.refund_for_user(direction, cidr_user_share, bytes);
}
pub fn try_consume(&self, direction: RateDirection, requested: u64) -> TrafficConsumeResult {
let reservation = self.try_reserve(direction, requested);
let result = reservation.result();
reservation.settle_written(result.granted);
result
}
pub fn observe_wait_ms(
@@ -82,6 +105,27 @@ impl TrafficLease {
}
}
impl TrafficReservation<'_> {
/// Returns the shaping decision associated with this reservation.
pub(crate) fn result(&self) -> TrafficConsumeResult {
self.result
}
/// Commits written bytes and refunds the uncommitted remainder.
pub(crate) fn settle_written(mut self, committed: u64) {
let committed = committed.min(self.result.granted);
if let Some(debit) = self.user.as_mut() {
debit.settle(committed);
}
if let Some(debit) = self.cidr.as_mut() {
debit.settle(committed);
}
if let Some(debit) = self.cidr_user.as_mut() {
debit.settle(committed);
}
}
}
impl Drop for TrafficLease {
fn drop(&mut self) {
if let Some(bucket) = self.user_bucket.as_ref() {
+12 -4
View File
@@ -5,6 +5,7 @@ impl TrafficLimiter {
pub fn new() -> Arc<Self> {
Arc::new(Self {
policy: ArcSwap::from_pointee(PolicySnapshot::default()),
policy_update: ParkingMutex::new(()),
user_buckets: ShardedRegistry::new(REGISTRY_SHARDS),
cidr_buckets: ShardedRegistry::new(REGISTRY_SHARDS),
user_scope: ScopeMetrics::default(),
@@ -18,6 +19,11 @@ impl TrafficLimiter {
user_limits: HashMap<String, RateLimitBps>,
cidr_limits: HashMap<CidrRateLimitKey, RateLimitBps>,
) {
let policy_update = self.policy_update.lock();
// Revision wrap could otherwise let an old lease restore stale rates.
let Some(revision) = self.policy.load().revision.checked_add(1) else {
return;
};
let filtered_users = user_limits
.into_iter()
.filter(|(_, limit)| limit.up_bps > 0 || limit.down_bps > 0)
@@ -78,6 +84,7 @@ impl TrafficLimiter {
.store(cidr_policy_entries as u64, Ordering::Relaxed);
self.policy.store(Arc::new(PolicySnapshot {
revision,
user_limits: filtered_users,
cidr_rules_v4,
cidr_rules_v6,
@@ -86,6 +93,7 @@ impl TrafficLimiter {
cidr_rule_keys,
}));
drop(policy_update);
self.maybe_cleanup();
}
@@ -99,12 +107,12 @@ impl TrafficLimiter {
if let Some(limit) = policy.user_limits.get(user).copied() {
let bucket = self.user_buckets.get_or_insert_with(
user,
|| UserBucket::new(limit),
|| UserBucket::new(policy.revision, limit),
|bucket| {
bucket.active_leases.fetch_add(1, Ordering::Relaxed);
},
);
bucket.set_rates(limit);
bucket.set_rates(policy.revision, limit);
self.user_scope
.active_leases
.fetch_add(1, Ordering::Relaxed);
@@ -121,12 +129,12 @@ impl TrafficLimiter {
};
let bucket = self.cidr_buckets.get_or_insert_with(
key,
|| CidrBucket::new(limits),
|| CidrBucket::new(policy.revision, limits),
|bucket| {
bucket.active_leases.fetch_add(1, Ordering::Relaxed);
},
);
bucket.set_rates(limits);
bucket.set_rates(policy.revision, limits);
self.cidr_scope
.active_leases
.fetch_add(1, Ordering::Relaxed);
+182
View File
@@ -74,3 +74,185 @@ fn auto_cidr_bucket_key_canonicalizes_network_address() {
"auto:6:2001:db8::/64"
);
}
#[test]
fn refund_from_an_old_epoch_does_not_reduce_the_current_epoch() {
let bucket = DirectionBucket::default();
let old_debit = bucket.try_reserve_at(7, 100, 80).unwrap();
let current_debit = bucket.try_reserve_at(8, 100, 60).unwrap();
drop(old_debit);
assert_eq!(bucket.used_at(8), Some(60));
drop(current_debit);
}
#[test]
fn concurrent_rollover_cannot_publish_multiple_epoch_budgets() {
const CONTENDERS: usize = 32;
let bucket = Arc::new(DirectionBucket::default());
let barrier = Arc::new(std::sync::Barrier::new(CONTENDERS));
let mut threads = Vec::with_capacity(CONTENDERS);
for _ in 0..CONTENDERS {
let bucket = Arc::clone(&bucket);
let barrier = Arc::clone(&barrier);
threads.push(std::thread::spawn(move || {
barrier.wait();
bucket
.try_reserve_at(9, 100, 100)
.map(|mut debit| debit.commit_all())
.unwrap_or(0)
}));
}
let granted = threads
.into_iter()
.map(|thread| thread.join().unwrap())
.sum::<u64>();
assert_eq!(granted, 100);
assert_eq!(bucket.used_at(9), Some(100));
}
#[test]
fn scheduler_pressure_never_exceeds_a_packed_epoch_budget() {
const CONTENDERS: usize = 4;
const EPOCHS: usize = 10_000;
let bucket = Arc::new(DirectionBucket::default());
let barrier = Arc::new(std::sync::Barrier::new(CONTENDERS));
let mut threads = Vec::with_capacity(CONTENDERS);
for _ in 0..CONTENDERS {
let bucket = Arc::clone(&bucket);
let barrier = Arc::clone(&barrier);
threads.push(std::thread::spawn(move || {
let mut grants = Vec::with_capacity(EPOCHS);
for epoch in 1..=EPOCHS as u64 {
barrier.wait();
let granted = bucket
.try_reserve_at(epoch, 100, 100)
.map(|mut debit| debit.commit_all())
.unwrap_or(0);
grants.push(granted);
barrier.wait();
}
grants
}));
}
let grants = threads
.into_iter()
.map(|thread| thread.join().unwrap())
.collect::<Vec<_>>();
for epoch_index in 0..EPOCHS {
let granted = grants
.iter()
.map(|thread_grants| thread_grants[epoch_index])
.sum::<u64>();
assert_eq!(granted, 100);
}
}
#[test]
fn stale_policy_revision_cannot_restore_an_old_rate() {
let bucket = UserBucket::new(2, rate(2_000, 3_000));
bucket.set_rates(3, rate(4_000, 5_000));
bucket.set_rates(2, rate(6_000, 7_000));
assert_eq!(bucket.rates.get(RateDirection::Up), 4_000);
assert_eq!(bucket.rates.get(RateDirection::Down), 5_000);
}
#[test]
fn dropped_debit_refunds_only_its_packed_epoch() {
let bucket = DirectionBucket::default();
let debit = bucket.try_reserve_at(11, 100, 80).unwrap();
drop(debit);
assert_eq!(bucket.used_at(11), Some(0));
assert!(
bucket
.try_reserve_at(PACKED_EPOCH_MAX + 1, 100, 1)
.is_none()
);
}
#[test]
fn concurrent_first_use_counts_one_active_cidr_user() {
const CONTENDERS: usize = 32;
let bucket = Arc::new(CidrDirectionBucket::default());
let user = Arc::new(CidrUserDirectionState::default());
let barrier = Arc::new(std::sync::Barrier::new(CONTENDERS));
let mut threads = Vec::with_capacity(CONTENDERS);
for _ in 0..CONTENDERS {
let bucket = Arc::clone(&bucket);
let user = Arc::clone(&user);
let barrier = Arc::clone(&barrier);
threads.push(std::thread::spawn(move || {
barrier.wait();
assert!(user.ensure_active(13, &bucket.active_users));
}));
}
for thread in threads {
thread.join().unwrap();
}
assert_eq!(bucket.active_users.used_at(13), Some(1));
}
#[test]
fn configured_rate_maximum_fits_the_packed_epoch_budget() {
assert_eq!(bytes_per_epoch(100_000_000_000), 250_000_000);
assert!(bytes_per_epoch(100_000_000_000) <= PACKED_USAGE_MASK);
}
#[test]
fn dropped_traffic_reservation_refunds_user_and_cidr_debits() {
let limiter = TrafficLimiter::new();
let mut user_limits = HashMap::new();
user_limits.insert("alice".to_string(), rate(400_000, 400_000));
let mut cidr_limits = HashMap::new();
cidr_limits.insert(
CidrRateLimitKey::Network("203.0.113.0/24".parse().unwrap()),
rate(400_000, 400_000),
);
limiter.apply_policy(user_limits, cidr_limits);
let lease = limiter
.acquire_lease("alice", "203.0.113.7".parse().unwrap())
.unwrap();
let reservation = lease.try_reserve(RateDirection::Down, 800);
assert_eq!(reservation.result().granted, 800);
let epoch = reservation.user.as_ref().unwrap().epoch;
drop(reservation);
let user_bucket = lease.user_bucket.as_ref().unwrap();
let cidr_bucket = lease.cidr_bucket.as_ref().unwrap();
let cidr_user = lease.cidr_user_share.as_ref().unwrap();
assert_eq!(user_bucket.down.used_at(epoch), Some(0));
assert_eq!(cidr_bucket.down.used.used_at(epoch), Some(0));
assert_eq!(cidr_user.down.used.used_at(epoch), Some(0));
}
#[test]
fn partial_traffic_settlement_charges_only_committed_bytes() {
let limiter = TrafficLimiter::new();
let mut user_limits = HashMap::new();
user_limits.insert("alice".to_string(), rate(400_000, 400_000));
limiter.apply_policy(user_limits, HashMap::new());
let lease = limiter
.acquire_lease("alice", "203.0.113.7".parse().unwrap())
.unwrap();
let reservation = lease.try_reserve(RateDirection::Down, 800);
let epoch = reservation.user.as_ref().unwrap().epoch;
reservation.settle_written(300);
assert_eq!(
lease.user_bucket.as_ref().unwrap().down.used_at(epoch),
Some(300)
);
}
+13 -6
View File
@@ -21,7 +21,7 @@ use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, AtomicUsize, Ordering};
use std::time::Instant;
pub(crate) use self::quota_store::QuotaStore;
pub(crate) use self::quota_store::{QuotaReservation, QuotaStore};
#[allow(unused_imports)]
pub use self::replay::{ReplayChecker, ReplayStats};
use self::telemetry::TelemetryPolicy;
@@ -392,11 +392,6 @@ impl UserStats {
self.quota.used()
}
#[inline]
pub(crate) fn refund_quota(&self, bytes: u64) {
self.quota.refund(bytes);
}
/// Attempts one CAS reservation step against the quota counter.
///
/// Callers control retry/yield policy. This primitive intentionally does
@@ -404,6 +399,18 @@ impl UserStats {
/// with their own contention strategy.
#[inline]
pub fn quota_try_reserve(&self, bytes: u64, limit: u64) -> Result<u64, QuotaReserveError> {
self.quota
.try_reserve(bytes, limit)
.map(QuotaReservation::commit)
}
/// Reserves quota until a direct I/O attempt is settled.
#[inline]
pub(crate) fn quota_reserve(
&self,
bytes: u64,
limit: u64,
) -> Result<QuotaReservation, QuotaReserveError> {
self.quota.try_reserve(bytes, limit)
}
}
+143 -36
View File
@@ -2,6 +2,7 @@ use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use arc_swap::ArcSwap;
use dashmap::DashMap;
use super::{QuotaReserveError, UserQuotaSnapshot};
@@ -12,11 +13,22 @@ pub struct QuotaStore {
users: DashMap<String, Arc<UserQuotaCounters>>,
}
/// Atomic quota state for one configured user.
#[derive(Default)]
/// Atomically replaceable quota state for one configured user.
pub(crate) struct UserQuotaCounters {
generation: ArcSwap<QuotaGeneration>,
}
struct QuotaGeneration {
used_bytes: AtomicU64,
last_reset_epoch_secs: AtomicU64,
last_reset_epoch_secs: u64,
}
/// Owns a quota debit until the corresponding I/O outcome is known.
#[must_use = "quota reservations must be committed or settled"]
pub(crate) struct QuotaReservation {
generation: Arc<QuotaGeneration>,
reserved_bytes: u64,
total_after_reserve: u64,
}
impl QuotaStore {
@@ -38,18 +50,12 @@ impl QuotaStore {
pub(crate) fn load(&self, user: &str, used_bytes: u64, last_reset_epoch_secs: u64) {
let state = self.user(user);
state.used_bytes.store(used_bytes, Ordering::Relaxed);
state
.last_reset_epoch_secs
.store(last_reset_epoch_secs, Ordering::Relaxed);
state.replace(used_bytes, last_reset_epoch_secs);
}
pub(crate) fn reset(&self, user: &str, now_epoch_secs: u64) -> UserQuotaSnapshot {
let state = self.user(user);
state.used_bytes.store(0, Ordering::Relaxed);
state
.last_reset_epoch_secs
.store(now_epoch_secs, Ordering::Relaxed);
state.replace(0, now_epoch_secs);
UserQuotaSnapshot {
used_bytes: 0,
last_reset_epoch_secs: now_epoch_secs,
@@ -64,8 +70,9 @@ impl QuotaStore {
let mut out = HashMap::new();
for entry in self.users.iter() {
let state = entry.value();
let used_bytes = state.used();
let last_reset_epoch_secs = state.last_reset_epoch_secs.load(Ordering::Relaxed);
let generation = state.generation.load_full();
let used_bytes = generation.used_bytes.load(Ordering::Relaxed);
let last_reset_epoch_secs = generation.last_reset_epoch_secs;
if used_bytes == 0 && last_reset_epoch_secs == 0 {
continue;
}
@@ -81,60 +88,121 @@ impl QuotaStore {
}
}
impl Default for UserQuotaCounters {
fn default() -> Self {
Self {
generation: ArcSwap::from_pointee(QuotaGeneration {
used_bytes: AtomicU64::new(0),
last_reset_epoch_secs: 0,
}),
}
}
}
impl UserQuotaCounters {
fn replace(&self, used_bytes: u64, last_reset_epoch_secs: u64) {
self.generation.store(Arc::new(QuotaGeneration {
used_bytes: AtomicU64::new(used_bytes),
last_reset_epoch_secs,
}));
}
#[inline]
pub(crate) fn used(&self) -> u64 {
self.used_bytes.load(Ordering::Relaxed)
self.generation.load().used_bytes.load(Ordering::Relaxed)
}
#[inline]
pub(crate) fn charge(&self, bytes: u64) -> u64 {
self.used_bytes
self.generation
.load_full()
.used_bytes
.fetch_add(bytes, Ordering::Relaxed)
.saturating_add(bytes)
}
#[inline]
pub(crate) fn refund(&self, bytes: u64) {
let mut current = self.used_bytes.load(Ordering::Relaxed);
loop {
let next = current.saturating_sub(bytes);
match self.used_bytes.compare_exchange_weak(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(observed) => current = observed,
}
}
}
#[inline]
pub(crate) fn try_reserve(&self, bytes: u64, limit: u64) -> Result<u64, QuotaReserveError> {
let current = self.used_bytes.load(Ordering::Relaxed);
pub(crate) fn try_reserve(
&self,
bytes: u64,
limit: u64,
) -> Result<QuotaReservation, QuotaReserveError> {
let generation = self.generation.load_full();
let current = generation.used_bytes.load(Ordering::Relaxed);
if bytes > limit.saturating_sub(current) {
return Err(QuotaReserveError::LimitExceeded);
}
let next = current.saturating_add(bytes);
match self.used_bytes.compare_exchange_weak(
match generation.used_bytes.compare_exchange_weak(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => Ok(next),
Ok(_) => Ok(QuotaReservation {
generation,
reserved_bytes: bytes,
total_after_reserve: next,
}),
Err(_) => Err(QuotaReserveError::Contended),
}
}
}
impl QuotaReservation {
/// Returns the number of bytes held by this reservation.
pub(crate) fn reserved_bytes(&self) -> u64 {
self.reserved_bytes
}
/// Commits the complete reservation and returns the generation-local total.
pub(crate) fn commit(mut self) -> u64 {
self.reserved_bytes = 0;
self.total_after_reserve
}
/// Commits part of the reservation and refunds the remainder.
pub(crate) fn settle(mut self, committed_bytes: u64) {
let committed_bytes = committed_bytes.min(self.reserved_bytes);
refund_generation(
self.generation.as_ref(),
self.reserved_bytes - committed_bytes,
);
self.reserved_bytes = 0;
}
}
impl Drop for QuotaReservation {
fn drop(&mut self) {
refund_generation(self.generation.as_ref(), self.reserved_bytes);
}
}
fn refund_generation(generation: &QuotaGeneration, bytes: u64) {
if bytes == 0 {
return;
}
let mut current = generation.used_bytes.load(Ordering::Relaxed);
loop {
let next = current.saturating_sub(bytes);
match generation.used_bytes.compare_exchange_weak(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(observed) => current = observed,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::stats::Stats;
use std::sync::Barrier;
#[test]
fn quota_counters_are_shared_across_stats_generations() {
@@ -148,4 +216,43 @@ mod tests {
second.reset_user_quota("alice");
assert_eq!(first.get_user_quota_used("alice"), 0);
}
#[test]
fn quota_snapshot_never_combines_different_generations() {
const ITERATIONS: u64 = 10_000;
let store = Arc::new(QuotaStore::default());
store.load("alice", 1, 1);
let barrier = Arc::new(Barrier::new(2));
let writer_store = Arc::clone(&store);
let writer_barrier = Arc::clone(&barrier);
let writer = std::thread::spawn(move || {
writer_barrier.wait();
for generation in 2..=ITERATIONS {
writer_store.load("alice", generation, generation);
}
});
barrier.wait();
for _ in 0..ITERATIONS {
let snapshot = store.snapshot().remove("alice").unwrap();
assert_eq!(snapshot.used_bytes, snapshot.last_reset_epoch_secs);
}
writer.join().unwrap();
}
#[test]
fn repeated_old_generation_refunds_leave_new_usage_intact() {
const ITERATIONS: u64 = 10_000;
let store = QuotaStore::default();
let state = store.user("alice");
for generation in 1..=ITERATIONS {
let reservation = state.try_reserve(80, 100).unwrap();
store.load("alice", 40, generation);
drop(reservation);
assert_eq!(store.used("alice"), 40);
store.reset("alice", generation);
}
}
}
+14
View File
@@ -289,6 +289,20 @@ fn test_quota_used_is_authoritative_and_independent_from_octets_telemetry() {
assert_eq!(stats.get_user_quota_used(user), 7);
}
#[test]
fn old_quota_reservation_refund_does_not_reduce_post_reset_usage() {
let stats = Stats::new();
let user = "quota-reset-generation-user";
let user_stats = stats.get_or_create_user_stats_handle(user);
let reservation = user_stats.quota_reserve(80, 100).unwrap();
stats.reset_user_quota(user);
stats.quota_charge_post_write(user_stats.as_ref(), 40);
drop(reservation);
assert_eq!(stats.get_user_quota_used(user), 40);
}
#[test]
fn test_cached_handle_survives_map_cleanup_until_last_drop() {
let stats = Stats::new();
+3
View File
@@ -33,6 +33,9 @@ mod pool_runtime_api;
mod pool_status;
mod pool_writer;
#[cfg(test)]
#[path = "tests/pool_writer_publication_tests.rs"]
mod pool_writer_publication_tests;
#[cfg(test)]
#[path = "tests/pool_writer_security_tests.rs"]
mod pool_writer_security_tests;
mod reader;
-8
View File
@@ -88,14 +88,6 @@ impl WritersState {
}
}
pub(super) async fn update<F, R>(&self, f: F) -> R
where
F: FnOnce(&mut Vec<MeWriter>) -> R,
{
let mut guard = self.write().await;
f(&mut guard)
}
fn debug_assert_store_guarded(&self) {
debug_assert!(
self.writers_write_guard.try_lock().is_err(),
@@ -1,4 +1,5 @@
use std::collections::HashMap;
use std::future::Future;
use std::io::ErrorKind;
use std::net::SocketAddr;
use std::sync::Arc;
@@ -19,6 +20,7 @@ use crate::protocol::constants::{RPC_CLOSE_EXT_U32, RPC_PING_U32};
use super::codec::{RpcWriter, WriterCommand, build_control_payload};
use super::pool::{MePool, MeWriter, WriterContour};
use super::pool_lifecycle::MeTaskRegistration;
use super::reader::reader_loop;
use super::wire::build_proxy_req_payload;
@@ -132,16 +132,6 @@ impl MePool {
drain_deadline_epoch_secs: drain_deadline_epoch_secs.clone(),
allow_drain_fallback: allow_drain_fallback.clone(),
};
self.writers
.update(|writers| writers.push(writer.clone()))
.await;
self.registry
.register_writer(writer_id, tx.clone(), byte_budget)
.await;
self.registry.mark_writer_idle(writer_id).await;
self.conn_count.fetch_add(1, Ordering::Relaxed);
self.notify_writer_epoch();
let reg = self.registry.clone();
let writers_arc = self.writers_arc();
let ping_tracker = Arc::new(tokio::sync::Mutex::new(HashMap::<i64, Instant>::new()));
@@ -177,8 +167,10 @@ impl MePool {
let route_fairshare_enabled = self.transport_policy.me_route_fairshare_enabled.clone();
let reader_route_data_wait_ms = self.transport_policy.me_reader_route_data_wait_ms.clone();
self.lifecycle
.spawn_registered_writer(task_registration, async move {
let writer_task = {
// Keep transport ownership behind a stable-size pointer while publication waits for
// locks.
Box::pin(async move {
// Reader MUST be the first branch in biased select! to avoid read starvation.
let exit = tokio::select! {
biased;
@@ -269,11 +261,41 @@ impl MePool {
let remaining = writers_arc.read().await.len();
debug!(writer_id, remaining, "ME writer lifecycle task finished");
});
})
};
self.publish_prepared_writer(writer, tx, byte_budget, task_registration, writer_task)
.await;
Ok(())
}
/// Commits writer visibility and lifecycle ownership after all cancellation points.
#[allow(clippy::too_many_arguments)]
pub(in crate::transport::middle_proxy) async fn publish_prepared_writer<F>(
self: &Arc<Self>,
writer: MeWriter,
tx: mpsc::Sender<WriterCommand>,
byte_budget: Arc<tokio::sync::Semaphore>,
task_registration: MeTaskRegistration<'_>,
writer_task: F,
) where
F: Future<Output = ()> + Send + 'static,
{
let (mut writers, mut registry_registration) = tokio::join!(
self.writers.write(),
self.registry.prepare_writer_registration()
);
registry_registration.install(writer.id, tx, byte_budget);
writers.push(writer);
self.conn_count.fetch_add(1, Ordering::Relaxed);
self.lifecycle
.spawn_registered_writer(task_registration, writer_task);
drop(writers);
drop(registry_registration);
self.notify_writer_epoch();
}
pub(crate) async fn remove_writer_and_close_clients(self: &Arc<Self>, writer_id: u64) {
// Full client cleanup now happens inside `registry.writer_lost` to keep
// writer reap/remove paths strictly non-blocking per connection.
@@ -337,7 +359,12 @@ impl MePool {
self.stats.increment_me_writer_removed_unexpected_total();
}
close_tx = Some(w.tx.clone());
self.conn_count.fetch_sub(1, Ordering::Relaxed);
// Teardown remains idempotent if the advisory count was already reconciled.
let _ =
self.conn_count
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |count| {
count.checked_sub(1)
});
removed = true;
}
}
+2
View File
@@ -18,6 +18,8 @@ const ROUTE_QUEUED_BYTE_PERMIT_UNIT: usize = 16 * 1024;
const ROUTE_QUEUED_PERMITS_PER_SLOT: usize = 4;
const ROUTE_QUEUED_MAX_FRAME_PERMITS: usize = 1024;
// Transactional writer registry publication.
mod publication;
mod writer;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -0,0 +1,50 @@
use std::sync::Arc;
use tokio::sync::{MutexGuard, Semaphore, mpsc};
use super::super::codec::WriterCommand;
use super::{BindingInner, ConnRegistry, WriterRoute};
/// Holds registry binding ownership until pool writer visibility is published.
pub(in crate::transport::middle_proxy) struct WriterRegistrationGuard<'a> {
registry: &'a ConnRegistry,
binding: MutexGuard<'a, BindingInner>,
}
impl ConnRegistry {
/// Acquires the only cancellation point required for writer registration.
pub(in crate::transport::middle_proxy) async fn prepare_writer_registration(
&self,
) -> WriterRegistrationGuard<'_> {
WriterRegistrationGuard {
registry: self,
binding: self.binding.inner.lock().await,
}
}
}
impl WriterRegistrationGuard<'_> {
/// Installs registry state while retaining binding ownership for pool publication.
pub(in crate::transport::middle_proxy) fn install(
&mut self,
writer_id: u64,
tx: mpsc::Sender<WriterCommand>,
byte_budget: Arc<Semaphore>,
) {
self.binding.conns_for_writer.entry(writer_id).or_default();
self.registry
.binding
.bound_clients_by_writer
.entry(writer_id)
.or_insert(0);
self.registry
.binding
.writer_idle_since_epoch_secs
.entry(writer_id)
.or_insert_with(ConnRegistry::now_epoch_secs);
self.registry
.writers
.map
.insert(writer_id, WriterRoute { tx, byte_budget });
}
}
+4 -30
View File
@@ -58,28 +58,16 @@ impl ConnRegistry {
}
/// Registers one writer command route and its matching memory budget atomically.
#[allow(dead_code)]
pub async fn register_writer(
&self,
writer_id: u64,
tx: mpsc::Sender<WriterCommand>,
byte_budget: Arc<tokio::sync::Semaphore>,
) {
let mut binding = self.binding.inner.lock().await;
binding
.conns_for_writer
.entry(writer_id)
.or_insert_with(HashSet::new);
self.binding
.bound_clients_by_writer
.entry(writer_id)
.or_insert(0);
self.binding
.writer_idle_since_epoch_secs
.entry(writer_id)
.or_insert_with(Self::now_epoch_secs);
self.writers
.map
.insert(writer_id, super::WriterRoute { tx, byte_budget });
self.prepare_writer_registration()
.await
.install(writer_id, tx, byte_budget);
}
/// Unregister connection, returning associated writer_id if any.
@@ -346,20 +334,6 @@ impl ConnRegistry {
true
}
pub async fn mark_writer_idle(&self, writer_id: u64) {
let mut binding = self.binding.inner.lock().await;
binding
.conns_for_writer
.entry(writer_id)
.or_insert_with(HashSet::new);
let count = binding
.conns_for_writer
.get(&writer_id)
.map(|set| set.len())
.unwrap_or(0);
self.set_writer_bound_count(writer_id, count);
}
pub async fn get_last_writer_meta(&self, writer_id: u64) -> Option<ConnMeta> {
self.binding
.last_meta_for_writer
@@ -0,0 +1,147 @@
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU32, AtomicU64, Ordering};
use std::time::{Duration, Instant};
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use super::codec::WriterCommand;
use super::pool::{MeWriter, WriterContour};
use super::pool_writer_security_tests::make_pool;
use super::registry::ConnMeta;
#[tokio::test]
async fn successful_writer_publication_is_fully_visible_and_removable() {
let pool = make_pool().await;
let writer_id = 76_002;
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443);
let (tx, _rx) = mpsc::channel::<WriterCommand>(8);
let byte_budget = pool.new_writer_byte_budget();
let cancel = CancellationToken::new();
let writer = MeWriter {
id: writer_id,
addr,
source_ip: addr.ip(),
writer_dc: 2,
generation: pool.current_generation(),
contour: Arc::new(AtomicU8::new(WriterContour::Active.as_u8())),
created_at: Instant::now(),
tx: tx.clone(),
byte_budget: byte_budget.clone(),
cancel: cancel.clone(),
degraded: Arc::new(AtomicBool::new(false)),
rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)),
draining: Arc::new(AtomicBool::new(false)),
draining_started_at_epoch_secs: Arc::new(AtomicU64::new(0)),
drain_deadline_epoch_secs: Arc::new(AtomicU64::new(0)),
allow_drain_fallback: Arc::new(AtomicBool::new(false)),
};
let task_started = Arc::new(AtomicBool::new(false));
let task_started_writer = Arc::clone(&task_started);
let writer_task = async move {
task_started_writer.store(true, Ordering::Release);
cancel.cancelled().await;
};
let task_registration = pool.lifecycle.try_register().unwrap();
pool.publish_prepared_writer(writer, tx, byte_budget, task_registration, writer_task)
.await;
assert_eq!(pool.writers.read().await.len(), 1);
assert_eq!(pool.conn_count.load(Ordering::Relaxed), 1);
tokio::time::timeout(Duration::from_secs(1), async {
while !task_started.load(Ordering::Acquire) {
tokio::task::yield_now().await;
}
})
.await
.expect("published writer task must start");
let (conn_id, _response_rx) = pool.registry.register().await;
assert!(
pool.registry
.bind_writer(
conn_id,
writer_id,
ConnMeta {
target_dc: 2,
client_addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 7300),
our_addr: addr,
proto_flags: 0,
},
)
.await
);
assert_eq!(
pool.registry.get_writer(conn_id).await.unwrap().writer_id,
writer_id
);
pool.remove_writer_and_close_clients(writer_id).await;
assert!(pool.writers.read().await.is_empty());
assert_eq!(pool.conn_count.load(Ordering::Relaxed), 0);
assert!(pool.registry.get_writer(conn_id).await.is_none());
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn published_writer_removal_race_preserves_count() {
let pool = make_pool().await;
for writer_id in 80_000..90_000 {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443);
let (tx, _rx) = mpsc::channel::<WriterCommand>(1);
let byte_budget = pool.new_writer_byte_budget();
let cancel = CancellationToken::new();
let writer = MeWriter {
id: writer_id,
addr,
source_ip: addr.ip(),
writer_dc: 2,
generation: pool.current_generation(),
contour: Arc::new(AtomicU8::new(WriterContour::Active.as_u8())),
created_at: Instant::now(),
tx: tx.clone(),
byte_budget: byte_budget.clone(),
cancel: cancel.clone(),
degraded: Arc::new(AtomicBool::new(false)),
rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)),
draining: Arc::new(AtomicBool::new(false)),
draining_started_at_epoch_secs: Arc::new(AtomicU64::new(0)),
drain_deadline_epoch_secs: Arc::new(AtomicU64::new(0)),
allow_drain_fallback: Arc::new(AtomicBool::new(false)),
};
let writer_task = async move {
cancel.cancelled().await;
};
let task_registration = pool.lifecycle.try_register().unwrap();
let remover_pool = Arc::clone(&pool);
let remover = tokio::spawn(async move {
loop {
if remover_pool
.writers
.snapshot()
.iter()
.any(|writer| writer.id == writer_id)
{
remover_pool
.remove_writer_and_close_clients(writer_id)
.await;
return;
}
tokio::task::yield_now().await;
}
});
pool.publish_prepared_writer(writer, tx, byte_budget, task_registration, writer_task)
.await;
tokio::time::timeout(Duration::from_secs(1), remover)
.await
.expect("published writer must become removable")
.expect("writer remover task must not panic");
assert!(pool.writers.read().await.is_empty());
assert_eq!(pool.conn_count.load(Ordering::Relaxed), 0);
}
}
@@ -15,7 +15,8 @@ use crate::crypto::SecureRandom;
use crate::network::probe::NetworkDecision;
use crate::stats::Stats;
async fn make_pool() -> Arc<MePool> {
/// Builds an isolated ME pool for writer-state tests.
pub(super) async fn make_pool() -> Arc<MePool> {
let general = GeneralConfig::default();
MePool::new(
@@ -154,6 +155,70 @@ async fn insert_writer(
pool.conn_count.fetch_add(1, Ordering::Relaxed);
}
#[tokio::test]
async fn cancelled_writer_publication_leaves_no_ghost_state() {
let pool = make_pool().await;
let writer_id = 76_001;
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443);
let (tx, mut rx) = mpsc::channel::<WriterCommand>(8);
let byte_budget = pool.new_writer_byte_budget();
let writer = MeWriter {
id: writer_id,
addr,
source_ip: addr.ip(),
writer_dc: 2,
generation: pool.current_generation(),
contour: Arc::new(AtomicU8::new(WriterContour::Active.as_u8())),
created_at: Instant::now(),
tx: tx.clone(),
byte_budget: byte_budget.clone(),
cancel: CancellationToken::new(),
degraded: Arc::new(AtomicBool::new(false)),
rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)),
draining: Arc::new(AtomicBool::new(false)),
draining_started_at_epoch_secs: Arc::new(AtomicU64::new(0)),
drain_deadline_epoch_secs: Arc::new(AtomicU64::new(0)),
allow_drain_fallback: Arc::new(AtomicBool::new(false)),
};
let task_started = Arc::new(AtomicBool::new(false));
let task_started_writer = Arc::clone(&task_started);
let writer_task = async move {
task_started_writer.store(true, Ordering::Release);
let _ = rx.recv().await;
};
let task_registration = pool.lifecycle.try_register().unwrap();
let held_registration = pool.registry.prepare_writer_registration().await;
let publication =
pool.publish_prepared_writer(writer, tx, byte_budget, task_registration, writer_task);
assert!(
tokio::time::timeout(Duration::from_millis(10), publication)
.await
.is_err()
);
drop(held_registration);
assert!(pool.writers.read().await.is_empty());
assert_eq!(pool.conn_count.load(Ordering::Relaxed), 0);
assert!(!task_started.load(Ordering::Acquire));
let (conn_id, _response_rx) = pool.registry.register().await;
assert!(
!pool
.registry
.bind_writer(
conn_id,
writer_id,
ConnMeta {
target_dc: 2,
client_addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 7300),
our_addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443),
proto_flags: 0,
},
)
.await
);
}
async fn current_writer_ids(pool: &Arc<MePool>) -> HashSet<u64> {
pool.writers
.read()