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(); let runtime_cfg = config_rx.borrow().clone();
data.user.in_runtime = runtime_cfg.access.users.contains_key(&data.user.username); data.user.in_runtime = runtime_cfg.access.users.contains_key(&data.user.username);
if let Some(enabled) = requested_enabled { if let Some(enabled) = requested_enabled {
shared let (_, cancelled) = shared
.proxy_shared .proxy_shared
.set_user_enabled(&data.user.username, enabled); .set_user_enabled(&data.user.username, enabled);
if !enabled { if !enabled {
let cancelled = shared
.proxy_shared
.cancel_user_sessions(&data.user.username);
if cancelled > 0 { if cancelled > 0 {
shared.runtime_events.record( shared.runtime_events.record(
"api.user.disable.runtime", "api.user.disable.runtime",
+2 -4
View File
@@ -104,8 +104,7 @@ pub(super) async fn handle(
}; };
let runtime_cfg = config_rx.borrow().clone(); let runtime_cfg = config_rx.borrow().clone();
data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); data.in_runtime = runtime_cfg.access.users.contains_key(&data.username);
let newly_disabled = shared.proxy_shared.set_user_enabled(base_user, false); let (newly_disabled, cancelled) = shared.proxy_shared.set_user_enabled(base_user, false);
let cancelled = shared.proxy_shared.cancel_user_sessions(base_user);
shared.runtime_events.record( shared.runtime_events.record(
"api.user.disable.ok", "api.user.disable.ok",
format!( format!(
@@ -290,11 +289,10 @@ pub(super) async fn handle(
let runtime_cfg = config_rx.borrow().clone(); let runtime_cfg = config_rx.borrow().clone();
data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); data.in_runtime = runtime_cfg.access.users.contains_key(&data.username);
if let Some(enabled) = enabled_update { if let Some(enabled) = enabled_update {
shared let (_, cancelled) = shared
.proxy_shared .proxy_shared
.set_user_enabled(&data.username, enabled); .set_user_enabled(&data.username, enabled);
if !enabled { if !enabled {
let cancelled = shared.proxy_shared.cancel_user_sessions(&data.username);
shared.runtime_events.record( shared.runtime_events.record(
"api.user.disable.runtime", "api.user.disable.runtime",
format!( format!(
+23 -496
View File
@@ -8,15 +8,22 @@
//! - `run [OPTIONS] [config.toml]` - Run in foreground (default behavior) //! - `run [OPTIONS] [config.toml]` - Run in foreground (default behavior)
//! - `healthcheck [OPTIONS] [config.toml]` - Run control-plane health probe //! - `healthcheck [OPTIONS] [config.toml]` - Run control-plane health probe
use rand::RngExt; use std::path::PathBuf;
use std::fs;
use std::path::{Path, PathBuf};
use std::process::Command;
use crate::healthcheck::{self, HealthcheckMode}; use crate::healthcheck::{self, HealthcheckMode};
#[cfg(unix)] #[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. /// CLI subcommand to execute.
#[derive(Debug, Clone, PartialEq, Eq)] #[derive(Debug, Clone, PartialEq, Eq)]
@@ -40,13 +47,20 @@ pub enum Subcommand {
/// Parsed subcommand with its options. /// Parsed subcommand with its options.
#[derive(Debug)] #[derive(Debug)]
pub struct ParsedCommand { pub struct ParsedCommand {
/// Selected command mode.
pub subcommand: Subcommand, pub subcommand: Subcommand,
/// PID file used by daemon-control commands.
pub pid_file: PathBuf, pub pid_file: PathBuf,
/// Configuration file passed to runtime or healthcheck.
pub config_path: String, pub config_path: String,
/// Requested healthcheck mode.
pub healthcheck_mode: HealthcheckMode, pub healthcheck_mode: HealthcheckMode,
/// Invalid healthcheck mode retained for command diagnostics.
pub healthcheck_mode_invalid: Option<String>, pub healthcheck_mode_invalid: Option<String>,
#[cfg(unix)] #[cfg(unix)]
/// Unix daemon lifecycle options.
pub daemon_opts: DaemonOptions, pub daemon_opts: DaemonOptions,
/// Fire-and-forget initialization options.
pub init_opts: Option<InitOptions>, pub init_opts: Option<InitOptions>,
} }
@@ -79,7 +93,6 @@ pub fn parse_command(args: &[String]) -> ParsedCommand {
return cmd; return cmd;
} }
// Check for subcommand as first argument
if let Some(first) = args.first() { if let Some(first) = args.first() {
match first.as_str() { match first.as_str() {
"start" => { "start" => {
@@ -120,11 +133,9 @@ pub fn parse_command(args: &[String]) -> ParsedCommand {
} }
} }
// Parse remaining options
let mut i = 0; let mut i = 0;
while i < args.len() { while i < args.len() {
match args[i].as_str() { match args[i].as_str() {
// Skip subcommand names
"start" | "stop" | "reload" | "status" | "run" | "healthcheck" => {} "start" | "stop" | "reload" | "status" | "run" | "healthcheck" => {}
"--mode" => { "--mode" => {
i += 1; i += 1;
@@ -154,7 +165,6 @@ pub fn parse_command(args: &[String]) -> ParsedCommand {
} }
} }
} }
// PID file option (for stop/reload/status)
"--pid-file" => { "--pid-file" => {
i += 1; i += 1;
if i < args.len() { if i < args.len() {
@@ -189,9 +199,9 @@ pub fn parse_command(args: &[String]) -> ParsedCommand {
#[cfg(unix)] #[cfg(unix)]
pub fn execute_subcommand(cmd: &ParsedCommand) -> Option<i32> { pub fn execute_subcommand(cmd: &ParsedCommand) -> Option<i32> {
match cmd.subcommand { match cmd.subcommand {
Subcommand::Stop => Some(cmd_stop(&cmd.pid_file)), Subcommand::Stop => Some(daemon_commands::stop(&cmd.pid_file)),
Subcommand::Reload => Some(cmd_reload(&cmd.pid_file)), Subcommand::Reload => Some(daemon_commands::reload(&cmd.pid_file)),
Subcommand::Status => Some(cmd_status(&cmd.pid_file)), Subcommand::Status => Some(daemon_commands::status(&cmd.pid_file)),
Subcommand::Healthcheck => { Subcommand::Healthcheck => {
if let Some(invalid_mode) = cmd.healthcheck_mode_invalid.as_ref() { if let Some(invalid_mode) = cmd.healthcheck_mode_invalid.as_ref() {
if invalid_mode.is_empty() { 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))] #[cfg(not(unix))]
pub fn execute_subcommand(cmd: &ParsedCommand) -> Option<i32> { pub fn execute_subcommand(cmd: &ParsedCommand) -> Option<i32> {
match cmd.subcommand { match cmd.subcommand {
@@ -261,487 +272,3 @@ pub fn execute_subcommand(cmd: &ParsedCommand) -> Option<i32> {
Subcommand::Run | Subcommand::Start => None, 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" "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 { 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" "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(); let mut cidr_auto_templates = HashSet::new();
for cidr in config.access.cidr_rate_limits.keys() { 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")); 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] #[test]
fn file_logging_requires_path() { fn file_logging_requires_path() {
let error = load_config_error_from_temp_toml( 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}; pub use access::{AccessConfig, CidrRateLimitKey, RateLimitBps};
#[allow(unused_imports)] #[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 api::{ApiConfig, ApiGrayAction};
pub use censorship::{ pub use censorship::{
AntiCensorshipConfig, ExclusiveMaskTarget, TlsFetchConfig, TlsFetchProfile, UnknownSniAction, AntiCensorshipConfig, ExclusiveMaskTarget, TlsFetchConfig, TlsFetchProfile, UnknownSniAction,
+5 -2
View File
@@ -1,5 +1,8 @@
use super::*; 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)] #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct AccessConfig { pub struct AccessConfig {
#[serde(default = "default_access_users")] #[serde(default = "default_access_users")]
@@ -260,10 +263,10 @@ fn parse_cidr_auto_prefix(
/// Transport rate limit in bits-per-second. /// Transport rate limit in bits-per-second.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)] #[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct RateLimitBps { 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)] #[serde(default)]
pub up_bps: u64, 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)] #[serde(default)]
pub down_bps: u64, pub down_bps: u64,
} }
+27 -278
View File
@@ -4,14 +4,20 @@
//! and privilege dropping for running telemt as a background service. //! and privilege dropping for running telemt as a background service.
use std::fs::{self, File, OpenOptions}; use std::fs::{self, File, OpenOptions};
use std::io::{self, Read, Write}; use std::io;
use std::os::unix::fs::OpenOptionsExt; use std::os::unix::fs::OpenOptionsExt;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use nix::errno::Errno; use nix::errno::Errno;
use nix::fcntl::{Flock, FlockArg}; use nix::unistd::{self, ForkResult, Gid, Uid, chdir, close, fork, getpid, setsid};
use nix::unistd::{self, ForkResult, Gid, Pid, Uid, chdir, close, fork, getpid, setsid}; use tracing::info;
use tracing::{debug, info, warn};
// 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. /// Default PID file location.
pub const DEFAULT_PID_FILE: &str = "/var/run/telemt.pid"; pub const DEFAULT_PID_FILE: &str = "/var/run/telemt.pid";
@@ -51,36 +57,47 @@ impl DaemonOptions {
/// Error types for daemon operations. /// Error types for daemon operations.
#[derive(Debug, thiserror::Error)] #[derive(Debug, thiserror::Error)]
pub enum DaemonError { pub enum DaemonError {
/// A daemonization fork failed.
#[error("fork failed: {0}")] #[error("fork failed: {0}")]
ForkFailed(#[source] nix::Error), ForkFailed(#[source] nix::Error),
/// Creation of the detached session failed.
#[error("setsid failed: {0}")] #[error("setsid failed: {0}")]
SetsidFailed(#[source] nix::Error), SetsidFailed(#[source] nix::Error),
/// Switching to the configured working directory failed.
#[error("chdir failed: {0}")] #[error("chdir failed: {0}")]
ChdirFailed(#[source] nix::Error), ChdirFailed(#[source] nix::Error),
/// Opening `/dev/null` for standard-stream redirection failed.
#[error("failed to open /dev/null: {0}")] #[error("failed to open /dev/null: {0}")]
DevNullFailed(#[source] io::Error), DevNullFailed(#[source] io::Error),
/// Redirecting a standard file descriptor failed.
#[error("failed to redirect stdio: {0}")] #[error("failed to redirect stdio: {0}")]
RedirectFailed(#[source] nix::Error), RedirectFailed(#[source] nix::Error),
/// A PID lifecycle operation failed.
#[error("PID file error: {0}")] #[error("PID file error: {0}")]
PidFile(String), PidFile(String),
/// Another process owns the daemon PID lifecycle.
#[error("another instance is already running (pid {0})")] #[error("another instance is already running (pid {0})")]
AlreadyRunning(i32), AlreadyRunning(i32),
/// The configured runtime user does not exist.
#[error("user '{0}' not found")] #[error("user '{0}' not found")]
UserNotFound(String), UserNotFound(String),
/// The configured runtime group does not exist.
#[error("group '{0}' not found")] #[error("group '{0}' not found")]
GroupNotFound(String), GroupNotFound(String),
/// Applying the configured runtime identity failed.
#[error("failed to set uid/gid: {0}")] #[error("failed to set uid/gid: {0}")]
PrivilegeDrop(#[source] nix::Error), PrivilegeDrop(#[source] nix::Error),
/// An underlying filesystem operation failed.
#[error("io error: {0}")] #[error("io error: {0}")]
Io(#[from] io::Error), Io(#[from] io::Error),
} }
@@ -106,38 +123,28 @@ pub enum DaemonizeResult {
/// Returns `DaemonizeResult::Parent` in the original parent (which should exit), /// Returns `DaemonizeResult::Parent` in the original parent (which should exit),
/// or `DaemonizeResult::Child` in the final daemon child. /// or `DaemonizeResult::Child` in the final daemon child.
pub fn daemonize(working_dir: Option<&Path>) -> Result<DaemonizeResult, DaemonError> { pub fn daemonize(working_dir: Option<&Path>) -> Result<DaemonizeResult, DaemonError> {
// First fork
match unsafe { fork() } { match unsafe { fork() } {
Ok(ForkResult::Parent { .. }) => { Ok(ForkResult::Parent { .. }) => {
// Parent exits
return Ok(DaemonizeResult::Parent); return Ok(DaemonizeResult::Parent);
} }
Ok(ForkResult::Child) => { Ok(ForkResult::Child) => {}
// Child continues
}
Err(e) => return Err(DaemonError::ForkFailed(e)), Err(e) => return Err(DaemonError::ForkFailed(e)),
} }
// Create new session, become session leader
setsid().map_err(DaemonError::SetsidFailed)?; setsid().map_err(DaemonError::SetsidFailed)?;
// Second fork to ensure we can never acquire a controlling terminal // Second fork to ensure we can never acquire a controlling terminal
match unsafe { fork() } { match unsafe { fork() } {
Ok(ForkResult::Parent { .. }) => { Ok(ForkResult::Parent { .. }) => {
// Intermediate parent exits
std::process::exit(0); std::process::exit(0);
} }
Ok(ForkResult::Child) => { Ok(ForkResult::Child) => {}
// Final daemon child continues
}
Err(e) => return Err(DaemonError::ForkFailed(e)), Err(e) => return Err(DaemonError::ForkFailed(e)),
} }
// Change working directory
let target_dir = working_dir.unwrap_or(Path::new("/")); let target_dir = working_dir.unwrap_or(Path::new("/"));
chdir(target_dir).map_err(DaemonError::ChdirFailed)?; chdir(target_dir).map_err(DaemonError::ChdirFailed)?;
// Redirect stdin, stdout, stderr to /dev/null
redirect_stdio_to_devnull()?; redirect_stdio_to_devnull()?;
Ok(DaemonizeResult::Child) Ok(DaemonizeResult::Child)
@@ -156,21 +163,17 @@ fn redirect_stdio_to_devnull() -> Result<(), DaemonError> {
// Use libc::dup2 directly for redirecting standard file descriptors // Use libc::dup2 directly for redirecting standard file descriptors
// nix 0.31's dup2 requires OwnedFd which doesn't work well with stdio fds // nix 0.31's dup2 requires OwnedFd which doesn't work well with stdio fds
unsafe { unsafe {
// Redirect stdin (fd 0)
if libc::dup2(devnull_fd, 0) < 0 { if libc::dup2(devnull_fd, 0) < 0 {
return Err(DaemonError::RedirectFailed(Errno::last())); return Err(DaemonError::RedirectFailed(Errno::last()));
} }
// Redirect stdout (fd 1)
if libc::dup2(devnull_fd, 1) < 0 { if libc::dup2(devnull_fd, 1) < 0 {
return Err(DaemonError::RedirectFailed(Errno::last())); return Err(DaemonError::RedirectFailed(Errno::last()));
} }
// Redirect stderr (fd 2)
if libc::dup2(devnull_fd, 2) < 0 { if libc::dup2(devnull_fd, 2) < 0 {
return Err(DaemonError::RedirectFailed(Errno::last())); return Err(DaemonError::RedirectFailed(Errno::last()));
} }
} }
// Close original devnull fd if it's not one of the standard fds
if devnull_fd > 2 { if devnull_fd > 2 {
let _ = close(devnull_fd); let _ = close(devnull_fd);
} }
@@ -178,166 +181,6 @@ fn redirect_stdio_to_devnull() -> Result<(), DaemonError> {
Ok(()) 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, // macOS gates nix::unistd::setgroups differently in the current dependency set,
// so call libc directly there while preserving the original nix path elsewhere. // so call libc directly there while preserving the original nix path elsewhere.
fn set_supplementary_groups(gid: Gid) -> Result<(), nix::Error> { fn set_supplementary_groups(gid: Gid) -> Result<(), nix::Error> {
@@ -383,10 +226,12 @@ pub fn drop_privileges(
}; };
if (target_uid.is_some() || target_gid.is_some()) 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
{ {
for file in pid_file.ownership_file_handles().into_iter().flatten() {
unistd::fchown(file, target_uid, target_gid).map_err(DaemonError::PrivilegeDrop)?; unistd::fchown(file, target_uid, target_gid).map_err(DaemonError::PrivilegeDrop)?;
} }
}
if let Some(gid) = target_gid { if let Some(gid) = target_gid {
unistd::setgid(gid).map_err(DaemonError::PrivilegeDrop)?; unistd::setgid(gid).map_err(DaemonError::PrivilegeDrop)?;
@@ -401,7 +246,7 @@ pub fn drop_privileges(
if uid.as_raw() != 0 if uid.as_raw() != 0
&& let Some(pid) = pid_file && 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!( let probe_path = parent.join(format!(
".telemt_pid_probe_{}_{}", ".telemt_pid_probe_{}_{}",
std::process::id(), std::process::id(),
@@ -436,7 +281,6 @@ pub fn drop_privileges(
/// Looks up a user by name and returns their UID. /// Looks up a user by name and returns their UID.
fn lookup_user(name: &str) -> Result<Uid, DaemonError> { fn lookup_user(name: &str) -> Result<Uid, DaemonError> {
// Use libc getpwnam
let c_name = let c_name =
std::ffi::CString::new(name).map_err(|_| DaemonError::UserNotFound(name.to_string()))?; 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)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
@@ -571,29 +345,4 @@ mod tests {
}; };
assert!(!opts.should_daemonize()); 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; break;
} }
let cfg = config_rx_user_enabled.borrow_and_update().clone(); let cfg = config_rx_user_enabled.borrow_and_update().clone();
for user in shared_user_enabled.apply_user_enabled_config(&cfg.access.user_enabled) { for (user, cancelled) in
let cancelled = shared_user_enabled.cancel_user_sessions(&user); shared_user_enabled.apply_user_enabled_config(&cfg.access.user_enabled)
{
if cancelled > 0 { if cancelled > 0 {
info!( info!(
user = %user, user = %user,
+17 -1
View File
@@ -79,7 +79,11 @@ where
let route_snapshot = deps.route_runtime.snapshot(); let route_snapshot = deps.route_runtime.snapshot();
let session_id = deps.rng.u64(); let session_id = deps.rng.u64();
let 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 session_cancel = user_session.token();
let selected_me_pool = if deps.config.general.use_middle_proxy let selected_me_pool = if deps.config.general.use_middle_proxy
&& matches!(route_snapshot.mode, RelayRouteMode::Middle) && matches!(route_snapshot.mode, RelayRouteMode::Middle)
@@ -246,6 +250,18 @@ impl UserConnectionReservation {
} }
self.stats.decrement_user_curr_connects(&self.user); 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 { impl Drop for UserConnectionReservation {
+40 -61
View File
@@ -16,10 +16,7 @@ mod quota;
pub(super) use self::combined::CombinedStream; pub(super) use self::combined::CombinedStream;
pub(super) use self::counters::SharedCounters; pub(super) use self::counters::SharedCounters;
pub(super) use self::quota::is_quota_io_error; pub(super) use self::quota::is_quota_io_error;
use self::quota::{ use self::quota::{QUOTA_RESERVE_MAX_ROUNDS, QUOTA_RESERVE_SPIN_RETRIES, quota_io_error};
QUOTA_RESERVE_MAX_ROUNDS, QUOTA_RESERVE_SPIN_RETRIES, quota_io_error,
refund_reserved_quota_bytes,
};
pub(super) use self::quota::{quota_adaptive_interval_bytes, should_immediate_quota_check}; pub(super) use self::quota::{quota_adaptive_interval_bytes, should_immediate_quota_check};
/// Transparent I/O wrapper that tracks per-user statistics and activity. /// 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 remaining_before = None;
let mut reserved_read_bytes = 0u64; let mut quota_reservation = None;
let mut read_limit = buf.remaining(); let mut read_limit = buf.remaining();
if let Some(limit) = this.quota_limit { if let Some(limit) = this.quota_limit {
let used_before = this.user_stats.quota_used(); let used_before = this.user_stats.quota_used();
@@ -231,11 +228,11 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
let desired = read_limit as u64; let desired = read_limit as u64;
let mut reserve_rounds = 0usize; let mut reserve_rounds = 0usize;
while reserved_read_bytes == 0 { while quota_reservation.is_none() {
for _ in 0..QUOTA_RESERVE_SPIN_RETRIES { for _ in 0..QUOTA_RESERVE_SPIN_RETRIES {
match this.user_stats.quota_try_reserve(desired, limit) { match this.user_stats.quota_reserve(desired, limit) {
Ok(_) => { Ok(reservation) => {
reserved_read_bytes = desired; quota_reservation = Some(reservation);
break; break;
} }
Err(crate::stats::QuotaReserveError::LimitExceeded) => { 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); reserve_rounds = reserve_rounds.saturating_add(1);
if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS { if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS {
this.stats.increment_quota_contention_timeout_total(); this.stats.increment_quota_contention_timeout_total();
@@ -287,9 +284,9 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
match read_result { match read_result {
Poll::Ready(Ok(n)) => { Poll::Ready(Ok(n)) => {
if reserved_read_bytes > n as u64 { if let Some(reservation) = quota_reservation.take() {
let refund_bytes = reserved_read_bytes - n as u64; let refund_bytes = reservation.reserved_bytes().saturating_sub(n as u64);
refund_reserved_quota_bytes(this.user_stats.as_ref(), refund_bytes); reservation.settle(n as u64);
this.stats.add_quota_refund_bytes_total(refund_bytes); this.stats.add_quota_refund_bytes_total(refund_bytes);
} }
if n > 0 { if n > 0 {
@@ -333,16 +330,16 @@ impl<S: AsyncRead + Unpin> AsyncRead for StatsIo<S> {
Poll::Ready(Ok(())) Poll::Ready(Ok(()))
} }
Poll::Pending => { Poll::Pending => {
if reserved_read_bytes > 0 { if let Some(reservation) = quota_reservation.take() {
refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_read_bytes); this.stats
this.stats.add_quota_refund_bytes_total(reserved_read_bytes); .add_quota_refund_bytes_total(reservation.reserved_bytes());
} }
Poll::Pending Poll::Pending
} }
Poll::Ready(Err(err)) => { Poll::Ready(Err(err)) => {
if reserved_read_bytes > 0 { if let Some(reservation) = quota_reservation.take() {
refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_read_bytes); this.stats
this.stats.add_quota_refund_bytes_total(reserved_read_bytes); .add_quota_refund_bytes_total(reservation.reserved_bytes());
} }
Poll::Ready(Err(err)) Poll::Ready(Err(err))
} }
@@ -361,14 +358,15 @@ impl<S: AsyncWrite + Unpin> AsyncWrite for StatsIo<S> {
return Poll::Ready(Err(quota_io_error())); return Poll::Ready(Err(quota_io_error()));
} }
let mut shaper_reserved_bytes = 0u64; let mut shaper_reservation = None;
let mut write_buf = buf; let mut write_buf = buf;
if let Some(lease) = this.traffic_lease.as_ref() { if let Some(lease) = this.traffic_lease.as_ref() {
if !buf.is_empty() { if !buf.is_empty() {
loop { 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 { if consume.granted > 0 {
shaper_reserved_bytes = consume.granted; shaper_reservation = Some(reservation);
if consume.granted < buf.len() as u64 { if consume.granted < buf.len() as u64 {
write_buf = &buf[..consume.granted as usize]; 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 remaining_before = None;
let mut reserved_bytes = 0u64; let mut quota_reservation = None;
if let Some(limit) = this.quota_limit { if let Some(limit) = this.quota_limit {
if !write_buf.is_empty() { if !write_buf.is_empty() {
let mut reserve_rounds = 0usize; let mut reserve_rounds = 0usize;
while reserved_bytes == 0 { while quota_reservation.is_none() {
let used_before = this.user_stats.quota_used(); let used_before = this.user_stats.quota_used();
let remaining = limit.saturating_sub(used_before); let remaining = limit.saturating_sub(used_before);
if remaining == 0 { 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); this.quota_exceeded.store(true, Ordering::Release);
return Poll::Ready(Err(quota_io_error())); 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 desired = remaining.min(write_buf.len() as u64);
let mut saw_contention = false; let mut saw_contention = false;
for _ in 0..QUOTA_RESERVE_SPIN_RETRIES { for _ in 0..QUOTA_RESERVE_SPIN_RETRIES {
match this.user_stats.quota_try_reserve(desired, limit) { match this.user_stats.quota_reserve(desired, limit) {
Ok(_) => { Ok(reservation) => {
reserved_bytes = desired; quota_reservation = Some(reservation);
write_buf = &write_buf[..desired as usize]; write_buf = &write_buf[..desired as usize];
break; 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); reserve_rounds = reserve_rounds.saturating_add(1);
if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS { if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS {
this.stats.increment_quota_contention_timeout_total(); this.stats.increment_quota_contention_timeout_total();
if let Some(lease) = this.traffic_lease.as_ref() { Self::arm_wait(&mut this.quota_wait, false, false);
lease.refund(RateDirection::Down, shaper_reserved_bytes); let _ =
} Self::poll_wait(&mut this.quota_wait, cx, None, RateDirection::Up);
let _ = this.arm_quota_wait(cx);
return Poll::Pending; return Poll::Pending;
} else if saw_contention { } else if saw_contention {
std::hint::spin_loop(); 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 used_before = this.user_stats.quota_used();
let remaining = limit.saturating_sub(used_before); let remaining = limit.saturating_sub(used_before);
if remaining == 0 { 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); this.quota_exceeded.store(true, Ordering::Release);
return Poll::Ready(Err(quota_io_error())); 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) { match Pin::new(&mut this.inner).poll_write(cx, write_buf) {
Poll::Ready(Ok(n)) => { Poll::Ready(Ok(n)) => {
if reserved_bytes > n as u64 { if let Some(reservation) = quota_reservation.take() {
let refund_bytes = reserved_bytes - n as u64; let refund_bytes = reservation.reserved_bytes().saturating_sub(n as u64);
refund_reserved_quota_bytes(this.user_stats.as_ref(), refund_bytes); reservation.settle(n as u64);
this.stats.add_quota_refund_bytes_total(refund_bytes); this.stats.add_quota_refund_bytes_total(refund_bytes);
} }
if shaper_reserved_bytes > n as u64 if let Some(reservation) = shaper_reservation.take() {
&& let Some(lease) = this.traffic_lease.as_ref() reservation.settle_written(n as u64);
{
lease.refund(RateDirection::Down, shaper_reserved_bytes - n as u64);
} }
if n > 0 { if n > 0 {
if let Some(lease) = this.traffic_lease.as_ref() { 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(Ok(n))
} }
Poll::Ready(Err(err)) => { Poll::Ready(Err(err)) => {
if reserved_bytes > 0 { if let Some(reservation) = quota_reservation.take() {
refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_bytes); this.stats
this.stats.add_quota_refund_bytes_total(reserved_bytes); .add_quota_refund_bytes_total(reservation.reserved_bytes());
}
if shaper_reserved_bytes > 0
&& let Some(lease) = this.traffic_lease.as_ref()
{
lease.refund(RateDirection::Down, shaper_reserved_bytes);
} }
Poll::Ready(Err(err)) Poll::Ready(Err(err))
} }
Poll::Pending => { Poll::Pending => {
if reserved_bytes > 0 { if let Some(reservation) = quota_reservation.take() {
refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_bytes); this.stats
this.stats.add_quota_refund_bytes_total(reserved_bytes); .add_quota_refund_bytes_total(reservation.reserved_bytes());
}
if shaper_reserved_bytes > 0
&& let Some(lease) = this.traffic_lease.as_ref()
{
lease.refund(RateDirection::Down, shaper_reserved_bytes);
} }
Poll::Pending Poll::Pending
} }
-8
View File
@@ -1,4 +1,3 @@
use crate::stats::UserStats;
use std::io; use std::io;
#[derive(Debug)] #[derive(Debug)]
@@ -46,10 +45,3 @@ pub(in crate::proxy::relay) fn should_immediate_quota_check(
) -> bool { ) -> bool {
remaining_before <= QUOTA_NEAR_LIMIT_BYTES || charge_bytes >= QUOTA_LARGE_CHARGE_BYTES 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);
}
+144 -37
View File
@@ -6,6 +6,7 @@ use std::sync::{Arc, Mutex};
use std::time::Instant; use std::time::Instant;
use dashmap::DashMap; use dashmap::DashMap;
use parking_lot::Mutex as ParkingMutex;
use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc}; use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc};
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
@@ -75,13 +76,18 @@ pub(crate) struct MiddleRelaySharedState {
pub(crate) relay_idle_mark_seq: AtomicU64, pub(crate) relay_idle_mark_seq: AtomicU64,
} }
#[derive(Default)]
struct UserAdmissionState {
disabled_users: HashSet<String>,
sessions_by_user: HashMap<String, HashMap<u64, CancellationToken>>,
}
pub(crate) struct ProxySharedState { pub(crate) struct ProxySharedState {
pub(crate) handshake: HandshakeSharedState, pub(crate) handshake: HandshakeSharedState,
pub(crate) middle_relay: MiddleRelaySharedState, pub(crate) middle_relay: MiddleRelaySharedState,
pub(crate) traffic_limiter: Arc<TrafficLimiter>, pub(crate) traffic_limiter: Arc<TrafficLimiter>,
pub(crate) direct_buffer_budget: Arc<DirectBufferBudget>, pub(crate) direct_buffer_budget: Arc<DirectBufferBudget>,
disabled_users: DashMap<String, ()>, user_admission: ParkingMutex<UserAdmissionState>,
active_user_sessions: DashMap<(String, u64), CancellationToken>,
pub(crate) conntrack_pressure_active: AtomicBool, pub(crate) conntrack_pressure_active: AtomicBool,
pub(crate) conntrack_close_tx: Mutex<Option<mpsc::Sender<ConntrackCloseEvent>>>, pub(crate) conntrack_close_tx: Mutex<Option<mpsc::Sender<ConntrackCloseEvent>>>,
masking_fallback_permits: Arc<Semaphore>, masking_fallback_permits: Arc<Semaphore>,
@@ -106,7 +112,18 @@ struct UserSessionGuard {
impl Drop for UserSessionGuard { impl Drop for UserSessionGuard {
fn drop(&mut self) { 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(), traffic_limiter: TrafficLimiter::new(),
direct_buffer_budget, direct_buffer_budget,
disabled_users: DashMap::new(), user_admission: ParkingMutex::new(UserAdmissionState::default()),
active_user_sessions: DashMap::new(),
conntrack_pressure_active: AtomicBool::new(false), conntrack_pressure_active: AtomicBool::new(false),
conntrack_close_tx: Mutex::new(None), conntrack_close_tx: Mutex::new(None),
masking_fallback_permits: Arc::new(Semaphore::new(MASKING_FALLBACK_MAX_CONCURRENT)), masking_fallback_permits: Arc::new(Semaphore::new(MASKING_FALLBACK_MAX_CONCURRENT)),
@@ -167,68 +183,102 @@ impl ProxySharedState {
} }
pub(crate) fn is_user_enabled(&self, user: &str) -> bool { 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 { 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 { if enabled {
self.disabled_users.remove(user); admission.disabled_users.remove(user);
false (false, Vec::new())
} else { } else {
self.disabled_users.insert(user.to_string(), ()).is_none() 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( pub(crate) fn apply_user_enabled_config(
&self, &self,
user_enabled: &HashMap<String, bool>, user_enabled: &HashMap<String, bool>,
) -> Vec<String> { ) -> Vec<(String, usize)> {
let desired_disabled = user_enabled let desired_disabled = user_enabled
.iter() .iter()
.filter_map(|(user, enabled)| (!*enabled).then_some(user.clone())) .filter_map(|(user, enabled)| (!*enabled).then_some(user.clone()))
.collect::<HashSet<_>>(); .collect::<HashSet<_>>();
let current_disabled = self let cancellations = {
.disabled_users let mut admission = self.user_admission.lock();
.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 let newly_disabled = desired_disabled
.difference(&current_disabled) .difference(&admission.disabled_users)
.cloned() .cloned()
.collect::<Vec<_>>(); .collect::<Vec<_>>();
for user in desired_disabled { admission.disabled_users = desired_disabled;
self.disabled_users.insert(user, ());
}
newly_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( pub(crate) fn register_user_session(
self: &Arc<Self>, self: &Arc<Self>,
user: &str, user: &str,
session_id: u64, session_id: u64,
) -> UserSessionRegistration { ) -> Option<UserSessionRegistration> {
let token = CancellationToken::new(); let token = CancellationToken::new();
let key = (user.to_string(), session_id); let key = (user.to_string(), session_id);
self.active_user_sessions.insert(key.clone(), token.clone()); let mut admission = self.user_admission.lock();
UserSessionRegistration { 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, token,
_guard: UserSessionGuard { _guard: UserSessionGuard {
shared: Arc::clone(self), shared: Arc::clone(self),
key, key,
}, },
} })
} }
pub(crate) fn cancel_user_sessions(&self, user: &str) -> usize { pub(crate) fn cancel_user_sessions(&self, user: &str) -> usize {
let tokens = self let tokens: Vec<CancellationToken> = self
.active_user_sessions .user_admission
.iter() .lock()
.filter_map(|entry| (entry.key().0 == user).then(|| entry.value().clone())) .sessions_by_user
.collect::<Vec<_>>(); .get(user)
.map(|sessions| sessions.values().cloned().collect())
.unwrap_or_default();
for token in &tokens { for token in &tokens {
token.cancel(); token.cancel();
} }
@@ -311,7 +361,7 @@ mod tests {
let mut newly_disabled = shared.apply_user_enabled_config(&user_enabled); let mut newly_disabled = shared.apply_user_enabled_config(&user_enabled);
newly_disabled.sort(); 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("alice"));
assert!(shared.is_user_enabled("bob")); assert!(shared.is_user_enabled("bob"));
@@ -325,9 +375,9 @@ mod tests {
#[test] #[test]
fn cancel_user_sessions_cancels_only_registered_matching_user() { fn cancel_user_sessions_cancels_only_registered_matching_user() {
let shared = ProxySharedState::new(); let shared = ProxySharedState::new();
let alice_1 = shared.register_user_session("alice", 1); let alice_1 = shared.register_user_session("alice", 1).unwrap();
let alice_2 = shared.register_user_session("alice", 2); let alice_2 = shared.register_user_session("alice", 2).unwrap();
let bob = shared.register_user_session("bob", 1); let bob = shared.register_user_session("bob", 1).unwrap();
let alice_1_token = alice_1.token(); let alice_1_token = alice_1.token();
let alice_2_token = alice_2.token(); let alice_2_token = alice_2.token();
let bob_token = bob.token(); let bob_token = bob.token();
@@ -339,4 +389,61 @@ mod tests {
assert!(alice_2_token.is_cancelled()); assert!(alice_2_token.is_cancelled());
assert!(!bob_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 arc_swap::ArcSwap;
use dashmap::DashMap; use dashmap::DashMap;
use ipnetwork::IpNetwork; use ipnetwork::IpNetwork;
use parking_lot::Mutex as ParkingMutex;
use crate::config::RateLimitBps; use crate::config::RateLimitBps;
@@ -32,6 +33,9 @@ const REGISTRY_SHARDS: usize = 64;
const FAIR_EPOCH_MS: u64 = 20; const FAIR_EPOCH_MS: u64 = 20;
const MAX_BORROW_CHUNK_BYTES: u64 = 32 * 1024; const MAX_BORROW_CHUNK_BYTES: u64 = 32 * 1024;
const CLEANUP_INTERVAL_SECS: u64 = 60; 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)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RateDirection { pub enum RateDirection {
@@ -76,12 +80,12 @@ struct ScopeMetrics {
struct AtomicRatePair { struct AtomicRatePair {
up_bps: AtomicU64, up_bps: AtomicU64,
down_bps: AtomicU64, down_bps: AtomicU64,
revision: ParkingMutex<u64>,
} }
#[derive(Default)] #[derive(Default)]
struct DirectionBucket { struct DirectionBucket {
epoch: AtomicU64, state: AtomicU64,
used: AtomicU64,
} }
struct UserBucket { struct UserBucket {
@@ -93,15 +97,13 @@ struct UserBucket {
#[derive(Default)] #[derive(Default)]
struct CidrDirectionBucket { struct CidrDirectionBucket {
epoch: AtomicU64, used: DirectionBucket,
used: AtomicU64, active_users: DirectionBucket,
active_users: AtomicU64,
} }
#[derive(Default)] #[derive(Default)]
struct CidrUserDirectionState { struct CidrUserDirectionState {
epoch: AtomicU64, used: DirectionBucket,
used: AtomicU64,
} }
struct CidrUserShare { struct CidrUserShare {
@@ -139,6 +141,7 @@ enum CidrPolicyMatch<'a> {
#[derive(Default)] #[derive(Default)]
struct PolicySnapshot { struct PolicySnapshot {
revision: u64,
user_limits: HashMap<String, RateLimitBps>, user_limits: HashMap<String, RateLimitBps>,
cidr_rules_v4: Vec<CidrRule>, cidr_rules_v4: Vec<CidrRule>,
cidr_rules_v6: Vec<CidrRule>, cidr_rules_v6: Vec<CidrRule>,
@@ -162,9 +165,25 @@ pub struct TrafficLease {
pub struct TrafficLimiter { pub struct TrafficLimiter {
policy: ArcSwap<PolicySnapshot>, policy: ArcSwap<PolicySnapshot>,
policy_update: ParkingMutex<()>,
user_buckets: ShardedRegistry<UserBucket>, user_buckets: ShardedRegistry<UserBucket>,
cidr_buckets: ShardedRegistry<CidrBucket>, cidr_buckets: ShardedRegistry<CidrBucket>,
user_scope: ScopeMetrics, user_scope: ScopeMetrics,
cidr_scope: ScopeMetrics, cidr_scope: ScopeMetrics,
last_cleanup_epoch_secs: AtomicU64, 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>>,
}
+215 -161
View File
@@ -26,202 +26,276 @@ impl ScopeMetrics {
} }
impl AtomicRatePair { impl AtomicRatePair {
pub(super) fn set(&self, limits: RateLimitBps) { pub(super) fn new(revision: u64, limits: RateLimitBps) -> Self {
self.up_bps.store(limits.up_bps, Ordering::Relaxed); let rates = Self::default();
self.down_bps.store(limits.down_bps, Ordering::Relaxed); 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 { pub(super) fn get(&self, direction: RateDirection) -> u64 {
match direction { match direction {
RateDirection::Up => self.up_bps.load(Ordering::Relaxed), RateDirection::Up => self.up_bps.load(Ordering::Acquire),
RateDirection::Down => self.down_bps.load(Ordering::Relaxed), RateDirection::Down => self.down_bps.load(Ordering::Acquire),
} }
} }
} }
impl DirectionBucket { impl DirectionBucket {
pub(super) fn sync_epoch(&self, epoch: u64) { fn unpack(state: u64) -> (u64, u64) {
let current = self.epoch.load(Ordering::Relaxed); (state >> PACKED_USAGE_BITS, state & PACKED_USAGE_MASK)
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);
}
} }
pub(super) fn try_consume(&self, cap_bps: u64, requested: u64) -> u64 { fn pack(epoch: u64, used: u64) -> Option<u64> {
if requested == 0 { if epoch > PACKED_EPOCH_MAX || used > PACKED_USAGE_MASK {
return 0; return None;
} }
if cap_bps == 0 { Some((epoch << PACKED_USAGE_BITS) | used)
return requested;
} }
let epoch = current_epoch(); pub(super) fn used_at(&self, epoch: u64) -> Option<u64> {
self.sync_epoch(epoch); if epoch > PACKED_EPOCH_MAX {
let cap_epoch = bytes_per_epoch(cap_bps); return None;
}
let (current_epoch, used) = Self::unpack(self.state.load(Ordering::Relaxed));
(current_epoch == epoch).then_some(used)
}
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 { loop {
let used = self.used.load(Ordering::Relaxed); let (observed_epoch, observed_used) = Self::unpack(observed);
if used >= cap_epoch { if observed_epoch > epoch {
return 0; 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); let grant = requested.min(remaining);
if grant == 0 { if grant == 0 {
return 0; return None;
} }
let next = used.saturating_add(grant); let next = Self::pack(epoch, used + grant)?;
if self match self.state.compare_exchange_weak(
.used observed,
.compare_exchange_weak(used, next, Ordering::Relaxed, Ordering::Relaxed) next,
.is_ok() Ordering::Relaxed,
{ Ordering::Relaxed,
return grant; ) {
Ok(_) => {
return Some(DirectionDebit {
bucket: self,
epoch,
refundable: grant,
});
}
Err(actual) => observed = actual,
} }
} }
} }
pub(super) fn refund(&self, bytes: u64) { fn refund_at(&self, epoch: u64, bytes: u64) {
if bytes == 0 { if bytes == 0 || epoch > PACKED_EPOCH_MAX {
return; 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 { impl UserBucket {
pub(super) fn new(limits: RateLimitBps) -> Self { pub(super) fn new(revision: u64, limits: RateLimitBps) -> Self {
let rates = AtomicRatePair::default();
rates.set(limits);
Self { Self {
rates, rates: AtomicRatePair::new(revision, limits),
up: DirectionBucket::default(), up: DirectionBucket::default(),
down: DirectionBucket::default(), down: DirectionBucket::default(),
active_leases: AtomicU64::new(0), active_leases: AtomicU64::new(0),
} }
} }
pub(super) fn set_rates(&self, limits: RateLimitBps) { pub(super) fn set_rates(&self, revision: u64, limits: RateLimitBps) {
self.rates.set(limits); 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); let cap_bps = self.rates.get(direction);
match direction { if cap_bps == 0 {
RateDirection::Up => self.up.try_consume(cap_bps, requested), return (requested, None);
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),
} }
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 { impl CidrDirectionBucket {
pub(super) fn sync_epoch(&self, epoch: u64) { pub(super) fn try_reserve<'a>(
let current = self.epoch.load(Ordering::Relaxed); &'a self,
if current == epoch { user_state: &'a CidrUserDirectionState,
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,
cap_epoch: u64, cap_epoch: u64,
requested: u64, requested: u64,
) -> u64 { ) -> (u64, Option<DirectionDebit<'a>>, Option<DirectionDebit<'a>>) {
if requested == 0 || cap_epoch == 0 { if requested == 0 || cap_epoch == 0 {
return 0; return (0, None, None);
} }
let epoch = current_epoch(); let epoch = current_epoch();
self.sync_epoch(epoch); if !user_state.ensure_active(epoch, &self.active_users) {
user_state.sync_epoch_and_mark_active(epoch, &self.active_users); return (0, None, None);
let active_users = self.active_users.load(Ordering::Relaxed).max(1); }
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); let fair_share = cap_epoch.saturating_div(active_users).max(1);
loop { loop {
let total_used = self.used.load(Ordering::Relaxed); let Some(user_used) = user_state.used.used_at(epoch) else {
if total_used >= cap_epoch { return (0, None, None);
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 guaranteed_remaining = fair_share.saturating_sub(user_used);
if grant == 0 { let (user_cap, desired) = if guaranteed_remaining > 0 {
return 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 next_total = total_used.saturating_add(grant); };
if self let user_granted = user_debit.granted();
.used let Some(aggregate_debit) = self.used.try_reserve_at(epoch, cap_epoch, user_granted)
.compare_exchange_weak(total_used, next_total, Ordering::Relaxed, Ordering::Relaxed) else {
.is_ok() return (0, None, None);
{ };
user_state.used.fetch_add(grant, Ordering::Relaxed); let granted = aggregate_debit.granted();
return grant; 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 { impl CidrUserDirectionState {
pub(super) fn sync_epoch_and_mark_active(&self, epoch: u64, active_users: &AtomicU64) { pub(super) fn ensure_active(&self, epoch: u64, active_users: &DirectionBucket) -> bool {
let current = self.epoch.load(Ordering::Relaxed); if epoch > PACKED_EPOCH_MAX {
if current == epoch { return false;
return;
} }
if current < epoch let mut observed = self.used.state.load(Ordering::Relaxed);
&& self loop {
.epoch let (observed_epoch, _) = DirectionBucket::unpack(observed);
.compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed) if observed_epoch == epoch {
.is_ok() return true;
{ }
self.used.store(0, Ordering::Relaxed); if observed_epoch > epoch {
active_users.fetch_add(1, Ordering::Relaxed); 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);
} }
} }
@@ -236,11 +310,9 @@ impl CidrUserShare {
} }
impl CidrBucket { impl CidrBucket {
pub(super) fn new(limits: RateLimitBps) -> Self { pub(super) fn new(revision: u64, limits: RateLimitBps) -> Self {
let rates = AtomicRatePair::default();
rates.set(limits);
Self { Self {
rates, rates: AtomicRatePair::new(revision, limits),
up: CidrDirectionBucket::default(), up: CidrDirectionBucket::default(),
down: CidrDirectionBucket::default(), down: CidrDirectionBucket::default(),
users: ShardedRegistry::new(REGISTRY_SHARDS), users: ShardedRegistry::new(REGISTRY_SHARDS),
@@ -248,8 +320,8 @@ impl CidrBucket {
} }
} }
pub(super) fn set_rates(&self, limits: RateLimitBps) { pub(super) fn set_rates(&self, revision: u64, limits: RateLimitBps) {
self.rates.set(limits); self.rates.set(revision, limits);
} }
pub(super) fn acquire_user_share(&self, user: &str) -> Arc<CidrUserShare> { pub(super) fn acquire_user_share(&self, user: &str) -> Arc<CidrUserShare> {
@@ -268,38 +340,20 @@ impl CidrBucket {
}); });
} }
pub(super) fn try_consume_for_user( pub(super) fn try_reserve_for_user<'a>(
&self, &'a self,
direction: RateDirection, direction: RateDirection,
share: &CidrUserShare, share: &'a CidrUserShare,
requested: u64, requested: u64,
) -> u64 { ) -> (u64, Option<DirectionDebit<'a>>, Option<DirectionDebit<'a>>) {
let cap_bps = self.rates.get(direction); let cap_bps = self.rates.get(direction);
if cap_bps == 0 { if cap_bps == 0 {
return requested; return (requested, None, None);
} }
let cap_epoch = bytes_per_epoch(cap_bps); let cap_epoch = bytes_per_epoch(cap_bps);
match direction { match direction {
RateDirection::Up => self.up.try_consume(&share.up, cap_epoch, requested), RateDirection::Up => self.up.try_reserve(&share.up, cap_epoch, requested),
RateDirection::Down => self.down.try_consume(&share.down, cap_epoch, requested), RateDirection::Down => self.down.try_reserve(&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);
}
} }
} }
+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 { pub(super) fn current_epoch() -> u64 {
let start = limiter_epoch_start(); let start = limiter_epoch_start();
let elapsed_ms = start.elapsed().as_millis() as u64; 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 { pub(super) fn limiter_epoch_start() -> &'static Instant {
+67 -23
View File
@@ -1,70 +1,93 @@
use super::*; use super::*;
impl TrafficLease { 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 { if requested == 0 {
return TrafficConsumeResult { return TrafficReservation {
result: TrafficConsumeResult {
granted: 0, granted: 0,
blocked_user: false, blocked_user: false,
blocked_cidr: false, blocked_cidr: false,
},
user: None,
cidr: None,
cidr_user: None,
}; };
} }
let mut granted = requested; let mut granted = requested;
let mut user_debit = None;
if let Some(user_bucket) = self.user_bucket.as_ref() { 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 { if user_granted == 0 {
self.limiter.observe_throttle(direction, true, false); self.limiter.observe_throttle(direction, true, false);
return TrafficConsumeResult { return TrafficReservation {
result: TrafficConsumeResult {
granted: 0, granted: 0,
blocked_user: true, blocked_user: true,
blocked_cidr: false, blocked_cidr: false,
},
user: user_debit,
cidr: None,
cidr_user: None,
}; };
} }
granted = user_granted; granted = user_granted;
} }
let mut cidr_debit = None;
let mut cidr_user_debit = None;
if let (Some(cidr_bucket), Some(cidr_user_share)) = if let (Some(cidr_bucket), Some(cidr_user_share)) =
(self.cidr_bucket.as_ref(), self.cidr_user_share.as_ref()) (self.cidr_bucket.as_ref(), self.cidr_user_share.as_ref())
{ {
let cidr_granted = let (cidr_granted, aggregate_debit, share_debit) =
cidr_bucket.try_consume_for_user(direction, cidr_user_share, granted); cidr_bucket.try_reserve_for_user(direction, cidr_user_share, granted);
cidr_debit = aggregate_debit;
cidr_user_debit = share_debit;
if cidr_granted < granted 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 { if cidr_granted == 0 {
self.limiter.observe_throttle(direction, false, true); self.limiter.observe_throttle(direction, false, true);
return TrafficConsumeResult { return TrafficReservation {
result: TrafficConsumeResult {
granted: 0, granted: 0,
blocked_user: false, blocked_user: false,
blocked_cidr: true, blocked_cidr: true,
},
user: user_debit,
cidr: cidr_debit,
cidr_user: cidr_user_debit,
}; };
} }
granted = cidr_granted; granted = cidr_granted;
} }
TrafficConsumeResult { TrafficReservation {
result: TrafficConsumeResult {
granted, granted,
blocked_user: false, blocked_user: false,
blocked_cidr: false, blocked_cidr: false,
},
user: user_debit,
cidr: cidr_debit,
cidr_user: cidr_user_debit,
} }
} }
pub fn refund(&self, direction: RateDirection, bytes: u64) { pub fn try_consume(&self, direction: RateDirection, requested: u64) -> TrafficConsumeResult {
if bytes == 0 { let reservation = self.try_reserve(direction, requested);
return; let result = reservation.result();
} reservation.settle_written(result.granted);
result
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 observe_wait_ms( 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 { impl Drop for TrafficLease {
fn drop(&mut self) { fn drop(&mut self) {
if let Some(bucket) = self.user_bucket.as_ref() { if let Some(bucket) = self.user_bucket.as_ref() {
+12 -4
View File
@@ -5,6 +5,7 @@ impl TrafficLimiter {
pub fn new() -> Arc<Self> { pub fn new() -> Arc<Self> {
Arc::new(Self { Arc::new(Self {
policy: ArcSwap::from_pointee(PolicySnapshot::default()), policy: ArcSwap::from_pointee(PolicySnapshot::default()),
policy_update: ParkingMutex::new(()),
user_buckets: ShardedRegistry::new(REGISTRY_SHARDS), user_buckets: ShardedRegistry::new(REGISTRY_SHARDS),
cidr_buckets: ShardedRegistry::new(REGISTRY_SHARDS), cidr_buckets: ShardedRegistry::new(REGISTRY_SHARDS),
user_scope: ScopeMetrics::default(), user_scope: ScopeMetrics::default(),
@@ -18,6 +19,11 @@ impl TrafficLimiter {
user_limits: HashMap<String, RateLimitBps>, user_limits: HashMap<String, RateLimitBps>,
cidr_limits: HashMap<CidrRateLimitKey, 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 let filtered_users = user_limits
.into_iter() .into_iter()
.filter(|(_, limit)| limit.up_bps > 0 || limit.down_bps > 0) .filter(|(_, limit)| limit.up_bps > 0 || limit.down_bps > 0)
@@ -78,6 +84,7 @@ impl TrafficLimiter {
.store(cidr_policy_entries as u64, Ordering::Relaxed); .store(cidr_policy_entries as u64, Ordering::Relaxed);
self.policy.store(Arc::new(PolicySnapshot { self.policy.store(Arc::new(PolicySnapshot {
revision,
user_limits: filtered_users, user_limits: filtered_users,
cidr_rules_v4, cidr_rules_v4,
cidr_rules_v6, cidr_rules_v6,
@@ -86,6 +93,7 @@ impl TrafficLimiter {
cidr_rule_keys, cidr_rule_keys,
})); }));
drop(policy_update);
self.maybe_cleanup(); self.maybe_cleanup();
} }
@@ -99,12 +107,12 @@ impl TrafficLimiter {
if let Some(limit) = policy.user_limits.get(user).copied() { if let Some(limit) = policy.user_limits.get(user).copied() {
let bucket = self.user_buckets.get_or_insert_with( let bucket = self.user_buckets.get_or_insert_with(
user, user,
|| UserBucket::new(limit), || UserBucket::new(policy.revision, limit),
|bucket| { |bucket| {
bucket.active_leases.fetch_add(1, Ordering::Relaxed); bucket.active_leases.fetch_add(1, Ordering::Relaxed);
}, },
); );
bucket.set_rates(limit); bucket.set_rates(policy.revision, limit);
self.user_scope self.user_scope
.active_leases .active_leases
.fetch_add(1, Ordering::Relaxed); .fetch_add(1, Ordering::Relaxed);
@@ -121,12 +129,12 @@ impl TrafficLimiter {
}; };
let bucket = self.cidr_buckets.get_or_insert_with( let bucket = self.cidr_buckets.get_or_insert_with(
key, key,
|| CidrBucket::new(limits), || CidrBucket::new(policy.revision, limits),
|bucket| { |bucket| {
bucket.active_leases.fetch_add(1, Ordering::Relaxed); bucket.active_leases.fetch_add(1, Ordering::Relaxed);
}, },
); );
bucket.set_rates(limits); bucket.set_rates(policy.revision, limits);
self.cidr_scope self.cidr_scope
.active_leases .active_leases
.fetch_add(1, Ordering::Relaxed); .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" "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::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, AtomicUsize, Ordering};
use std::time::Instant; use std::time::Instant;
pub(crate) use self::quota_store::QuotaStore; pub(crate) use self::quota_store::{QuotaReservation, QuotaStore};
#[allow(unused_imports)] #[allow(unused_imports)]
pub use self::replay::{ReplayChecker, ReplayStats}; pub use self::replay::{ReplayChecker, ReplayStats};
use self::telemetry::TelemetryPolicy; use self::telemetry::TelemetryPolicy;
@@ -392,11 +392,6 @@ impl UserStats {
self.quota.used() 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. /// Attempts one CAS reservation step against the quota counter.
/// ///
/// Callers control retry/yield policy. This primitive intentionally does /// Callers control retry/yield policy. This primitive intentionally does
@@ -404,6 +399,18 @@ impl UserStats {
/// with their own contention strategy. /// with their own contention strategy.
#[inline] #[inline]
pub fn quota_try_reserve(&self, bytes: u64, limit: u64) -> Result<u64, QuotaReserveError> { 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) self.quota.try_reserve(bytes, limit)
} }
} }
+145 -38
View File
@@ -2,6 +2,7 @@ use std::collections::HashMap;
use std::sync::Arc; use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::atomic::{AtomicU64, Ordering};
use arc_swap::ArcSwap;
use dashmap::DashMap; use dashmap::DashMap;
use super::{QuotaReserveError, UserQuotaSnapshot}; use super::{QuotaReserveError, UserQuotaSnapshot};
@@ -12,11 +13,22 @@ pub struct QuotaStore {
users: DashMap<String, Arc<UserQuotaCounters>>, users: DashMap<String, Arc<UserQuotaCounters>>,
} }
/// Atomic quota state for one configured user. /// Atomically replaceable quota state for one configured user.
#[derive(Default)]
pub(crate) struct UserQuotaCounters { pub(crate) struct UserQuotaCounters {
generation: ArcSwap<QuotaGeneration>,
}
struct QuotaGeneration {
used_bytes: AtomicU64, 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 { impl QuotaStore {
@@ -38,18 +50,12 @@ impl QuotaStore {
pub(crate) fn load(&self, user: &str, used_bytes: u64, last_reset_epoch_secs: u64) { pub(crate) fn load(&self, user: &str, used_bytes: u64, last_reset_epoch_secs: u64) {
let state = self.user(user); let state = self.user(user);
state.used_bytes.store(used_bytes, Ordering::Relaxed); state.replace(used_bytes, last_reset_epoch_secs);
state
.last_reset_epoch_secs
.store(last_reset_epoch_secs, Ordering::Relaxed);
} }
pub(crate) fn reset(&self, user: &str, now_epoch_secs: u64) -> UserQuotaSnapshot { pub(crate) fn reset(&self, user: &str, now_epoch_secs: u64) -> UserQuotaSnapshot {
let state = self.user(user); let state = self.user(user);
state.used_bytes.store(0, Ordering::Relaxed); state.replace(0, now_epoch_secs);
state
.last_reset_epoch_secs
.store(now_epoch_secs, Ordering::Relaxed);
UserQuotaSnapshot { UserQuotaSnapshot {
used_bytes: 0, used_bytes: 0,
last_reset_epoch_secs: now_epoch_secs, last_reset_epoch_secs: now_epoch_secs,
@@ -64,8 +70,9 @@ impl QuotaStore {
let mut out = HashMap::new(); let mut out = HashMap::new();
for entry in self.users.iter() { for entry in self.users.iter() {
let state = entry.value(); let state = entry.value();
let used_bytes = state.used(); let generation = state.generation.load_full();
let last_reset_epoch_secs = state.last_reset_epoch_secs.load(Ordering::Relaxed); 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 { if used_bytes == 0 && last_reset_epoch_secs == 0 {
continue; continue;
} }
@@ -81,25 +88,105 @@ 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 { 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] #[inline]
pub(crate) fn used(&self) -> u64 { pub(crate) fn used(&self) -> u64 {
self.used_bytes.load(Ordering::Relaxed) self.generation.load().used_bytes.load(Ordering::Relaxed)
} }
#[inline] #[inline]
pub(crate) fn charge(&self, bytes: u64) -> u64 { pub(crate) fn charge(&self, bytes: u64) -> u64 {
self.used_bytes self.generation
.load_full()
.used_bytes
.fetch_add(bytes, Ordering::Relaxed) .fetch_add(bytes, Ordering::Relaxed)
.saturating_add(bytes) .saturating_add(bytes)
} }
#[inline] #[inline]
pub(crate) fn refund(&self, bytes: u64) { pub(crate) fn try_reserve(
let mut current = self.used_bytes.load(Ordering::Relaxed); &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 generation.used_bytes.compare_exchange_weak(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
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 { loop {
let next = current.saturating_sub(bytes); let next = current.saturating_sub(bytes);
match self.used_bytes.compare_exchange_weak( match generation.used_bytes.compare_exchange_weak(
current, current,
next, next,
Ordering::Relaxed, Ordering::Relaxed,
@@ -109,32 +196,13 @@ impl UserQuotaCounters {
Err(observed) => current = observed, 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);
if bytes > limit.saturating_sub(current) {
return Err(QuotaReserveError::LimitExceeded);
}
let next = current.saturating_add(bytes);
match self.used_bytes.compare_exchange_weak(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => Ok(next),
Err(_) => Err(QuotaReserveError::Contended),
}
}
} }
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::stats::Stats; use crate::stats::Stats;
use std::sync::Barrier;
#[test] #[test]
fn quota_counters_are_shared_across_stats_generations() { fn quota_counters_are_shared_across_stats_generations() {
@@ -148,4 +216,43 @@ mod tests {
second.reset_user_quota("alice"); second.reset_user_quota("alice");
assert_eq!(first.get_user_quota_used("alice"), 0); 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); 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] #[test]
fn test_cached_handle_survives_map_cleanup_until_last_drop() { fn test_cached_handle_survives_map_cleanup_until_last_drop() {
let stats = Stats::new(); let stats = Stats::new();
+3
View File
@@ -33,6 +33,9 @@ mod pool_runtime_api;
mod pool_status; mod pool_status;
mod pool_writer; mod pool_writer;
#[cfg(test)] #[cfg(test)]
#[path = "tests/pool_writer_publication_tests.rs"]
mod pool_writer_publication_tests;
#[cfg(test)]
#[path = "tests/pool_writer_security_tests.rs"] #[path = "tests/pool_writer_security_tests.rs"]
mod pool_writer_security_tests; mod pool_writer_security_tests;
mod reader; 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) { fn debug_assert_store_guarded(&self) {
debug_assert!( debug_assert!(
self.writers_write_guard.try_lock().is_err(), self.writers_write_guard.try_lock().is_err(),
@@ -1,4 +1,5 @@
use std::collections::HashMap; use std::collections::HashMap;
use std::future::Future;
use std::io::ErrorKind; use std::io::ErrorKind;
use std::net::SocketAddr; use std::net::SocketAddr;
use std::sync::Arc; 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::codec::{RpcWriter, WriterCommand, build_control_payload};
use super::pool::{MePool, MeWriter, WriterContour}; use super::pool::{MePool, MeWriter, WriterContour};
use super::pool_lifecycle::MeTaskRegistration;
use super::reader::reader_loop; use super::reader::reader_loop;
use super::wire::build_proxy_req_payload; use super::wire::build_proxy_req_payload;
@@ -132,16 +132,6 @@ impl MePool {
drain_deadline_epoch_secs: drain_deadline_epoch_secs.clone(), drain_deadline_epoch_secs: drain_deadline_epoch_secs.clone(),
allow_drain_fallback: allow_drain_fallback.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 reg = self.registry.clone();
let writers_arc = self.writers_arc(); let writers_arc = self.writers_arc();
let ping_tracker = Arc::new(tokio::sync::Mutex::new(HashMap::<i64, Instant>::new())); 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 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(); let reader_route_data_wait_ms = self.transport_policy.me_reader_route_data_wait_ms.clone();
self.lifecycle let writer_task = {
.spawn_registered_writer(task_registration, async move { // 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. // Reader MUST be the first branch in biased select! to avoid read starvation.
let exit = tokio::select! { let exit = tokio::select! {
biased; biased;
@@ -269,11 +261,41 @@ impl MePool {
let remaining = writers_arc.read().await.len(); let remaining = writers_arc.read().await.len();
debug!(writer_id, remaining, "ME writer lifecycle task finished"); debug!(writer_id, remaining, "ME writer lifecycle task finished");
}); })
};
self.publish_prepared_writer(writer, tx, byte_budget, task_registration, writer_task)
.await;
Ok(()) 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) { 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 // Full client cleanup now happens inside `registry.writer_lost` to keep
// writer reap/remove paths strictly non-blocking per connection. // writer reap/remove paths strictly non-blocking per connection.
@@ -337,7 +359,12 @@ impl MePool {
self.stats.increment_me_writer_removed_unexpected_total(); self.stats.increment_me_writer_removed_unexpected_total();
} }
close_tx = Some(w.tx.clone()); 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; 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_PERMITS_PER_SLOT: usize = 4;
const ROUTE_QUEUED_MAX_FRAME_PERMITS: usize = 1024; const ROUTE_QUEUED_MAX_FRAME_PERMITS: usize = 1024;
// Transactional writer registry publication.
mod publication;
mod writer; mod writer;
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[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. /// Registers one writer command route and its matching memory budget atomically.
#[allow(dead_code)]
pub async fn register_writer( pub async fn register_writer(
&self, &self,
writer_id: u64, writer_id: u64,
tx: mpsc::Sender<WriterCommand>, tx: mpsc::Sender<WriterCommand>,
byte_budget: Arc<tokio::sync::Semaphore>, byte_budget: Arc<tokio::sync::Semaphore>,
) { ) {
let mut binding = self.binding.inner.lock().await; self.prepare_writer_registration()
binding .await
.conns_for_writer .install(writer_id, tx, byte_budget);
.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 });
} }
/// Unregister connection, returning associated writer_id if any. /// Unregister connection, returning associated writer_id if any.
@@ -346,20 +334,6 @@ impl ConnRegistry {
true 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> { pub async fn get_last_writer_meta(&self, writer_id: u64) -> Option<ConnMeta> {
self.binding self.binding
.last_meta_for_writer .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::network::probe::NetworkDecision;
use crate::stats::Stats; 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(); let general = GeneralConfig::default();
MePool::new( MePool::new(
@@ -154,6 +155,70 @@ async fn insert_writer(
pool.conn_count.fetch_add(1, Ordering::Relaxed); 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> { async fn current_writer_ids(pool: &Arc<MePool>) -> HashSet<u64> {
pool.writers pool.writers
.read() .read()