Redesign runtime w/ include-aware config + Module Split + Listener Lifecycle + Atomic Reload

This commit is contained in:
Alexey
2026-08-22 16:13:25 +03:00
parent f9910fb29e
commit 7f4b87bea4
72 changed files with 14433 additions and 12976 deletions
+384
View File
@@ -0,0 +1,384 @@
use std::path::PathBuf;
use std::sync::Arc;
use std::time::{Instant, SystemTime, UNIX_EPOCH};
use tracing::{info, warn};
use tracing_subscriber::{EnvFilter, fmt, prelude::*, reload as tracing_reload};
use crate::config::{LogLevel, ProxyConfig};
use crate::startup::{
COMPONENT_CONFIG_LOAD, COMPONENT_TRACING_INIT, StartupTracker,
};
use super::helpers::{
parse_cli, print_maestro_line, resolve_runtime_base_dir, resolve_runtime_config_path,
set_maestro_colors_enabled,
};
use super::runtime_tasks;
use super::validate_synlimit_privilege_drop;
pub(super) struct BootstrapState {
pub(super) process_started_at: Instant,
pub(super) process_started_at_epoch_secs: u64,
pub(super) startup_tracker: Arc<StartupTracker>,
pub(super) config: ProxyConfig,
pub(super) config_path: PathBuf,
pub(super) has_rust_log: bool,
pub(super) effective_log_level: LogLevel,
pub(super) runtime_log_filter: runtime_tasks::RuntimeLogFilter,
pub(super) logging_guard: Option<crate::logging::LoggingGuard>,
}
pub(super) async fn bootstrap(
privilege_drop_requested: bool,
) -> std::result::Result<BootstrapState, Box<dyn std::error::Error>> {
let process_started_at = Instant::now();
let process_started_at_epoch_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let startup_tracker = Arc::new(StartupTracker::new(process_started_at_epoch_secs));
startup_tracker
.start_component(
COMPONENT_CONFIG_LOAD,
Some("load and validate config".to_string()),
)
.await;
let cli_args = parse_cli();
let config_path_cli = cli_args.config_path;
let config_path_explicit = cli_args.config_path_explicit;
let data_path = cli_args.data_path;
let cli_silent = cli_args.silent;
let cli_log_level = cli_args.log_level;
let log_cli_options = cli_args.log_cli_options;
let startup_cwd = match std::env::current_dir() {
Ok(cwd) => cwd,
Err(e) => {
eprintln!("[telemt] Can't read current_dir: {}", e);
std::process::exit(1);
}
};
if let Some(ref data_path) = data_path
&& !data_path.is_absolute()
{
eprintln!(
"[telemt] data_path must be absolute: {}",
data_path.display()
);
std::process::exit(1);
}
let mut config_path =
resolve_runtime_config_path(&config_path_cli, &startup_cwd, config_path_explicit);
let runtime_base_dir = resolve_runtime_base_dir(
&config_path,
&startup_cwd,
config_path_explicit,
data_path.as_deref(),
);
if !runtime_base_dir.exists()
&& let Err(e) = std::fs::create_dir_all(&runtime_base_dir)
{
eprintln!(
"[telemt] Can't create runtime directory {}: {}",
runtime_base_dir.display(),
e
);
std::process::exit(1);
}
if !runtime_base_dir.is_dir() {
eprintln!(
"[telemt] Runtime path exists but is not a directory: {}",
runtime_base_dir.display()
);
std::process::exit(1);
}
if let Err(e) = std::env::set_current_dir(&runtime_base_dir) {
eprintln!(
"[telemt] Can't use runtime directory {}: {}",
runtime_base_dir.display(),
e
);
std::process::exit(1);
}
let mut config = match ProxyConfig::load(&config_path) {
Ok(c) => c,
Err(e) => {
if config_path.exists() {
eprintln!("[telemt] Error: {}", e);
std::process::exit(1);
} else {
let default = ProxyConfig::default();
let serialized =
match toml::to_string_pretty(&default).or_else(|_| toml::to_string(&default)) {
Ok(value) => Some(value),
Err(serialize_error) => {
eprintln!(
"[telemt] Warning: failed to serialize default config: {}",
serialize_error
);
None
}
};
if config_path_explicit {
if let Some(serialized) = serialized.as_ref() {
if let Err(write_error) = std::fs::write(&config_path, serialized) {
eprintln!(
"[telemt] Error: failed to create explicit config at {}: {}",
config_path.display(),
write_error
);
std::process::exit(1);
}
eprintln!(
"[telemt] Created default config at {}",
config_path.display()
);
} else {
eprintln!(
"[telemt] Warning: running with in-memory default config without writing to disk"
);
}
} else {
let runtime_config_path = runtime_base_dir.join("telemt.toml");
let fallback_config_path = runtime_base_dir.join("config.toml");
let mut persisted = false;
if let Some(serialized) = serialized.as_ref() {
match std::fs::create_dir_all(&runtime_base_dir) {
Ok(()) => match std::fs::write(&runtime_config_path, serialized) {
Ok(()) => {
config_path = runtime_config_path;
eprintln!(
"[telemt] Created default config at {}",
config_path.display()
);
persisted = true;
}
Err(write_error) => {
eprintln!(
"[telemt] Warning: failed to write default config at {}: {}",
runtime_config_path.display(),
write_error
);
}
},
Err(create_error) => {
eprintln!(
"[telemt] Warning: failed to create {}: {}",
runtime_base_dir.display(),
create_error
);
}
}
if !persisted {
match std::fs::write(&fallback_config_path, serialized) {
Ok(()) => {
config_path = fallback_config_path;
eprintln!(
"[telemt] Created default config at {}",
config_path.display()
);
persisted = true;
}
Err(write_error) => {
eprintln!(
"[telemt] Warning: failed to write default config at {}: {}",
fallback_config_path.display(),
write_error
);
}
}
}
}
if !persisted {
eprintln!(
"[telemt] Warning: running with in-memory default config without writing to disk"
);
}
}
default
}
}
};
if let Err(e) = config.validate() {
eprintln!("[telemt] Invalid config: {}", e);
std::process::exit(1);
}
validate_synlimit_privilege_drop(&config, privilege_drop_requested)?;
if let Some(p) = data_path {
config.general.data_path = Some(p);
}
if let Some(ref data_path) = config.general.data_path {
if !data_path.is_absolute() {
eprintln!(
"[telemt] data_path must be absolute: {}",
data_path.display()
);
std::process::exit(1);
}
if data_path.exists() {
if !data_path.is_dir() {
eprintln!(
"[telemt] data_path exists but is not a directory: {}",
data_path.display()
);
std::process::exit(1);
}
} else if let Err(e) = std::fs::create_dir_all(data_path) {
eprintln!(
"[telemt] Can't create data_path {}: {}",
data_path.display(),
e
);
std::process::exit(1);
}
if let Err(e) = std::env::set_current_dir(data_path) {
eprintln!(
"[telemt] Can't use data_path {}: {}",
data_path.display(),
e
);
std::process::exit(1);
}
}
if let Err(e) = crate::network::dns_overrides::install_entries(&config.network.dns_overrides) {
eprintln!("[telemt] Invalid network.dns_overrides: {}", e);
std::process::exit(1);
}
set_maestro_colors_enabled(!config.general.disable_colors);
startup_tracker
.complete_component(COMPONENT_CONFIG_LOAD, Some("config is ready".to_string()))
.await;
let has_rust_log = std::env::var("RUST_LOG").is_ok();
let effective_log_level = if cli_silent {
LogLevel::Silent
} else if let Some(ref s) = cli_log_level {
LogLevel::from_str_loose(s)
} else {
config.general.log_level.clone()
};
let initial_filter_spec = runtime_tasks::log_filter_spec(has_rust_log, &effective_log_level);
let log_destination =
match crate::logging::resolve_log_destination(&config.logging, &log_cli_options) {
Ok(destination) => destination,
Err(error) => {
eprintln!("[telemt] {error}");
std::process::exit(1);
}
};
let (filter_layer, filter_handle) =
tracing_reload::Layer::new(EnvFilter::new(initial_filter_spec.clone()));
startup_tracker
.start_component(
COMPONENT_TRACING_INIT,
Some("initialize tracing subscriber".to_string()),
)
.await;
let logging_guard: Option<crate::logging::LoggingGuard>;
match log_destination {
crate::logging::LogDestination::Stderr => {
let fmt_layer = if config.general.disable_colors {
fmt::Layer::default().with_ansi(false)
} else {
fmt::Layer::default().with_ansi(true)
};
tracing_subscriber::registry()
.with(filter_layer)
.with(fmt_layer)
.init();
logging_guard = None;
}
#[cfg(unix)]
crate::logging::LogDestination::Syslog => {
let logging_opts = crate::logging::LoggingOptions {
destination: log_destination,
disable_colors: true,
};
let (_, guard) = crate::logging::init_logging(&logging_opts, &initial_filter_spec);
logging_guard = Some(guard);
}
crate::logging::LogDestination::File { .. } => {
let logging_opts = crate::logging::LoggingOptions {
destination: log_destination,
disable_colors: true,
};
let (_, guard) = crate::logging::init_logging(&logging_opts, &initial_filter_spec);
logging_guard = Some(guard);
}
}
let runtime_log_filter = runtime_tasks::RuntimeLogFilter::new(filter_handle);
startup_tracker
.complete_component(
COMPONENT_TRACING_INIT,
Some("tracing initialized".to_string()),
)
.await;
print_maestro_line(format!("Telemt MTProxy v{}", env!("CARGO_PKG_VERSION")));
info!("Log level: {}", effective_log_level);
if config.general.disable_colors {
info!("Colors: disabled");
}
info!(
"Modes: classic={} secure={} tls={}",
config.general.modes.classic, config.general.modes.secure, config.general.modes.tls
);
if config.general.modes.classic {
warn!("Classic mode is vulnerable to DPI detection; enable only for legacy clients");
}
info!("TLS domain: {}", config.censorship.tls_domain);
if let Some(ref sock) = config.censorship.mask_unix_sock {
info!("Mask: {} -> unix:{}", config.censorship.mask, sock);
if !std::path::Path::new(sock).exists() {
warn!(
"Unix socket '{}' does not exist yet. Masking will fail until it appears.",
sock
);
}
} else {
info!(
"Mask: {} -> {}:{}",
config.censorship.mask,
config
.censorship
.mask_host
.as_deref()
.unwrap_or(&config.censorship.tls_domain),
config.censorship.mask_port
);
}
if config.censorship.tls_domain == "www.google.com" {
warn!("Using default tls_domain. Consider setting a custom domain.");
}
Ok(BootstrapState {
process_started_at,
process_started_at_epoch_secs,
startup_tracker,
config,
config_path,
has_rust_log,
effective_log_level,
runtime_log_filter,
logging_guard,
})
}
+15 -492
View File
@@ -1,64 +1,22 @@
use std::error::Error;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use std::time::Duration;
use arc_swap::ArcSwap;
use tokio::net::TcpListener;
#[cfg(unix)]
use tokio::net::UnixListener;
use tracing::{debug, error, info, warn};
use crate::config::{ProxyConfig, RstOnCloseMode};
use crate::proxy::ClientHandler;
use crate::startup::{COMPONENT_LISTENERS_BIND, StartupTracker};
use crate::transport::socket::set_linger_zero;
use crate::transport::{ListenOptions, create_listener, find_listener_processes};
use super::generation::RuntimeGeneration;
use super::helpers::{
expected_handshake_close_description, is_expected_handshake_eof, peer_close_description,
print_proxy_links,
};
//! Client listener planning, binding, lifecycle control, and accept loops.
//!
//! Submodules keep process-owned socket state separate from generation-owned
//! runtime state:
//! - `plan` derives deterministic bind intent from validated configuration.
//! - `bind` prepares and activates sockets without partial startup binding.
//! - `accept` runs cancellation-aware TCP accept loops.
//! - `control` coordinates reversible listener transitions and shutdown.
mod accept;
mod bind;
mod control;
mod plan;
#[cfg(unix)]
mod unix;
#[cfg(unix)]
/// Runs the Unix socket accept loop against the active runtime generation.
pub(crate) use unix::spawn_unix_accept_loop;
/// Owns the sockets bound during process startup.
pub(crate) struct BoundListeners {
/// TCP listeners and their immutable bind-time settings.
pub(crate) listeners: Vec<BoundTcpListener>,
#[cfg(unix)]
/// The optional Unix listener transferred to its accept loop after startup.
pub(crate) unix_listener: Option<UnixListener>,
}
/// A TCP listener and the connection settings fixed when it was bound.
pub(crate) struct BoundTcpListener {
listener: TcpListener,
proxy_protocol: bool,
tls_response_fragment_size: Option<u16>,
}
fn listener_port_or_legacy(listener: &crate::config::ListenerConfig, config: &ProxyConfig) -> u16 {
listener.port.unwrap_or(config.server.port)
}
fn default_link_port(config: &ProxyConfig) -> u16 {
config
.server
.listeners
.first()
.and_then(|listener| listener.port)
.unwrap_or(config.server.port)
}
fn mss_segment_multiplier(client_mss: u16) -> u16 {
1460u16.div_ceil(client_mss)
}
pub(crate) use bind::bind_listeners;
pub(crate) use control::{ListenerManager, PreparedListenerTransition};
pub(crate) use plan::listener_rebind_supported;
#[cfg(any(target_os = "linux", test))]
fn tcp_mss_runtime_profile(
@@ -72,441 +30,6 @@ fn tcp_mss_runtime_profile(
}
}
#[allow(clippy::too_many_arguments)]
/// Binds configured TCP and Unix listeners without starting accept loops.
pub(crate) async fn bind_listeners(
config: &Arc<ProxyConfig>,
decision_ipv4_dc: bool,
decision_ipv6_dc: bool,
detected_ip_v4: Option<IpAddr>,
detected_ip_v6: Option<IpAddr>,
startup_tracker: &Arc<StartupTracker>,
) -> Result<BoundListeners, Box<dyn Error>> {
startup_tracker
.start_component(
COMPONENT_LISTENERS_BIND,
Some("bind TCP/Unix listeners".to_string()),
)
.await;
let mut listeners = Vec::new();
let bulk_client_mss = match config.server.client_mss_bulk_value() {
Ok(value) => value,
Err(error) => {
warn!(
error = %error,
"Invalid bulk client MSS after config validation; disabling bulk MSS"
);
None
}
};
for listener_conf in &config.server.listeners {
let listener_port = listener_port_or_legacy(listener_conf, config);
let addr = SocketAddr::new(listener_conf.ip, listener_port);
if addr.is_ipv4() && !decision_ipv4_dc {
warn!(%addr, "Skipping IPv4 listener: IPv4 disabled by [network]");
continue;
}
if addr.is_ipv6() && !decision_ipv6_dc {
warn!(%addr, "Skipping IPv6 listener: IPv6 disabled by [network]");
continue;
}
let configured_client_mss = match listener_conf.effective_client_mss(&config.server) {
Ok(value) => value,
Err(error) => {
warn!(
%addr,
error = %error,
"Invalid listener client MSS after config validation; using kernel default"
);
None
}
};
#[cfg(target_os = "linux")]
let (client_mss, tls_response_fragment_size) =
tcp_mss_runtime_profile(configured_client_mss, bulk_client_mss);
#[cfg(not(target_os = "linux"))]
let (client_mss, tls_response_fragment_size) = (configured_client_mss, None);
let options = ListenOptions {
reuse_port: listener_conf.reuse_allow,
ipv6_only: listener_conf.ip.is_ipv6(),
backlog: config.server.listen_backlog,
client_mss,
..Default::default()
};
match create_listener(addr, &options) {
Ok(socket) => {
let listener = TcpListener::from_std(socket.into())?;
info!("Listening on {}", addr);
if let Some(client_mss) = client_mss {
info!(
%addr,
client_mss,
segment_multiplier = mss_segment_multiplier(client_mss),
"Client-facing TCP MSS configured"
);
}
if let Some(fragment_size) = tls_response_fragment_size {
info!(
%addr,
fragment_size,
bulk_mss = client_mss,
"Initial FakeTLS response best-effort chunking configured"
);
}
let listener_proxy_protocol = listener_conf
.proxy_protocol
.unwrap_or(config.server.proxy_protocol);
let public_host = if let Some(ref announce) = listener_conf.announce {
announce.clone()
} else if listener_conf.ip.is_unspecified() {
if listener_conf.ip.is_ipv4() {
detected_ip_v4
.map(|ip| ip.to_string())
.unwrap_or_else(|| listener_conf.ip.to_string())
} else {
detected_ip_v6
.map(|ip| ip.to_string())
.unwrap_or_else(|| listener_conf.ip.to_string())
}
} else {
listener_conf.ip.to_string()
};
if config.general.links.public_host.is_none()
&& !config.general.links.show.is_empty()
{
let link_port = config.general.links.public_port.unwrap_or(listener_port);
print_proxy_links(&public_host, link_port, config);
}
listeners.push(BoundTcpListener {
listener,
proxy_protocol: listener_proxy_protocol,
tls_response_fragment_size,
});
}
Err(e) => {
if e.kind() == std::io::ErrorKind::AddrInUse {
let owners = find_listener_processes(addr);
if owners.is_empty() {
error!(
%addr,
"Failed to bind: address already in use (owner process unresolved)"
);
} else {
for owner in owners {
error!(
%addr,
pid = owner.pid,
process = %owner.process,
"Failed to bind: address already in use"
);
}
}
if !listener_conf.reuse_allow {
error!(
%addr,
"reuse_allow=false; set [[server.listeners]].reuse_allow=true to allow multi-instance listening"
);
}
} else {
error!("Failed to bind to {}: {}", addr, e);
}
}
}
}
if !config.general.links.show.is_empty()
&& (config.general.links.public_host.is_some() || listeners.is_empty())
{
let (host, port) = if let Some(ref h) = config.general.links.public_host {
(
h.clone(),
config
.general
.links
.public_port
.unwrap_or(default_link_port(config)),
)
} else {
let ip = detected_ip_v4.or(detected_ip_v6).map(|ip| ip.to_string());
if ip.is_none() {
warn!(
"show_link is configured but public IP could not be detected. Set public_host in config."
);
}
(
ip.unwrap_or_else(|| "UNKNOWN".to_string()),
config
.general
.links
.public_port
.unwrap_or(default_link_port(config)),
)
};
print_proxy_links(&host, port, config);
}
#[cfg(unix)]
let mut unix_listener_out = None;
#[cfg(unix)]
if let Some(ref unix_path) = config.server.listen_unix_sock {
let _ = tokio::fs::remove_file(unix_path).await;
let unix_listener = UnixListener::bind(unix_path)?;
if let Some(ref perm_str) = config.server.listen_unix_sock_perm {
match u32::from_str_radix(perm_str.trim_start_matches('0'), 8) {
Ok(mode) => {
use std::os::unix::fs::PermissionsExt;
let perms = std::fs::Permissions::from_mode(mode);
if let Err(e) = std::fs::set_permissions(unix_path, perms) {
error!(
"Failed to set unix socket permissions to {}: {}",
perm_str, e
);
} else {
info!("Listening on unix:{} (mode {})", unix_path, perm_str);
}
}
Err(e) => {
warn!(
"Invalid listen_unix_sock_perm '{}': {}. Ignoring.",
perm_str, e
);
info!("Listening on unix:{}", unix_path);
}
}
} else {
info!("Listening on unix:{}", unix_path);
}
unix_listener_out = Some(unix_listener);
}
#[cfg(unix)]
let has_unix_listener = unix_listener_out.is_some();
#[cfg(not(unix))]
let has_unix_listener = false;
startup_tracker
.complete_component(
COMPONENT_LISTENERS_BIND,
Some(format!(
"listeners configured tcp={} unix={}",
listeners.len(),
has_unix_listener
)),
)
.await;
Ok(BoundListeners {
listeners,
#[cfg(unix)]
unix_listener: unix_listener_out,
})
}
/// Starts one TCP accept loop per bound listener.
pub(crate) fn spawn_tcp_accept_loops(
listeners: Vec<BoundTcpListener>,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
) {
for bound_listener in listeners {
let listener = bound_listener.listener;
let listener_proxy_protocol = bound_listener.proxy_protocol;
let tls_response_fragment_size = bound_listener.tls_response_fragment_size;
let active_runtime = active_runtime.clone();
tokio::spawn(async move {
loop {
match listener.accept().await {
Ok((stream, peer_addr)) => {
let runtime = active_runtime.load_full();
let config = runtime.config();
let rst_mode = config.general.rst_on_close;
#[cfg(unix)]
let raw_fd = {
use std::os::unix::io::AsRawFd;
stream.as_raw_fd()
};
if matches!(rst_mode, RstOnCloseMode::Errors | RstOnCloseMode::Always) {
let _ = set_linger_zero(&stream);
}
if !*runtime.admission_rx.borrow() {
debug!(peer = %peer_addr, "Admission gate closed, dropping connection");
drop(stream);
continue;
}
let accept_permit_timeout_ms = config.server.accept_permit_timeout_ms;
let permit = if accept_permit_timeout_ms == 0 {
match runtime.max_connections.clone().acquire_owned().await {
Ok(permit) => permit,
Err(_) => {
error!("Connection limiter is closed");
break;
}
}
} else {
match tokio::time::timeout(
Duration::from_millis(accept_permit_timeout_ms),
runtime.max_connections.clone().acquire_owned(),
)
.await
{
Ok(Ok(permit)) => permit,
Ok(Err(_)) => {
error!("Connection limiter is closed");
break;
}
Err(_) => {
runtime.stats.increment_accept_permit_timeout_total();
debug!(
peer = %peer_addr,
timeout_ms = accept_permit_timeout_ms,
"Dropping accepted connection: permit wait timeout"
);
drop(stream);
continue;
}
}
};
let stats = runtime.stats.clone();
let upstream_manager = runtime.upstream_manager.clone();
let replay_checker = runtime.replay_checker.clone();
let buffer_pool = runtime.buffer_pool.clone();
let rng = runtime.rng.clone();
let me_pool = runtime.me_pool.clone();
let me_pool_runtime = runtime.me_pool_runtime.clone();
let route_runtime = runtime.route_runtime.clone();
let tls_cache = runtime.tls_cache.clone();
let ip_tracker = runtime.ip_tracker.clone();
let beobachten = runtime.beobachten.clone();
let shared = runtime.proxy_shared.clone();
let proxy_protocol_enabled = listener_proxy_protocol;
let real_peer_report = Arc::new(std::sync::Mutex::new(None));
let real_peer_report_for_handler = real_peer_report.clone();
let _ = runtime.spawn_session(async move {
let _permit = permit;
if let Err(e) = ClientHandler::new_with_shared(
stream,
peer_addr,
config,
stats,
upstream_manager,
replay_checker,
buffer_pool,
rng,
me_pool,
Some(me_pool_runtime),
route_runtime,
tls_cache,
ip_tracker,
beobachten,
shared,
proxy_protocol_enabled,
real_peer_report_for_handler,
#[cfg(unix)]
raw_fd,
rst_mode,
tls_response_fragment_size,
)
.run()
.await
{
let real_peer = match real_peer_report.lock() {
Ok(guard) => *guard,
Err(_) => None,
};
let peer_close_reason = peer_close_description(&e);
let handshake_close_reason =
expected_handshake_close_description(&e);
let me_closed =
matches!(&e, crate::error::ProxyError::MiddleConnectionLost);
let route_switched =
matches!(&e, crate::error::ProxyError::RouteSwitched);
match (peer_close_reason, me_closed) {
(Some(reason), _) => {
if let Some(real_peer) = real_peer {
debug!(
peer = %peer_addr,
real_peer = %real_peer,
error = %e,
close_reason = reason,
"Connection closed by peer"
);
} else {
debug!(
peer = %peer_addr,
error = %e,
close_reason = reason,
"Connection closed by peer"
);
}
}
(_, true) => {
if let Some(real_peer) = real_peer {
warn!(peer = %peer_addr, real_peer = %real_peer, error = %e, "Connection closed: Middle-End dropped session");
} else {
warn!(peer = %peer_addr, error = %e, "Connection closed: Middle-End dropped session");
}
}
_ if route_switched => {
if let Some(real_peer) = real_peer {
info!(peer = %peer_addr, real_peer = %real_peer, error = %e, "Connection closed by controlled route cutover");
} else {
info!(peer = %peer_addr, error = %e, "Connection closed by controlled route cutover");
}
}
_ if is_expected_handshake_eof(&e) => {
let reason = handshake_close_reason
.unwrap_or("Peer closed during initial handshake");
if let Some(real_peer) = real_peer {
info!(
peer = %peer_addr,
real_peer = %real_peer,
error = %e,
close_reason = reason,
"Connection closed during initial handshake"
);
} else {
info!(
peer = %peer_addr,
error = %e,
close_reason = reason,
"Connection closed during initial handshake"
);
}
}
_ => {
if let Some(real_peer) = real_peer {
warn!(peer = %peer_addr, real_peer = %real_peer, error = %e, "Connection closed with error");
} else {
warn!(peer = %peer_addr, error = %e, "Connection closed with error");
}
}
}
}
});
}
Err(e) => {
error!("Accept error: {}", e);
tokio::time::sleep(Duration::from_millis(100)).await;
}
}
}
});
}
}
#[cfg(test)]
mod tests {
use super::tcp_mss_runtime_profile;
+269
View File
@@ -0,0 +1,269 @@
use std::sync::Arc;
use std::time::Duration;
use arc_swap::ArcSwap;
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::OwnedSemaphorePermit;
use tokio::task::JoinHandle;
use tokio_util::sync::CancellationToken;
use tracing::{debug, error, info, warn};
use crate::config::RstOnCloseMode;
use crate::proxy::ClientHandler;
use crate::transport::socket::set_linger_zero;
use super::bind::BoundTcpListener;
use super::plan::ListenerBindSpec;
use crate::maestro::generation::RuntimeGeneration;
use crate::maestro::helpers::{
expected_handshake_close_description, is_expected_handshake_eof, peer_close_description,
};
pub(super) struct ListenerSlot {
pub(super) spec: ListenerBindSpec,
listener: Arc<TcpListener>,
cancellation: CancellationToken,
task: Option<JoinHandle<()>>,
}
enum PermitWait {
Acquired(OwnedSemaphorePermit),
TimedOut,
Closed,
Cancelled,
}
async fn wait_for_permit(
runtime: &Arc<RuntimeGeneration>,
cancellation: &CancellationToken,
) -> PermitWait {
let timeout_ms = runtime.config().server.accept_permit_timeout_ms;
let acquire = runtime.max_connections.clone().acquire_owned();
if timeout_ms == 0 {
return tokio::select! {
biased;
_ = cancellation.cancelled() => PermitWait::Cancelled,
permit = acquire => match permit {
Ok(permit) => PermitWait::Acquired(permit),
Err(_) => PermitWait::Closed,
},
};
}
tokio::select! {
biased;
_ = cancellation.cancelled() => PermitWait::Cancelled,
result = tokio::time::timeout(Duration::from_millis(timeout_ms), acquire) => {
match result {
Ok(Ok(permit)) => PermitWait::Acquired(permit),
Ok(Err(_)) => PermitWait::Closed,
Err(_) => PermitWait::TimedOut,
}
}
}
}
fn spawn_client_session(
stream: TcpStream,
peer_addr: std::net::SocketAddr,
runtime: Arc<RuntimeGeneration>,
permit: OwnedSemaphorePermit,
spec: &ListenerBindSpec,
) {
let config = runtime.config();
let rst_mode = config.general.rst_on_close;
#[cfg(unix)]
let raw_fd = {
use std::os::unix::io::AsRawFd;
stream.as_raw_fd()
};
if matches!(rst_mode, RstOnCloseMode::Errors | RstOnCloseMode::Always) {
let _ = set_linger_zero(&stream);
}
let stats = runtime.stats.clone();
let upstream_manager = runtime.upstream_manager.clone();
let replay_checker = runtime.replay_checker.clone();
let buffer_pool = runtime.buffer_pool.clone();
let rng = runtime.rng.clone();
let me_pool = runtime.me_pool.clone();
let me_pool_runtime = runtime.me_pool_runtime.clone();
let route_runtime = runtime.route_runtime.clone();
let tls_cache = runtime.tls_cache.clone();
let ip_tracker = runtime.ip_tracker.clone();
let beobachten = runtime.beobachten.clone();
let shared = runtime.proxy_shared.clone();
let proxy_protocol_enabled = spec.proxy_protocol;
let tls_response_fragment_size = spec.tls_response_fragment_size;
let real_peer_report = Arc::new(std::sync::Mutex::new(None));
let real_peer_report_for_handler = real_peer_report.clone();
let _ = runtime.spawn_session(async move {
let _permit = permit;
if let Err(error_value) = ClientHandler::new_with_shared(
stream,
peer_addr,
config,
stats,
upstream_manager,
replay_checker,
buffer_pool,
rng,
me_pool,
Some(me_pool_runtime),
route_runtime,
tls_cache,
ip_tracker,
beobachten,
shared,
proxy_protocol_enabled,
real_peer_report_for_handler,
#[cfg(unix)]
raw_fd,
rst_mode,
tls_response_fragment_size,
)
.run()
.await
{
let real_peer = real_peer_report.lock().ok().and_then(|guard| *guard);
let peer_close_reason = peer_close_description(&error_value);
let handshake_close_reason = expected_handshake_close_description(&error_value);
let me_closed = matches!(
&error_value,
crate::error::ProxyError::MiddleConnectionLost
);
let route_switched =
matches!(&error_value, crate::error::ProxyError::RouteSwitched);
match (peer_close_reason, me_closed) {
(Some(reason), _) => {
if let Some(real_peer) = real_peer {
debug!(peer = %peer_addr, real_peer = %real_peer, error = %error_value, close_reason = reason, "Connection closed by peer");
} else {
debug!(peer = %peer_addr, error = %error_value, close_reason = reason, "Connection closed by peer");
}
}
(_, true) => {
if let Some(real_peer) = real_peer {
warn!(peer = %peer_addr, real_peer = %real_peer, error = %error_value, "Connection closed: Middle-End dropped session");
} else {
warn!(peer = %peer_addr, error = %error_value, "Connection closed: Middle-End dropped session");
}
}
_ if route_switched => {
if let Some(real_peer) = real_peer {
info!(peer = %peer_addr, real_peer = %real_peer, error = %error_value, "Connection closed by controlled route cutover");
} else {
info!(peer = %peer_addr, error = %error_value, "Connection closed by controlled route cutover");
}
}
_ if is_expected_handshake_eof(&error_value) => {
let reason = handshake_close_reason
.unwrap_or("Peer closed during initial handshake");
if let Some(real_peer) = real_peer {
info!(peer = %peer_addr, real_peer = %real_peer, error = %error_value, close_reason = reason, "Connection closed during initial handshake");
} else {
info!(peer = %peer_addr, error = %error_value, close_reason = reason, "Connection closed during initial handshake");
}
}
_ => {
if let Some(real_peer) = real_peer {
warn!(peer = %peer_addr, real_peer = %real_peer, error = %error_value, "Connection closed with error");
} else {
warn!(peer = %peer_addr, error = %error_value, "Connection closed with error");
}
}
}
}
});
}
async fn run_accept_loop(
listener: Arc<TcpListener>,
spec: ListenerBindSpec,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
cancellation: CancellationToken,
) {
loop {
let accepted = tokio::select! {
biased;
_ = cancellation.cancelled() => return,
accepted = listener.accept() => accepted,
};
match accepted {
Ok((stream, peer_addr)) => {
let runtime = active_runtime.load_full();
if !*runtime.admission_rx.borrow() {
debug!(peer = %peer_addr, "Admission gate closed, dropping connection");
drop(stream);
continue;
}
match wait_for_permit(&runtime, &cancellation).await {
PermitWait::Acquired(permit) => {
spawn_client_session(stream, peer_addr, runtime, permit, &spec);
}
PermitWait::TimedOut => {
runtime.stats.increment_accept_permit_timeout_total();
debug!(
peer = %peer_addr,
timeout_ms = runtime.config().server.accept_permit_timeout_ms,
"Dropping accepted connection: permit wait timeout"
);
}
PermitWait::Closed => {
error!(addr = %spec.addr, "Connection limiter is closed");
return;
}
PermitWait::Cancelled => return,
}
}
Err(error_value) => {
error!(addr = %spec.addr, error = %error_value, "TCP accept error");
tokio::select! {
biased;
_ = cancellation.cancelled() => return,
_ = tokio::time::sleep(Duration::from_millis(100)) => {}
}
}
}
}
}
impl ListenerSlot {
pub(super) fn start(
bound: BoundTcpListener,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
) -> Self {
let cancellation = CancellationToken::new();
let task = tokio::spawn(run_accept_loop(
bound.listener.clone(),
bound.spec.clone(),
active_runtime,
cancellation.clone(),
));
Self {
spec: bound.spec,
listener: bound.listener,
cancellation,
task: Some(task),
}
}
pub(super) async fn stop(&mut self) -> Result<(), String> {
self.cancellation.cancel();
if let Some(task) = self.task.take() {
task.await
.map_err(|error_value| format!("listener {} task failed: {error_value}", self.spec.addr))?;
}
Ok(())
}
pub(super) fn restart(&mut self, active_runtime: Arc<ArcSwap<RuntimeGeneration>>) {
self.cancellation = CancellationToken::new();
self.task = Some(tokio::spawn(run_accept_loop(
self.listener.clone(),
self.spec.clone(),
active_runtime,
self.cancellation.clone(),
)));
}
}
+260
View File
@@ -0,0 +1,260 @@
use std::error::Error;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use socket2::Socket;
use tokio::net::TcpListener;
#[cfg(unix)]
use tokio::net::UnixListener;
use tracing::{error, info, warn};
use crate::config::ProxyConfig;
use crate::startup::{COMPONENT_LISTENERS_BIND, StartupTracker};
use crate::transport::socket::{activate_listener_socket, bind_listener_socket};
use crate::transport::find_listener_processes;
use super::plan::{ListenerBindSpec, listener_bind_plan};
use crate::maestro::helpers::print_proxy_links;
/// Owns sockets bound before process accept loops start.
pub(crate) struct BoundListeners {
pub(super) listeners: Vec<BoundTcpListener>,
#[cfg(unix)]
pub(super) unix_listener: Option<UnixListener>,
}
impl BoundListeners {
pub(crate) fn is_empty(&self) -> bool {
let tcp_empty = self.listeners.is_empty();
#[cfg(unix)]
{
tcp_empty && self.unix_listener.is_none()
}
#[cfg(not(unix))]
{
tcp_empty
}
}
}
/// Active socket and immutable connection policy for one endpoint.
pub(super) struct BoundTcpListener {
pub(super) listener: Arc<TcpListener>,
pub(super) spec: ListenerBindSpec,
}
/// Socket bound for a candidate transition but not yet listening.
pub(super) struct PreparedTcpListener {
socket: Socket,
spec: ListenerBindSpec,
}
fn mss_segment_multiplier(client_mss: u16) -> u16 {
1460u16.div_ceil(client_mss)
}
fn default_link_port(config: &ProxyConfig) -> u16 {
config
.server
.listeners
.first()
.and_then(|listener| listener.port)
.unwrap_or(config.server.port)
}
fn log_bind_error(addr: SocketAddr, reuse_allow: bool, error_value: &std::io::Error) {
if error_value.kind() == std::io::ErrorKind::AddrInUse {
let owners = find_listener_processes(addr);
if owners.is_empty() {
error!(%addr, "Failed to bind: address already in use (owner process unresolved)");
} else {
for owner in owners {
error!(
%addr,
pid = owner.pid,
process = %owner.process,
"Failed to bind: address already in use"
);
}
}
if !reuse_allow {
error!(
%addr,
"reuse_allow=false; set [[server.listeners]].reuse_allow=true to allow multi-instance listening"
);
}
} else {
error!(%addr, error = %error_value, "Failed to bind listener");
}
}
pub(super) fn prepare_listener(
spec: ListenerBindSpec,
) -> std::io::Result<PreparedTcpListener> {
match bind_listener_socket(spec.addr, &spec.options) {
Ok(socket) => Ok(PreparedTcpListener { socket, spec }),
Err(error_value) => {
log_bind_error(spec.addr, spec.options.reuse_port, &error_value);
Err(error_value)
}
}
}
impl PreparedTcpListener {
pub(super) fn activate(self) -> std::io::Result<BoundTcpListener> {
activate_listener_socket(&self.socket, self.spec.options.backlog)?;
let listener = TcpListener::from_std(self.socket.into())?;
Ok(BoundTcpListener {
listener: Arc::new(listener),
spec: self.spec,
})
}
}
fn log_listener_profile(spec: &ListenerBindSpec) {
info!(addr = %spec.addr, "Listening on TCP endpoint");
if let Some(client_mss) = spec.options.client_mss {
info!(
addr = %spec.addr,
client_mss,
segment_multiplier = mss_segment_multiplier(client_mss),
"Client-facing TCP MSS configured"
);
}
if let Some(fragment_size) = spec.tls_response_fragment_size {
info!(
addr = %spec.addr,
fragment_size,
bulk_mss = spec.options.client_mss,
"Initial FakeTLS response best-effort chunking configured"
);
}
}
fn print_configured_links(
config: &ProxyConfig,
plan: &std::collections::BTreeMap<SocketAddr, ListenerBindSpec>,
detected_ip_v4: Option<IpAddr>,
detected_ip_v6: Option<IpAddr>,
) {
for listener in &config.server.listeners {
let port = listener.port.unwrap_or(config.server.port);
let addr = SocketAddr::new(listener.ip, port);
if !plan.contains_key(&addr) || config.general.links.public_host.is_some() {
continue;
}
let public_host = if let Some(announce) = &listener.announce {
announce.clone()
} else if listener.ip.is_unspecified() {
if listener.ip.is_ipv4() {
detected_ip_v4
} else {
detected_ip_v6
}
.map(|ip| ip.to_string())
.unwrap_or_else(|| listener.ip.to_string())
} else {
listener.ip.to_string()
};
if !config.general.links.show.is_empty() {
let link_port = config.general.links.public_port.unwrap_or(port);
print_proxy_links(&public_host, link_port, config);
}
}
if config.general.links.show.is_empty() || config.general.links.public_host.is_none() {
return;
}
let host = config.general.links.public_host.as_deref().unwrap_or_default();
let port = config
.general
.links
.public_port
.unwrap_or_else(|| default_link_port(config));
print_proxy_links(host, port, config);
}
/// Binds every eligible configured listener or fails without a partial inventory.
pub(crate) async fn bind_listeners(
config: &Arc<ProxyConfig>,
detected_ip_v4: Option<IpAddr>,
detected_ip_v6: Option<IpAddr>,
startup_tracker: &Arc<StartupTracker>,
) -> Result<BoundListeners, Box<dyn Error>> {
startup_tracker
.start_component(
COMPONENT_LISTENERS_BIND,
Some("bind TCP/Unix listeners".to_string()),
)
.await;
let plan = listener_bind_plan(config).map_err(std::io::Error::other)?;
let mut prepared = Vec::with_capacity(plan.len());
for spec in plan.values().cloned() {
prepared.push(prepare_listener(spec)?);
}
let mut listeners = Vec::with_capacity(prepared.len());
for candidate in prepared {
let bound = candidate.activate()?;
log_listener_profile(&bound.spec);
listeners.push(bound);
}
print_configured_links(config, &plan, detected_ip_v4, detected_ip_v6);
#[cfg(unix)]
let mut unix_listener_out = None;
#[cfg(unix)]
if let Some(unix_path) = &config.server.listen_unix_sock {
let _ = tokio::fs::remove_file(unix_path).await;
let unix_listener = UnixListener::bind(unix_path)?;
if let Some(perm_str) = &config.server.listen_unix_sock_perm {
match u32::from_str_radix(perm_str.trim_start_matches('0'), 8) {
Ok(mode) => {
use std::os::unix::fs::PermissionsExt;
let permissions = std::fs::Permissions::from_mode(mode);
if let Err(error_value) = std::fs::set_permissions(unix_path, permissions) {
error!(
path = %unix_path,
permissions = %perm_str,
error = %error_value,
"Failed to set Unix socket permissions"
);
} else {
info!(path = %unix_path, permissions = %perm_str, "Listening on Unix socket");
}
}
Err(error_value) => {
warn!(
path = %unix_path,
permissions = %perm_str,
error = %error_value,
"Invalid Unix socket permissions; keeping umask-derived mode"
);
}
}
} else {
info!(path = %unix_path, "Listening on Unix socket");
}
unix_listener_out = Some(unix_listener);
}
#[cfg(unix)]
let has_unix_listener = unix_listener_out.is_some();
#[cfg(not(unix))]
let has_unix_listener = false;
startup_tracker
.complete_component(
COMPONENT_LISTENERS_BIND,
Some(format!(
"listeners configured tcp={} unix={}",
listeners.len(),
has_unix_listener
)),
)
.await;
Ok(BoundListeners {
listeners,
#[cfg(unix)]
unix_listener: unix_listener_out,
})
}
+307
View File
@@ -0,0 +1,307 @@
use std::collections::{BTreeMap, BTreeSet};
use std::net::SocketAddr;
use std::sync::Arc;
use arc_swap::ArcSwap;
use crate::config::ProxyConfig;
use crate::maestro::generation::RuntimeGeneration;
use super::accept::ListenerSlot;
use super::bind::{
BoundListeners, BoundTcpListener, PreparedTcpListener, prepare_listener,
};
use super::plan::{ListenerBindSpec, listener_bind_plan};
#[cfg(unix)]
use super::unix::UnixAcceptHandle;
/// Process-owned listener inventory and accept-task lifecycle controller.
pub(crate) struct ListenerManager {
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
slots: BTreeMap<SocketAddr, ListenerSlot>,
#[cfg(unix)]
unix: Option<UnixAcceptHandle>,
}
pub(crate) struct PreparedListenerTransition {
target_specs: BTreeMap<SocketAddr, ListenerBindSpec>,
additions: Vec<PreparedTcpListener>,
removals: Vec<SocketAddr>,
}
pub(crate) struct PendingListenerTransition {
target_specs: BTreeMap<SocketAddr, ListenerBindSpec>,
additions: Vec<BoundTcpListener>,
removals: Vec<SocketAddr>,
}
impl ListenerManager {
/// Starts accept loops for the complete startup-bound inventory.
pub(crate) fn start(
bound: BoundListeners,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
) -> Self {
let mut slots = BTreeMap::new();
for listener in bound.listeners {
let addr = listener.spec.addr;
slots.insert(addr, ListenerSlot::start(listener, active_runtime.clone()));
}
#[cfg(unix)]
let unix = bound
.unix_listener
.map(|listener| UnixAcceptHandle::start(listener, active_runtime.clone()));
Self {
active_runtime,
slots,
#[cfg(unix)]
unix,
}
}
#[cfg(test)]
pub(crate) fn empty(active_runtime: Arc<ArcSwap<RuntimeGeneration>>) -> Self {
Self {
active_runtime,
slots: BTreeMap::new(),
#[cfg(unix)]
unix: None,
}
}
/// Binds added endpoints without calling `listen(2)` or changing active tasks.
pub(crate) fn prepare_transition(
&self,
desired: &ProxyConfig,
) -> Result<Option<PreparedListenerTransition>, String> {
let target_specs = listener_bind_plan(desired)?;
let current_addresses: BTreeSet<_> = self.slots.keys().copied().collect();
let target_addresses: BTreeSet<_> = target_specs.keys().copied().collect();
if current_addresses == target_addresses
&& self
.slots
.iter()
.all(|(addr, slot)| target_specs.get(addr) == Some(&slot.spec))
{
return Ok(None);
}
for addr in current_addresses.intersection(&target_addresses) {
let current = &self.slots[addr].spec;
let desired_spec = &target_specs[addr];
if current != desired_spec {
return Err(format!(
"listener {addr} bind policy changed at the same endpoint; process restart required"
));
}
}
let mut additions = Vec::new();
for addr in target_addresses.difference(&current_addresses) {
let spec = target_specs
.get(addr)
.expect("address originated from target listener plan")
.clone();
additions.push(prepare_listener(spec).map_err(|error_value| {
format!("failed to prepare listener {addr}: {error_value}")
})?);
}
let removals = current_addresses
.difference(&target_addresses)
.copied()
.collect();
Ok(Some(PreparedListenerTransition {
target_specs,
additions,
removals,
}))
}
/// Activates additions and stops removed acceptors before the runtime swap.
pub(crate) async fn begin_transition(
&mut self,
prepared: PreparedListenerTransition,
) -> Result<PendingListenerTransition, String> {
let mut additions = Vec::with_capacity(prepared.additions.len());
for candidate in prepared.additions {
additions.push(candidate.activate().map_err(|error_value| {
format!("failed to activate prepared listener: {error_value}")
})?);
}
let mut stopped = Vec::new();
for addr in &prepared.removals {
let stop_result = self
.slots
.get_mut(addr)
.expect("removal originated from active listener inventory")
.stop()
.await;
if let Err(error_value) = stop_result {
for stopped_addr in stopped {
if let Some(stopped_slot) = self.slots.get_mut(&stopped_addr) {
stopped_slot.restart(self.active_runtime.clone());
}
}
self.slots
.get_mut(addr)
.expect("failed slot remains in active listener inventory")
.restart(self.active_runtime.clone());
return Err(error_value);
}
stopped.push(*addr);
}
Ok(PendingListenerTransition {
target_specs: prepared.target_specs,
additions,
removals: prepared.removals,
})
}
/// Publishes new acceptors after the runtime generation has been swapped.
pub(crate) fn finish_transition(&mut self, pending: PendingListenerTransition) {
for addr in pending.removals {
self.slots.remove(&addr);
}
for listener in pending.additions {
let addr = listener.spec.addr;
self.slots.insert(
addr,
ListenerSlot::start(listener, self.active_runtime.clone()),
);
}
debug_assert_eq!(
self.slots
.iter()
.map(|(addr, slot)| (*addr, slot.spec.clone()))
.collect::<BTreeMap<_, _>>(),
pending.target_specs
);
}
/// Stops and joins every accept task before sockets are released.
pub(crate) async fn shutdown(&mut self) -> Result<(), String> {
let mut errors = Vec::new();
for slot in self.slots.values_mut() {
if let Err(error_value) = slot.stop().await {
errors.push(error_value);
}
}
#[cfg(unix)]
if let Some(unix) = &mut self.unix
&& let Err(error_value) = unix.stop().await
{
errors.push(error_value);
}
self.slots.clear();
#[cfg(unix)]
{
self.unix = None;
}
if errors.is_empty() {
Ok(())
} else {
Err(errors.join("; "))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::{ListenerConfig, SynLimitMode};
use crate::maestro::generation::test_runtime_generation;
use crate::transport::ListenOptions;
use tokio::net::{TcpListener, TcpStream};
fn listener_config(addr: SocketAddr) -> ListenerConfig {
ListenerConfig {
ip: addr.ip(),
port: Some(addr.port()),
client_mss: None,
synlimit: SynLimitMode::Off,
synlimit_seconds: 60,
synlimit_hitcount: 48,
synlimit_burst: 24,
synlimit_ios_seconds: 1,
synlimit_ios_hitcount: 12,
synlimit_ios_burst: 24,
synlimit_hashlimit_expire_ms: 60_000,
synlimit_hashlimit_size: 32_768,
announce: None,
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
}
}
async fn bound_listener() -> (BoundTcpListener, SocketAddr) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let spec = ListenerBindSpec {
addr,
options: ListenOptions {
reuse_port: false,
..Default::default()
},
proxy_protocol: false,
tls_response_fragment_size: None,
};
(
BoundTcpListener {
listener: Arc::new(listener),
spec,
},
addr,
)
}
#[tokio::test]
async fn candidate_bind_failure_keeps_old_listener_accepting() {
let runtime = test_runtime_generation(1, ProxyConfig::default());
let active_runtime = Arc::new(ArcSwap::from(runtime.clone()));
let (old_listener, old_addr) = bound_listener().await;
let bound = BoundListeners {
listeners: vec![old_listener],
#[cfg(unix)]
unix_listener: None,
};
let mut manager = ListenerManager::start(bound, active_runtime);
let blocker = TcpListener::bind("127.0.0.1:0").await.unwrap();
let blocked_addr = blocker.local_addr().unwrap();
let mut desired = ProxyConfig::default();
desired.server.listeners = vec![listener_config(blocked_addr)];
assert!(manager.prepare_transition(&desired).is_err());
TcpStream::connect(old_addr).await.unwrap();
manager.shutdown().await.unwrap();
runtime.stop_sessions().await;
}
#[tokio::test]
async fn added_listener_is_dormant_until_transition_begins() {
let runtime = test_runtime_generation(1, ProxyConfig::default());
let active_runtime = Arc::new(ArcSwap::from(runtime.clone()));
let (old_listener, _old_addr) = bound_listener().await;
let bound = BoundListeners {
listeners: vec![old_listener],
#[cfg(unix)]
unix_listener: None,
};
let mut manager = ListenerManager::start(bound, active_runtime);
let reservation = TcpListener::bind("127.0.0.1:0").await.unwrap();
let new_addr = reservation.local_addr().unwrap();
drop(reservation);
let mut desired = ProxyConfig::default();
desired.server.listeners = vec![listener_config(new_addr)];
let prepared = manager.prepare_transition(&desired).unwrap().unwrap();
assert!(TcpStream::connect(new_addr).await.is_err());
let pending = manager.begin_transition(prepared).await.unwrap();
manager.finish_transition(pending);
TcpStream::connect(new_addr).await.unwrap();
manager.shutdown().await.unwrap();
runtime.stop_sessions().await;
}
}
+174
View File
@@ -0,0 +1,174 @@
use std::collections::{BTreeMap, BTreeSet};
use std::net::SocketAddr;
use crate::config::{ProxyConfig, ServerConfig, SynLimitMode};
use crate::transport::ListenOptions;
use super::tcp_mss_runtime_profile;
/// Immutable socket and connection policy for one listener endpoint.
#[derive(Debug, Clone, PartialEq, Eq)]
pub(super) struct ListenerBindSpec {
pub(super) addr: SocketAddr,
pub(super) options: ListenOptions,
pub(super) proxy_protocol: bool,
pub(super) tls_response_fragment_size: Option<u16>,
}
fn listener_port_or_legacy(
listener: &crate::config::ListenerConfig,
server: &ServerConfig,
) -> u16 {
listener.port.unwrap_or(server.port)
}
/// Derives inbound listener intent without consulting transient outbound probes.
pub(crate) fn listener_bind_plan(
config: &ProxyConfig,
) -> Result<BTreeMap<SocketAddr, ListenerBindSpec>, String> {
let mut plan = BTreeMap::new();
let bulk_client_mss = config
.server
.client_mss_bulk_value()
.map_err(|error| format!("invalid server.client_mss_bulk: {error}"))?;
for listener in &config.server.listeners {
let addr = SocketAddr::new(
listener.ip,
listener_port_or_legacy(listener, &config.server),
);
if addr.is_ipv4() && !config.network.ipv4 {
continue;
}
if addr.is_ipv6() && config.network.ipv6 == Some(false) {
continue;
}
let configured_client_mss = listener
.effective_client_mss(&config.server)
.map_err(|error| format!("invalid client MSS for listener {addr}: {error}"))?;
#[cfg(target_os = "linux")]
let (client_mss, tls_response_fragment_size) =
tcp_mss_runtime_profile(configured_client_mss, bulk_client_mss);
#[cfg(not(target_os = "linux"))]
let (client_mss, tls_response_fragment_size) = (configured_client_mss, None);
let spec = ListenerBindSpec {
addr,
options: ListenOptions {
reuse_port: listener.reuse_allow,
ipv6_only: listener.ip.is_ipv6(),
backlog: config.server.listen_backlog,
client_mss,
..Default::default()
},
proxy_protocol: listener
.proxy_protocol
.unwrap_or(config.server.proxy_protocol),
tls_response_fragment_size,
};
if plan.insert(addr, spec).is_some() {
return Err(format!("duplicate effective listener endpoint: {addr}"));
}
}
Ok(plan)
}
fn any_synlimit_enabled(config: &ProxyConfig) -> bool {
config
.server
.listeners
.iter()
.any(|listener| listener.synlimit != SynLimitMode::Off)
}
/// Returns whether an endpoint-only change can use coordinated process rebind.
pub(crate) fn listener_rebind_supported(old: &ProxyConfig, desired: &ProxyConfig) -> bool {
if any_synlimit_enabled(old) || any_synlimit_enabled(desired) {
return false;
}
let Ok(old_plan) = listener_bind_plan(old) else {
return false;
};
let Ok(desired_plan) = listener_bind_plan(desired) else {
return false;
};
let retained: BTreeSet<_> = old_plan
.keys()
.filter(|addr| desired_plan.contains_key(addr))
.copied()
.collect();
retained
.iter()
.all(|addr| old_plan.get(addr) == desired_plan.get(addr))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::ListenerConfig;
fn listener(ip: &str, port: u16) -> ListenerConfig {
ListenerConfig {
ip: ip.parse().unwrap(),
port: Some(port),
client_mss: None,
synlimit: SynLimitMode::Off,
synlimit_seconds: 60,
synlimit_hitcount: 48,
synlimit_burst: 24,
synlimit_ios_seconds: 1,
synlimit_ios_hitcount: 12,
synlimit_ios_burst: 24,
synlimit_hashlimit_expire_ms: 60_000,
synlimit_hashlimit_size: 32_768,
announce: None,
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
}
}
#[test]
fn plan_depends_on_inbound_family_policy_only() {
let mut config = ProxyConfig::default();
config.server.listeners = vec![listener("0.0.0.0", 443), listener("::", 443)];
config.network.ipv4 = true;
config.network.ipv6 = None;
let plan = listener_bind_plan(&config).unwrap();
assert_eq!(plan.len(), 2);
config.network.ipv6 = Some(false);
let plan = listener_bind_plan(&config).unwrap();
assert_eq!(plan.len(), 1);
assert!(plan.keys().all(SocketAddr::is_ipv4));
}
#[test]
fn duplicate_effective_endpoint_is_rejected() {
let mut config = ProxyConfig::default();
config.server.listeners = vec![listener("127.0.0.1", 443), listener("127.0.0.1", 443)];
assert!(listener_bind_plan(&config).is_err());
}
#[test]
fn retained_policy_change_is_not_rebindable() {
let mut old = ProxyConfig::default();
old.server.listeners = vec![listener("127.0.0.1", 443)];
let mut desired = old.clone();
desired.server.listeners[0].proxy_protocol = Some(true);
assert!(!listener_rebind_supported(&old, &desired));
}
#[test]
fn endpoint_move_without_synlimit_is_rebindable() {
let mut old = ProxyConfig::default();
old.server.listeners = vec![listener("127.0.0.1", 443)];
let mut desired = old.clone();
desired.server.listeners[0].port = Some(444);
assert!(listener_rebind_supported(&old, &desired));
}
}
+136 -94
View File
@@ -5,113 +5,155 @@ use std::time::Duration;
use arc_swap::ArcSwap;
use tokio::net::UnixListener;
use tokio::task::JoinHandle;
use tokio_util::sync::CancellationToken;
use tracing::{debug, error};
use super::RuntimeGeneration;
use crate::maestro::generation::RuntimeGeneration;
pub(crate) fn spawn_unix_accept_loop(
listener: Option<UnixListener>,
pub(super) struct UnixAcceptHandle {
_listener: Arc<UnixListener>,
cancellation: CancellationToken,
task: Option<JoinHandle<()>>,
}
async fn run_unix_accept_loop(
listener: Arc<UnixListener>,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
cancellation: CancellationToken,
) {
let Some(listener) = listener else {
return;
};
tokio::spawn(async move {
let connection_counter = AtomicU64::new(1);
loop {
match listener.accept().await {
Ok((stream, _)) => {
let runtime = active_runtime.load_full();
if !*runtime.admission_rx.borrow() {
drop(stream);
continue;
}
let config = runtime.config();
let timeout_ms = config.server.accept_permit_timeout_ms;
let permit = if timeout_ms == 0 {
match runtime.max_connections.clone().acquire_owned().await {
let connection_counter = AtomicU64::new(1);
loop {
let accepted = tokio::select! {
biased;
_ = cancellation.cancelled() => return,
accepted = listener.accept() => accepted,
};
match accepted {
Ok((stream, _)) => {
let runtime = active_runtime.load_full();
if !*runtime.admission_rx.borrow() {
drop(stream);
continue;
}
let config = runtime.config();
let timeout_ms = config.server.accept_permit_timeout_ms;
let acquire = runtime.max_connections.clone().acquire_owned();
let permit = if timeout_ms == 0 {
tokio::select! {
biased;
_ = cancellation.cancelled() => return,
permit = acquire => match permit {
Ok(permit) => permit,
Err(_) => {
error!("Connection limiter is closed");
break;
return;
}
}
} else {
match tokio::time::timeout(
Duration::from_millis(timeout_ms),
runtime.max_connections.clone().acquire_owned(),
}
} else {
match tokio::select! {
biased;
_ = cancellation.cancelled() => return,
result = tokio::time::timeout(Duration::from_millis(timeout_ms), acquire) => result,
} {
Ok(Ok(permit)) => permit,
Ok(Err(_)) => {
error!("Connection limiter is closed");
return;
}
Err(_) => {
runtime.stats.increment_accept_permit_timeout_total();
debug!(
timeout_ms,
"Dropping accepted Unix connection: permit wait timeout"
);
continue;
}
}
};
let connection_id = connection_counter.fetch_add(1, Ordering::Relaxed);
let fake_peer =
SocketAddr::from(([127, 0, 0, 1], (connection_id % 65535) as u16));
let stats = runtime.stats.clone();
let upstream_manager = runtime.upstream_manager.clone();
let replay_checker = runtime.replay_checker.clone();
let buffer_pool = runtime.buffer_pool.clone();
let rng = runtime.rng.clone();
let me_pool = runtime.me_pool.clone();
let me_pool_runtime = runtime.me_pool_runtime.clone();
let route_runtime = runtime.route_runtime.clone();
let tls_cache = runtime.tls_cache.clone();
let ip_tracker = runtime.ip_tracker.clone();
let beobachten = runtime.beobachten.clone();
let shared = runtime.proxy_shared.clone();
let proxy_protocol_enabled = config.server.proxy_protocol;
let _ = runtime.spawn_session(async move {
let _permit = permit;
if let Err(error_value) =
crate::proxy::client::handle_client_stream_with_shared_and_pool_runtime(
stream,
fake_peer,
config,
stats,
upstream_manager,
replay_checker,
buffer_pool,
rng,
me_pool,
Some(me_pool_runtime),
route_runtime,
tls_cache,
ip_tracker,
beobachten,
shared,
proxy_protocol_enabled,
)
.await
{
Ok(Ok(permit)) => permit,
Ok(Err(_)) => {
error!("Connection limiter is closed");
break;
}
Err(_) => {
runtime.stats.increment_accept_permit_timeout_total();
debug!(
timeout_ms,
"Dropping accepted unix connection: permit wait timeout"
);
drop(stream);
continue;
}
}
};
let connection_id = connection_counter.fetch_add(1, Ordering::Relaxed);
let fake_peer =
SocketAddr::from(([127, 0, 0, 1], (connection_id % 65535) as u16));
let stats = runtime.stats.clone();
let upstream_manager = runtime.upstream_manager.clone();
let replay_checker = runtime.replay_checker.clone();
let buffer_pool = runtime.buffer_pool.clone();
let rng = runtime.rng.clone();
let me_pool = runtime.me_pool.clone();
let me_pool_runtime = runtime.me_pool_runtime.clone();
let route_runtime = runtime.route_runtime.clone();
let tls_cache = runtime.tls_cache.clone();
let ip_tracker = runtime.ip_tracker.clone();
let beobachten = runtime.beobachten.clone();
let shared = runtime.proxy_shared.clone();
let proxy_protocol_enabled = config.server.proxy_protocol;
let _ = runtime.spawn_session(async move {
let _permit = permit;
if let Err(error) =
crate::proxy::client::handle_client_stream_with_shared_and_pool_runtime(
stream,
fake_peer,
config,
stats,
upstream_manager,
replay_checker,
buffer_pool,
rng,
me_pool,
Some(me_pool_runtime),
route_runtime,
tls_cache,
ip_tracker,
beobachten,
shared,
proxy_protocol_enabled,
)
.await
{
debug!(error = %error, "Unix socket connection error");
}
});
}
Err(error) => {
error!(error = %error, "Unix socket accept error");
tokio::time::sleep(Duration::from_millis(100)).await;
{
debug!(error = %error_value, "Unix socket connection error");
}
});
}
Err(error_value) => {
error!(error = %error_value, "Unix socket accept error");
tokio::select! {
biased;
_ = cancellation.cancelled() => return,
_ = tokio::time::sleep(Duration::from_millis(100)) => {}
}
}
}
});
}
}
impl UnixAcceptHandle {
pub(super) fn start(
listener: UnixListener,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
) -> Self {
let listener = Arc::new(listener);
let cancellation = CancellationToken::new();
let task = tokio::spawn(run_unix_accept_loop(
listener.clone(),
active_runtime,
cancellation.clone(),
));
Self {
_listener: listener,
cancellation,
task: Some(task),
}
}
pub(super) async fn stop(&mut self) -> Result<(), String> {
self.cancellation.cancel();
if let Some(task) = self.task.take() {
task.await
.map_err(|error_value| format!("Unix listener task failed: {error_value}"))?;
}
Ok(())
}
}
+29 -981
View File
File diff suppressed because it is too large Load Diff
+342
View File
@@ -0,0 +1,342 @@
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use arc_swap::ArcSwap;
use tokio::sync::{RwLock, watch};
use tracing::{error, info, warn};
use crate::api;
use crate::ip_tracker::UserIpTracker;
use crate::network::probe::{decide_network_capabilities, log_probe_result, run_probe};
use crate::proxy::direct_buffer_budget::{
DirectBufferBudget, resolve_direct_buffer_hard_limit,
};
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::shared_state::ProxySharedState;
use crate::startup::{COMPONENT_API_BOOTSTRAP, COMPONENT_NETWORK_PROBE};
use crate::stats::telemetry::TelemetryPolicy;
use crate::stats::{QuotaStore, Stats};
use crate::synlimit_control;
use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool;
use super::{
bootstrap, generation, listeners, reload, reload_supervisor, runtime_startup, runtime_tasks,
shutdown, tls_bootstrap,
};
// Shared maestro startup and main loop. `drop_after_bind` runs on Unix after listeners are bound
// and privileged firewall setup completes; it is a no-op on other platforms.
pub(super) async fn run_telemt_core(
privilege_drop_requested: bool,
drop_after_bind: impl FnOnce(),
) -> std::result::Result<(), Box<dyn std::error::Error>> {
let bootstrap::BootstrapState {
process_started_at,
process_started_at_epoch_secs,
startup_tracker,
config,
config_path,
has_rust_log,
effective_log_level,
runtime_log_filter,
logging_guard: _logging_guard,
} = bootstrap::bootstrap(privilege_drop_requested).await?;
let quota_store = Arc::new(QuotaStore::default());
let stats = Arc::new(Stats::with_quota_store(quota_store.clone()));
let runtime_task_scope = generation::RuntimeTaskScope::new();
stats.apply_telemetry_policy(TelemetryPolicy::from_config(&config.general.telemetry));
let quota_state_path = config.general.quota_state_path.clone();
crate::quota_state::load_quota_state(&quota_state_path, stats.as_ref()).await;
let upstream_manager = Arc::new(
UpstreamManager::new(
config.upstreams.clone(),
config.general.upstream_connect_retry_attempts,
config.general.upstream_connect_retry_backoff_ms,
config.general.upstream_connect_budget_ms,
config.general.tg_connect,
config.general.upstream_unhealthy_fail_threshold,
config.general.upstream_connect_failfast_hard_errors,
stats.clone(),
)
.with_dns_overrides(&config.network.dns_overrides)?,
);
let ip_tracker = Arc::new(UserIpTracker::new());
ip_tracker
.load_limits(
config.access.user_max_unique_ips_global_each,
&config.access.user_max_unique_ips,
)
.await;
ip_tracker
.set_limit_policy(
config.access.user_max_unique_ips_mode,
config.access.user_max_unique_ips_window_secs,
)
.await;
if config.access.user_max_unique_ips_global_each > 0
|| !config.access.user_max_unique_ips.is_empty()
{
info!(
global_each_limit = config.access.user_max_unique_ips_global_each,
explicit_user_limits = config.access.user_max_unique_ips.len(),
"User unique IP limits configured"
);
}
if !config.network.dns_overrides.is_empty() {
info!(
"Runtime DNS overrides configured: {} entries",
config.network.dns_overrides.len()
);
}
let direct_buffer_hard_limit =
resolve_direct_buffer_hard_limit(config.general.direct_relay_buffer_budget_max_bytes).await;
let direct_buffer_budget = DirectBufferBudget::new(direct_buffer_hard_limit);
info!(
hard_limit_bytes = direct_buffer_hard_limit,
configured_override_bytes = config.general.direct_relay_buffer_budget_max_bytes,
"Direct relay buffer budget initialized"
);
let shared_state =
ProxySharedState::new_with_direct_buffer_budget(direct_buffer_budget.clone());
shared_state.apply_user_enabled_config(&config.access.user_enabled);
shared_state.traffic_limiter.apply_policy(
config.access.user_rate_limits.clone(),
config.access.cidr_rate_limits.clone(),
);
let (detected_ips_tx, detected_ips_rx) = watch::channel((None::<IpAddr>, None::<IpAddr>));
let initial_direct_first = config.general.use_middle_proxy && config.general.me2dc_fallback;
let initial_admission_open = !config.general.use_middle_proxy || initial_direct_first;
let (admission_tx, admission_rx) = watch::channel(initial_admission_open);
let (reload_control, reload_commands) = reload::ReloadControl::channel(1);
let (active_runtime_tx, active_runtime_rx) =
watch::channel(None::<Arc<ArcSwap<generation::RuntimeGeneration>>>);
let (runtime_watch_tx, runtime_watch_rx) =
watch::channel(None::<generation::RuntimeWatchState>);
let initial_route_mode = if !config.general.use_middle_proxy || initial_direct_first {
RelayRouteMode::Direct
} else {
RelayRouteMode::Middle
};
let route_runtime = Arc::new(RouteRuntimeController::new(initial_route_mode));
let api_me_pool = Arc::new(RwLock::new(None::<Arc<MePool>>));
startup_tracker
.start_component(
COMPONENT_API_BOOTSTRAP,
Some("spawn API listener task".to_string()),
)
.await;
if config.server.api.enabled {
let listen = match config.server.api.listen.parse::<SocketAddr>() {
Ok(listen) => listen,
Err(error) => {
warn!(
error = %error,
listen = %config.server.api.listen,
"Invalid server.api.listen; API is disabled"
);
SocketAddr::from(([127, 0, 0, 1], 0))
}
};
if listen.port() != 0 {
let stats_api = stats.clone();
let ip_tracker_api = ip_tracker.clone();
let me_pool_api = api_me_pool.clone();
let upstream_manager_api = upstream_manager.clone();
let route_runtime_api = route_runtime.clone();
let proxy_shared_api = shared_state.clone();
let config_path_api = config_path.clone();
let quota_state_path_api = quota_state_path.clone();
let startup_tracker_api = startup_tracker.clone();
let detected_ips_rx_api = detected_ips_rx.clone();
let reload_control_api = reload_control.clone();
let active_runtime_rx_api = active_runtime_rx.clone();
let runtime_watch_rx_api = runtime_watch_rx.clone();
tokio::spawn(async move {
api::serve(
listen,
stats_api,
ip_tracker_api,
me_pool_api,
route_runtime_api,
proxy_shared_api,
upstream_manager_api,
config_path_api,
quota_state_path_api,
detected_ips_rx_api,
process_started_at_epoch_secs,
startup_tracker_api,
reload_control_api,
active_runtime_rx_api,
runtime_watch_rx_api,
)
.await;
});
startup_tracker
.complete_component(
COMPONENT_API_BOOTSTRAP,
Some(format!("api task spawned on {}", listen)),
)
.await;
} else {
startup_tracker
.skip_component(
COMPONENT_API_BOOTSTRAP,
Some("server.api.listen has zero port".to_string()),
)
.await;
}
} else {
startup_tracker
.skip_component(
COMPONENT_API_BOOTSTRAP,
Some("server.api.enabled is false".to_string()),
)
.await;
}
let mut tls_domains = Vec::with_capacity(1 + config.censorship.tls_domains.len());
tls_domains.push(config.censorship.tls_domain.clone());
for domain in &config.censorship.tls_domains {
if !tls_domains.contains(domain) {
tls_domains.push(domain.clone());
}
}
let tls_cache = tls_bootstrap::bootstrap_tls_front(
&config,
&tls_domains,
upstream_manager.clone(),
&startup_tracker,
runtime_task_scope.clone(),
tls_bootstrap::TlsBootstrapPolicy::BestEffort,
)
.await?;
startup_tracker
.start_component(
COMPONENT_NETWORK_PROBE,
Some("probe network capabilities".to_string()),
)
.await;
let probe = run_probe(
&config.network,
&config.upstreams,
config.general.middle_proxy_nat_probe,
config.general.stun_nat_probe_concurrency,
)
.await?;
detected_ips_tx.send_replace((
probe.detected_ipv4.map(IpAddr::V4),
probe.detected_ipv6.map(IpAddr::V6),
));
let decision =
decide_network_capabilities(&config.network, &probe, config.general.middle_proxy_nat_ip);
log_probe_result(&probe, &decision);
startup_tracker
.complete_component(
COMPONENT_NETWORK_PROBE,
Some("network capabilities determined".to_string()),
)
.await;
let runtime = runtime_startup::prepare_runtime(
config,
&config_path,
&probe,
&decision,
process_started_at,
&startup_tracker,
stats.clone(),
upstream_manager.clone(),
ip_tracker.clone(),
shared_state.clone(),
direct_buffer_budget,
route_runtime.clone(),
api_me_pool.clone(),
runtime_task_scope.clone(),
admission_tx,
&runtime_log_filter,
has_rust_log,
&effective_log_level,
)
.await;
let _admission_tx_hold = runtime.admission_tx;
let runtime_generation = generation::RuntimeGeneration::new(
1,
runtime.config_rx.clone(),
admission_rx,
stats.clone(),
upstream_manager.clone(),
runtime.replay_checker,
runtime.buffer_pool,
runtime.rng,
runtime.me_pool,
api_me_pool,
route_runtime,
tls_cache,
ip_tracker,
runtime.beobachten,
shared_state,
runtime.max_connections,
runtime_task_scope,
);
let active_runtime = Arc::new(ArcSwap::from(runtime_generation));
let bound = listeners::bind_listeners(
&runtime.config,
runtime.detected_ip_v4,
runtime.detected_ip_v6,
&startup_tracker,
)
.await?;
if bound.is_empty() {
error!("No listeners. Exiting.");
std::process::exit(1);
}
synlimit_control::reconcile_synlimit_rules(&runtime.config)
.await
.map_err(std::io::Error::other)?;
drop_after_bind();
runtime_tasks::spawn_metrics_if_configured(
&runtime.config,
&startup_tracker,
active_runtime.clone(),
)
.await;
runtime_watch_tx.send_replace(Some(active_runtime.load_full().watch_state()));
active_runtime_tx.send_replace(Some(active_runtime.clone()));
runtime_tasks::mark_runtime_ready(&startup_tracker).await;
let listener_manager = listeners::ListenerManager::start(bound, active_runtime.clone());
let reload_supervisor = reload_supervisor::ReloadSupervisor::spawn(
active_runtime.clone(),
reload_control,
reload_commands,
config_path,
quota_store,
detected_ips_tx,
runtime_log_filter,
runtime_watch_tx,
listener_manager,
);
shutdown::spawn_signal_handlers(active_runtime.clone(), process_started_at);
shutdown::wait_for_shutdown(
process_started_at,
active_runtime,
quota_state_path,
reload_supervisor,
)
.await;
Ok(())
}
+99 -6
View File
@@ -4,12 +4,14 @@ use std::time::Duration;
use arc_swap::ArcSwap;
use tokio::sync::watch;
use tokio::sync::Mutex;
use tokio_util::sync::CancellationToken;
use tracing::{info, warn};
use crate::stats::QuotaStore;
use super::generation::{RuntimeGeneration, RuntimeWatchState};
use super::listeners::{ListenerManager, PreparedListenerTransition};
use super::reload::{
ReloadCommand, ReloadCommandReceiver, ReloadControl, ReloadFailurePolicy, ReloadMode,
ReloadPhase,
@@ -26,6 +28,7 @@ pub(crate) struct ReloadSupervisor {
detected_ips_tx: watch::Sender<(Option<std::net::IpAddr>, Option<std::net::IpAddr>)>,
runtime_log_filter: RuntimeLogFilter,
runtime_watch_tx: watch::Sender<Option<RuntimeWatchState>>,
listener_manager: Arc<Mutex<ListenerManager>>,
}
/// Process-owned handle that quiesces reloads before shutdown snapshots the runtime.
@@ -33,16 +36,18 @@ pub(crate) struct ReloadSupervisorHandle {
control: ReloadControl,
shutdown: CancellationToken,
join: tokio::task::JoinHandle<()>,
listener_manager: Arc<Mutex<ListenerManager>>,
}
impl ReloadSupervisorHandle {
/// Stops new submissions and waits for the accepted reload to finish.
pub(crate) async fn quiesce(self) {
pub(crate) async fn quiesce(self) -> Arc<Mutex<ListenerManager>> {
self.control.begin_shutdown().await;
self.shutdown.cancel();
if let Err(error) = self.join.await {
warn!(error = %error, "Reload supervisor failed while quiescing");
}
self.listener_manager
}
}
@@ -99,7 +104,9 @@ impl ReloadSupervisor {
detected_ips_tx: watch::Sender<(Option<std::net::IpAddr>, Option<std::net::IpAddr>)>,
runtime_log_filter: RuntimeLogFilter,
runtime_watch_tx: watch::Sender<Option<RuntimeWatchState>>,
listener_manager: ListenerManager,
) -> ReloadSupervisorHandle {
let listener_manager = Arc::new(Mutex::new(listener_manager));
let supervisor = Self {
active_runtime,
control,
@@ -109,6 +116,7 @@ impl ReloadSupervisor {
detected_ips_tx,
runtime_log_filter,
runtime_watch_tx,
listener_manager: listener_manager.clone(),
};
let control = supervisor.control.clone();
let shutdown = CancellationToken::new();
@@ -117,6 +125,7 @@ impl ReloadSupervisor {
control,
shutdown,
join,
listener_manager,
}
}
@@ -171,18 +180,41 @@ impl ReloadSupervisor {
}
};
let listener_transition = match self
.listener_manager
.lock()
.await
.prepare_transition(prepared.generation.config().as_ref())
{
Ok(transition) => transition,
Err(error) => {
let _ = cleanup_candidate(&prepared.generation).await;
self.runtime_log_filter
.apply_reload(&old_runtime.config().general.log_level);
self.control.fail(command.reload_id, error).await;
return;
}
};
let revision_action = revision_gate_action(
&command.config_revision,
crate::api::config_store::current_revision_for_maestro(&self.config_path).await,
command.request.failure_policy,
);
self.activate_prepared(command, old_runtime, prepared, revision_action, |entries| {
crate::network::dns_overrides::install_entries(entries)
.map_err(|error| error.to_string())
})
self.activate_prepared_with_transition(
command,
old_runtime,
prepared,
listener_transition,
revision_action,
|entries| {
crate::network::dns_overrides::install_entries(entries)
.map_err(|error| error.to_string())
},
)
.await;
}
#[cfg(test)]
async fn activate_prepared<InstallDns>(
&self,
command: ReloadCommand,
@@ -192,6 +224,41 @@ impl ReloadSupervisor {
install_dns: InstallDns,
) where
InstallDns: FnOnce(&[String]) -> Result<(), String>,
{
let listener_transition = match self
.listener_manager
.lock()
.await
.prepare_transition(prepared.generation.config().as_ref())
{
Ok(transition) => transition,
Err(error) => {
let _ = cleanup_candidate(&prepared.generation).await;
self.control.fail(command.reload_id, error).await;
return;
}
};
self.activate_prepared_with_transition(
command,
old_runtime,
prepared,
listener_transition,
revision_action,
install_dns,
)
.await;
}
async fn activate_prepared_with_transition<InstallDns>(
&self,
command: ReloadCommand,
old_runtime: Arc<RuntimeGeneration>,
prepared: PreparedRuntime,
listener_transition: Option<PreparedListenerTransition>,
revision_action: RevisionGateAction,
install_dns: InstallDns,
) where
InstallDns: FnOnce(&[String]) -> Result<(), String>,
{
match revision_action {
RevisionGateAction::Proceed => {}
@@ -211,7 +278,6 @@ impl ReloadSupervisor {
.mark_phase(command.reload_id, ReloadPhase::Activating)
.await;
let new_runtime = prepared.generation;
old_runtime.stop_accepting_sessions();
if let Err(error) = install_dns(&new_runtime.config().network.dns_overrides) {
let message = format!("runtime DNS activation failed: {}", error);
if command.request.failure_policy == ReloadFailurePolicy::Rollback {
@@ -224,7 +290,34 @@ impl ReloadSupervisor {
}
self.control.add_warning(command.reload_id, message).await;
}
let pending_listener_transition = if let Some(listener_transition) = listener_transition {
match self
.listener_manager
.lock()
.await
.begin_transition(listener_transition)
.await
{
Ok(pending) => Some(pending),
Err(error) => {
let _ = cleanup_candidate(&new_runtime).await;
self.runtime_log_filter
.apply_reload(&old_runtime.config().general.log_level);
self.control.fail(command.reload_id, error).await;
return;
}
}
} else {
None
};
old_runtime.stop_accepting_sessions();
let replaced = self.active_runtime.swap(new_runtime.clone());
if let Some(pending) = pending_listener_transition {
self.listener_manager
.lock()
.await
.finish_transition(pending);
}
self.detected_ips_tx.send_replace(prepared.detected_ips);
self.runtime_log_filter
.apply_reload(&new_runtime.config().general.log_level);
+4
View File
@@ -33,6 +33,7 @@ async fn fixture(request: ReloadRequest) -> ReloadFixture {
.unwrap();
let (detected_ips_tx, _detected_ips_rx) = watch::channel((None, None));
let (runtime_watch_tx, runtime_watch_rx) = watch::channel(Some(old_runtime.watch_state()));
let listener_manager = Arc::new(Mutex::new(ListenerManager::empty(active_runtime.clone())));
let supervisor = Arc::new(ReloadSupervisor {
active_runtime,
control: control.clone(),
@@ -42,6 +43,7 @@ async fn fixture(request: ReloadRequest) -> ReloadFixture {
detected_ips_tx,
runtime_log_filter: runtime_log_filter(),
runtime_watch_tx,
listener_manager,
});
let command = ReloadCommand {
reload_id: accepted.reload_id,
@@ -293,6 +295,7 @@ async fn quiesce_joins_idle_supervisor_and_rejects_later_submissions() {
let (control, commands) = ReloadControl::channel(runtime.id);
let (detected_ips_tx, _detected_ips_rx) = watch::channel((None, None));
let (runtime_watch_tx, _runtime_watch_rx) = watch::channel(Some(runtime.watch_state()));
let listener_manager = ListenerManager::empty(active_runtime.clone());
let handle = ReloadSupervisor::spawn(
active_runtime,
control.clone(),
@@ -302,6 +305,7 @@ async fn quiesce_joins_idle_supervisor_and_rejects_later_submissions() {
detected_ips_tx,
runtime_log_filter(),
runtime_watch_tx,
listener_manager,
);
tokio::time::timeout(Duration::from_secs(1), handle.quiesce())
+8 -6
View File
@@ -24,6 +24,7 @@ use crate::transport::middle_proxy::MePool;
use super::admission;
use super::generation::{RuntimeGeneration, RuntimeTaskScope};
use super::listeners::listener_rebind_supported;
use super::runtime_tasks::RuntimeLogFilter;
use super::{me_startup, runtime_tasks, tls_bootstrap};
@@ -326,18 +327,19 @@ pub(crate) fn resolve_reload_config(
let mut effective = desired.clone();
let mut fields = Vec::new();
let listener_identity_matches = listeners_have_same_bind_identity(&old.server, &desired.server);
let listener_process_fields_changed = !listener_identity_matches
|| !listener_process_fields_equal(&old.server, &desired.server);
if old.server.port != desired.server.port
let global_listener_policy_changed = old.server.port != desired.server.port
|| old.server.listen_addr_ipv4 != desired.server.listen_addr_ipv4
|| old.server.listen_addr_ipv6 != desired.server.listen_addr_ipv6
|| old.server.listen_tcp != desired.server.listen_tcp
|| old.server.client_mss != desired.server.client_mss
|| old.server.client_mss_bulk != desired.server.client_mss_bulk
|| old.server.proxy_protocol != desired.server.proxy_protocol
|| old.server.listen_backlog != desired.server.listen_backlog
|| listener_process_fields_changed
{
|| old.server.listen_backlog != desired.server.listen_backlog;
let listener_policy_changed =
listener_identity_matches && !listener_process_fields_equal(&old.server, &desired.server);
let unsupported_identity_change = !listener_identity_matches
&& !listener_rebind_supported(old, desired);
if global_listener_policy_changed || listener_policy_changed || unsupported_identity_change {
fields.push("server.listeners".to_string());
effective.server.port = old.server.port;
effective.server.listen_addr_ipv4 = old.server.listen_addr_ipv4.clone();
+53
View File
@@ -1,5 +1,26 @@
use super::*;
fn test_listener(port: u16) -> crate::config::ListenerConfig {
crate::config::ListenerConfig {
ip: "127.0.0.1".parse().unwrap(),
port: Some(port),
client_mss: None,
synlimit: crate::config::SynLimitMode::Off,
synlimit_seconds: 60,
synlimit_hitcount: 48,
synlimit_burst: 24,
synlimit_ios_seconds: 1,
synlimit_ios_hitcount: 12,
synlimit_ios_burst: 24,
synlimit_hashlimit_expire_ms: 60_000,
synlimit_hashlimit_size: 32_768,
announce: None,
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
}
}
#[test]
fn process_socket_and_logging_changes_are_deferred() {
let old = ProxyConfig::default();
@@ -129,3 +150,35 @@ fn strict_middle_proxy_requires_a_prepared_pool() {
assert!(!strict_middle_proxy_unavailable(true, true, false));
assert!(!strict_middle_proxy_unavailable(false, false, false));
}
#[test]
fn endpoint_only_listener_move_is_runtime_rebindable() {
let mut old = ProxyConfig::default();
old.server.listeners = vec![test_listener(443)];
let mut desired = old.clone();
desired.server.listeners[0].port = Some(8443);
let resolved = resolve_reload_config(&old, &desired);
assert!(resolved.deferred_process_fields.is_empty());
assert_eq!(resolved.effective.server.listeners[0].port, Some(8443));
assert!(resolved.runtime_changed);
}
#[test]
fn synlimited_endpoint_move_remains_restart_only() {
let mut old = ProxyConfig::default();
old.server.listeners = vec![test_listener(443)];
old.server.listeners[0].synlimit = crate::config::SynLimitMode::Nftables;
let mut desired = old.clone();
desired.server.listeners[0].port = Some(8443);
let resolved = resolve_reload_config(&old, &desired);
assert_eq!(
resolved.deferred_process_fields,
vec!["server.listeners".to_string()]
);
assert_eq!(resolved.effective.server.listeners[0].port, Some(443));
assert!(!resolved.runtime_changed);
}
+369
View File
@@ -0,0 +1,369 @@
use std::net::IpAddr;
use std::path::Path;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::{RwLock, Semaphore, watch};
use tracing::{info, warn};
use crate::config::{LogLevel, ProxyConfig};
use crate::conntrack_control;
use crate::crypto::SecureRandom;
use crate::ip_tracker::UserIpTracker;
use crate::network::probe::{NetworkDecision, NetworkProbe};
use crate::proxy::direct_buffer_budget::{
DirectBufferBudget, run_direct_buffer_budget_controller,
};
use crate::proxy::route_mode::RouteRuntimeController;
use crate::proxy::shared_state::ProxySharedState;
use crate::startup::{
COMPONENT_DC_CONNECTIVITY_PING, COMPONENT_ME_CONNECTIVITY_PING,
COMPONENT_ME_POOL_CONSTRUCT, COMPONENT_ME_POOL_INIT_STAGE1, COMPONENT_ME_PROXY_CONFIG_V4,
COMPONENT_ME_PROXY_CONFIG_V6, COMPONENT_ME_SECRET_FETCH, StartupMeStatus, StartupTracker,
};
use crate::stats::beobachten::BeobachtenStore;
use crate::stats::{ReplayChecker, Stats};
use crate::stream::BufferPool;
use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool;
use super::admission;
use super::generation::RuntimeTaskScope;
use super::{connectivity, me_startup, runtime_tasks};
pub(super) struct RuntimeStartupState {
pub(super) config: Arc<ProxyConfig>,
pub(super) beobachten: Arc<BeobachtenStore>,
pub(super) rng: Arc<SecureRandom>,
pub(super) max_connections: Arc<Semaphore>,
pub(super) me_pool: Option<Arc<MePool>>,
pub(super) replay_checker: Arc<ReplayChecker>,
pub(super) buffer_pool: Arc<BufferPool>,
pub(super) config_rx: watch::Receiver<Arc<ProxyConfig>>,
pub(super) detected_ip_v4: Option<IpAddr>,
pub(super) detected_ip_v6: Option<IpAddr>,
pub(super) admission_tx: watch::Sender<bool>,
}
#[allow(clippy::too_many_arguments)]
pub(super) async fn prepare_runtime(
mut config: ProxyConfig,
config_path: &Path,
probe: &NetworkProbe,
decision: &NetworkDecision,
process_started_at: Instant,
startup_tracker: &Arc<StartupTracker>,
stats: Arc<Stats>,
upstream_manager: Arc<UpstreamManager>,
ip_tracker: Arc<UserIpTracker>,
shared_state: Arc<ProxySharedState>,
direct_buffer_budget: Arc<DirectBufferBudget>,
route_runtime: Arc<RouteRuntimeController>,
api_me_pool: Arc<RwLock<Option<Arc<MePool>>>>,
runtime_task_scope: RuntimeTaskScope,
admission_tx: watch::Sender<bool>,
runtime_log_filter: &runtime_tasks::RuntimeLogFilter,
has_rust_log: bool,
effective_log_level: &LogLevel,
) -> RuntimeStartupState {
let prefer_ipv6 = decision.prefer_ipv6();
let mut use_middle_proxy = config.general.use_middle_proxy;
let beobachten = Arc::new(BeobachtenStore::new());
let rng = Arc::new(SecureRandom::new());
let max_connections_limit = if config.server.max_connections == 0 {
Semaphore::MAX_PERMITS
} else {
config.server.max_connections as usize
};
let max_connections = Arc::new(Semaphore::new(max_connections_limit));
let me2dc_fallback = config.general.me2dc_fallback;
let me_init_retry_attempts = config.general.me_init_retry_attempts;
if use_middle_proxy && !decision.ipv4_me && !decision.ipv6_me {
if me2dc_fallback {
warn!(
"No usable IP family for Middle Proxy detected; Direct-DC startup fallback is active while ME init retries continue"
);
} else {
warn!(
"No usable IP family for Middle Proxy detected; me2dc_fallback=false, ME init retries stay active"
);
}
}
if use_middle_proxy {
startup_tracker
.set_me_status(StartupMeStatus::Initializing, COMPONENT_ME_SECRET_FETCH)
.await;
startup_tracker
.start_component(
COMPONENT_ME_SECRET_FETCH,
Some("fetch proxy-secret from source/cache".to_string()),
)
.await;
startup_tracker
.set_me_retry_limit(if !me2dc_fallback || me_init_retry_attempts == 0 {
"unlimited".to_string()
} else {
me_init_retry_attempts.to_string()
})
.await;
} else {
startup_tracker
.set_me_status(StartupMeStatus::Skipped, "skipped")
.await;
startup_tracker
.skip_component(
COMPONENT_ME_SECRET_FETCH,
Some("middle proxy mode disabled".to_string()),
)
.await;
startup_tracker
.skip_component(
COMPONENT_ME_PROXY_CONFIG_V4,
Some("middle proxy mode disabled".to_string()),
)
.await;
startup_tracker
.skip_component(
COMPONENT_ME_PROXY_CONFIG_V6,
Some("middle proxy mode disabled".to_string()),
)
.await;
startup_tracker
.skip_component(
COMPONENT_ME_POOL_CONSTRUCT,
Some("middle proxy mode disabled".to_string()),
)
.await;
startup_tracker
.skip_component(
COMPONENT_ME_POOL_INIT_STAGE1,
Some("middle proxy mode disabled".to_string()),
)
.await;
}
let (me_ready_tx, me_ready_rx) = watch::channel(0_u64);
let direct_first_startup = use_middle_proxy && me2dc_fallback;
let me_pool: Option<Arc<MePool>> = if direct_first_startup {
None
} else {
me_startup::initialize_me_pool(
use_middle_proxy,
&config,
decision,
probe,
startup_tracker,
upstream_manager.clone(),
rng.clone(),
stats.clone(),
api_me_pool.clone(),
me_ready_tx.clone(),
runtime_task_scope.clone(),
)
.await
};
if direct_first_startup {
startup_tracker.set_transport_mode("direct").await;
startup_tracker.set_degraded(true).await;
info!(
"Transport: Direct DC startup fallback active; Middle-End bootstrap continues in background"
);
} else if me_pool.is_some() {
startup_tracker.set_transport_mode("middle_proxy").await;
startup_tracker.set_degraded(false).await;
info!("Transport: Middle-End Proxy - all DC-over-RPC");
} else {
let _ = use_middle_proxy;
use_middle_proxy = false;
config.general.use_middle_proxy = false;
startup_tracker.set_transport_mode("direct").await;
startup_tracker.set_degraded(true).await;
if me2dc_fallback {
startup_tracker
.set_me_status(StartupMeStatus::Failed, "fallback_to_direct")
.await;
} else {
startup_tracker
.set_me_status(StartupMeStatus::Skipped, "skipped")
.await;
}
info!("Transport: Direct DC - TCP - standard DC-over-TCP");
}
let config = Arc::new(config);
let replay_checker = Arc::new(ReplayChecker::new(
config.access.replay_check_len,
Duration::from_secs(config.access.replay_window_secs),
));
let buffer_pool = Arc::new(BufferPool::with_config(64 * 1024, 4096));
if direct_first_startup {
startup_tracker
.skip_component(
COMPONENT_ME_CONNECTIVITY_PING,
Some("deferred by direct-first startup".to_string()),
)
.await;
startup_tracker
.skip_component(
COMPONENT_DC_CONNECTIVITY_PING,
Some("background health checks active".to_string()),
)
.await;
} else {
connectivity::run_startup_connectivity(
&config,
&me_pool,
rng.clone(),
startup_tracker,
upstream_manager.clone(),
prefer_ipv6,
decision,
process_started_at,
api_me_pool.clone(),
)
.await;
}
let runtime_watches = runtime_tasks::spawn_runtime_tasks(
&config,
config_path,
probe,
prefer_ipv6,
decision.ipv4_dc,
decision.ipv6_dc,
startup_tracker,
stats.clone(),
upstream_manager.clone(),
replay_checker.clone(),
me_pool.clone(),
rng.clone(),
ip_tracker.clone(),
beobachten.clone(),
me_pool.clone(),
shared_state.clone(),
me_ready_tx.clone(),
runtime_task_scope.clone(),
)
.await;
let config_rx = runtime_watches.config_rx;
let log_level_rx = runtime_watches.log_level_rx;
let detected_ip_v4 = runtime_watches.detected_ip_v4;
let detected_ip_v6 = runtime_watches.detected_ip_v6;
runtime_log_filter.start(
has_rust_log,
effective_log_level,
log_level_rx,
runtime_task_scope.clone(),
);
if direct_first_startup {
let config_bg = config.clone();
let decision_bg = decision.clone();
let probe_bg = probe.clone();
let startup_tracker_bg = startup_tracker.clone();
let upstream_manager_bg = upstream_manager.clone();
let rng_bg = rng.clone();
let stats_bg = stats.clone();
let api_me_pool_bg = api_me_pool.clone();
let me_ready_tx_bg = me_ready_tx.clone();
let config_rx_bg = config_rx.clone();
let task_scope_bg = runtime_task_scope.clone();
runtime_task_scope.spawn(async move {
let mut bootstrap_attempt: u32 = 0;
loop {
bootstrap_attempt = bootstrap_attempt.saturating_add(1);
let pool = me_startup::initialize_me_pool(
true,
config_bg.as_ref(),
&decision_bg,
&probe_bg,
&startup_tracker_bg,
upstream_manager_bg.clone(),
rng_bg.clone(),
stats_bg.clone(),
api_me_pool_bg.clone(),
me_ready_tx_bg.clone(),
task_scope_bg.clone(),
)
.await;
if let Some(pool) = pool {
runtime_tasks::spawn_middle_proxy_runtime_tasks(
config_bg.as_ref(),
config_rx_bg,
pool,
rng_bg,
me_ready_tx_bg,
task_scope_bg,
);
break;
}
if me_init_retry_attempts > 0 && bootstrap_attempt >= me_init_retry_attempts {
break;
}
tokio::time::sleep(Duration::from_secs(2)).await;
}
});
let startup_tracker_ready = startup_tracker.clone();
let api_me_pool_ready = api_me_pool.clone();
let mut me_ready_rx_transport = me_ready_tx.subscribe();
runtime_task_scope.spawn(async move {
if me_ready_rx_transport.changed().await.is_ok() {
if let Some(pool) = api_me_pool_ready.read().await.as_ref() {
pool.set_runtime_ready(true);
}
startup_tracker_ready
.set_transport_mode("middle_proxy")
.await;
startup_tracker_ready.set_degraded(false).await;
info!("Transport: Middle-End Proxy restored for new sessions");
}
});
}
admission::configure_admission_gate(
&config,
me_pool.clone(),
api_me_pool,
route_runtime,
&admission_tx,
config_rx.clone(),
me_ready_rx,
runtime_task_scope.clone(),
)
.await;
let conntrack_scope = runtime_task_scope.clone();
runtime_task_scope.spawn(conntrack_control::run_conntrack_controller(
config_rx.clone(),
stats.clone(),
shared_state.clone(),
conntrack_scope.cancellation_token(),
));
runtime_task_scope.spawn(run_direct_buffer_budget_controller(
direct_buffer_budget,
buffer_pool.clone(),
stats,
shared_state,
config.server.max_connections,
));
RuntimeStartupState {
config,
beobachten,
rng,
max_connections,
me_pool,
replay_checker,
buffer_pool,
config_rx,
detected_ip_v4,
detected_ip_v6,
admission_tx,
}
}
+5 -1
View File
@@ -95,7 +95,7 @@ async fn perform_shutdown(
let shutdown_started_at = Instant::now();
info!(signal = %signal, "Received shutdown signal");
reload_supervisor.quiesce().await;
let listener_manager = reload_supervisor.quiesce().await;
let runtime = active_runtime.load_full();
let stats = runtime.stats.as_ref();
@@ -108,6 +108,10 @@ async fn perform_shutdown(
let uptime_secs = process_started_at.elapsed().as_secs();
info!("Uptime: {}", format_uptime(uptime_secs));
if let Err(error) = listener_manager.lock().await.shutdown().await {
warn!(error = %error, "Failed to stop one or more listener tasks cleanly");
}
// Graceful ME pool shutdown
runtime.stop_sessions().await;
runtime.stop_background_tasks().await;