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