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