diff --git a/src/api/handler/fixed_routes.rs b/src/api/handler/fixed_routes.rs index 57cdf99..96ae7c6 100644 --- a/src/api/handler/fixed_routes.rs +++ b/src/api/handler/fixed_routes.rs @@ -35,13 +35,10 @@ pub(super) async fn create_user_route( let runtime_cfg = config_rx.borrow().clone(); data.user.in_runtime = runtime_cfg.access.users.contains_key(&data.user.username); if let Some(enabled) = requested_enabled { - shared + let (_, cancelled) = shared .proxy_shared .set_user_enabled(&data.user.username, enabled); if !enabled { - let cancelled = shared - .proxy_shared - .cancel_user_sessions(&data.user.username); if cancelled > 0 { shared.runtime_events.record( "api.user.disable.runtime", diff --git a/src/api/handler/user_routes.rs b/src/api/handler/user_routes.rs index 004f01d..72ec55d 100644 --- a/src/api/handler/user_routes.rs +++ b/src/api/handler/user_routes.rs @@ -104,8 +104,7 @@ pub(super) async fn handle( }; let runtime_cfg = config_rx.borrow().clone(); data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); - let newly_disabled = shared.proxy_shared.set_user_enabled(base_user, false); - let cancelled = shared.proxy_shared.cancel_user_sessions(base_user); + let (newly_disabled, cancelled) = shared.proxy_shared.set_user_enabled(base_user, false); shared.runtime_events.record( "api.user.disable.ok", format!( @@ -290,11 +289,10 @@ pub(super) async fn handle( let runtime_cfg = config_rx.borrow().clone(); data.in_runtime = runtime_cfg.access.users.contains_key(&data.username); if let Some(enabled) = enabled_update { - shared + let (_, cancelled) = shared .proxy_shared .set_user_enabled(&data.username, enabled); if !enabled { - let cancelled = shared.proxy_shared.cancel_user_sessions(&data.username); shared.runtime_events.record( "api.user.disable.runtime", format!( diff --git a/src/cli.rs b/src/cli.rs index e740b2b..d9a4e15 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -8,15 +8,22 @@ //! - `run [OPTIONS] [config.toml]` - Run in foreground (default behavior) //! - `healthcheck [OPTIONS] [config.toml]` - Run control-plane health probe -use rand::RngExt; -use std::fs; -use std::path::{Path, PathBuf}; -use std::process::Command; +use std::path::PathBuf; use crate::healthcheck::{self, HealthcheckMode}; #[cfg(unix)] -use crate::daemon::{self, DEFAULT_PID_FILE, DaemonOptions}; +use crate::daemon::{DEFAULT_PID_FILE, DaemonOptions}; + +// Unix daemon control and argument parsing. +#[cfg(unix)] +mod daemon_commands; +// Fire-and-forget installation workflow. +mod init; + +#[cfg(unix)] +pub use daemon_commands::parse_daemon_args; +pub use init::{InitOptions, parse_init_args, run_init}; /// CLI subcommand to execute. #[derive(Debug, Clone, PartialEq, Eq)] @@ -40,13 +47,20 @@ pub enum Subcommand { /// Parsed subcommand with its options. #[derive(Debug)] pub struct ParsedCommand { + /// Selected command mode. pub subcommand: Subcommand, + /// PID file used by daemon-control commands. pub pid_file: PathBuf, + /// Configuration file passed to runtime or healthcheck. pub config_path: String, + /// Requested healthcheck mode. pub healthcheck_mode: HealthcheckMode, + /// Invalid healthcheck mode retained for command diagnostics. pub healthcheck_mode_invalid: Option, #[cfg(unix)] + /// Unix daemon lifecycle options. pub daemon_opts: DaemonOptions, + /// Fire-and-forget initialization options. pub init_opts: Option, } @@ -79,7 +93,6 @@ pub fn parse_command(args: &[String]) -> ParsedCommand { return cmd; } - // Check for subcommand as first argument if let Some(first) = args.first() { match first.as_str() { "start" => { @@ -120,11 +133,9 @@ pub fn parse_command(args: &[String]) -> ParsedCommand { } } - // Parse remaining options let mut i = 0; while i < args.len() { match args[i].as_str() { - // Skip subcommand names "start" | "stop" | "reload" | "status" | "run" | "healthcheck" => {} "--mode" => { i += 1; @@ -154,7 +165,6 @@ pub fn parse_command(args: &[String]) -> ParsedCommand { } } } - // PID file option (for stop/reload/status) "--pid-file" => { i += 1; if i < args.len() { @@ -189,9 +199,9 @@ pub fn parse_command(args: &[String]) -> ParsedCommand { #[cfg(unix)] pub fn execute_subcommand(cmd: &ParsedCommand) -> Option { match cmd.subcommand { - Subcommand::Stop => Some(cmd_stop(&cmd.pid_file)), - Subcommand::Reload => Some(cmd_reload(&cmd.pid_file)), - Subcommand::Status => Some(cmd_status(&cmd.pid_file)), + Subcommand::Stop => Some(daemon_commands::stop(&cmd.pid_file)), + Subcommand::Reload => Some(daemon_commands::reload(&cmd.pid_file)), + Subcommand::Status => Some(daemon_commands::status(&cmd.pid_file)), Subcommand::Healthcheck => { if let Some(invalid_mode) = cmd.healthcheck_mode_invalid.as_ref() { if invalid_mode.is_empty() { @@ -224,6 +234,7 @@ pub fn execute_subcommand(cmd: &ParsedCommand) -> Option { } } +/// Executes a non-server subcommand on platforms without daemon support. #[cfg(not(unix))] pub fn execute_subcommand(cmd: &ParsedCommand) -> Option { match cmd.subcommand { @@ -261,487 +272,3 @@ pub fn execute_subcommand(cmd: &ParsedCommand) -> Option { 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, - 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 { - 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> { - 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 = (0..16).map(|_| rng.random::()).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!("==================="); -} diff --git a/src/cli/daemon_commands.rs b/src/cli/daemon_commands.rs new file mode 100644 index 0000000..ee6fcca --- /dev/null +++ b/src/cli/daemon_commands.rs @@ -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()); + } +} diff --git a/src/cli/init.rs b/src/cli/init.rs new file mode 100644 index 0000000..fa36218 --- /dev/null +++ b/src/cli/init.rs @@ -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, + /// 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 { + 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> { + 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 = (0..16).map(|_| rng.random::()).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!("==================="); +} diff --git a/src/config/load/validate_runtime.rs b/src/config/load/validate_runtime.rs index ae05757..6cae0ed 100644 --- a/src/config/load/validate_runtime.rs +++ b/src/config/load/validate_runtime.rs @@ -196,6 +196,13 @@ pub(super) fn validate(config: &mut ProxyConfig) -> Result<()> { "access.user_rate_limits.{user} must set at least one non-zero direction" ))); } + for (direction, value) in [("up_bps", limit.up_bps), ("down_bps", limit.down_bps)] { + if value > MAX_RATE_LIMIT_BPS { + return Err(ProxyError::Config(format!( + "access.user_rate_limits.{user}.{direction} must be within [0, {MAX_RATE_LIMIT_BPS}]" + ))); + } + } } for (cidr, limit) in &config.access.cidr_rate_limits { @@ -204,6 +211,13 @@ pub(super) fn validate(config: &mut ProxyConfig) -> Result<()> { "access.cidr_rate_limits.{cidr} must set at least one non-zero direction" ))); } + for (direction, value) in [("up_bps", limit.up_bps), ("down_bps", limit.down_bps)] { + if value > MAX_RATE_LIMIT_BPS { + return Err(ProxyError::Config(format!( + "access.cidr_rate_limits.{cidr}.{direction} must be within [0, {MAX_RATE_LIMIT_BPS}]" + ))); + } + } } let mut cidr_auto_templates = HashSet::new(); for cidr in config.access.cidr_rate_limits.keys() { diff --git a/src/config/tests/load_basic_tests/defaults_access_tests.rs b/src/config/tests/load_basic_tests/defaults_access_tests.rs index 7d3d8de..4f6908e 100644 --- a/src/config/tests/load_basic_tests/defaults_access_tests.rs +++ b/src/config/tests/load_basic_tests/defaults_access_tests.rs @@ -288,6 +288,68 @@ fn cidr_rate_limits_reject_duplicate_normalized_auto_templates() { assert!(error.contains("duplicates normalized auto-template *6/128")); } +#[test] +fn rate_limits_accept_the_packed_counter_maximum() { + let cfg = load_config_from_temp_toml( + r#" + [censorship] + tls_domain = "example.com" + + [access.users] + user = "00000000000000000000000000000000" + + [access.user_rate_limits] + user = { up_bps = 100000000000, down_bps = 0 } + + [access.cidr_rate_limits] + "203.0.113.0/24" = { up_bps = 0, down_bps = 100000000000 } + "#, + ); + + assert_eq!(cfg.access.user_rate_limits["user"].up_bps, 100_000_000_000); + assert_eq!( + cfg.access.cidr_rate_limits[&CidrRateLimitKey::Network("203.0.113.0/24".parse().unwrap())] + .down_bps, + 100_000_000_000 + ); +} + +#[test] +fn user_rate_limits_reject_values_above_the_packed_counter_maximum() { + let error = load_config_error_from_temp_toml( + r#" + [censorship] + tls_domain = "example.com" + + [access.users] + user = "00000000000000000000000000000000" + + [access.user_rate_limits] + user = { up_bps = 100000000001, down_bps = 0 } + "#, + ); + + assert!(error.contains("access.user_rate_limits.user.up_bps must be within")); +} + +#[test] +fn cidr_rate_limits_reject_values_above_the_packed_counter_maximum() { + let error = load_config_error_from_temp_toml( + r#" + [censorship] + tls_domain = "example.com" + + [access.users] + user = "00000000000000000000000000000000" + + [access.cidr_rate_limits] + "203.0.113.0/24" = { up_bps = 0, down_bps = 100000000001 } + "#, + ); + + assert!(error.contains("access.cidr_rate_limits.203.0.113.0/24.down_bps must be within")); +} + #[test] fn file_logging_requires_path() { let error = load_config_error_from_temp_toml( diff --git a/src/config/types.rs b/src/config/types.rs index dfc0026..e769ef4 100644 --- a/src/config/types.rs +++ b/src/config/types.rs @@ -31,7 +31,7 @@ mod web_debug; pub use access::{AccessConfig, CidrRateLimitKey, RateLimitBps}; #[allow(unused_imports)] -pub(crate) use access::{CidrAutoTemplate, CidrAutoTemplateFamily}; +pub(crate) use access::{CidrAutoTemplate, CidrAutoTemplateFamily, MAX_RATE_LIMIT_BPS}; pub use api::{ApiConfig, ApiGrayAction}; pub use censorship::{ AntiCensorshipConfig, ExclusiveMaskTarget, TlsFetchConfig, TlsFetchProfile, UnknownSniAction, diff --git a/src/config/types/access.rs b/src/config/types/access.rs index 3f68d60..ac35dff 100644 --- a/src/config/types/access.rs +++ b/src/config/types/access.rs @@ -1,5 +1,8 @@ use super::*; +/// Highest rate that fits one packed 20 ms shaping epoch. +pub(crate) const MAX_RATE_LIMIT_BPS: u64 = 100_000_000_000; + #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct AccessConfig { #[serde(default = "default_access_users")] @@ -260,10 +263,10 @@ fn parse_cidr_auto_prefix( /// Transport rate limit in bits-per-second. #[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)] pub struct RateLimitBps { - /// Upload direction limit in bits-per-second; `0` means unlimited. + /// Upload limit in bits-per-second within `0..=100_000_000_000`; `0` means unlimited. #[serde(default)] pub up_bps: u64, - /// Download direction limit in bits-per-second; `0` means unlimited. + /// Download limit in bits-per-second within `0..=100_000_000_000`; `0` means unlimited. #[serde(default)] pub down_bps: u64, } diff --git a/src/daemon/mod.rs b/src/daemon/mod.rs index 66fca0f..15d274c 100644 --- a/src/daemon/mod.rs +++ b/src/daemon/mod.rs @@ -4,14 +4,20 @@ //! and privilege dropping for running telemt as a background service. use std::fs::{self, File, OpenOptions}; -use std::io::{self, Read, Write}; +use std::io; use std::os::unix::fs::OpenOptionsExt; use std::path::{Path, PathBuf}; use nix::errno::Errno; -use nix::fcntl::{Flock, FlockArg}; -use nix::unistd::{self, ForkResult, Gid, Pid, Uid, chdir, close, fork, getpid, setsid}; -use tracing::{debug, info, warn}; +use nix::unistd::{self, ForkResult, Gid, Uid, chdir, close, fork, getpid, setsid}; +use tracing::info; + +// PID file ownership and process-control helpers. +mod pid_file; + +pub use pid_file::DaemonStatus; +#[allow(unused_imports)] +pub use pid_file::{PidFile, check_status, read_pid_file, signal_pid_file}; /// Default PID file location. pub const DEFAULT_PID_FILE: &str = "/var/run/telemt.pid"; @@ -51,36 +57,47 @@ impl DaemonOptions { /// Error types for daemon operations. #[derive(Debug, thiserror::Error)] pub enum DaemonError { + /// A daemonization fork failed. #[error("fork failed: {0}")] ForkFailed(#[source] nix::Error), + /// Creation of the detached session failed. #[error("setsid failed: {0}")] SetsidFailed(#[source] nix::Error), + /// Switching to the configured working directory failed. #[error("chdir failed: {0}")] ChdirFailed(#[source] nix::Error), + /// Opening `/dev/null` for standard-stream redirection failed. #[error("failed to open /dev/null: {0}")] DevNullFailed(#[source] io::Error), + /// Redirecting a standard file descriptor failed. #[error("failed to redirect stdio: {0}")] RedirectFailed(#[source] nix::Error), + /// A PID lifecycle operation failed. #[error("PID file error: {0}")] PidFile(String), + /// Another process owns the daemon PID lifecycle. #[error("another instance is already running (pid {0})")] AlreadyRunning(i32), + /// The configured runtime user does not exist. #[error("user '{0}' not found")] UserNotFound(String), + /// The configured runtime group does not exist. #[error("group '{0}' not found")] GroupNotFound(String), + /// Applying the configured runtime identity failed. #[error("failed to set uid/gid: {0}")] PrivilegeDrop(#[source] nix::Error), + /// An underlying filesystem operation failed. #[error("io error: {0}")] Io(#[from] io::Error), } @@ -106,38 +123,28 @@ pub enum DaemonizeResult { /// Returns `DaemonizeResult::Parent` in the original parent (which should exit), /// or `DaemonizeResult::Child` in the final daemon child. pub fn daemonize(working_dir: Option<&Path>) -> Result { - // First fork match unsafe { fork() } { Ok(ForkResult::Parent { .. }) => { - // Parent exits return Ok(DaemonizeResult::Parent); } - Ok(ForkResult::Child) => { - // Child continues - } + Ok(ForkResult::Child) => {} Err(e) => return Err(DaemonError::ForkFailed(e)), } - // Create new session, become session leader setsid().map_err(DaemonError::SetsidFailed)?; // Second fork to ensure we can never acquire a controlling terminal match unsafe { fork() } { Ok(ForkResult::Parent { .. }) => { - // Intermediate parent exits std::process::exit(0); } - Ok(ForkResult::Child) => { - // Final daemon child continues - } + Ok(ForkResult::Child) => {} Err(e) => return Err(DaemonError::ForkFailed(e)), } - // Change working directory let target_dir = working_dir.unwrap_or(Path::new("/")); chdir(target_dir).map_err(DaemonError::ChdirFailed)?; - // Redirect stdin, stdout, stderr to /dev/null redirect_stdio_to_devnull()?; Ok(DaemonizeResult::Child) @@ -156,21 +163,17 @@ fn redirect_stdio_to_devnull() -> Result<(), DaemonError> { // Use libc::dup2 directly for redirecting standard file descriptors // nix 0.31's dup2 requires OwnedFd which doesn't work well with stdio fds unsafe { - // Redirect stdin (fd 0) if libc::dup2(devnull_fd, 0) < 0 { return Err(DaemonError::RedirectFailed(Errno::last())); } - // Redirect stdout (fd 1) if libc::dup2(devnull_fd, 1) < 0 { return Err(DaemonError::RedirectFailed(Errno::last())); } - // Redirect stderr (fd 2) if libc::dup2(devnull_fd, 2) < 0 { return Err(DaemonError::RedirectFailed(Errno::last())); } } - // Close original devnull fd if it's not one of the standard fds if devnull_fd > 2 { let _ = close(devnull_fd); } @@ -178,166 +181,6 @@ fn redirect_stdio_to_devnull() -> Result<(), DaemonError> { Ok(()) } -/// PID file manager with flock-based locking. -pub struct PidFile { - path: PathBuf, - file: Option, - locked: bool, -} - -impl PidFile { - /// Creates a new PID file manager for the given path. - pub fn new>(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, DaemonError> { - if !self.path.exists() { - return Ok(None); - } - - // Try to read existing PID - let mut contents = String::new(); - File::open(&self.path) - .and_then(|mut f| f.read_to_string(&mut contents)) - .map_err(|e| { - DaemonError::PidFile(format!("cannot read {}: {}", self.path.display(), e)) - })?; - - let pid: i32 = contents - .trim() - .parse() - .map_err(|_| DaemonError::PidFile(format!("invalid PID in {}", self.path.display())))?; - - // Check if process is still running - if is_process_running(pid) { - Ok(Some(pid)) - } else { - // Stale PID file - debug!(pid, path = %self.path.display(), "Removing stale PID file"); - let _ = fs::remove_file(&self.path); - Ok(None) - } - } - - /// Acquires the PID file lock and writes the current PID. - /// - /// Fails if another instance is already running. - pub fn acquire(&mut self) -> Result<(), DaemonError> { - // Check for running instance first - if let Some(pid) = self.check_running()? { - return Err(DaemonError::AlreadyRunning(pid)); - } - - // Ensure parent directory exists - if let Some(parent) = self.path.parent() { - if !parent.exists() { - fs::create_dir_all(parent).map_err(|e| { - DaemonError::PidFile(format!( - "cannot create directory {}: {}", - parent.display(), - e - )) - })?; - } - } - - // Open/create PID file with exclusive lock - let file = OpenOptions::new() - .write(true) - .create(true) - .truncate(true) - .mode(0o644) - .open(&self.path) - .map_err(|e| { - DaemonError::PidFile(format!("cannot open {}: {}", self.path.display(), e)) - })?; - - // Try to acquire exclusive lock (non-blocking) - let flock = Flock::lock(file, FlockArg::LockExclusiveNonblock).map_err(|(_, errno)| { - // Check if another instance grabbed the lock - if let Some(pid) = self.check_running().ok().flatten() { - DaemonError::AlreadyRunning(pid) - } else { - DaemonError::PidFile(format!("cannot lock {}: {}", self.path.display(), errno)) - } - })?; - - // Write our PID - let pid = getpid(); - let mut file = flock - .unlock() - .map_err(|(_, errno)| DaemonError::PidFile(format!("unlock failed: {}", errno)))?; - - writeln!(file, "{}", pid).map_err(|e| { - DaemonError::PidFile(format!( - "cannot write PID to {}: {}", - self.path.display(), - e - )) - })?; - - // Re-acquire lock and keep it - let flock = Flock::lock(file, FlockArg::LockExclusiveNonblock).map_err(|(_, errno)| { - DaemonError::PidFile(format!("cannot re-lock {}: {}", self.path.display(), errno)) - })?; - - self.file = Some(flock.unlock().map_err(|(_, errno)| { - DaemonError::PidFile(format!("unlock for storage failed: {}", errno)) - })?); - self.locked = true; - - info!(pid = pid.as_raw(), path = %self.path.display(), "PID file created"); - Ok(()) - } - - /// Releases the PID file lock and removes the file. - pub fn release(&mut self) -> Result<(), DaemonError> { - if let Some(file) = self.file.take() { - drop(file); - } - self.locked = false; - - if self.path.exists() { - fs::remove_file(&self.path).map_err(|e| { - DaemonError::PidFile(format!("cannot remove {}: {}", self.path.display(), e)) - })?; - debug!(path = %self.path.display(), "PID file removed"); - } - - Ok(()) - } - - /// Returns the path to this PID file. - #[allow(dead_code)] - pub fn path(&self) -> &Path { - &self.path - } -} - -impl Drop for PidFile { - fn drop(&mut self) { - if self.locked { - if let Err(e) = self.release() { - warn!(error = %e, "Failed to clean up PID file on drop"); - } - } - } -} - -/// Checks if a process with the given PID is running. -fn is_process_running(pid: i32) -> bool { - // kill(pid, 0) checks if process exists without sending a signal - nix::sys::signal::kill(Pid::from_raw(pid), None).is_ok() -} - // macOS gates nix::unistd::setgroups differently in the current dependency set, // so call libc directly there while preserving the original nix path elsewhere. fn set_supplementary_groups(gid: Gid) -> Result<(), nix::Error> { @@ -383,9 +226,11 @@ pub fn drop_privileges( }; if (target_uid.is_some() || target_gid.is_some()) - && let Some(file) = pid_file.and_then(|pid| pid.file.as_ref()) + && let Some(pid_file) = pid_file { - unistd::fchown(file, target_uid, target_gid).map_err(DaemonError::PrivilegeDrop)?; + for file in pid_file.ownership_file_handles().into_iter().flatten() { + unistd::fchown(file, target_uid, target_gid).map_err(DaemonError::PrivilegeDrop)?; + } } if let Some(gid) = target_gid { @@ -401,7 +246,7 @@ pub fn drop_privileges( if uid.as_raw() != 0 && let Some(pid) = pid_file { - let parent = pid.path.parent().unwrap_or(Path::new(".")); + let parent = pid.path().parent().unwrap_or(Path::new(".")); let probe_path = parent.join(format!( ".telemt_pid_probe_{}_{}", std::process::id(), @@ -436,7 +281,6 @@ pub fn drop_privileges( /// Looks up a user by name and returns their UID. fn lookup_user(name: &str) -> Result { - // Use libc getpwnam let c_name = std::ffi::CString::new(name).map_err(|_| DaemonError::UserNotFound(name.to_string()))?; @@ -480,76 +324,6 @@ fn lookup_group(name: &str) -> Result { } } -/// Reads PID from a PID file. -#[allow(dead_code)] -pub fn read_pid_file>(path: P) -> Result { - 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>( - 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>(path: P) -> DaemonStatus { - let path = path.as_ref(); - - if !path.exists() { - return DaemonStatus::NotRunning; - } - - match read_pid_file(path) { - Ok(pid) => { - if is_process_running(pid) { - DaemonStatus::Running(pid) - } else { - DaemonStatus::Stale(pid) - } - } - Err(_) => DaemonStatus::NotRunning, - } -} - #[cfg(test)] mod tests { use super::*; @@ -571,29 +345,4 @@ mod tests { }; assert!(!opts.should_daemonize()); } - - #[test] - fn test_check_status_not_running() { - let path = "/tmp/telemt_test_nonexistent.pid"; - assert_eq!(check_status(path), DaemonStatus::NotRunning); - } - - #[test] - fn test_pid_file_basic() { - let path = "/tmp/telemt_test_pidfile.pid"; - let _ = fs::remove_file(path); - - let mut pf = PidFile::new(path); - assert!(pf.check_running().unwrap().is_none()); - - pf.acquire().unwrap(); - assert!(Path::new(path).exists()); - - // Read it back - let pid = read_pid_file(path).unwrap(); - assert_eq!(pid, std::process::id() as i32); - - pf.release().unwrap(); - assert!(!Path::new(path).exists()); - } } diff --git a/src/daemon/pid_file.rs b/src/daemon/pid_file.rs new file mode 100644 index 0000000..7f3e95b --- /dev/null +++ b/src/daemon/pid_file.rs @@ -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, + lock_file: Option>, +} + +impl PidFile { + /// Creates a new PID file manager for the given path. + pub fn new>(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, 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, 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>(path: P) -> Result { + 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>( + 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>(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 { + 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() {} + + assert_send_sync::(); + } + + #[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(); + } +} diff --git a/src/maestro/runtime_tasks.rs b/src/maestro/runtime_tasks.rs index cc2a124..c8a302a 100644 --- a/src/maestro/runtime_tasks.rs +++ b/src/maestro/runtime_tasks.rs @@ -288,8 +288,9 @@ pub(crate) async fn spawn_runtime_tasks( break; } let cfg = config_rx_user_enabled.borrow_and_update().clone(); - for user in shared_user_enabled.apply_user_enabled_config(&cfg.access.user_enabled) { - let cancelled = shared_user_enabled.cancel_user_sessions(&user); + for (user, cancelled) in + shared_user_enabled.apply_user_enabled_config(&cfg.access.user_enabled) + { if cancelled > 0 { info!( user = %user, diff --git a/src/proxy/authenticated.rs b/src/proxy/authenticated.rs index 60315fb..77e0c9d 100644 --- a/src/proxy/authenticated.rs +++ b/src/proxy/authenticated.rs @@ -79,7 +79,11 @@ where let route_snapshot = deps.route_runtime.snapshot(); let session_id = deps.rng.u64(); - let user_session = deps.shared.register_user_session(&user, session_id); + let Some(user_session) = deps.shared.register_user_session(&user, session_id) else { + user_reservation.release_deferred(); + warn!(user = %user, "Disabled user rejected during final admission"); + return Err(ProxyError::UserDisabled { user }); + }; let session_cancel = user_session.token(); let selected_me_pool = if deps.config.general.use_middle_proxy && matches!(route_snapshot.mode, RelayRouteMode::Middle) @@ -246,6 +250,18 @@ impl UserConnectionReservation { } self.stats.decrement_user_curr_connects(&self.user); } + + /// Defers IP cleanup when admission fails after the asynchronous reservation step. + pub(crate) fn release_deferred(mut self) { + if !self.active { + return; + } + self.active = false; + self.stats.decrement_user_curr_connects(&self.user); + if self.tracks_ip { + self.ip_tracker.enqueue_cleanup(self.user.clone(), self.ip); + } + } } impl Drop for UserConnectionReservation { diff --git a/src/proxy/relay/io.rs b/src/proxy/relay/io.rs index fb30f3f..e2d5f78 100644 --- a/src/proxy/relay/io.rs +++ b/src/proxy/relay/io.rs @@ -16,10 +16,7 @@ mod quota; pub(super) use self::combined::CombinedStream; pub(super) use self::counters::SharedCounters; pub(super) use self::quota::is_quota_io_error; -use self::quota::{ - QUOTA_RESERVE_MAX_ROUNDS, QUOTA_RESERVE_SPIN_RETRIES, quota_io_error, - refund_reserved_quota_bytes, -}; +use self::quota::{QUOTA_RESERVE_MAX_ROUNDS, QUOTA_RESERVE_SPIN_RETRIES, quota_io_error}; pub(super) use self::quota::{quota_adaptive_interval_bytes, should_immediate_quota_check}; /// Transparent I/O wrapper that tracks per-user statistics and activity. @@ -213,7 +210,7 @@ impl AsyncRead for StatsIo { } let mut remaining_before = None; - let mut reserved_read_bytes = 0u64; + let mut quota_reservation = None; let mut read_limit = buf.remaining(); if let Some(limit) = this.quota_limit { let used_before = this.user_stats.quota_used(); @@ -231,11 +228,11 @@ impl AsyncRead for StatsIo { let desired = read_limit as u64; let mut reserve_rounds = 0usize; - while reserved_read_bytes == 0 { + while quota_reservation.is_none() { for _ in 0..QUOTA_RESERVE_SPIN_RETRIES { - match this.user_stats.quota_try_reserve(desired, limit) { - Ok(_) => { - reserved_read_bytes = desired; + match this.user_stats.quota_reserve(desired, limit) { + Ok(reservation) => { + quota_reservation = Some(reservation); break; } Err(crate::stats::QuotaReserveError::LimitExceeded) => { @@ -248,7 +245,7 @@ impl AsyncRead for StatsIo { } } - if reserved_read_bytes == 0 { + if quota_reservation.is_none() { reserve_rounds = reserve_rounds.saturating_add(1); if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS { this.stats.increment_quota_contention_timeout_total(); @@ -287,9 +284,9 @@ impl AsyncRead for StatsIo { match read_result { Poll::Ready(Ok(n)) => { - if reserved_read_bytes > n as u64 { - let refund_bytes = reserved_read_bytes - n as u64; - refund_reserved_quota_bytes(this.user_stats.as_ref(), refund_bytes); + if let Some(reservation) = quota_reservation.take() { + let refund_bytes = reservation.reserved_bytes().saturating_sub(n as u64); + reservation.settle(n as u64); this.stats.add_quota_refund_bytes_total(refund_bytes); } if n > 0 { @@ -333,16 +330,16 @@ impl AsyncRead for StatsIo { Poll::Ready(Ok(())) } Poll::Pending => { - if reserved_read_bytes > 0 { - refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_read_bytes); - this.stats.add_quota_refund_bytes_total(reserved_read_bytes); + if let Some(reservation) = quota_reservation.take() { + this.stats + .add_quota_refund_bytes_total(reservation.reserved_bytes()); } Poll::Pending } Poll::Ready(Err(err)) => { - if reserved_read_bytes > 0 { - refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_read_bytes); - this.stats.add_quota_refund_bytes_total(reserved_read_bytes); + if let Some(reservation) = quota_reservation.take() { + this.stats + .add_quota_refund_bytes_total(reservation.reserved_bytes()); } Poll::Ready(Err(err)) } @@ -361,14 +358,15 @@ impl AsyncWrite for StatsIo { return Poll::Ready(Err(quota_io_error())); } - let mut shaper_reserved_bytes = 0u64; + let mut shaper_reservation = None; let mut write_buf = buf; if let Some(lease) = this.traffic_lease.as_ref() { if !buf.is_empty() { loop { - let consume = lease.try_consume(RateDirection::Down, buf.len() as u64); + let reservation = lease.try_reserve(RateDirection::Down, buf.len() as u64); + let consume = reservation.result(); if consume.granted > 0 { - shaper_reserved_bytes = consume.granted; + shaper_reservation = Some(reservation); if consume.granted < buf.len() as u64 { write_buf = &buf[..consume.granted as usize]; } @@ -398,17 +396,14 @@ impl AsyncWrite for StatsIo { } let mut remaining_before = None; - let mut reserved_bytes = 0u64; + let mut quota_reservation = None; if let Some(limit) = this.quota_limit { if !write_buf.is_empty() { let mut reserve_rounds = 0usize; - while reserved_bytes == 0 { + while quota_reservation.is_none() { let used_before = this.user_stats.quota_used(); let remaining = limit.saturating_sub(used_before); if remaining == 0 { - if let Some(lease) = this.traffic_lease.as_ref() { - lease.refund(RateDirection::Down, shaper_reserved_bytes); - } this.quota_exceeded.store(true, Ordering::Release); return Poll::Ready(Err(quota_io_error())); } @@ -417,9 +412,9 @@ impl AsyncWrite for StatsIo { let desired = remaining.min(write_buf.len() as u64); let mut saw_contention = false; for _ in 0..QUOTA_RESERVE_SPIN_RETRIES { - match this.user_stats.quota_try_reserve(desired, limit) { - Ok(_) => { - reserved_bytes = desired; + match this.user_stats.quota_reserve(desired, limit) { + Ok(reservation) => { + quota_reservation = Some(reservation); write_buf = &write_buf[..desired as usize]; break; } @@ -433,14 +428,13 @@ impl AsyncWrite for StatsIo { } } - if reserved_bytes == 0 { + if quota_reservation.is_none() { reserve_rounds = reserve_rounds.saturating_add(1); if reserve_rounds >= QUOTA_RESERVE_MAX_ROUNDS { this.stats.increment_quota_contention_timeout_total(); - if let Some(lease) = this.traffic_lease.as_ref() { - lease.refund(RateDirection::Down, shaper_reserved_bytes); - } - let _ = this.arm_quota_wait(cx); + Self::arm_wait(&mut this.quota_wait, false, false); + let _ = + Self::poll_wait(&mut this.quota_wait, cx, None, RateDirection::Up); return Poll::Pending; } else if saw_contention { std::hint::spin_loop(); @@ -451,9 +445,6 @@ impl AsyncWrite for StatsIo { let used_before = this.user_stats.quota_used(); let remaining = limit.saturating_sub(used_before); if remaining == 0 { - if let Some(lease) = this.traffic_lease.as_ref() { - lease.refund(RateDirection::Down, shaper_reserved_bytes); - } this.quota_exceeded.store(true, Ordering::Release); return Poll::Ready(Err(quota_io_error())); } @@ -463,15 +454,13 @@ impl AsyncWrite for StatsIo { match Pin::new(&mut this.inner).poll_write(cx, write_buf) { Poll::Ready(Ok(n)) => { - if reserved_bytes > n as u64 { - let refund_bytes = reserved_bytes - n as u64; - refund_reserved_quota_bytes(this.user_stats.as_ref(), refund_bytes); + if let Some(reservation) = quota_reservation.take() { + let refund_bytes = reservation.reserved_bytes().saturating_sub(n as u64); + reservation.settle(n as u64); this.stats.add_quota_refund_bytes_total(refund_bytes); } - if shaper_reserved_bytes > n as u64 - && let Some(lease) = this.traffic_lease.as_ref() - { - lease.refund(RateDirection::Down, shaper_reserved_bytes - n as u64); + if let Some(reservation) = shaper_reservation.take() { + reservation.settle_written(n as u64); } if n > 0 { if let Some(lease) = this.traffic_lease.as_ref() { @@ -513,26 +502,16 @@ impl AsyncWrite for StatsIo { Poll::Ready(Ok(n)) } Poll::Ready(Err(err)) => { - if reserved_bytes > 0 { - refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_bytes); - this.stats.add_quota_refund_bytes_total(reserved_bytes); - } - if shaper_reserved_bytes > 0 - && let Some(lease) = this.traffic_lease.as_ref() - { - lease.refund(RateDirection::Down, shaper_reserved_bytes); + if let Some(reservation) = quota_reservation.take() { + this.stats + .add_quota_refund_bytes_total(reservation.reserved_bytes()); } Poll::Ready(Err(err)) } Poll::Pending => { - if reserved_bytes > 0 { - refund_reserved_quota_bytes(this.user_stats.as_ref(), reserved_bytes); - this.stats.add_quota_refund_bytes_total(reserved_bytes); - } - if shaper_reserved_bytes > 0 - && let Some(lease) = this.traffic_lease.as_ref() - { - lease.refund(RateDirection::Down, shaper_reserved_bytes); + if let Some(reservation) = quota_reservation.take() { + this.stats + .add_quota_refund_bytes_total(reservation.reserved_bytes()); } Poll::Pending } diff --git a/src/proxy/relay/io/quota.rs b/src/proxy/relay/io/quota.rs index 82e83f5..b5f87e0 100644 --- a/src/proxy/relay/io/quota.rs +++ b/src/proxy/relay/io/quota.rs @@ -1,4 +1,3 @@ -use crate::stats::UserStats; use std::io; #[derive(Debug)] @@ -46,10 +45,3 @@ pub(in crate::proxy::relay) fn should_immediate_quota_check( ) -> bool { remaining_before <= QUOTA_NEAR_LIMIT_BYTES || charge_bytes >= QUOTA_LARGE_CHARGE_BYTES } - -pub(super) fn refund_reserved_quota_bytes(user_stats: &UserStats, reserved_bytes: u64) { - if reserved_bytes == 0 { - return; - } - user_stats.refund_quota(reserved_bytes); -} diff --git a/src/proxy/shared_state.rs b/src/proxy/shared_state.rs index 6a8774c..9623b0c 100644 --- a/src/proxy/shared_state.rs +++ b/src/proxy/shared_state.rs @@ -6,6 +6,7 @@ use std::sync::{Arc, Mutex}; use std::time::Instant; use dashmap::DashMap; +use parking_lot::Mutex as ParkingMutex; use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc}; use tokio_util::sync::CancellationToken; @@ -75,13 +76,18 @@ pub(crate) struct MiddleRelaySharedState { pub(crate) relay_idle_mark_seq: AtomicU64, } +#[derive(Default)] +struct UserAdmissionState { + disabled_users: HashSet, + sessions_by_user: HashMap>, +} + pub(crate) struct ProxySharedState { pub(crate) handshake: HandshakeSharedState, pub(crate) middle_relay: MiddleRelaySharedState, pub(crate) traffic_limiter: Arc, pub(crate) direct_buffer_budget: Arc, - disabled_users: DashMap, - active_user_sessions: DashMap<(String, u64), CancellationToken>, + user_admission: ParkingMutex, pub(crate) conntrack_pressure_active: AtomicBool, pub(crate) conntrack_close_tx: Mutex>>, masking_fallback_permits: Arc, @@ -106,7 +112,18 @@ struct UserSessionGuard { impl Drop for UserSessionGuard { fn drop(&mut self) { - self.shared.active_user_sessions.remove(&self.key); + let mut admission = self.shared.user_admission.lock(); + let remove_user = admission + .sessions_by_user + .get_mut(&self.key.0) + .map(|sessions| { + sessions.remove(&self.key.1); + sessions.is_empty() + }) + .unwrap_or(false); + if remove_user { + admission.sessions_by_user.remove(&self.key.0); + } } } @@ -150,8 +167,7 @@ impl ProxySharedState { }, traffic_limiter: TrafficLimiter::new(), direct_buffer_budget, - disabled_users: DashMap::new(), - active_user_sessions: DashMap::new(), + user_admission: ParkingMutex::new(UserAdmissionState::default()), conntrack_pressure_active: AtomicBool::new(false), conntrack_close_tx: Mutex::new(None), masking_fallback_permits: Arc::new(Semaphore::new(MASKING_FALLBACK_MAX_CONCURRENT)), @@ -167,68 +183,102 @@ impl ProxySharedState { } pub(crate) fn is_user_enabled(&self, user: &str) -> bool { - !self.disabled_users.contains_key(user) + !self.user_admission.lock().disabled_users.contains(user) } - pub(crate) fn set_user_enabled(&self, user: &str, enabled: bool) -> bool { - if enabled { - self.disabled_users.remove(user); - false - } else { - self.disabled_users.insert(user.to_string(), ()).is_none() + pub(crate) fn set_user_enabled(&self, user: &str, enabled: bool) -> (bool, usize) { + let (newly_disabled, tokens) = { + let mut admission = self.user_admission.lock(); + if enabled { + admission.disabled_users.remove(user); + (false, Vec::new()) + } else { + let newly_disabled = admission.disabled_users.insert(user.to_string()); + let tokens = admission + .sessions_by_user + .get(user) + .map(|sessions| sessions.values().cloned().collect()) + .unwrap_or_default(); + (newly_disabled, tokens) + } + }; + for token in &tokens { + token.cancel(); } + (newly_disabled, tokens.len()) } pub(crate) fn apply_user_enabled_config( &self, user_enabled: &HashMap, - ) -> Vec { + ) -> Vec<(String, usize)> { let desired_disabled = user_enabled .iter() .filter_map(|(user, enabled)| (!*enabled).then_some(user.clone())) .collect::>(); - let current_disabled = self - .disabled_users - .iter() - .map(|entry| entry.key().clone()) - .collect::>(); - - for user in current_disabled.difference(&desired_disabled) { - self.disabled_users.remove(user); - } - let newly_disabled = desired_disabled - .difference(¤t_disabled) - .cloned() - .collect::>(); - for user in desired_disabled { - self.disabled_users.insert(user, ()); - } - newly_disabled + let cancellations = { + let mut admission = self.user_admission.lock(); + let newly_disabled = desired_disabled + .difference(&admission.disabled_users) + .cloned() + .collect::>(); + admission.disabled_users = desired_disabled; + newly_disabled + .into_iter() + .map(|user| { + let tokens = admission + .sessions_by_user + .get(&user) + .map(|sessions| sessions.values().cloned().collect()) + .unwrap_or_default(); + (user, tokens) + }) + .collect::)>>() + }; + cancellations + .into_iter() + .map(|(user, tokens)| { + for token in &tokens { + token.cancel(); + } + (user, tokens.len()) + }) + .collect() } pub(crate) fn register_user_session( self: &Arc, user: &str, session_id: u64, - ) -> UserSessionRegistration { + ) -> Option { let token = CancellationToken::new(); let key = (user.to_string(), session_id); - self.active_user_sessions.insert(key.clone(), token.clone()); - UserSessionRegistration { + let mut admission = self.user_admission.lock(); + if admission.disabled_users.contains(user) { + return None; + } + admission + .sessions_by_user + .entry(key.0.clone()) + .or_default() + .insert(session_id, token.clone()); + Some(UserSessionRegistration { token, _guard: UserSessionGuard { shared: Arc::clone(self), key, }, - } + }) } pub(crate) fn cancel_user_sessions(&self, user: &str) -> usize { - let tokens = self - .active_user_sessions - .iter() - .filter_map(|entry| (entry.key().0 == user).then(|| entry.value().clone())) - .collect::>(); + let tokens: Vec = self + .user_admission + .lock() + .sessions_by_user + .get(user) + .map(|sessions| sessions.values().cloned().collect()) + .unwrap_or_default(); for token in &tokens { token.cancel(); } @@ -311,7 +361,7 @@ mod tests { let mut newly_disabled = shared.apply_user_enabled_config(&user_enabled); newly_disabled.sort(); - assert_eq!(newly_disabled, vec!["alice".to_string()]); + assert_eq!(newly_disabled, vec![("alice".to_string(), 0)]); assert!(!shared.is_user_enabled("alice")); assert!(shared.is_user_enabled("bob")); @@ -325,9 +375,9 @@ mod tests { #[test] fn cancel_user_sessions_cancels_only_registered_matching_user() { let shared = ProxySharedState::new(); - let alice_1 = shared.register_user_session("alice", 1); - let alice_2 = shared.register_user_session("alice", 2); - let bob = shared.register_user_session("bob", 1); + let alice_1 = shared.register_user_session("alice", 1).unwrap(); + let alice_2 = shared.register_user_session("alice", 2).unwrap(); + let bob = shared.register_user_session("bob", 1).unwrap(); let alice_1_token = alice_1.token(); let alice_2_token = alice_2.token(); let bob_token = bob.token(); @@ -339,4 +389,61 @@ mod tests { assert!(alice_2_token.is_cancelled()); assert!(!bob_token.is_cancelled()); } + + #[test] + fn disabled_user_cannot_register_after_the_cancellation_snapshot() { + let shared = ProxySharedState::new(); + + assert_eq!(shared.set_user_enabled("alice", false), (true, 0)); + assert_eq!(shared.cancel_user_sessions("alice"), 0); + + let late = shared.register_user_session("alice", 1); + assert!( + late.is_none(), + "a session registered after disable returned must be rejected" + ); + } + + #[test] + fn disabling_user_cancels_existing_sessions_before_return() { + let shared = ProxySharedState::new(); + let registration = shared.register_user_session("alice", 1).unwrap(); + let token = registration.token(); + + assert_eq!(shared.set_user_enabled("alice", false), (true, 1)); + assert!(token.is_cancelled()); + assert!(shared.register_user_session("alice", 2).is_none()); + + assert_eq!(shared.set_user_enabled("alice", true), (false, 0)); + assert!(shared.register_user_session("alice", 3).is_some()); + } + + #[test] + fn concurrent_disable_and_registration_never_leave_a_live_session() { + const ITERATIONS: usize = 10_000; + + let shared = ProxySharedState::new(); + let barrier = Arc::new(std::sync::Barrier::new(2)); + let register_shared = Arc::clone(&shared); + let register_barrier = Arc::clone(&barrier); + let register = std::thread::spawn(move || { + let mut registrations = Vec::with_capacity(ITERATIONS); + for session_id in 0..ITERATIONS as u64 { + let user = format!("user-{session_id}"); + register_barrier.wait(); + registrations.push(register_shared.register_user_session(&user, session_id)); + } + registrations + }); + + for session_id in 0..ITERATIONS as u64 { + let user = format!("user-{session_id}"); + barrier.wait(); + shared.set_user_enabled(&user, false); + } + + for registration in register.join().unwrap().into_iter().flatten() { + assert!(registration.token().is_cancelled()); + } + } } diff --git a/src/proxy/traffic_limiter.rs b/src/proxy/traffic_limiter.rs index a7a6068..ce454de 100644 --- a/src/proxy/traffic_limiter.rs +++ b/src/proxy/traffic_limiter.rs @@ -6,6 +6,7 @@ use std::sync::atomic::{AtomicU64, Ordering}; use arc_swap::ArcSwap; use dashmap::DashMap; use ipnetwork::IpNetwork; +use parking_lot::Mutex as ParkingMutex; use crate::config::RateLimitBps; @@ -32,6 +33,9 @@ const REGISTRY_SHARDS: usize = 64; const FAIR_EPOCH_MS: u64 = 20; const MAX_BORROW_CHUNK_BYTES: u64 = 32 * 1024; const CLEANUP_INTERVAL_SECS: u64 = 60; +const PACKED_USAGE_BITS: u32 = 28; +const PACKED_USAGE_MASK: u64 = (1u64 << PACKED_USAGE_BITS) - 1; +const PACKED_EPOCH_MAX: u64 = (1u64 << (u64::BITS - PACKED_USAGE_BITS)) - 1; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum RateDirection { @@ -76,12 +80,12 @@ struct ScopeMetrics { struct AtomicRatePair { up_bps: AtomicU64, down_bps: AtomicU64, + revision: ParkingMutex, } #[derive(Default)] struct DirectionBucket { - epoch: AtomicU64, - used: AtomicU64, + state: AtomicU64, } struct UserBucket { @@ -93,15 +97,13 @@ struct UserBucket { #[derive(Default)] struct CidrDirectionBucket { - epoch: AtomicU64, - used: AtomicU64, - active_users: AtomicU64, + used: DirectionBucket, + active_users: DirectionBucket, } #[derive(Default)] struct CidrUserDirectionState { - epoch: AtomicU64, - used: AtomicU64, + used: DirectionBucket, } struct CidrUserShare { @@ -139,6 +141,7 @@ enum CidrPolicyMatch<'a> { #[derive(Default)] struct PolicySnapshot { + revision: u64, user_limits: HashMap, cidr_rules_v4: Vec, cidr_rules_v6: Vec, @@ -162,9 +165,25 @@ pub struct TrafficLease { pub struct TrafficLimiter { policy: ArcSwap, + policy_update: ParkingMutex<()>, user_buckets: ShardedRegistry, cidr_buckets: ShardedRegistry, user_scope: ScopeMetrics, cidr_scope: ScopeMetrics, last_cleanup_epoch_secs: AtomicU64, } + +struct DirectionDebit<'a> { + bucket: &'a DirectionBucket, + epoch: u64, + refundable: u64, +} + +/// Refunds uncommitted shaping budget when an I/O attempt is cancelled. +#[must_use = "traffic reservations must be settled after the I/O attempt"] +pub(crate) struct TrafficReservation<'a> { + result: TrafficConsumeResult, + user: Option>, + cidr: Option>, + cidr_user: Option>, +} diff --git a/src/proxy/traffic_limiter/buckets.rs b/src/proxy/traffic_limiter/buckets.rs index f6fe592..748ba53 100644 --- a/src/proxy/traffic_limiter/buckets.rs +++ b/src/proxy/traffic_limiter/buckets.rs @@ -26,203 +26,277 @@ impl ScopeMetrics { } impl AtomicRatePair { - pub(super) fn set(&self, limits: RateLimitBps) { - self.up_bps.store(limits.up_bps, Ordering::Relaxed); - self.down_bps.store(limits.down_bps, Ordering::Relaxed); + pub(super) fn new(revision: u64, limits: RateLimitBps) -> Self { + let rates = Self::default(); + rates.set(revision, limits); + rates + } + + pub(super) fn set(&self, revision: u64, limits: RateLimitBps) { + let mut current_revision = self.revision.lock(); + if revision < *current_revision { + return; + } + self.up_bps.store(limits.up_bps, Ordering::Release); + self.down_bps.store(limits.down_bps, Ordering::Release); + *current_revision = revision; } pub(super) fn get(&self, direction: RateDirection) -> u64 { match direction { - RateDirection::Up => self.up_bps.load(Ordering::Relaxed), - RateDirection::Down => self.down_bps.load(Ordering::Relaxed), + RateDirection::Up => self.up_bps.load(Ordering::Acquire), + RateDirection::Down => self.down_bps.load(Ordering::Acquire), } } } impl DirectionBucket { - pub(super) fn sync_epoch(&self, epoch: u64) { - let current = self.epoch.load(Ordering::Relaxed); - if current == epoch { - return; - } - if current < epoch - && self - .epoch - .compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed) - .is_ok() - { - self.used.store(0, Ordering::Relaxed); - } + fn unpack(state: u64) -> (u64, u64) { + (state >> PACKED_USAGE_BITS, state & PACKED_USAGE_MASK) } - pub(super) fn try_consume(&self, cap_bps: u64, requested: u64) -> u64 { - if requested == 0 { - return 0; + fn pack(epoch: u64, used: u64) -> Option { + if epoch > PACKED_EPOCH_MAX || used > PACKED_USAGE_MASK { + return None; } - if cap_bps == 0 { - return requested; + Some((epoch << PACKED_USAGE_BITS) | used) + } + + pub(super) fn used_at(&self, epoch: u64) -> Option { + if epoch > PACKED_EPOCH_MAX { + return None; } + let (current_epoch, used) = Self::unpack(self.state.load(Ordering::Relaxed)); + (current_epoch == epoch).then_some(used) + } - let epoch = current_epoch(); - self.sync_epoch(epoch); - let cap_epoch = bytes_per_epoch(cap_bps); + pub(super) fn try_reserve_at( + &self, + epoch: u64, + cap: u64, + requested: u64, + ) -> Option> { + if requested == 0 || cap == 0 || epoch > PACKED_EPOCH_MAX { + return None; + } + let cap = cap.min(PACKED_USAGE_MASK); + let mut observed = self.state.load(Ordering::Relaxed); loop { - let used = self.used.load(Ordering::Relaxed); - if used >= cap_epoch { - return 0; + let (observed_epoch, observed_used) = Self::unpack(observed); + if observed_epoch > epoch { + return None; } - let remaining = cap_epoch.saturating_sub(used); + let used = if observed_epoch == epoch { + observed_used + } else { + 0 + }; + if used >= cap { + return None; + } + let remaining = cap - used; let grant = requested.min(remaining); if grant == 0 { - return 0; + return None; } - let next = used.saturating_add(grant); - if self - .used - .compare_exchange_weak(used, next, Ordering::Relaxed, Ordering::Relaxed) - .is_ok() - { - return grant; + let next = Self::pack(epoch, used + grant)?; + match self.state.compare_exchange_weak( + observed, + next, + Ordering::Relaxed, + Ordering::Relaxed, + ) { + Ok(_) => { + return Some(DirectionDebit { + bucket: self, + epoch, + refundable: grant, + }); + } + Err(actual) => observed = actual, } } } - pub(super) fn refund(&self, bytes: u64) { - if bytes == 0 { + fn refund_at(&self, epoch: u64, bytes: u64) { + if bytes == 0 || epoch > PACKED_EPOCH_MAX { return; } - decrement_atomic_saturating(&self.used, bytes); + + let mut observed = self.state.load(Ordering::Relaxed); + loop { + let (observed_epoch, used) = Self::unpack(observed); + if observed_epoch != epoch || used == 0 { + return; + } + let next = Self::pack(epoch, used.saturating_sub(bytes)).unwrap_or(observed); + match self.state.compare_exchange_weak( + observed, + next, + Ordering::Relaxed, + Ordering::Relaxed, + ) { + Ok(_) => return, + Err(actual) => observed = actual, + } + } + } +} + +impl DirectionDebit<'_> { + fn granted(&self) -> u64 { + self.refundable + } + + pub(super) fn shrink_to(&mut self, retained: u64) { + let retained = retained.min(self.refundable); + self.bucket + .refund_at(self.epoch, self.refundable - retained); + self.refundable = retained; + } + + pub(super) fn settle(&mut self, committed: u64) { + self.shrink_to(committed); + self.refundable = 0; + } + + pub(super) fn commit_all(&mut self) -> u64 { + let committed = self.refundable; + self.refundable = 0; + committed + } +} + +impl Drop for DirectionDebit<'_> { + fn drop(&mut self) { + self.bucket.refund_at(self.epoch, self.refundable); } } impl UserBucket { - pub(super) fn new(limits: RateLimitBps) -> Self { - let rates = AtomicRatePair::default(); - rates.set(limits); + pub(super) fn new(revision: u64, limits: RateLimitBps) -> Self { Self { - rates, + rates: AtomicRatePair::new(revision, limits), up: DirectionBucket::default(), down: DirectionBucket::default(), active_leases: AtomicU64::new(0), } } - pub(super) fn set_rates(&self, limits: RateLimitBps) { - self.rates.set(limits); + pub(super) fn set_rates(&self, revision: u64, limits: RateLimitBps) { + self.rates.set(revision, limits); } - pub(super) fn try_consume(&self, direction: RateDirection, requested: u64) -> u64 { + pub(super) fn try_reserve( + &self, + direction: RateDirection, + requested: u64, + ) -> (u64, Option>) { let cap_bps = self.rates.get(direction); - match direction { - RateDirection::Up => self.up.try_consume(cap_bps, requested), - RateDirection::Down => self.down.try_consume(cap_bps, requested), - } - } - - pub(super) fn refund(&self, direction: RateDirection, bytes: u64) { - match direction { - RateDirection::Up => self.up.refund(bytes), - RateDirection::Down => self.down.refund(bytes), + if cap_bps == 0 { + return (requested, None); } + let cap = bytes_per_epoch(cap_bps); + let debit = match direction { + RateDirection::Up => self.up.try_reserve_at(current_epoch(), cap, requested), + RateDirection::Down => self.down.try_reserve_at(current_epoch(), cap, requested), + }; + let granted = debit.as_ref().map(DirectionDebit::granted).unwrap_or(0); + (granted, debit) } } impl CidrDirectionBucket { - pub(super) fn sync_epoch(&self, epoch: u64) { - let current = self.epoch.load(Ordering::Relaxed); - if current == epoch { - return; - } - if current < epoch - && self - .epoch - .compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed) - .is_ok() - { - self.used.store(0, Ordering::Relaxed); - self.active_users.store(0, Ordering::Relaxed); - } - } - - pub(super) fn try_consume( - &self, - user_state: &CidrUserDirectionState, + pub(super) fn try_reserve<'a>( + &'a self, + user_state: &'a CidrUserDirectionState, cap_epoch: u64, requested: u64, - ) -> u64 { + ) -> (u64, Option>, Option>) { if requested == 0 || cap_epoch == 0 { - return 0; + return (0, None, None); } let epoch = current_epoch(); - self.sync_epoch(epoch); - user_state.sync_epoch_and_mark_active(epoch, &self.active_users); - let active_users = self.active_users.load(Ordering::Relaxed).max(1); + if !user_state.ensure_active(epoch, &self.active_users) { + return (0, None, None); + } + let Some(active_users) = self.active_users.used_at(epoch) else { + return (0, None, None); + }; + let active_users = active_users.max(1); let fair_share = cap_epoch.saturating_div(active_users).max(1); loop { - let total_used = self.used.load(Ordering::Relaxed); - if total_used >= cap_epoch { - return 0; - } - let total_remaining = cap_epoch.saturating_sub(total_used); - let user_used = user_state.used.load(Ordering::Relaxed); - let guaranteed_remaining = fair_share.saturating_sub(user_used); - - let grant = if guaranteed_remaining > 0 { - requested.min(guaranteed_remaining).min(total_remaining) - } else { - requested.min(total_remaining).min(MAX_BORROW_CHUNK_BYTES) + let Some(user_used) = user_state.used.used_at(epoch) else { + return (0, None, None); }; - - if grant == 0 { - return 0; - } - - let next_total = total_used.saturating_add(grant); - if self - .used - .compare_exchange_weak(total_used, next_total, Ordering::Relaxed, Ordering::Relaxed) - .is_ok() - { - user_state.used.fetch_add(grant, Ordering::Relaxed); - return grant; + let guaranteed_remaining = fair_share.saturating_sub(user_used); + let (user_cap, desired) = if guaranteed_remaining > 0 { + (fair_share, requested.min(guaranteed_remaining)) + } else { + (PACKED_USAGE_MASK, requested.min(MAX_BORROW_CHUNK_BYTES)) + }; + let Some(mut user_debit) = user_state.used.try_reserve_at(epoch, user_cap, desired) + else { + if guaranteed_remaining > 0 { + continue; + } + return (0, None, None); + }; + let user_granted = user_debit.granted(); + let Some(aggregate_debit) = self.used.try_reserve_at(epoch, cap_epoch, user_granted) + else { + return (0, None, None); + }; + let granted = aggregate_debit.granted(); + if granted < user_granted { + user_debit.shrink_to(granted); } + return (granted, Some(aggregate_debit), Some(user_debit)); } } - - pub(super) fn refund(&self, bytes: u64) { - if bytes == 0 { - return; - } - decrement_atomic_saturating(&self.used, bytes); - } } impl CidrUserDirectionState { - pub(super) fn sync_epoch_and_mark_active(&self, epoch: u64, active_users: &AtomicU64) { - let current = self.epoch.load(Ordering::Relaxed); - if current == epoch { - return; + pub(super) fn ensure_active(&self, epoch: u64, active_users: &DirectionBucket) -> bool { + if epoch > PACKED_EPOCH_MAX { + return false; } - if current < epoch - && self - .epoch - .compare_exchange(current, epoch, Ordering::Relaxed, Ordering::Relaxed) - .is_ok() - { - self.used.store(0, Ordering::Relaxed); - active_users.fetch_add(1, Ordering::Relaxed); + let mut observed = self.used.state.load(Ordering::Relaxed); + loop { + let (observed_epoch, _) = DirectionBucket::unpack(observed); + if observed_epoch == epoch { + return true; + } + if observed_epoch > epoch { + return false; + } + let Some(mut active_debit) = active_users.try_reserve_at(epoch, PACKED_USAGE_MASK, 1) + else { + return false; + }; + let Some(next) = DirectionBucket::pack(epoch, 0) else { + return false; + }; + match self.used.state.compare_exchange( + observed, + next, + Ordering::Relaxed, + Ordering::Relaxed, + ) { + Ok(_) => { + active_debit.commit_all(); + return true; + } + Err(actual) => { + drop(active_debit); + observed = actual; + } + } } } - - pub(super) fn refund(&self, bytes: u64) { - if bytes == 0 { - return; - } - decrement_atomic_saturating(&self.used, bytes); - } } impl CidrUserShare { @@ -236,11 +310,9 @@ impl CidrUserShare { } impl CidrBucket { - pub(super) fn new(limits: RateLimitBps) -> Self { - let rates = AtomicRatePair::default(); - rates.set(limits); + pub(super) fn new(revision: u64, limits: RateLimitBps) -> Self { Self { - rates, + rates: AtomicRatePair::new(revision, limits), up: CidrDirectionBucket::default(), down: CidrDirectionBucket::default(), users: ShardedRegistry::new(REGISTRY_SHARDS), @@ -248,8 +320,8 @@ impl CidrBucket { } } - pub(super) fn set_rates(&self, limits: RateLimitBps) { - self.rates.set(limits); + pub(super) fn set_rates(&self, revision: u64, limits: RateLimitBps) { + self.rates.set(revision, limits); } pub(super) fn acquire_user_share(&self, user: &str) -> Arc { @@ -268,38 +340,20 @@ impl CidrBucket { }); } - pub(super) fn try_consume_for_user( - &self, + pub(super) fn try_reserve_for_user<'a>( + &'a self, direction: RateDirection, - share: &CidrUserShare, + share: &'a CidrUserShare, requested: u64, - ) -> u64 { + ) -> (u64, Option>, Option>) { let cap_bps = self.rates.get(direction); if cap_bps == 0 { - return requested; + return (requested, None, None); } let cap_epoch = bytes_per_epoch(cap_bps); match direction { - RateDirection::Up => self.up.try_consume(&share.up, cap_epoch, requested), - RateDirection::Down => self.down.try_consume(&share.down, cap_epoch, requested), - } - } - - pub(super) fn refund_for_user( - &self, - direction: RateDirection, - share: &CidrUserShare, - bytes: u64, - ) { - match direction { - RateDirection::Up => { - self.up.refund(bytes); - share.up.refund(bytes); - } - RateDirection::Down => { - self.down.refund(bytes); - share.down.refund(bytes); - } + RateDirection::Up => self.up.try_reserve(&share.up, cap_epoch, requested), + RateDirection::Down => self.down.try_reserve(&share.down, cap_epoch, requested), } } diff --git a/src/proxy/traffic_limiter/helpers.rs b/src/proxy/traffic_limiter/helpers.rs index 272fe03..a25eada 100644 --- a/src/proxy/traffic_limiter/helpers.rs +++ b/src/proxy/traffic_limiter/helpers.rs @@ -56,7 +56,7 @@ pub(super) fn auto_cidr_bucket_key(ip: IpAddr, prefix_len: u8) -> Option pub(super) fn current_epoch() -> u64 { let start = limiter_epoch_start(); let elapsed_ms = start.elapsed().as_millis() as u64; - elapsed_ms / FAIR_EPOCH_MS + elapsed_ms / FAIR_EPOCH_MS + 1 } pub(super) fn limiter_epoch_start() -> &'static Instant { diff --git a/src/proxy/traffic_limiter/lease.rs b/src/proxy/traffic_limiter/lease.rs index ad136af..99f3d26 100644 --- a/src/proxy/traffic_limiter/lease.rs +++ b/src/proxy/traffic_limiter/lease.rs @@ -1,70 +1,93 @@ use super::*; impl TrafficLease { - pub fn try_consume(&self, direction: RateDirection, requested: u64) -> TrafficConsumeResult { + /// Reserves shaping budget until the associated I/O result is settled. + pub(crate) fn try_reserve( + &self, + direction: RateDirection, + requested: u64, + ) -> TrafficReservation<'_> { if requested == 0 { - return TrafficConsumeResult { - granted: 0, - blocked_user: false, - blocked_cidr: false, + return TrafficReservation { + result: TrafficConsumeResult { + granted: 0, + blocked_user: false, + blocked_cidr: false, + }, + user: None, + cidr: None, + cidr_user: None, }; } let mut granted = requested; + let mut user_debit = None; if let Some(user_bucket) = self.user_bucket.as_ref() { - let user_granted = user_bucket.try_consume(direction, granted); + let (user_granted, debit) = user_bucket.try_reserve(direction, granted); + user_debit = debit; if user_granted == 0 { self.limiter.observe_throttle(direction, true, false); - return TrafficConsumeResult { - granted: 0, - blocked_user: true, - blocked_cidr: false, + return TrafficReservation { + result: TrafficConsumeResult { + granted: 0, + blocked_user: true, + blocked_cidr: false, + }, + user: user_debit, + cidr: None, + cidr_user: None, }; } granted = user_granted; } + let mut cidr_debit = None; + let mut cidr_user_debit = None; if let (Some(cidr_bucket), Some(cidr_user_share)) = (self.cidr_bucket.as_ref(), self.cidr_user_share.as_ref()) { - let cidr_granted = - cidr_bucket.try_consume_for_user(direction, cidr_user_share, granted); + let (cidr_granted, aggregate_debit, share_debit) = + cidr_bucket.try_reserve_for_user(direction, cidr_user_share, granted); + cidr_debit = aggregate_debit; + cidr_user_debit = share_debit; if cidr_granted < granted - && let Some(user_bucket) = self.user_bucket.as_ref() + && let Some(debit) = user_debit.as_mut() { - user_bucket.refund(direction, granted.saturating_sub(cidr_granted)); + debit.shrink_to(cidr_granted); } if cidr_granted == 0 { self.limiter.observe_throttle(direction, false, true); - return TrafficConsumeResult { - granted: 0, - blocked_user: false, - blocked_cidr: true, + return TrafficReservation { + result: TrafficConsumeResult { + granted: 0, + blocked_user: false, + blocked_cidr: true, + }, + user: user_debit, + cidr: cidr_debit, + cidr_user: cidr_user_debit, }; } granted = cidr_granted; } - TrafficConsumeResult { - granted, - blocked_user: false, - blocked_cidr: false, + TrafficReservation { + result: TrafficConsumeResult { + granted, + blocked_user: false, + blocked_cidr: false, + }, + user: user_debit, + cidr: cidr_debit, + cidr_user: cidr_user_debit, } } - pub fn refund(&self, direction: RateDirection, bytes: u64) { - if bytes == 0 { - return; - } - - if let Some(user_bucket) = self.user_bucket.as_ref() { - user_bucket.refund(direction, bytes); - } - if let (Some(cidr_bucket), Some(cidr_user_share)) = - (self.cidr_bucket.as_ref(), self.cidr_user_share.as_ref()) - { - cidr_bucket.refund_for_user(direction, cidr_user_share, bytes); - } + pub fn try_consume(&self, direction: RateDirection, requested: u64) -> TrafficConsumeResult { + let reservation = self.try_reserve(direction, requested); + let result = reservation.result(); + reservation.settle_written(result.granted); + result } pub fn observe_wait_ms( @@ -82,6 +105,27 @@ impl TrafficLease { } } +impl TrafficReservation<'_> { + /// Returns the shaping decision associated with this reservation. + pub(crate) fn result(&self) -> TrafficConsumeResult { + self.result + } + + /// Commits written bytes and refunds the uncommitted remainder. + pub(crate) fn settle_written(mut self, committed: u64) { + let committed = committed.min(self.result.granted); + if let Some(debit) = self.user.as_mut() { + debit.settle(committed); + } + if let Some(debit) = self.cidr.as_mut() { + debit.settle(committed); + } + if let Some(debit) = self.cidr_user.as_mut() { + debit.settle(committed); + } + } +} + impl Drop for TrafficLease { fn drop(&mut self) { if let Some(bucket) = self.user_bucket.as_ref() { diff --git a/src/proxy/traffic_limiter/limiter.rs b/src/proxy/traffic_limiter/limiter.rs index 4aaf411..e52df40 100644 --- a/src/proxy/traffic_limiter/limiter.rs +++ b/src/proxy/traffic_limiter/limiter.rs @@ -5,6 +5,7 @@ impl TrafficLimiter { pub fn new() -> Arc { Arc::new(Self { policy: ArcSwap::from_pointee(PolicySnapshot::default()), + policy_update: ParkingMutex::new(()), user_buckets: ShardedRegistry::new(REGISTRY_SHARDS), cidr_buckets: ShardedRegistry::new(REGISTRY_SHARDS), user_scope: ScopeMetrics::default(), @@ -18,6 +19,11 @@ impl TrafficLimiter { user_limits: HashMap, cidr_limits: HashMap, ) { + let policy_update = self.policy_update.lock(); + // Revision wrap could otherwise let an old lease restore stale rates. + let Some(revision) = self.policy.load().revision.checked_add(1) else { + return; + }; let filtered_users = user_limits .into_iter() .filter(|(_, limit)| limit.up_bps > 0 || limit.down_bps > 0) @@ -78,6 +84,7 @@ impl TrafficLimiter { .store(cidr_policy_entries as u64, Ordering::Relaxed); self.policy.store(Arc::new(PolicySnapshot { + revision, user_limits: filtered_users, cidr_rules_v4, cidr_rules_v6, @@ -86,6 +93,7 @@ impl TrafficLimiter { cidr_rule_keys, })); + drop(policy_update); self.maybe_cleanup(); } @@ -99,12 +107,12 @@ impl TrafficLimiter { if let Some(limit) = policy.user_limits.get(user).copied() { let bucket = self.user_buckets.get_or_insert_with( user, - || UserBucket::new(limit), + || UserBucket::new(policy.revision, limit), |bucket| { bucket.active_leases.fetch_add(1, Ordering::Relaxed); }, ); - bucket.set_rates(limit); + bucket.set_rates(policy.revision, limit); self.user_scope .active_leases .fetch_add(1, Ordering::Relaxed); @@ -121,12 +129,12 @@ impl TrafficLimiter { }; let bucket = self.cidr_buckets.get_or_insert_with( key, - || CidrBucket::new(limits), + || CidrBucket::new(policy.revision, limits), |bucket| { bucket.active_leases.fetch_add(1, Ordering::Relaxed); }, ); - bucket.set_rates(limits); + bucket.set_rates(policy.revision, limits); self.cidr_scope .active_leases .fetch_add(1, Ordering::Relaxed); diff --git a/src/proxy/traffic_limiter/tests.rs b/src/proxy/traffic_limiter/tests.rs index b9147da..8bfe6e3 100644 --- a/src/proxy/traffic_limiter/tests.rs +++ b/src/proxy/traffic_limiter/tests.rs @@ -74,3 +74,185 @@ fn auto_cidr_bucket_key_canonicalizes_network_address() { "auto:6:2001:db8::/64" ); } + +#[test] +fn refund_from_an_old_epoch_does_not_reduce_the_current_epoch() { + let bucket = DirectionBucket::default(); + let old_debit = bucket.try_reserve_at(7, 100, 80).unwrap(); + let current_debit = bucket.try_reserve_at(8, 100, 60).unwrap(); + + drop(old_debit); + + assert_eq!(bucket.used_at(8), Some(60)); + drop(current_debit); +} + +#[test] +fn concurrent_rollover_cannot_publish_multiple_epoch_budgets() { + const CONTENDERS: usize = 32; + + let bucket = Arc::new(DirectionBucket::default()); + let barrier = Arc::new(std::sync::Barrier::new(CONTENDERS)); + let mut threads = Vec::with_capacity(CONTENDERS); + for _ in 0..CONTENDERS { + let bucket = Arc::clone(&bucket); + let barrier = Arc::clone(&barrier); + threads.push(std::thread::spawn(move || { + barrier.wait(); + bucket + .try_reserve_at(9, 100, 100) + .map(|mut debit| debit.commit_all()) + .unwrap_or(0) + })); + } + + let granted = threads + .into_iter() + .map(|thread| thread.join().unwrap()) + .sum::(); + 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::>(); + for epoch_index in 0..EPOCHS { + let granted = grants + .iter() + .map(|thread_grants| thread_grants[epoch_index]) + .sum::(); + 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) + ); +} diff --git a/src/stats/mod.rs b/src/stats/mod.rs index 170ad08..90a1c8f 100644 --- a/src/stats/mod.rs +++ b/src/stats/mod.rs @@ -21,7 +21,7 @@ use std::sync::Arc; use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, AtomicUsize, Ordering}; use std::time::Instant; -pub(crate) use self::quota_store::QuotaStore; +pub(crate) use self::quota_store::{QuotaReservation, QuotaStore}; #[allow(unused_imports)] pub use self::replay::{ReplayChecker, ReplayStats}; use self::telemetry::TelemetryPolicy; @@ -392,11 +392,6 @@ impl UserStats { self.quota.used() } - #[inline] - pub(crate) fn refund_quota(&self, bytes: u64) { - self.quota.refund(bytes); - } - /// Attempts one CAS reservation step against the quota counter. /// /// Callers control retry/yield policy. This primitive intentionally does @@ -404,6 +399,18 @@ impl UserStats { /// with their own contention strategy. #[inline] pub fn quota_try_reserve(&self, bytes: u64, limit: u64) -> Result { + 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 { self.quota.try_reserve(bytes, limit) } } diff --git a/src/stats/quota_store.rs b/src/stats/quota_store.rs index e338b76..5911fa1 100644 --- a/src/stats/quota_store.rs +++ b/src/stats/quota_store.rs @@ -2,6 +2,7 @@ use std::collections::HashMap; use std::sync::Arc; use std::sync::atomic::{AtomicU64, Ordering}; +use arc_swap::ArcSwap; use dashmap::DashMap; use super::{QuotaReserveError, UserQuotaSnapshot}; @@ -12,11 +13,22 @@ pub struct QuotaStore { users: DashMap>, } -/// Atomic quota state for one configured user. -#[derive(Default)] +/// Atomically replaceable quota state for one configured user. pub(crate) struct UserQuotaCounters { + generation: ArcSwap, +} + +struct QuotaGeneration { used_bytes: AtomicU64, - last_reset_epoch_secs: AtomicU64, + last_reset_epoch_secs: u64, +} + +/// Owns a quota debit until the corresponding I/O outcome is known. +#[must_use = "quota reservations must be committed or settled"] +pub(crate) struct QuotaReservation { + generation: Arc, + reserved_bytes: u64, + total_after_reserve: u64, } impl QuotaStore { @@ -38,18 +50,12 @@ impl QuotaStore { pub(crate) fn load(&self, user: &str, used_bytes: u64, last_reset_epoch_secs: u64) { let state = self.user(user); - state.used_bytes.store(used_bytes, Ordering::Relaxed); - state - .last_reset_epoch_secs - .store(last_reset_epoch_secs, Ordering::Relaxed); + state.replace(used_bytes, last_reset_epoch_secs); } pub(crate) fn reset(&self, user: &str, now_epoch_secs: u64) -> UserQuotaSnapshot { let state = self.user(user); - state.used_bytes.store(0, Ordering::Relaxed); - state - .last_reset_epoch_secs - .store(now_epoch_secs, Ordering::Relaxed); + state.replace(0, now_epoch_secs); UserQuotaSnapshot { used_bytes: 0, last_reset_epoch_secs: now_epoch_secs, @@ -64,8 +70,9 @@ impl QuotaStore { let mut out = HashMap::new(); for entry in self.users.iter() { let state = entry.value(); - let used_bytes = state.used(); - let last_reset_epoch_secs = state.last_reset_epoch_secs.load(Ordering::Relaxed); + let generation = state.generation.load_full(); + let used_bytes = generation.used_bytes.load(Ordering::Relaxed); + let last_reset_epoch_secs = generation.last_reset_epoch_secs; if used_bytes == 0 && last_reset_epoch_secs == 0 { continue; } @@ -81,60 +88,121 @@ impl QuotaStore { } } +impl Default for UserQuotaCounters { + fn default() -> Self { + Self { + generation: ArcSwap::from_pointee(QuotaGeneration { + used_bytes: AtomicU64::new(0), + last_reset_epoch_secs: 0, + }), + } + } +} + impl UserQuotaCounters { + fn replace(&self, used_bytes: u64, last_reset_epoch_secs: u64) { + self.generation.store(Arc::new(QuotaGeneration { + used_bytes: AtomicU64::new(used_bytes), + last_reset_epoch_secs, + })); + } + #[inline] pub(crate) fn used(&self) -> u64 { - self.used_bytes.load(Ordering::Relaxed) + self.generation.load().used_bytes.load(Ordering::Relaxed) } #[inline] pub(crate) fn charge(&self, bytes: u64) -> u64 { - self.used_bytes + self.generation + .load_full() + .used_bytes .fetch_add(bytes, Ordering::Relaxed) .saturating_add(bytes) } #[inline] - pub(crate) fn refund(&self, bytes: u64) { - let mut current = self.used_bytes.load(Ordering::Relaxed); - loop { - let next = current.saturating_sub(bytes); - match self.used_bytes.compare_exchange_weak( - current, - next, - Ordering::Relaxed, - Ordering::Relaxed, - ) { - Ok(_) => return, - Err(observed) => current = observed, - } - } - } - - #[inline] - pub(crate) fn try_reserve(&self, bytes: u64, limit: u64) -> Result { - let current = self.used_bytes.load(Ordering::Relaxed); + pub(crate) fn try_reserve( + &self, + bytes: u64, + limit: u64, + ) -> Result { + let generation = self.generation.load_full(); + let current = generation.used_bytes.load(Ordering::Relaxed); if bytes > limit.saturating_sub(current) { return Err(QuotaReserveError::LimitExceeded); } let next = current.saturating_add(bytes); - match self.used_bytes.compare_exchange_weak( + match generation.used_bytes.compare_exchange_weak( current, next, Ordering::Relaxed, Ordering::Relaxed, ) { - Ok(_) => Ok(next), + Ok(_) => Ok(QuotaReservation { + generation, + reserved_bytes: bytes, + total_after_reserve: next, + }), Err(_) => Err(QuotaReserveError::Contended), } } } +impl QuotaReservation { + /// Returns the number of bytes held by this reservation. + pub(crate) fn reserved_bytes(&self) -> u64 { + self.reserved_bytes + } + + /// Commits the complete reservation and returns the generation-local total. + pub(crate) fn commit(mut self) -> u64 { + self.reserved_bytes = 0; + self.total_after_reserve + } + + /// Commits part of the reservation and refunds the remainder. + pub(crate) fn settle(mut self, committed_bytes: u64) { + let committed_bytes = committed_bytes.min(self.reserved_bytes); + refund_generation( + self.generation.as_ref(), + self.reserved_bytes - committed_bytes, + ); + self.reserved_bytes = 0; + } +} + +impl Drop for QuotaReservation { + fn drop(&mut self) { + refund_generation(self.generation.as_ref(), self.reserved_bytes); + } +} + +fn refund_generation(generation: &QuotaGeneration, bytes: u64) { + if bytes == 0 { + return; + } + let mut current = generation.used_bytes.load(Ordering::Relaxed); + loop { + let next = current.saturating_sub(bytes); + match generation.used_bytes.compare_exchange_weak( + current, + next, + Ordering::Relaxed, + Ordering::Relaxed, + ) { + Ok(_) => return, + Err(observed) => current = observed, + } + } +} + #[cfg(test)] mod tests { use super::*; use crate::stats::Stats; + use std::sync::Barrier; #[test] fn quota_counters_are_shared_across_stats_generations() { @@ -148,4 +216,43 @@ mod tests { second.reset_user_quota("alice"); assert_eq!(first.get_user_quota_used("alice"), 0); } + + #[test] + fn quota_snapshot_never_combines_different_generations() { + const ITERATIONS: u64 = 10_000; + + let store = Arc::new(QuotaStore::default()); + store.load("alice", 1, 1); + let barrier = Arc::new(Barrier::new(2)); + let writer_store = Arc::clone(&store); + let writer_barrier = Arc::clone(&barrier); + let writer = std::thread::spawn(move || { + writer_barrier.wait(); + for generation in 2..=ITERATIONS { + writer_store.load("alice", generation, generation); + } + }); + + barrier.wait(); + for _ in 0..ITERATIONS { + let snapshot = store.snapshot().remove("alice").unwrap(); + assert_eq!(snapshot.used_bytes, snapshot.last_reset_epoch_secs); + } + writer.join().unwrap(); + } + + #[test] + fn repeated_old_generation_refunds_leave_new_usage_intact() { + const ITERATIONS: u64 = 10_000; + + let store = QuotaStore::default(); + let state = store.user("alice"); + for generation in 1..=ITERATIONS { + let reservation = state.try_reserve(80, 100).unwrap(); + store.load("alice", 40, generation); + drop(reservation); + assert_eq!(store.used("alice"), 40); + store.reset("alice", generation); + } + } } diff --git a/src/stats/tests.rs b/src/stats/tests.rs index 650b123..f90e601 100644 --- a/src/stats/tests.rs +++ b/src/stats/tests.rs @@ -289,6 +289,20 @@ fn test_quota_used_is_authoritative_and_independent_from_octets_telemetry() { assert_eq!(stats.get_user_quota_used(user), 7); } +#[test] +fn old_quota_reservation_refund_does_not_reduce_post_reset_usage() { + let stats = Stats::new(); + let user = "quota-reset-generation-user"; + let user_stats = stats.get_or_create_user_stats_handle(user); + let reservation = user_stats.quota_reserve(80, 100).unwrap(); + + stats.reset_user_quota(user); + stats.quota_charge_post_write(user_stats.as_ref(), 40); + drop(reservation); + + assert_eq!(stats.get_user_quota_used(user), 40); +} + #[test] fn test_cached_handle_survives_map_cleanup_until_last_drop() { let stats = Stats::new(); diff --git a/src/transport/middle_proxy/mod.rs b/src/transport/middle_proxy/mod.rs index 2c38fb5..8adc45c 100644 --- a/src/transport/middle_proxy/mod.rs +++ b/src/transport/middle_proxy/mod.rs @@ -33,6 +33,9 @@ mod pool_runtime_api; mod pool_status; mod pool_writer; #[cfg(test)] +#[path = "tests/pool_writer_publication_tests.rs"] +mod pool_writer_publication_tests; +#[cfg(test)] #[path = "tests/pool_writer_security_tests.rs"] mod pool_writer_security_tests; mod reader; diff --git a/src/transport/middle_proxy/pool.rs b/src/transport/middle_proxy/pool.rs index 232b556..cd4451a 100644 --- a/src/transport/middle_proxy/pool.rs +++ b/src/transport/middle_proxy/pool.rs @@ -88,14 +88,6 @@ impl WritersState { } } - pub(super) async fn update(&self, f: F) -> R - where - F: FnOnce(&mut Vec) -> R, - { - let mut guard = self.write().await; - f(&mut guard) - } - fn debug_assert_store_guarded(&self) { debug_assert!( self.writers_write_guard.try_lock().is_err(), diff --git a/src/transport/middle_proxy/pool_writer.rs b/src/transport/middle_proxy/pool_writer.rs index 6841a29..f116225 100644 --- a/src/transport/middle_proxy/pool_writer.rs +++ b/src/transport/middle_proxy/pool_writer.rs @@ -1,4 +1,5 @@ use std::collections::HashMap; +use std::future::Future; use std::io::ErrorKind; use std::net::SocketAddr; use std::sync::Arc; @@ -19,6 +20,7 @@ use crate::protocol::constants::{RPC_CLOSE_EXT_U32, RPC_PING_U32}; use super::codec::{RpcWriter, WriterCommand, build_control_payload}; use super::pool::{MePool, MeWriter, WriterContour}; +use super::pool_lifecycle::MeTaskRegistration; use super::reader::reader_loop; use super::wire::build_proxy_req_payload; diff --git a/src/transport/middle_proxy/pool_writer/runtime.rs b/src/transport/middle_proxy/pool_writer/runtime.rs index d59218d..427a197 100644 --- a/src/transport/middle_proxy/pool_writer/runtime.rs +++ b/src/transport/middle_proxy/pool_writer/runtime.rs @@ -132,16 +132,6 @@ impl MePool { drain_deadline_epoch_secs: drain_deadline_epoch_secs.clone(), allow_drain_fallback: allow_drain_fallback.clone(), }; - self.writers - .update(|writers| writers.push(writer.clone())) - .await; - self.registry - .register_writer(writer_id, tx.clone(), byte_budget) - .await; - self.registry.mark_writer_idle(writer_id).await; - self.conn_count.fetch_add(1, Ordering::Relaxed); - self.notify_writer_epoch(); - let reg = self.registry.clone(); let writers_arc = self.writers_arc(); let ping_tracker = Arc::new(tokio::sync::Mutex::new(HashMap::::new())); @@ -177,8 +167,10 @@ impl MePool { let route_fairshare_enabled = self.transport_policy.me_route_fairshare_enabled.clone(); let reader_route_data_wait_ms = self.transport_policy.me_reader_route_data_wait_ms.clone(); - self.lifecycle - .spawn_registered_writer(task_registration, async move { + let writer_task = { + // Keep transport ownership behind a stable-size pointer while publication waits for + // locks. + Box::pin(async move { // Reader MUST be the first branch in biased select! to avoid read starvation. let exit = tokio::select! { biased; @@ -269,11 +261,41 @@ impl MePool { let remaining = writers_arc.read().await.len(); debug!(writer_id, remaining, "ME writer lifecycle task finished"); - }); + }) + }; + + self.publish_prepared_writer(writer, tx, byte_budget, task_registration, writer_task) + .await; Ok(()) } + /// Commits writer visibility and lifecycle ownership after all cancellation points. + #[allow(clippy::too_many_arguments)] + pub(in crate::transport::middle_proxy) async fn publish_prepared_writer( + self: &Arc, + writer: MeWriter, + tx: mpsc::Sender, + byte_budget: Arc, + task_registration: MeTaskRegistration<'_>, + writer_task: F, + ) where + F: Future + 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, writer_id: u64) { // Full client cleanup now happens inside `registry.writer_lost` to keep // writer reap/remove paths strictly non-blocking per connection. @@ -337,7 +359,12 @@ impl MePool { self.stats.increment_me_writer_removed_unexpected_total(); } close_tx = Some(w.tx.clone()); - self.conn_count.fetch_sub(1, Ordering::Relaxed); + // Teardown remains idempotent if the advisory count was already reconciled. + let _ = + self.conn_count + .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |count| { + count.checked_sub(1) + }); removed = true; } } diff --git a/src/transport/middle_proxy/registry.rs b/src/transport/middle_proxy/registry.rs index e97a233..62ca021 100644 --- a/src/transport/middle_proxy/registry.rs +++ b/src/transport/middle_proxy/registry.rs @@ -18,6 +18,8 @@ const ROUTE_QUEUED_BYTE_PERMIT_UNIT: usize = 16 * 1024; const ROUTE_QUEUED_PERMITS_PER_SLOT: usize = 4; const ROUTE_QUEUED_MAX_FRAME_PERMITS: usize = 1024; +// Transactional writer registry publication. +mod publication; mod writer; #[derive(Debug, Clone, Copy, PartialEq, Eq)] diff --git a/src/transport/middle_proxy/registry/publication.rs b/src/transport/middle_proxy/registry/publication.rs new file mode 100644 index 0000000..8e29a55 --- /dev/null +++ b/src/transport/middle_proxy/registry/publication.rs @@ -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, + byte_budget: Arc, + ) { + 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 }); + } +} diff --git a/src/transport/middle_proxy/registry/writer.rs b/src/transport/middle_proxy/registry/writer.rs index a3d375f..ca0a2e7 100644 --- a/src/transport/middle_proxy/registry/writer.rs +++ b/src/transport/middle_proxy/registry/writer.rs @@ -58,28 +58,16 @@ impl ConnRegistry { } /// Registers one writer command route and its matching memory budget atomically. + #[allow(dead_code)] pub async fn register_writer( &self, writer_id: u64, tx: mpsc::Sender, byte_budget: Arc, ) { - let mut binding = self.binding.inner.lock().await; - binding - .conns_for_writer - .entry(writer_id) - .or_insert_with(HashSet::new); - self.binding - .bound_clients_by_writer - .entry(writer_id) - .or_insert(0); - self.binding - .writer_idle_since_epoch_secs - .entry(writer_id) - .or_insert_with(Self::now_epoch_secs); - self.writers - .map - .insert(writer_id, super::WriterRoute { tx, byte_budget }); + self.prepare_writer_registration() + .await + .install(writer_id, tx, byte_budget); } /// Unregister connection, returning associated writer_id if any. @@ -346,20 +334,6 @@ impl ConnRegistry { true } - pub async fn mark_writer_idle(&self, writer_id: u64) { - let mut binding = self.binding.inner.lock().await; - binding - .conns_for_writer - .entry(writer_id) - .or_insert_with(HashSet::new); - let count = binding - .conns_for_writer - .get(&writer_id) - .map(|set| set.len()) - .unwrap_or(0); - self.set_writer_bound_count(writer_id, count); - } - pub async fn get_last_writer_meta(&self, writer_id: u64) -> Option { self.binding .last_meta_for_writer diff --git a/src/transport/middle_proxy/tests/pool_writer_publication_tests.rs b/src/transport/middle_proxy/tests/pool_writer_publication_tests.rs new file mode 100644 index 0000000..0eb8f01 --- /dev/null +++ b/src/transport/middle_proxy/tests/pool_writer_publication_tests.rs @@ -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::(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::(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); + } +} diff --git a/src/transport/middle_proxy/tests/pool_writer_security_tests.rs b/src/transport/middle_proxy/tests/pool_writer_security_tests.rs index 1d57257..f468b47 100644 --- a/src/transport/middle_proxy/tests/pool_writer_security_tests.rs +++ b/src/transport/middle_proxy/tests/pool_writer_security_tests.rs @@ -15,7 +15,8 @@ use crate::crypto::SecureRandom; use crate::network::probe::NetworkDecision; use crate::stats::Stats; -async fn make_pool() -> Arc { +/// Builds an isolated ME pool for writer-state tests. +pub(super) async fn make_pool() -> Arc { let general = GeneralConfig::default(); MePool::new( @@ -154,6 +155,70 @@ async fn insert_writer( pool.conn_count.fetch_add(1, Ordering::Relaxed); } +#[tokio::test] +async fn cancelled_writer_publication_leaves_no_ghost_state() { + let pool = make_pool().await; + let writer_id = 76_001; + let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); + let (tx, mut rx) = mpsc::channel::(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) -> HashSet { pool.writers .read()