Merge branch 'flow/3.5.0' into flow-bulk-mss

This commit is contained in:
Alexey
2026-08-12 21:05:14 +03:00
committed by GitHub
39 changed files with 4551 additions and 920 deletions
+4 -1
View File
@@ -8,6 +8,8 @@ use crate::config::ProxyConfig;
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::transport::middle_proxy::MePool;
use super::generation::RuntimeTaskScope;
const STARTUP_FALLBACK_AFTER: Duration = Duration::from_secs(80);
const RUNTIME_FALLBACK_AFTER: Duration = Duration::from_secs(6);
@@ -19,6 +21,7 @@ pub(crate) async fn configure_admission_gate(
admission_tx: &watch::Sender<bool>,
config_rx: watch::Receiver<Arc<ProxyConfig>>,
me_ready_rx: watch::Receiver<u64>,
task_scope: RuntimeTaskScope,
) {
if config.general.use_middle_proxy {
if me_pool.is_some() || config.general.me2dc_fallback {
@@ -64,7 +67,7 @@ pub(crate) async fn configure_admission_gate(
let mut config_rx_gate = config_rx.clone();
let mut me_ready_rx_gate = me_ready_rx;
let mut admission_poll_ms = config.general.me_admission_poll_ms.max(1);
tokio::spawn(async move {
task_scope.spawn(async move {
let mut gate_open = initial_gate_open;
let mut route_mode = initial_route_mode;
let mut ready_observed = initial_ready;
+435
View File
@@ -0,0 +1,435 @@
use std::future::Future;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use tokio::sync::{RwLock, Semaphore, watch};
use tokio_util::sync::CancellationToken;
use tokio_util::task::TaskTracker;
use crate::config::ProxyConfig;
use crate::crypto::SecureRandom;
use crate::ip_tracker::UserIpTracker;
#[cfg(test)]
use crate::proxy::route_mode::RelayRouteMode;
use crate::proxy::route_mode::RouteRuntimeController;
use crate::proxy::shared_state::ProxySharedState;
use crate::stats::beobachten::BeobachtenStore;
use crate::stats::{ReplayChecker, Stats};
use crate::stream::BufferPool;
use crate::tls_front::TlsFrontCache;
use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool;
const SESSION_STOP_TIMEOUT: Duration = Duration::from_secs(5);
const BACKGROUND_STOP_TIMEOUT: Duration = Duration::from_secs(5);
const SESSION_ADMISSION_CLOSED: usize = 1 << (usize::BITS - 1);
const SESSION_REGISTRATION_COUNT: usize = SESSION_ADMISSION_CLOSED - 1;
struct SessionAdmission {
state: AtomicUsize,
}
struct SessionRegistration<'a> {
admission: &'a SessionAdmission,
}
impl SessionAdmission {
fn new() -> Self {
Self {
state: AtomicUsize::new(0),
}
}
fn try_register(&self) -> Option<SessionRegistration<'_>> {
let mut state = self.state.load(Ordering::Acquire);
loop {
if state & SESSION_ADMISSION_CLOSED != 0
|| state & SESSION_REGISTRATION_COUNT == SESSION_REGISTRATION_COUNT
{
return None;
}
match self.state.compare_exchange_weak(
state,
state + 1,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return Some(SessionRegistration { admission: self }),
Err(observed) => state = observed,
}
}
}
fn close(&self) {
self.state
.fetch_or(SESSION_ADMISSION_CLOSED, Ordering::AcqRel);
}
fn reopen(&self) {
self.state
.fetch_and(!SESSION_ADMISSION_CLOSED, Ordering::AcqRel);
}
async fn wait_for_registrations(&self) {
while self.state.load(Ordering::Acquire) & SESSION_REGISTRATION_COUNT != 0 {
tokio::task::yield_now().await;
}
}
}
impl Drop for SessionRegistration<'_> {
fn drop(&mut self) {
self.admission.state.fetch_sub(1, Ordering::Release);
}
}
/// Process-visible control-plane receivers for one active runtime generation.
#[derive(Clone)]
pub(crate) struct RuntimeWatchState {
pub(crate) generation_id: u64,
pub(crate) config_rx: watch::Receiver<Arc<ProxyConfig>>,
pub(crate) admission_rx: watch::Receiver<bool>,
}
/// Cancellation and join ownership for one generation's background tasks.
#[derive(Clone)]
pub(crate) struct RuntimeTaskScope {
tracker: TaskTracker,
cancel: CancellationToken,
}
impl RuntimeTaskScope {
/// Creates an open generation-owned task scope.
pub(crate) fn new() -> Self {
Self {
tracker: TaskTracker::new(),
cancel: CancellationToken::new(),
}
}
/// Spawns one task that is cancelled when the generation stops.
pub(crate) fn spawn<F>(&self, future: F)
where
F: Future<Output = ()> + Send + 'static,
{
let cancel = self.cancel.clone();
self.tracker.spawn(async move {
tokio::select! {
_ = cancel.cancelled() => {}
_ = future => {}
}
});
}
/// Returns the cancellation signal shared by generation-owned controllers.
pub(crate) fn cancellation_token(&self) -> CancellationToken {
self.cancel.clone()
}
/// Cancels the scope and waits within the bounded background-task budget.
pub(crate) async fn stop(&self) {
self.cancel.cancel();
self.tracker.close();
let _ = tokio::time::timeout(BACKGROUND_STOP_TIMEOUT, self.tracker.wait()).await;
}
}
/// Runtime-owned data plane and control-plane dependencies for one generation.
pub(crate) struct RuntimeGeneration {
pub(crate) id: u64,
pub(crate) config_rx: watch::Receiver<Arc<ProxyConfig>>,
pub(crate) admission_rx: watch::Receiver<bool>,
pub(crate) stats: Arc<Stats>,
pub(crate) upstream_manager: Arc<UpstreamManager>,
pub(crate) replay_checker: Arc<ReplayChecker>,
pub(crate) buffer_pool: Arc<BufferPool>,
pub(crate) rng: Arc<SecureRandom>,
pub(crate) me_pool: Option<Arc<MePool>>,
pub(crate) me_pool_runtime: Arc<RwLock<Option<Arc<MePool>>>>,
pub(crate) route_runtime: Arc<RouteRuntimeController>,
pub(crate) tls_cache: Option<Arc<TlsFrontCache>>,
pub(crate) ip_tracker: Arc<UserIpTracker>,
pub(crate) beobachten: Arc<BeobachtenStore>,
pub(crate) proxy_shared: Arc<ProxySharedState>,
pub(crate) max_connections: Arc<Semaphore>,
background_tasks: RuntimeTaskScope,
sessions: TaskTracker,
session_cancel: CancellationToken,
session_admission: SessionAdmission,
}
impl RuntimeGeneration {
#[allow(clippy::too_many_arguments)]
/// Builds one fully owned runtime generation.
pub(crate) fn new(
id: u64,
config_rx: watch::Receiver<Arc<ProxyConfig>>,
admission_rx: watch::Receiver<bool>,
stats: Arc<Stats>,
upstream_manager: Arc<UpstreamManager>,
replay_checker: Arc<ReplayChecker>,
buffer_pool: Arc<BufferPool>,
rng: Arc<SecureRandom>,
me_pool: Option<Arc<MePool>>,
me_pool_runtime: Arc<RwLock<Option<Arc<MePool>>>>,
route_runtime: Arc<RouteRuntimeController>,
tls_cache: Option<Arc<TlsFrontCache>>,
ip_tracker: Arc<UserIpTracker>,
beobachten: Arc<BeobachtenStore>,
proxy_shared: Arc<ProxySharedState>,
max_connections: Arc<Semaphore>,
background_tasks: RuntimeTaskScope,
) -> Arc<Self> {
Arc::new(Self {
id,
config_rx,
admission_rx,
stats,
upstream_manager,
replay_checker,
buffer_pool,
rng,
me_pool,
me_pool_runtime,
route_runtime,
tls_cache,
ip_tracker,
beobachten,
proxy_shared,
max_connections,
background_tasks,
sessions: TaskTracker::new(),
session_cancel: CancellationToken::new(),
session_admission: SessionAdmission::new(),
})
}
/// Returns the latest hot-reloaded configuration for this generation.
pub(crate) fn config(&self) -> Arc<ProxyConfig> {
self.config_rx.borrow().clone()
}
/// Returns receivers used by process-scoped observers of this generation.
pub(crate) fn watch_state(&self) -> RuntimeWatchState {
RuntimeWatchState {
generation_id: self.id,
config_rx: self.config_rx.clone(),
admission_rx: self.admission_rx.clone(),
}
}
/// Returns the initial or asynchronously published Middle-End pool.
pub(crate) async fn current_me_pool(&self) -> Option<Arc<MePool>> {
if let Some(pool) = &self.me_pool {
return Some(pool.clone());
}
self.me_pool_runtime.read().await.clone()
}
/// Registers a session only while admission remains open.
pub(crate) fn spawn_session<F>(&self, future: F) -> bool
where
F: Future<Output = ()> + Send + 'static,
{
let Some(_registration) = self.session_admission.try_register() else {
return false;
};
let cancel = self.session_cancel.clone();
self.sessions.spawn(async move {
tokio::select! {
_ = cancel.cancelled() => {}
_ = future => {}
}
});
true
}
/// Closes admission while preserving already registered sessions.
pub(crate) fn stop_accepting_sessions(&self) {
self.session_admission.close();
}
/// Reopens admission after a candidate activation rolls back.
pub(crate) fn resume_accepting_sessions(&self) {
self.session_admission.reopen();
}
/// Waits for registered sessions and cancels them when the deadline expires.
pub(crate) async fn drain_sessions(&self, timeout: Duration) -> bool {
self.stop_accepting_sessions();
self.session_admission.wait_for_registrations().await;
self.sessions.close();
if tokio::time::timeout(timeout, self.sessions.wait())
.await
.is_ok()
{
return true;
}
self.stop_sessions().await;
false
}
/// Cancels all sessions and waits within the bounded session-stop budget.
pub(crate) async fn stop_sessions(&self) {
self.stop_accepting_sessions();
self.session_admission.wait_for_registrations().await;
self.session_cancel.cancel();
self.sessions.close();
let _ = tokio::time::timeout(SESSION_STOP_TIMEOUT, self.sessions.wait()).await;
}
/// Stops all background tasks owned by this generation.
pub(crate) async fn stop_background_tasks(&self) {
self.background_tasks.stop().await;
}
}
#[cfg(test)]
/// Builds a lightweight runtime generation without network startup tasks.
pub(super) fn test_runtime_generation(id: u64, config: ProxyConfig) -> Arc<RuntimeGeneration> {
let (config_tx, config_rx) = watch::channel(Arc::new(config.clone()));
let (_admission_tx, admission_rx) = watch::channel(true);
let stats = Arc::new(Stats::new());
let upstream_manager = Arc::new(UpstreamManager::new(
config.upstreams,
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(),
));
let _config_tx = config_tx;
RuntimeGeneration::new(
id,
config_rx,
admission_rx,
stats,
upstream_manager,
Arc::new(ReplayChecker::new(128, Duration::from_secs(60))),
Arc::new(BufferPool::with_config(4096, 16)),
Arc::new(SecureRandom::new()),
None,
Arc::new(RwLock::new(None)),
Arc::new(RouteRuntimeController::new(RelayRouteMode::Direct)),
None,
Arc::new(UserIpTracker::new()),
Arc::new(BeobachtenStore::new()),
ProxySharedState::new(),
Arc::new(Semaphore::new(64)),
RuntimeTaskScope::new(),
)
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::sync::Barrier;
#[tokio::test]
async fn stop_sessions_cancels_tracked_future() {
let generation = test_runtime_generation(1, ProxyConfig::default());
let started = Arc::new(tokio::sync::Notify::new());
let dropped = Arc::new(tokio::sync::Notify::new());
let started_task = started.clone();
let dropped_task = dropped.clone();
assert!(generation.spawn_session(async move {
struct DropSignal(Arc<tokio::sync::Notify>);
impl Drop for DropSignal {
fn drop(&mut self) {
self.0.notify_one();
}
}
let _drop_signal = DropSignal(dropped_task);
started_task.notify_one();
std::future::pending::<()>().await;
}));
started.notified().await;
generation.stop_sessions().await;
tokio::time::timeout(Duration::from_secs(1), dropped.notified())
.await
.unwrap();
assert!(!generation.spawn_session(async {}));
}
#[tokio::test]
async fn runtime_task_scope_joins_cancelled_background_task() {
let scope = RuntimeTaskScope::new();
scope.spawn(std::future::pending());
tokio::time::timeout(Duration::from_secs(1), scope.stop())
.await
.unwrap();
}
#[tokio::test]
async fn session_admission_waits_for_registration_started_before_cutover() {
let admission = Arc::new(SessionAdmission::new());
let registration = admission.try_register().unwrap();
admission.close();
assert!(admission.try_register().is_none());
let wait_admission = admission.clone();
let waiter = tokio::spawn(async move {
wait_admission.wait_for_registrations().await;
});
tokio::task::yield_now().await;
assert!(!waiter.is_finished());
drop(registration);
tokio::time::timeout(Duration::from_secs(1), waiter)
.await
.unwrap()
.unwrap();
admission.reopen();
assert!(admission.try_register().is_some());
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn cutover_never_leaves_late_session_registrations() {
const ATTEMPTS: usize = 10_000;
let admission = Arc::new(SessionAdmission::new());
let tracker = TaskTracker::new();
let cancel = CancellationToken::new();
let start = Arc::new(Barrier::new(ATTEMPTS + 1));
let live = Arc::new(AtomicUsize::new(0));
let mut attempts = tokio::task::JoinSet::new();
for _ in 0..ATTEMPTS {
let admission = admission.clone();
let tracker = tracker.clone();
let cancel = cancel.clone();
let start = start.clone();
let live = live.clone();
attempts.spawn(async move {
start.wait().await;
let Some(_registration) = admission.try_register() else {
return;
};
tracker.spawn(async move {
live.fetch_add(1, Ordering::AcqRel);
cancel.cancelled().await;
live.fetch_sub(1, Ordering::AcqRel);
});
});
}
start.wait().await;
admission.close();
admission.wait_for_registrations().await;
tracker.close();
cancel.cancel();
while attempts.join_next().await.is_some() {}
tokio::time::timeout(Duration::from_secs(1), tracker.wait())
.await
.unwrap();
assert_eq!(live.load(Ordering::Acquire), 0);
assert!(admission.try_register().is_none());
}
}
+117
View File
@@ -0,0 +1,117 @@
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use arc_swap::ArcSwap;
use tokio::net::UnixListener;
use tracing::{debug, error};
use super::RuntimeGeneration;
pub(crate) fn spawn_unix_accept_loop(
listener: Option<UnixListener>,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
) {
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 {
Ok(permit) => permit,
Err(_) => {
error!("Connection limiter is closed");
break;
}
}
} else {
match tokio::time::timeout(
Duration::from_millis(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!(
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;
}
}
}
});
}
+184 -200
View File
@@ -1,9 +1,11 @@
#![allow(clippy::too_many_arguments)]
use std::future::Future;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::{RwLock, watch};
use tokio_util::task::AbortOnDropHandle;
use tracing::{error, info, warn};
use crate::config::ProxyConfig;
@@ -17,8 +19,61 @@ use crate::stats::Stats;
use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool;
use super::generation::RuntimeTaskScope;
use super::helpers::load_startup_proxy_config_snapshot;
async fn supervise_me_task<F, Fut>(task_name: &'static str, mut task: F)
where
F: FnMut() -> Fut,
Fut: Future<Output = ()> + Send + 'static,
{
loop {
let result = AbortOnDropHandle::new(tokio::spawn(task())).await;
match result {
Ok(()) => warn!(
task = task_name,
"Middle-End supervisor task exited unexpectedly, restarting"
),
Err(error) => {
error!(task = task_name, error = %error, "Middle-End supervisor task panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
}
fn spawn_me_supervisors(
task_scope: RuntimeTaskScope,
pool: Arc<MePool>,
rng: Arc<SecureRandom>,
min_connections: usize,
) {
let health_pool = pool.clone();
let health_rng = rng;
task_scope.spawn(supervise_me_task("health_monitor", move || {
let pool = health_pool.clone();
let rng = health_rng.clone();
async move {
crate::transport::middle_proxy::me_health_monitor(pool, rng, min_connections).await;
}
}));
let drain_pool = pool.clone();
task_scope.spawn(supervise_me_task("drain_timeout_enforcer", move || {
let pool = drain_pool.clone();
async move {
crate::transport::middle_proxy::me_drain_timeout_enforcer(pool).await;
}
}));
task_scope.spawn(supervise_me_task("zombie_writer_watchdog", move || {
let pool = pool.clone();
async move {
crate::transport::middle_proxy::me_zombie_writer_watchdog(pool).await;
}
}));
}
pub(crate) async fn initialize_me_pool(
use_middle_proxy: bool,
config: &ProxyConfig,
@@ -30,6 +85,7 @@ pub(crate) async fn initialize_me_pool(
stats: Arc<Stats>,
api_me_pool: Arc<RwLock<Option<Arc<MePool>>>>,
me_ready_tx: watch::Sender<u64>,
task_scope: RuntimeTaskScope,
) -> Option<Arc<MePool>> {
if !use_middle_proxy {
return None;
@@ -52,15 +108,8 @@ pub(crate) async fn initialize_me_pool(
.as_ref()
.map(|tag| hex::decode(tag).expect("general.ad_tag must be validated before startup"));
// =============================================================
// CRITICAL: Download Telegram proxy-secret (NOT user secret!)
//
// C MTProxy uses TWO separate secrets:
// -S flag = 16-byte user secret for client obfuscation
// --aes-pwd = 32-512 byte binary file for ME RPC auth
//
// proxy-secret is from: https://core.telegram.org/getProxySecret
// =============================================================
// The Telegram proxy-secret authenticates ME RPC and is distinct from client secrets.
// It corresponds to the C MTProxy --aes-pwd input and may be fetched from Telegram.
let proxy_secret_path = config.general.proxy_secret_path.as_deref();
let pool_size = config.general.middle_proxy_pool_size.max(1);
let proxy_secret = loop {
@@ -319,143 +368,70 @@ pub(crate) async fn initialize_me_pool(
let rng_bg = rng.clone();
let startup_tracker_bg = startup_tracker.clone();
let me_ready_tx_bg = me_ready_tx.clone();
let task_scope_bg = task_scope.clone();
let retry_limit = if me_init_retry_attempts == 0 {
String::from("unlimited")
} else {
me_init_retry_attempts.to_string()
};
std::thread::spawn(move || {
let runtime = match tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
{
Ok(runtime) => runtime,
Err(error) => {
error!(error = %error, "Failed to build background runtime for ME initialization");
return;
}
};
runtime.block_on(async move {
let mut init_attempt: u32 = 0;
loop {
init_attempt = init_attempt.saturating_add(1);
startup_tracker_bg.set_me_init_attempt(init_attempt).await;
match pool_bg.init(pool_size, &rng_bg).await {
Ok(()) => {
startup_tracker_bg.set_me_last_error(None).await;
startup_tracker_bg
.complete_component(
COMPONENT_ME_POOL_INIT_STAGE1,
Some("ME pool initialized".to_string()),
)
.await;
startup_tracker_bg
.set_me_status(StartupMeStatus::Ready, "ready")
.await;
me_ready_tx_bg.send_modify(|version| {
*version = version.saturating_add(1);
});
info!(
task_scope.spawn(async move {
let mut init_attempt: u32 = 0;
loop {
init_attempt = init_attempt.saturating_add(1);
startup_tracker_bg.set_me_init_attempt(init_attempt).await;
match pool_bg.init(pool_size, &rng_bg).await {
Ok(()) => {
startup_tracker_bg.set_me_last_error(None).await;
startup_tracker_bg
.complete_component(
COMPONENT_ME_POOL_INIT_STAGE1,
Some("ME pool initialized".to_string()),
)
.await;
startup_tracker_bg
.set_me_status(StartupMeStatus::Ready, "ready")
.await;
me_ready_tx_bg.send_modify(|version| {
*version = version.saturating_add(1);
});
info!(
attempt = init_attempt,
"Middle-End pool initialized successfully"
);
spawn_me_supervisors(
task_scope_bg,
pool_bg.clone(),
rng_bg.clone(),
pool_size,
);
break;
}
Err(e) => {
startup_tracker_bg
.set_me_last_error(Some(e.to_string()))
.await;
if init_attempt >= me_init_warn_after_attempts {
warn!(
error = %e,
attempt = init_attempt,
"Middle-End pool initialized successfully"
retry_limit = %retry_limit,
retry_in_secs = 2,
"ME pool is not ready yet; retrying background initialization"
);
} else {
info!(
error = %e,
attempt = init_attempt,
retry_limit = %retry_limit,
retry_in_secs = 2,
"ME pool startup warmup: retrying background initialization"
);
// ── Supervised background tasks ──────────────────
// Each task runs inside a nested tokio::spawn so
// that a panic is caught via JoinHandle and the
// outer loop restarts the task automatically.
let pool_health = pool_bg.clone();
let rng_health = rng_bg.clone();
let min_conns = pool_size;
tokio::spawn(async move {
loop {
let p = pool_health.clone();
let r = rng_health.clone();
let res = tokio::spawn(async move {
crate::transport::middle_proxy::me_health_monitor(
p, r, min_conns,
)
.await;
})
.await;
match res {
Ok(()) => warn!("me_health_monitor exited unexpectedly, restarting"),
Err(e) => {
error!(error = %e, "me_health_monitor panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
});
let pool_drain_enforcer = pool_bg.clone();
tokio::spawn(async move {
loop {
let p = pool_drain_enforcer.clone();
let res = tokio::spawn(async move {
crate::transport::middle_proxy::me_drain_timeout_enforcer(p).await;
})
.await;
match res {
Ok(()) => warn!("me_drain_timeout_enforcer exited unexpectedly, restarting"),
Err(e) => {
error!(error = %e, "me_drain_timeout_enforcer panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
});
let pool_watchdog = pool_bg.clone();
tokio::spawn(async move {
loop {
let p = pool_watchdog.clone();
let res = tokio::spawn(async move {
crate::transport::middle_proxy::me_zombie_writer_watchdog(p).await;
})
.await;
match res {
Ok(()) => warn!("me_zombie_writer_watchdog exited unexpectedly, restarting"),
Err(e) => {
error!(error = %e, "me_zombie_writer_watchdog panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
});
// CRITICAL: keep the current-thread runtime
// alive. Without this, block_on() returns,
// the Runtime is dropped, and ALL spawned
// background tasks (health monitor, drain
// enforcer, zombie watchdog) are silently
// cancelled — causing the draining-writer
// leak that brought us here.
std::future::pending::<()>().await;
unreachable!();
}
Err(e) => {
startup_tracker_bg.set_me_last_error(Some(e.to_string())).await;
if init_attempt >= me_init_warn_after_attempts {
warn!(
error = %e,
attempt = init_attempt,
retry_limit = %retry_limit,
retry_in_secs = 2,
"ME pool is not ready yet; retrying background initialization"
);
} else {
info!(
error = %e,
attempt = init_attempt,
retry_limit = %retry_limit,
retry_in_secs = 2,
"ME pool startup warmup: retrying background initialization"
);
}
pool_bg.reset_stun_state();
tokio::time::sleep(Duration::from_secs(2)).await;
}
pool_bg.reset_stun_state();
tokio::time::sleep(Duration::from_secs(2)).await;
}
}
});
}
});
startup_tracker
.set_me_status(StartupMeStatus::Initializing, "background_init")
@@ -490,70 +466,12 @@ pub(crate) async fn initialize_me_pool(
"Middle-End pool initialized successfully"
);
// ── Supervised background tasks ──────────────────
let pool_clone = pool.clone();
let rng_clone = rng.clone();
let min_conns = pool_size;
tokio::spawn(async move {
loop {
let p = pool_clone.clone();
let r = rng_clone.clone();
let res = tokio::spawn(async move {
crate::transport::middle_proxy::me_health_monitor(
p, r, min_conns,
)
.await;
})
.await;
match res {
Ok(()) => warn!(
"me_health_monitor exited unexpectedly, restarting"
),
Err(e) => {
error!(error = %e, "me_health_monitor panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
});
let pool_drain_enforcer = pool.clone();
tokio::spawn(async move {
loop {
let p = pool_drain_enforcer.clone();
let res = tokio::spawn(async move {
crate::transport::middle_proxy::me_drain_timeout_enforcer(p).await;
})
.await;
match res {
Ok(()) => warn!(
"me_drain_timeout_enforcer exited unexpectedly, restarting"
),
Err(e) => {
error!(error = %e, "me_drain_timeout_enforcer panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
});
let pool_watchdog = pool.clone();
tokio::spawn(async move {
loop {
let p = pool_watchdog.clone();
let res = tokio::spawn(async move {
crate::transport::middle_proxy::me_zombie_writer_watchdog(p).await;
})
.await;
match res {
Ok(()) => warn!(
"me_zombie_writer_watchdog exited unexpectedly, restarting"
),
Err(e) => {
error!(error = %e, "me_zombie_writer_watchdog panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
});
spawn_me_supervisors(
task_scope.clone(),
pool.clone(),
rng.clone(),
pool_size,
);
break Some(pool);
}
@@ -666,3 +584,69 @@ pub(crate) async fn initialize_me_pool(
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::Notify;
struct DropSignal(Arc<Notify>);
impl Drop for DropSignal {
fn drop(&mut self) {
self.0.notify_one();
}
}
#[tokio::test]
async fn scoped_supervisor_aborts_its_current_child() {
let scope = RuntimeTaskScope::new();
let dropped = Arc::new(Notify::new());
let dropped_for_task = dropped.clone();
scope.spawn(supervise_me_task("test", move || {
let dropped = dropped_for_task.clone();
async move {
let _signal = DropSignal(dropped);
std::future::pending::<()>().await;
}
}));
tokio::task::yield_now().await;
scope.stop().await;
tokio::time::timeout(Duration::from_secs(1), dropped.notified())
.await
.unwrap();
}
#[tokio::test]
async fn supervisor_restarts_exited_child_and_stops_with_runtime_scope() {
let scope = RuntimeTaskScope::new();
let starts = Arc::new(AtomicUsize::new(0));
let restarted = Arc::new(Notify::new());
let starts_task = starts.clone();
let restarted_task = restarted.clone();
scope.spawn(supervise_me_task("restart_test", move || {
let starts = starts_task.clone();
let restarted = restarted_task.clone();
async move {
if starts.fetch_add(1, Ordering::AcqRel) + 1 >= 3 {
restarted.notify_one();
}
}
}));
tokio::time::timeout(Duration::from_secs(1), restarted.notified())
.await
.unwrap();
scope.stop().await;
let stopped_at = starts.load(Ordering::Acquire);
for _ in 0..100 {
tokio::task::yield_now().await;
}
assert!(stopped_at >= 3);
assert_eq!(starts.load(Ordering::Acquire), stopped_at);
}
}
+113 -84
View File
@@ -13,19 +13,24 @@
// - shutdown: graceful shutdown sequence and uptime logging.
mod admission;
mod connectivity;
pub(crate) mod generation;
mod helpers;
mod listeners;
mod me_startup;
pub(crate) mod reload;
mod reload_supervisor;
pub(crate) mod runtime_build;
mod runtime_tasks;
mod shutdown;
mod tls_bootstrap;
use arc_swap::ArcSwap;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use tokio::sync::{RwLock, Semaphore, watch};
use tracing::{error, info, warn};
use tracing_subscriber::{EnvFilter, fmt, prelude::*, reload};
use tracing_subscriber::{EnvFilter, fmt, prelude::*, reload as tracing_reload};
use crate::api;
use crate::config::{LogLevel, ProxyConfig};
@@ -34,7 +39,7 @@ use crate::crypto::SecureRandom;
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, spawn_direct_buffer_budget_controller,
DirectBufferBudget, resolve_direct_buffer_hard_limit, run_direct_buffer_budget_controller,
};
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::shared_state::ProxySharedState;
@@ -46,7 +51,7 @@ use crate::startup::{
};
use crate::stats::beobachten::BeobachtenStore;
use crate::stats::telemetry::TelemetryPolicy;
use crate::stats::{ReplayChecker, Stats};
use crate::stats::{QuotaStore, ReplayChecker, Stats};
use crate::stream::BufferPool;
use crate::synlimit_control;
use crate::transport::UpstreamManager;
@@ -343,7 +348,7 @@ async fn run_telemt_core(
}
};
let (filter_layer, filter_handle) =
reload::Layer::new(EnvFilter::new(initial_filter_spec.clone()));
tracing_reload::Layer::new(EnvFilter::new(initial_filter_spec.clone()));
startup_tracker
.start_component(
COMPONENT_TRACING_INIT,
@@ -387,6 +392,7 @@ async fn run_telemt_core(
_logging_guard = Some(guard);
}
}
let runtime_log_filter = runtime_tasks::RuntimeLogFilter::new(filter_handle);
startup_tracker
.complete_component(
@@ -433,21 +439,26 @@ async fn run_telemt_core(
warn!("Using default tls_domain. Consider setting a custom domain.");
}
let stats = Arc::new(Stats::new());
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(),
));
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(
@@ -492,11 +503,15 @@ async fn run_telemt_core(
config.access.cidr_rate_limits.clone(),
);
let (api_config_tx, api_config_rx) = watch::channel(Arc::new(config.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 {
@@ -530,12 +545,13 @@ async fn run_telemt_core(
let upstream_manager_api = upstream_manager.clone();
let route_runtime_api = route_runtime.clone();
let proxy_shared_api = shared_state.clone();
let config_rx_api = api_config_rx.clone();
let admission_rx_api = admission_rx.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,
@@ -545,13 +561,14 @@ async fn run_telemt_core(
route_runtime_api,
proxy_shared_api,
upstream_manager_api,
config_rx_api,
admission_rx_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;
});
@@ -591,8 +608,10 @@ async fn run_telemt_core(
&tls_domains,
upstream_manager.clone(),
&startup_tracker,
runtime_task_scope.clone(),
tls_bootstrap::TlsBootstrapPolicy::BestEffort,
)
.await;
.await?;
startup_tracker
.start_component(
@@ -718,6 +737,7 @@ async fn run_telemt_core(
stats.clone(),
api_me_pool.clone(),
me_ready_tx.clone(),
runtime_task_scope.clone(),
)
.await
};
@@ -805,16 +825,22 @@ async fn run_telemt_core(
rng.clone(),
ip_tracker.clone(),
beobachten.clone(),
api_config_tx.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();
@@ -827,7 +853,8 @@ async fn run_telemt_core(
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();
tokio::spawn(async move {
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);
@@ -842,6 +869,7 @@ async fn run_telemt_core(
stats_bg.clone(),
api_me_pool_bg.clone(),
me_ready_tx_bg.clone(),
task_scope_bg.clone(),
)
.await;
if let Some(pool) = pool {
@@ -851,6 +879,7 @@ async fn run_telemt_core(
pool,
rng_bg,
me_ready_tx_bg,
task_scope_bg,
);
break;
}
@@ -864,7 +893,7 @@ async fn run_telemt_core(
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();
tokio::spawn(async move {
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);
@@ -886,20 +915,54 @@ async fn run_telemt_core(
&admission_tx,
config_rx.clone(),
me_ready_rx,
runtime_task_scope.clone(),
)
.await;
let _admission_tx_hold = admission_tx;
conntrack_control::spawn_conntrack_controller(
let conntrack_scope = runtime_task_scope.clone();
runtime_task_scope.spawn(conntrack_control::run_conntrack_controller(
config_rx.clone(),
stats.clone(),
shared_state.clone(),
);
spawn_direct_buffer_budget_controller(
conntrack_scope.cancellation_token(),
));
runtime_task_scope.spawn(run_direct_buffer_budget_controller(
direct_buffer_budget,
buffer_pool.clone(),
stats.clone(),
shared_state.clone(),
config.server.max_connections,
));
let runtime_generation = generation::RuntimeGeneration::new(
1,
config_rx.clone(),
admission_rx.clone(),
stats.clone(),
upstream_manager.clone(),
replay_checker.clone(),
buffer_pool.clone(),
rng.clone(),
me_pool.clone(),
api_me_pool.clone(),
route_runtime.clone(),
tls_cache.clone(),
ip_tracker.clone(),
beobachten.clone(),
shared_state.clone(),
max_connections.clone(),
runtime_task_scope.clone(),
);
let active_runtime = Arc::new(ArcSwap::from(runtime_generation));
let reload_supervisor = reload_supervisor::ReloadSupervisor::spawn(
active_runtime.clone(),
reload_control,
reload_commands,
config_path.clone(),
quota_store,
detected_ips_tx,
runtime_log_filter,
runtime_watch_tx.clone(),
);
let bound = listeners::bind_listeners(
@@ -909,25 +972,15 @@ async fn run_telemt_core(
detected_ip_v4,
detected_ip_v6,
&startup_tracker,
config_rx.clone(),
admission_rx.clone(),
stats.clone(),
upstream_manager.clone(),
replay_checker.clone(),
buffer_pool.clone(),
rng.clone(),
me_pool.clone(),
api_me_pool.clone(),
route_runtime.clone(),
tls_cache.clone(),
ip_tracker.clone(),
beobachten.clone(),
shared_state.clone(),
max_connections.clone(),
)
.await?;
let listeners = bound.listeners;
let has_unix_listener = bound.has_unix_listener;
#[cfg(unix)]
let unix_listener = bound.unix_listener;
#[cfg(unix)]
let has_unix_listener = unix_listener.is_some();
#[cfg(not(unix))]
let has_unix_listener = false;
if listeners.is_empty() && !has_unix_listener {
error!("No listeners. Exiting.");
@@ -937,54 +990,30 @@ async fn run_telemt_core(
// On Unix, caller supplies privilege drop after bind (may require root for port < 1024).
drop_after_bind();
synlimit_control::reconcile_synlimit_rules(&config).await;
synlimit_control::spawn_synlimit_controller(config_rx.clone());
let synlimit_controller = synlimit_control::spawn_synlimit_controller(runtime_watch_rx);
runtime_tasks::apply_runtime_log_filter(
has_rust_log,
&effective_log_level,
filter_handle,
log_level_rx,
)
.await;
runtime_tasks::spawn_metrics_if_configured(
&config,
&startup_tracker,
stats.clone(),
beobachten.clone(),
shared_state.clone(),
ip_tracker.clone(),
tls_cache.clone(),
config_rx.clone(),
)
.await;
runtime_tasks::spawn_metrics_if_configured(&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;
// Spawn signal handlers for SIGUSR1/SIGUSR2 (non-shutdown signals)
shutdown::spawn_signal_handlers(stats.clone(), process_started_at);
shutdown::spawn_signal_handlers(active_runtime.clone(), process_started_at);
listeners::spawn_tcp_accept_loops(
listeners,
config_rx.clone(),
admission_rx.clone(),
stats.clone(),
upstream_manager.clone(),
replay_checker.clone(),
buffer_pool.clone(),
rng.clone(),
me_pool.clone(),
api_me_pool.clone(),
route_runtime.clone(),
tls_cache.clone(),
ip_tracker.clone(),
beobachten.clone(),
shared_state,
max_connections.clone(),
);
listeners::spawn_tcp_accept_loops(listeners, active_runtime.clone());
#[cfg(unix)]
listeners::spawn_unix_accept_loop(unix_listener, active_runtime.clone());
shutdown::wait_for_shutdown(process_started_at, me_pool, stats, quota_state_path).await;
shutdown::wait_for_shutdown(
process_started_at,
active_runtime,
quota_state_path,
synlimit_controller,
reload_supervisor,
)
.await;
Ok(())
}
+449
View File
@@ -0,0 +1,449 @@
use std::collections::VecDeque;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
use serde::{Deserialize, Serialize};
use tokio::sync::{Mutex, mpsc};
use crate::config::ProxyConfig;
const RELOAD_HISTORY_CAPACITY: usize = 32;
const RELOAD_COMMAND_CAPACITY: usize = 1;
const MAX_DRAIN_TIMEOUT_SECS: u64 = 3_600;
/// Session handling policy for an in-process runtime reload.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize, Serialize)]
#[serde(rename_all = "snake_case")]
pub(crate) enum ReloadMode {
#[default]
Instant,
Drain,
}
/// Failure policy applied during the activation barrier.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize, Serialize)]
#[serde(rename_all = "snake_case")]
pub(crate) enum ReloadFailurePolicy {
#[default]
KeepNew,
Rollback,
}
/// Request body accepted by the maestro reload endpoint.
#[derive(Debug, Clone, Default, PartialEq, Eq, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub(crate) struct ReloadRequest {
#[serde(default)]
pub(crate) mode: ReloadMode,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub(crate) timeout_secs: Option<u64>,
#[serde(default)]
pub(crate) failure_policy: ReloadFailurePolicy,
}
impl ReloadRequest {
/// Validates mode-specific request parameters.
pub(crate) fn validate(&self) -> Result<(), &'static str> {
match (self.mode, self.timeout_secs) {
(ReloadMode::Instant, None) => Ok(()),
(ReloadMode::Instant, Some(_)) => Err("timeout_secs is only valid when mode is drain"),
(ReloadMode::Drain, Some(1..=MAX_DRAIN_TIMEOUT_SECS)) => Ok(()),
(ReloadMode::Drain, Some(_)) => Err("timeout_secs must be within 1..=3600"),
(ReloadMode::Drain, None) => Err("timeout_secs is required when mode is drain"),
}
}
/// Parses optional PATCH query parameters into a reload request.
pub(crate) fn from_query(query: Option<&str>) -> Result<Option<Self>, String> {
let Some(query) = query.filter(|query| !query.is_empty()) else {
return Ok(None);
};
let mut mode = None;
let mut timeout_secs = None;
let mut failure_policy = None;
for (key, value) in url::form_urlencoded::parse(query.as_bytes()) {
match key.as_ref() {
"reload" if mode.is_none() => {
mode = Some(match value.as_ref() {
"instant" => ReloadMode::Instant,
"drain" => ReloadMode::Drain,
_ => return Err("reload must be instant or drain".to_string()),
});
}
"timeout_secs" if timeout_secs.is_none() => {
timeout_secs = Some(
value
.parse::<u64>()
.map_err(|_| "timeout_secs must be an integer".to_string())?,
);
}
"failure_policy" if failure_policy.is_none() => {
failure_policy = Some(match value.as_ref() {
"keep_new" => ReloadFailurePolicy::KeepNew,
"rollback" => ReloadFailurePolicy::Rollback,
_ => {
return Err("failure_policy must be keep_new or rollback".to_string());
}
});
}
"reload" | "timeout_secs" | "failure_policy" => {
return Err(format!("duplicate query parameter: {}", key));
}
_ => return Err(format!("unknown query parameter: {}", key)),
}
}
let mode = mode.ok_or_else(|| "reload query parameter is required".to_string())?;
let request = Self {
mode,
timeout_secs,
failure_policy: failure_policy.unwrap_or_default(),
};
request.validate().map_err(str::to_string)?;
Ok(Some(request))
}
}
/// Observable phase of one reload operation.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(rename_all = "snake_case")]
pub(crate) enum ReloadPhase {
Accepted,
Preparing,
Activating,
Draining,
Succeeded,
RolledBack,
Failed,
}
impl ReloadPhase {
fn is_terminal(self) -> bool {
matches!(
self,
ReloadPhase::Succeeded | ReloadPhase::RolledBack | ReloadPhase::Failed
)
}
}
/// Bounded public status for one reload operation.
#[derive(Debug, Clone, Serialize)]
pub(crate) struct ReloadStatus {
pub(crate) reload_id: u64,
pub(crate) target_generation: u64,
pub(crate) config_revision: String,
pub(crate) state: ReloadPhase,
pub(crate) mode: ReloadMode,
pub(crate) failure_policy: ReloadFailurePolicy,
pub(crate) requested_at_epoch_secs: u64,
#[serde(skip_serializing_if = "Option::is_none")]
pub(crate) started_at_epoch_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub(crate) finished_at_epoch_secs: Option<u64>,
#[serde(
rename = "deferred_process_fields",
default,
skip_serializing_if = "Vec::is_empty"
)]
pub(crate) deferred_fields: Vec<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub(crate) warnings: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub(crate) error: Option<String>,
}
/// Accepted operation metadata returned before asynchronous preparation starts.
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub(crate) struct ReloadAccepted {
pub(crate) reload_id: u64,
pub(crate) target_generation: u64,
pub(crate) config_revision: String,
pub(crate) state: ReloadPhase,
pub(crate) mode: ReloadMode,
pub(crate) failure_policy: ReloadFailurePolicy,
}
pub(crate) struct ReloadCommand {
pub(crate) reload_id: u64,
pub(crate) target_generation: u64,
pub(crate) config: Arc<ProxyConfig>,
pub(crate) config_revision: String,
pub(crate) request: ReloadRequest,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum ReloadSubmitError {
InProgress(u64),
MaestroUnavailable,
}
#[derive(Clone)]
pub(crate) struct ReloadControl {
command_tx: mpsc::Sender<ReloadCommand>,
status_store: Arc<ReloadStatusStore>,
active_generation: Arc<AtomicU64>,
}
pub(crate) struct ReloadCommandReceiver {
command_rx: mpsc::Receiver<ReloadCommand>,
}
struct ReloadStatusState {
next_reload_id: u64,
active_reload_id: Option<u64>,
statuses: VecDeque<ReloadStatus>,
accepting_commands: bool,
}
impl Default for ReloadStatusState {
fn default() -> Self {
Self {
next_reload_id: 0,
active_reload_id: None,
statuses: VecDeque::new(),
accepting_commands: true,
}
}
}
#[derive(Default)]
struct ReloadStatusStore {
state: Mutex<ReloadStatusState>,
}
impl ReloadControl {
/// Creates the process-scoped coordinator channel and status store.
pub(crate) fn channel(initial_generation: u64) -> (Self, ReloadCommandReceiver) {
let (command_tx, command_rx) = mpsc::channel(RELOAD_COMMAND_CAPACITY);
(
Self {
command_tx,
status_store: Arc::new(ReloadStatusStore::default()),
active_generation: Arc::new(AtomicU64::new(initial_generation)),
},
ReloadCommandReceiver { command_rx },
)
}
/// Atomically reserves and enqueues one reload operation.
pub(crate) async fn submit(
&self,
config: Arc<ProxyConfig>,
config_revision: String,
request: ReloadRequest,
) -> Result<ReloadAccepted, ReloadSubmitError> {
let target_generation = self
.active_generation
.load(Ordering::Acquire)
.saturating_add(1);
let status = self
.status_store
.reserve(target_generation, config_revision, request.clone())
.await?;
let command = ReloadCommand {
reload_id: status.reload_id,
target_generation,
config,
config_revision: status.config_revision.clone(),
request,
};
if self.command_tx.try_send(command).is_err() {
self.status_store
.finish(
status.reload_id,
ReloadPhase::Failed,
Some("maestro command channel is closed".to_string()),
)
.await;
return Err(ReloadSubmitError::MaestroUnavailable);
}
Ok(ReloadAccepted {
reload_id: status.reload_id,
target_generation,
config_revision: status.config_revision,
state: ReloadPhase::Accepted,
mode: status.mode,
failure_policy: status.failure_policy,
})
}
/// Returns a retained reload status by identifier.
pub(crate) async fn status(&self, reload_id: u64) -> Option<ReloadStatus> {
self.status_store.get(reload_id).await
}
/// Returns the identifier of the currently active reload.
pub(crate) async fn in_progress(&self) -> Option<u64> {
self.status_store.state.lock().await.active_reload_id
}
/// Rejects new commands while preserving an already accepted operation.
pub(crate) async fn begin_shutdown(&self) {
self.status_store.state.lock().await.accepting_commands = false;
}
/// Records a non-terminal lifecycle phase.
pub(crate) async fn mark_phase(&self, reload_id: u64, phase: ReloadPhase) {
self.status_store.mark_phase(reload_id, phase).await;
}
/// Records process-owned fields deferred until the next process restart.
pub(crate) async fn set_deferred_fields(&self, reload_id: u64, fields: Vec<String>) {
self.status_store
.update(reload_id, |status| status.deferred_fields = fields)
.await;
}
/// Commits the active generation and completes the matching reload.
pub(crate) async fn succeed(&self, reload_id: u64, generation: u64) {
self.status_store
.finish_success(reload_id, generation, &self.active_generation)
.await;
}
/// Marks the matching reload as failed.
pub(crate) async fn fail(&self, reload_id: u64, error: impl Into<String>) {
self.status_store
.finish(reload_id, ReloadPhase::Failed, Some(error.into()))
.await;
}
/// Marks the matching reload as rolled back.
pub(crate) async fn rolled_back(&self, reload_id: u64, error: impl Into<String>) {
self.status_store
.finish(reload_id, ReloadPhase::RolledBack, Some(error.into()))
.await;
}
/// Appends a non-fatal warning to the matching reload status.
pub(crate) async fn add_warning(&self, reload_id: u64, warning: impl Into<String>) {
let warning = warning.into();
self.status_store
.update(reload_id, |status| status.warnings.push(warning))
.await;
}
}
impl ReloadCommandReceiver {
/// Receives the next accepted reload command.
pub(crate) async fn recv(&mut self) -> Option<ReloadCommand> {
self.command_rx.recv().await
}
}
impl ReloadStatusStore {
async fn reserve(
&self,
target_generation: u64,
config_revision: String,
request: ReloadRequest,
) -> Result<ReloadStatus, ReloadSubmitError> {
let mut state = self.state.lock().await;
if !state.accepting_commands {
return Err(ReloadSubmitError::MaestroUnavailable);
}
if let Some(reload_id) = state.active_reload_id {
return Err(ReloadSubmitError::InProgress(reload_id));
}
state.next_reload_id = state.next_reload_id.saturating_add(1).max(1);
let reload_id = state.next_reload_id;
let status = ReloadStatus {
reload_id,
target_generation,
config_revision,
state: ReloadPhase::Accepted,
mode: request.mode,
failure_policy: request.failure_policy,
requested_at_epoch_secs: now_epoch_secs(),
started_at_epoch_secs: None,
finished_at_epoch_secs: None,
deferred_fields: Vec::new(),
warnings: Vec::new(),
error: None,
};
state.active_reload_id = Some(reload_id);
state.statuses.push_back(status.clone());
while state.statuses.len() > RELOAD_HISTORY_CAPACITY {
state.statuses.pop_front();
}
Ok(status)
}
async fn get(&self, reload_id: u64) -> Option<ReloadStatus> {
self.state
.lock()
.await
.statuses
.iter()
.find(|status| status.reload_id == reload_id)
.cloned()
}
async fn mark_phase(&self, reload_id: u64, phase: ReloadPhase) {
self.update(reload_id, |status| {
status.state = phase;
if status.started_at_epoch_secs.is_none() && phase != ReloadPhase::Accepted {
status.started_at_epoch_secs = Some(now_epoch_secs());
}
})
.await;
}
async fn finish(&self, reload_id: u64, phase: ReloadPhase, error: Option<String>) {
debug_assert!(phase.is_terminal());
let mut state = self.state.lock().await;
if let Some(status) = state
.statuses
.iter_mut()
.find(|status| status.reload_id == reload_id)
{
status.state = phase;
status.error = error;
status.finished_at_epoch_secs = Some(now_epoch_secs());
}
if state.active_reload_id == Some(reload_id) {
state.active_reload_id = None;
}
}
async fn finish_success(&self, reload_id: u64, generation: u64, active_generation: &AtomicU64) {
let mut state = self.state.lock().await;
if state.active_reload_id != Some(reload_id) {
return;
}
let Some(status) = state
.statuses
.iter_mut()
.find(|status| status.reload_id == reload_id)
else {
return;
};
status.state = ReloadPhase::Succeeded;
status.error = None;
status.finished_at_epoch_secs = Some(now_epoch_secs());
active_generation.store(generation, Ordering::Release);
state.active_reload_id = None;
}
async fn update(&self, reload_id: u64, update: impl FnOnce(&mut ReloadStatus)) {
let mut state = self.state.lock().await;
if let Some(status) = state
.statuses
.iter_mut()
.find(|status| status.reload_id == reload_id)
{
update(status);
}
}
}
fn now_epoch_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
}
#[cfg(test)]
#[path = "reload_tests.rs"]
mod tests;
+280
View File
@@ -0,0 +1,280 @@
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use arc_swap::ArcSwap;
use tokio::sync::watch;
use tokio_util::sync::CancellationToken;
use tracing::{info, warn};
use crate::stats::QuotaStore;
use super::generation::{RuntimeGeneration, RuntimeWatchState};
use super::reload::{
ReloadCommand, ReloadCommandReceiver, ReloadControl, ReloadFailurePolicy, ReloadMode,
ReloadPhase,
};
use super::runtime_build::{PreparedRuntime, deferred_process_fields, prepare_runtime};
use super::runtime_tasks::RuntimeLogFilter;
pub(crate) struct ReloadSupervisor {
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
control: ReloadControl,
commands: ReloadCommandReceiver,
config_path: PathBuf,
quota_store: Arc<QuotaStore>,
detected_ips_tx: watch::Sender<(Option<std::net::IpAddr>, Option<std::net::IpAddr>)>,
runtime_log_filter: RuntimeLogFilter,
runtime_watch_tx: watch::Sender<Option<RuntimeWatchState>>,
}
/// Process-owned handle that quiesces reloads before shutdown snapshots the runtime.
pub(crate) struct ReloadSupervisorHandle {
control: ReloadControl,
shutdown: CancellationToken,
join: tokio::task::JoinHandle<()>,
}
impl ReloadSupervisorHandle {
/// Stops new submissions and waits for the accepted reload to finish.
pub(crate) async fn quiesce(self) {
self.control.begin_shutdown().await;
self.shutdown.cancel();
if let Err(error) = self.join.await {
warn!(error = %error, "Reload supervisor failed while quiescing");
}
}
}
#[derive(Debug, PartialEq, Eq)]
enum RevisionGateAction {
Proceed,
Warn(String),
Rollback(String),
}
fn revision_gate_action(
accepted_revision: &str,
current_revision: Result<String, String>,
failure_policy: ReloadFailurePolicy,
) -> RevisionGateAction {
let warning = match current_revision {
Ok(current) if current == accepted_revision => return RevisionGateAction::Proceed,
Ok(current) => format!(
"config revision changed during preparation: accepted={} current={}",
accepted_revision, current
),
Err(error) => format!("config revision verification failed: {}", error),
};
match failure_policy {
ReloadFailurePolicy::KeepNew => RevisionGateAction::Warn(warning),
ReloadFailurePolicy::Rollback => RevisionGateAction::Rollback(warning),
}
}
async fn stop_background_and_middle_end(generation: &RuntimeGeneration) -> bool {
generation.stop_background_tasks().await;
let Some(pool) = generation.current_me_pool().await else {
return false;
};
tokio::time::timeout(Duration::from_secs(2), pool.shutdown_send_close_conn_all())
.await
.is_err()
}
async fn cleanup_candidate(generation: &RuntimeGeneration) -> bool {
generation.stop_sessions().await;
stop_background_and_middle_end(generation).await
}
impl ReloadSupervisor {
#[allow(clippy::too_many_arguments)]
/// Starts the process-scoped reload supervisor and returns its shutdown owner.
pub(crate) fn spawn(
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
control: ReloadControl,
commands: ReloadCommandReceiver,
config_path: PathBuf,
quota_store: Arc<QuotaStore>,
detected_ips_tx: watch::Sender<(Option<std::net::IpAddr>, Option<std::net::IpAddr>)>,
runtime_log_filter: RuntimeLogFilter,
runtime_watch_tx: watch::Sender<Option<RuntimeWatchState>>,
) -> ReloadSupervisorHandle {
let supervisor = Self {
active_runtime,
control,
commands,
config_path,
quota_store,
detected_ips_tx,
runtime_log_filter,
runtime_watch_tx,
};
let control = supervisor.control.clone();
let shutdown = CancellationToken::new();
let join = tokio::spawn(supervisor.run(shutdown.clone()));
ReloadSupervisorHandle {
control,
shutdown,
join,
}
}
async fn run(mut self, shutdown: CancellationToken) {
loop {
tokio::select! {
biased;
_ = shutdown.cancelled() => {
if self.control.in_progress().await.is_some()
&& let Some(command) = self.commands.recv().await
{
self.reload(command).await;
}
break;
}
command = self.commands.recv() => {
let Some(command) = command else {
break;
};
self.reload(command).await;
}
}
}
}
async fn reload(&self, command: ReloadCommand) {
self.control
.mark_phase(command.reload_id, ReloadPhase::Preparing)
.await;
let old_runtime = self.active_runtime.load_full();
let deferred = deferred_process_fields(&old_runtime.config(), &command.config);
self.control
.set_deferred_fields(command.reload_id, deferred)
.await;
let prepared = match prepare_runtime(
command.target_generation,
command.config.as_ref().clone(),
&self.config_path,
self.quota_store.clone(),
self.runtime_log_filter.clone(),
)
.await
{
Ok(prepared) => prepared,
Err(error) => {
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())
})
.await;
}
async fn activate_prepared<InstallDns>(
&self,
command: ReloadCommand,
old_runtime: Arc<RuntimeGeneration>,
prepared: PreparedRuntime,
revision_action: RevisionGateAction,
install_dns: InstallDns,
) where
InstallDns: FnOnce(&[String]) -> Result<(), String>,
{
match revision_action {
RevisionGateAction::Proceed => {}
RevisionGateAction::Warn(warning) => {
self.control.add_warning(command.reload_id, warning).await;
}
RevisionGateAction::Rollback(warning) => {
let _ = cleanup_candidate(&prepared.generation).await;
self.runtime_log_filter
.apply_reload(&old_runtime.config().general.log_level);
self.control.rolled_back(command.reload_id, warning).await;
return;
}
}
self.control
.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 {
old_runtime.resume_accepting_sessions();
let _ = cleanup_candidate(&new_runtime).await;
self.runtime_log_filter
.apply_reload(&old_runtime.config().general.log_level);
self.control.rolled_back(command.reload_id, message).await;
return;
}
self.control.add_warning(command.reload_id, message).await;
}
let replaced = self.active_runtime.swap(new_runtime.clone());
self.detected_ips_tx.send_replace(prepared.detected_ips);
self.runtime_log_filter
.apply_reload(&new_runtime.config().general.log_level);
self.runtime_watch_tx
.send_replace(Some(new_runtime.watch_state()));
info!(
reload_id = command.reload_id,
old_generation = replaced.id,
new_generation = new_runtime.id,
config_revision = %command.config_revision,
"Runtime generation activated"
);
match command.request.mode {
ReloadMode::Instant => {
replaced.stop_sessions().await;
}
ReloadMode::Drain => {
self.control
.mark_phase(command.reload_id, ReloadPhase::Draining)
.await;
let timeout = Duration::from_secs(
command
.request
.timeout_secs
.expect("validated drain request must carry timeout_secs"),
);
if !replaced.drain_sessions(timeout).await {
let warning = format!(
"generation {} exceeded drain timeout; remaining sessions were cancelled",
replaced.id
);
warn!(reload_id = command.reload_id, warning = %warning);
self.control.add_warning(command.reload_id, warning).await;
}
}
}
if stop_background_and_middle_end(&replaced).await {
let warning = format!(
"generation {} Middle-End close broadcast timed out",
replaced.id
);
warn!(reload_id = command.reload_id, warning = %warning);
self.control.add_warning(command.reload_id, warning).await;
}
self.control
.succeed(command.reload_id, new_runtime.id)
.await;
}
}
#[cfg(test)]
#[path = "reload_supervisor_tests.rs"]
mod tests;
+320
View File
@@ -0,0 +1,320 @@
use super::*;
use crate::config::ProxyConfig;
use crate::maestro::generation::test_runtime_generation;
use crate::maestro::reload::{ReloadRequest, ReloadSubmitError};
use crate::stats::QuotaStore;
use tokio::sync::Notify;
use tracing_subscriber::{EnvFilter, Registry};
struct ReloadFixture {
supervisor: Arc<ReloadSupervisor>,
control: ReloadControl,
command: ReloadCommand,
old_runtime: Arc<RuntimeGeneration>,
new_runtime: Arc<RuntimeGeneration>,
runtime_watch_rx: watch::Receiver<Option<RuntimeWatchState>>,
}
fn runtime_log_filter() -> RuntimeLogFilter {
let (_layer, handle) =
tracing_subscriber::reload::Layer::<EnvFilter, Registry>::new(EnvFilter::new("info"));
RuntimeLogFilter::new(handle)
}
async fn fixture(request: ReloadRequest) -> ReloadFixture {
let old_runtime = test_runtime_generation(1, ProxyConfig::default());
let new_config = Arc::new(ProxyConfig::default());
let new_runtime = test_runtime_generation(2, new_config.as_ref().clone());
let active_runtime = Arc::new(ArcSwap::from(old_runtime.clone()));
let (control, commands) = ReloadControl::channel(old_runtime.id);
let accepted = control
.submit(new_config.clone(), "revision".to_string(), request.clone())
.await
.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 supervisor = Arc::new(ReloadSupervisor {
active_runtime,
control: control.clone(),
commands,
config_path: PathBuf::new(),
quota_store: Arc::new(QuotaStore::default()),
detected_ips_tx,
runtime_log_filter: runtime_log_filter(),
runtime_watch_tx,
});
let command = ReloadCommand {
reload_id: accepted.reload_id,
target_generation: accepted.target_generation,
config: new_config,
config_revision: accepted.config_revision,
request,
};
ReloadFixture {
supervisor,
control,
command,
old_runtime,
new_runtime,
runtime_watch_rx,
}
}
struct DropSignal(Arc<Notify>);
impl Drop for DropSignal {
fn drop(&mut self) {
self.0.notify_one();
}
}
#[test]
fn revision_gate_proceeds_only_on_verified_match() {
assert_eq!(
revision_gate_action(
"accepted",
Ok("accepted".to_string()),
ReloadFailurePolicy::Rollback,
),
RevisionGateAction::Proceed
);
}
#[test]
fn revision_gate_applies_failure_policy_to_mismatch_and_read_error() {
for result in [Ok("changed".to_string()), Err("read failed".to_string())] {
assert!(matches!(
revision_gate_action("accepted", result.clone(), ReloadFailurePolicy::KeepNew,),
RevisionGateAction::Warn(_)
));
assert!(matches!(
revision_gate_action("accepted", result, ReloadFailurePolicy::Rollback),
RevisionGateAction::Rollback(_)
));
}
}
#[tokio::test]
async fn revision_rollback_keeps_old_generation_and_cleans_candidate() {
let fixture = fixture(ReloadRequest {
failure_policy: ReloadFailurePolicy::Rollback,
..ReloadRequest::default()
})
.await;
let candidate_dropped = Arc::new(Notify::new());
let candidate_drop = candidate_dropped.clone();
assert!(fixture.new_runtime.spawn_session(async move {
let _drop_signal = DropSignal(candidate_drop);
std::future::pending::<()>().await;
}));
tokio::task::yield_now().await;
fixture
.supervisor
.activate_prepared(
fixture.command,
fixture.old_runtime.clone(),
PreparedRuntime {
generation: fixture.new_runtime,
detected_ips: (None, None),
},
RevisionGateAction::Rollback("revision changed".to_string()),
|_| -> Result<(), String> { panic!("DNS activation must not run on rollback") },
)
.await;
tokio::time::timeout(Duration::from_secs(1), candidate_dropped.notified())
.await
.unwrap();
assert_eq!(fixture.supervisor.active_runtime.load().id, 1);
assert_eq!(
fixture
.runtime_watch_rx
.borrow()
.as_ref()
.unwrap()
.generation_id,
1
);
assert!(fixture.old_runtime.spawn_session(async {}));
let status = fixture.control.status(1).await.unwrap();
assert_eq!(status.state, ReloadPhase::RolledBack);
fixture.old_runtime.stop_sessions().await;
}
#[tokio::test]
async fn dns_failure_policy_controls_rollback_or_keep_new() {
for policy in [ReloadFailurePolicy::Rollback, ReloadFailurePolicy::KeepNew] {
let fixture = fixture(ReloadRequest {
failure_policy: policy,
..ReloadRequest::default()
})
.await;
fixture
.supervisor
.activate_prepared(
fixture.command,
fixture.old_runtime.clone(),
PreparedRuntime {
generation: fixture.new_runtime.clone(),
detected_ips: (None, None),
},
RevisionGateAction::Proceed,
|_| Err("invalid DNS entry".to_string()),
)
.await;
let status = fixture.control.status(1).await.unwrap();
match policy {
ReloadFailurePolicy::Rollback => {
assert_eq!(fixture.supervisor.active_runtime.load().id, 1);
assert_eq!(status.state, ReloadPhase::RolledBack);
assert!(fixture.old_runtime.spawn_session(async {}));
fixture.old_runtime.stop_sessions().await;
}
ReloadFailurePolicy::KeepNew => {
assert_eq!(fixture.supervisor.active_runtime.load().id, 2);
assert_eq!(status.state, ReloadPhase::Succeeded);
assert_eq!(status.warnings.len(), 1);
assert!(!fixture.old_runtime.spawn_session(async {}));
fixture.new_runtime.stop_sessions().await;
}
}
}
}
#[tokio::test]
async fn drain_publishes_new_generation_before_old_sessions_finish() {
let mut fixture = fixture(ReloadRequest {
mode: ReloadMode::Drain,
timeout_secs: Some(30),
..ReloadRequest::default()
})
.await;
let old_started = Arc::new(Notify::new());
let old_release = Arc::new(Notify::new());
let started = old_started.clone();
let release = old_release.clone();
assert!(fixture.old_runtime.spawn_session(async move {
started.notify_one();
release.notified().await;
}));
old_started.notified().await;
let supervisor = fixture.supervisor.clone();
let old_runtime = fixture.old_runtime.clone();
let new_runtime = fixture.new_runtime.clone();
let activation = tokio::spawn(async move {
supervisor
.activate_prepared(
fixture.command,
old_runtime,
PreparedRuntime {
generation: new_runtime,
detected_ips: (None, None),
},
RevisionGateAction::Proceed,
|_| Ok(()),
)
.await;
});
fixture.runtime_watch_rx.changed().await.unwrap();
assert_eq!(
fixture
.runtime_watch_rx
.borrow()
.as_ref()
.unwrap()
.generation_id,
2
);
assert!(!activation.is_finished());
assert!(!fixture.old_runtime.spawn_session(async {}));
old_release.notify_one();
activation.await.unwrap();
assert_eq!(
fixture.control.status(1).await.unwrap().state,
ReloadPhase::Succeeded
);
fixture.new_runtime.stop_sessions().await;
}
#[tokio::test(start_paused = true)]
async fn drain_timeout_cancels_old_sessions_and_records_one_warning() {
let mut fixture = fixture(ReloadRequest {
mode: ReloadMode::Drain,
timeout_secs: Some(1),
..ReloadRequest::default()
})
.await;
let dropped = Arc::new(Notify::new());
let drop_signal = dropped.clone();
assert!(fixture.old_runtime.spawn_session(async move {
let _drop_signal = DropSignal(drop_signal);
std::future::pending::<()>().await;
}));
tokio::task::yield_now().await;
let supervisor = fixture.supervisor.clone();
let old_runtime = fixture.old_runtime.clone();
let new_runtime = fixture.new_runtime.clone();
let activation = tokio::spawn(async move {
supervisor
.activate_prepared(
fixture.command,
old_runtime,
PreparedRuntime {
generation: new_runtime,
detected_ips: (None, None),
},
RevisionGateAction::Proceed,
|_| Ok(()),
)
.await;
});
fixture.runtime_watch_rx.changed().await.unwrap();
tokio::task::yield_now().await;
tokio::time::advance(Duration::from_secs(1)).await;
activation.await.unwrap();
dropped.notified().await;
let status = fixture.control.status(1).await.unwrap();
assert_eq!(status.state, ReloadPhase::Succeeded);
assert_eq!(status.warnings.len(), 1);
assert!(status.warnings[0].contains("exceeded drain timeout"));
fixture.new_runtime.stop_sessions().await;
}
#[tokio::test]
async fn quiesce_joins_idle_supervisor_and_rejects_later_submissions() {
let runtime = test_runtime_generation(1, ProxyConfig::default());
let active_runtime = Arc::new(ArcSwap::from(runtime.clone()));
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 handle = ReloadSupervisor::spawn(
active_runtime,
control.clone(),
commands,
PathBuf::new(),
Arc::new(QuotaStore::default()),
detected_ips_tx,
runtime_log_filter(),
runtime_watch_tx,
);
tokio::time::timeout(Duration::from_secs(1), handle.quiesce())
.await
.unwrap();
let result = control
.submit(
Arc::new(ProxyConfig::default()),
"revision".to_string(),
ReloadRequest::default(),
)
.await;
assert_eq!(result, Err(ReloadSubmitError::MaestroUnavailable));
runtime.stop_sessions().await;
}
+256
View File
@@ -0,0 +1,256 @@
use super::*;
#[test]
fn request_defaults_to_instant_keep_new() {
let request: ReloadRequest = serde_json::from_str("{}").unwrap();
assert_eq!(request, ReloadRequest::default());
assert_eq!(request.validate(), Ok(()));
}
#[test]
fn drain_requires_bounded_timeout() {
let missing = ReloadRequest {
mode: ReloadMode::Drain,
..ReloadRequest::default()
};
assert!(missing.validate().is_err());
let valid = ReloadRequest {
mode: ReloadMode::Drain,
timeout_secs: Some(30),
..ReloadRequest::default()
};
assert_eq!(valid.validate(), Ok(()));
}
#[test]
fn patch_query_parses_reload_policy() {
let request =
ReloadRequest::from_query(Some("reload=drain&timeout_secs=30&failure_policy=rollback"))
.unwrap()
.unwrap();
assert_eq!(request.mode, ReloadMode::Drain);
assert_eq!(request.timeout_secs, Some(30));
assert_eq!(request.failure_policy, ReloadFailurePolicy::Rollback);
assert!(ReloadRequest::from_query(Some("timeout_secs=30")).is_err());
}
#[test]
fn status_uses_documented_deferred_process_fields_key() {
let status = ReloadStatus {
reload_id: 1,
target_generation: 2,
config_revision: "revision".to_string(),
state: ReloadPhase::Succeeded,
mode: ReloadMode::Instant,
failure_policy: ReloadFailurePolicy::KeepNew,
requested_at_epoch_secs: 10,
started_at_epoch_secs: Some(11),
finished_at_epoch_secs: Some(12),
deferred_fields: vec!["server.listeners".to_string()],
warnings: Vec::new(),
error: None,
};
let value = serde_json::to_value(status).unwrap();
assert_eq!(
value["deferred_process_fields"],
serde_json::json!(["server.listeners"])
);
assert!(value.get("deferred_fields").is_none());
}
#[tokio::test]
async fn coordinator_rejects_concurrent_reload_and_releases_terminal_slot() {
let (control, mut receiver) = ReloadControl::channel(1);
let first = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-1".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
let second = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-2".to_string(),
ReloadRequest::default(),
)
.await;
assert_eq!(second, Err(ReloadSubmitError::InProgress(first.reload_id)));
control
.succeed(first.reload_id, first.target_generation)
.await;
let third = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-3".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
assert_eq!(third.reload_id, first.reload_id + 1);
}
#[tokio::test]
async fn terminal_outcomes_release_slot_and_only_success_advances_generation() {
let (control, mut receiver) = ReloadControl::channel(7);
let failed = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-failed".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
control
.mark_phase(failed.reload_id, ReloadPhase::Preparing)
.await;
control.fail(failed.reload_id, "prepare failed").await;
let failed_status = control.status(failed.reload_id).await.unwrap();
assert_eq!(failed_status.state, ReloadPhase::Failed);
assert_eq!(failed_status.error.as_deref(), Some("prepare failed"));
assert!(failed_status.started_at_epoch_secs.is_some());
assert!(failed_status.finished_at_epoch_secs.is_some());
let rolled_back = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-rollback".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
assert_eq!(rolled_back.target_generation, 8);
control
.rolled_back(rolled_back.reload_id, "revision changed")
.await;
let succeeded = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-success".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
assert_eq!(succeeded.target_generation, 8);
control
.succeed(succeeded.reload_id, succeeded.target_generation)
.await;
let next = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-next".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
assert_eq!(next.target_generation, 9);
}
#[tokio::test]
async fn stale_success_cannot_advance_generation_or_release_active_reload() {
let (control, mut receiver) = ReloadControl::channel(3);
let active = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-active".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
control.succeed(active.reload_id + 100, 99).await;
assert_eq!(control.in_progress().await, Some(active.reload_id));
control.fail(active.reload_id, "expected failure").await;
let next = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-next".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
assert_eq!(next.target_generation, 4);
}
#[tokio::test]
async fn status_history_retains_only_the_latest_entries() {
let (control, mut receiver) = ReloadControl::channel(1);
let mut reload_ids = Vec::new();
for index in 0..=RELOAD_HISTORY_CAPACITY {
let accepted = control
.submit(
Arc::new(ProxyConfig::default()),
format!("rev-{index}"),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
reload_ids.push(accepted.reload_id);
control.fail(accepted.reload_id, "expected failure").await;
}
assert!(control.status(reload_ids[0]).await.is_none());
assert!(control.status(reload_ids[1]).await.is_some());
assert!(control.status(*reload_ids.last().unwrap()).await.is_some());
}
#[tokio::test]
async fn closed_command_channel_marks_reload_failed_and_releases_slot() {
let (control, receiver) = ReloadControl::channel(1);
drop(receiver);
let result = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-closed".to_string(),
ReloadRequest::default(),
)
.await;
assert_eq!(result, Err(ReloadSubmitError::MaestroUnavailable));
assert_eq!(control.in_progress().await, None);
let status = control.status(1).await.unwrap();
assert_eq!(status.state, ReloadPhase::Failed);
assert_eq!(
status.error.as_deref(),
Some("maestro command channel is closed")
);
}
#[tokio::test]
async fn shutdown_gate_rejects_new_commands_without_disturbing_active_status() {
let (control, mut receiver) = ReloadControl::channel(4);
let active = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-active".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
control.begin_shutdown().await;
let rejected = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-rejected".to_string(),
ReloadRequest::default(),
)
.await;
assert_eq!(rejected, Err(ReloadSubmitError::MaestroUnavailable));
assert_eq!(control.in_progress().await, Some(active.reload_id));
control.fail(active.reload_id, "shutdown test").await;
}
+383
View File
@@ -0,0 +1,383 @@
use std::net::IpAddr;
use std::path::Path;
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use tokio::sync::{RwLock, Semaphore, watch};
use crate::config::ProxyConfig;
use crate::crypto::SecureRandom;
use crate::ip_tracker::UserIpTracker;
use crate::network::probe::{decide_network_capabilities, run_probe};
use crate::proxy::direct_buffer_budget::{
DirectBufferBudget, resolve_direct_buffer_hard_limit, run_direct_buffer_budget_controller,
};
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::shared_state::ProxySharedState;
use crate::startup::StartupTracker;
use crate::stats::beobachten::BeobachtenStore;
use crate::stats::telemetry::TelemetryPolicy;
use crate::stats::{QuotaStore, ReplayChecker, Stats};
use crate::stream::BufferPool;
use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool;
use super::admission;
use super::generation::{RuntimeGeneration, RuntimeTaskScope};
use super::runtime_tasks::RuntimeLogFilter;
use super::{me_startup, runtime_tasks, tls_bootstrap};
pub(crate) struct PreparedRuntime {
pub(crate) generation: Arc<RuntimeGeneration>,
pub(crate) detected_ips: (Option<IpAddr>, Option<IpAddr>),
}
pub(crate) async fn prepare_runtime(
generation_id: u64,
config: ProxyConfig,
config_path: &Path,
quota_store: Arc<QuotaStore>,
runtime_log_filter: RuntimeLogFilter,
) -> Result<PreparedRuntime, String> {
let started_at_epoch_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let startup_tracker = Arc::new(StartupTracker::new(started_at_epoch_secs));
let task_scope = RuntimeTaskScope::new();
let stats = Arc::new(Stats::with_quota_store(quota_store));
stats.apply_telemetry_policy(TelemetryPolicy::from_config(&config.general.telemetry));
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)
.map_err(|error| format!("DNS override preparation failed: {}", error))?,
);
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;
let hard_limit =
resolve_direct_buffer_hard_limit(config.general.direct_relay_buffer_budget_max_bytes).await;
let direct_buffer_budget = DirectBufferBudget::new(hard_limit);
let proxy_shared =
ProxySharedState::new_with_direct_buffer_budget(direct_buffer_budget.clone());
proxy_shared.apply_user_enabled_config(&config.access.user_enabled);
proxy_shared.traffic_limiter.apply_policy(
config.access.user_rate_limits.clone(),
config.access.cidr_rate_limits.clone(),
);
let probe = run_probe(
&config.network,
&config.upstreams,
config.general.middle_proxy_nat_probe,
config.general.stun_nat_probe_concurrency,
)
.await
.map_err(|error| format!("network probe failed: {}", error))?;
let decision =
decide_network_capabilities(&config.network, &probe, config.general.middle_proxy_nat_ip);
let prefer_ipv6 = decision.prefer_ipv6();
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,
task_scope.clone(),
tls_bootstrap::TlsBootstrapPolicy::RequireReady,
)
.await
.map_err(|error| error.to_string())?;
let beobachten = Arc::new(BeobachtenStore::new());
let rng = Arc::new(SecureRandom::new());
let route_mode = if !config.general.use_middle_proxy || config.general.me2dc_fallback {
RelayRouteMode::Direct
} else {
RelayRouteMode::Middle
};
let route_runtime = Arc::new(RouteRuntimeController::new(route_mode));
let me_pool_runtime = Arc::new(RwLock::new(None::<Arc<MePool>>));
let (me_ready_tx, me_ready_rx) = watch::channel(0_u64);
let direct_first_startup = config.general.use_middle_proxy && config.general.me2dc_fallback;
let me_pool = if direct_first_startup {
None
} else {
me_startup::initialize_me_pool(
config.general.use_middle_proxy,
&config,
&decision,
&probe,
&startup_tracker,
upstream_manager.clone(),
rng.clone(),
stats.clone(),
me_pool_runtime.clone(),
me_ready_tx.clone(),
task_scope.clone(),
)
.await
};
if strict_middle_proxy_unavailable(
config.general.use_middle_proxy,
direct_first_startup,
me_pool.is_some(),
) {
task_scope.stop().await;
return Err(
"Middle-End pool is required but did not become ready during reload preparation"
.to_string(),
);
}
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));
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 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(),
proxy_shared.clone(),
me_ready_tx.clone(),
task_scope.clone(),
)
.await;
let config_rx = watches.config_rx;
runtime_log_filter.spawn_watcher(watches.log_level_rx, task_scope.clone());
let initial_admission_open = !config.general.use_middle_proxy || me_pool.is_some();
let (admission_tx, admission_rx) = watch::channel(initial_admission_open);
admission::configure_admission_gate(
&config,
me_pool.clone(),
me_pool_runtime.clone(),
route_runtime.clone(),
&admission_tx,
config_rx.clone(),
me_ready_rx,
task_scope.clone(),
)
.await;
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 me_pool_runtime_bg = me_pool_runtime.clone();
let me_ready_tx_bg = me_ready_tx.clone();
let config_rx_bg = config_rx.clone();
let task_scope_bg = task_scope.clone();
let retry_limit = config.general.me_init_retry_attempts;
task_scope.spawn(async move {
let mut attempt = 0_u32;
loop {
attempt = 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(),
me_pool_runtime_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 retry_limit > 0 && attempt >= retry_limit {
break;
}
tokio::time::sleep(Duration::from_secs(2)).await;
}
});
}
let conntrack_scope = task_scope.clone();
task_scope.spawn(crate::conntrack_control::run_conntrack_controller(
config_rx.clone(),
stats.clone(),
proxy_shared.clone(),
conntrack_scope.cancellation_token(),
));
task_scope.spawn(run_direct_buffer_budget_controller(
direct_buffer_budget,
buffer_pool.clone(),
stats.clone(),
proxy_shared.clone(),
config.server.max_connections,
));
let generation = RuntimeGeneration::new(
generation_id,
config_rx,
admission_rx,
stats,
upstream_manager,
replay_checker,
buffer_pool,
rng,
me_pool,
me_pool_runtime,
route_runtime,
tls_cache,
ip_tracker,
beobachten,
proxy_shared,
max_connections,
task_scope,
);
drop(admission_tx);
Ok(PreparedRuntime {
generation,
detected_ips: (
probe.detected_ipv4.map(IpAddr::V4),
probe.detected_ipv6.map(IpAddr::V6),
),
})
}
fn strict_middle_proxy_unavailable(
use_middle_proxy: bool,
direct_first_startup: bool,
pool_available: bool,
) -> bool {
use_middle_proxy && !direct_first_startup && !pool_available
}
pub(crate) fn deferred_process_fields(old: &ProxyConfig, new: &ProxyConfig) -> Vec<String> {
let mut fields = Vec::new();
if old.server.port != new.server.port
|| old.server.proxy_protocol != new.server.proxy_protocol
|| old.server.listen_backlog != new.server.listen_backlog
|| serde_json::to_value(&old.server.listeners).ok()
!= serde_json::to_value(&new.server.listeners).ok()
{
fields.push("server.listeners".to_string());
}
if old.server.listen_unix_sock != new.server.listen_unix_sock
|| old.server.listen_unix_sock_perm != new.server.listen_unix_sock_perm
{
fields.push("server.listen_unix_sock".to_string());
}
if old.server.api.listen != new.server.api.listen
|| old.server.api.enabled != new.server.api.enabled
{
fields.push("server.api.listen".to_string());
}
if old.server.metrics_listen != new.server.metrics_listen
|| old.server.metrics_port != new.server.metrics_port
{
fields.push("server.metrics_listen".to_string());
}
if old.general.quota_state_path != new.general.quota_state_path {
fields.push("general.quota_state_path".to_string());
}
if old.general.disable_colors != new.general.disable_colors {
fields.push("general.disable_colors".to_string());
}
if old.general.data_path != new.general.data_path {
fields.push("general.data_path".to_string());
}
if serde_json::to_value(&old.logging).ok() != serde_json::to_value(&new.logging).ok() {
fields.push("logging".to_string());
}
fields
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn process_socket_and_logging_changes_are_deferred() {
let old = ProxyConfig::default();
let mut new = old.clone();
new.server.listen_backlog = new.server.listen_backlog.saturating_add(1);
new.general.disable_colors = !new.general.disable_colors;
let fields = deferred_process_fields(&old, &new);
assert!(fields.contains(&"server.listeners".to_string()));
assert!(fields.contains(&"general.disable_colors".to_string()));
}
#[test]
fn runtime_only_change_does_not_require_process_rebind() {
let old = ProxyConfig::default();
let mut new = old.clone();
new.censorship.tls_domain = "reload.example".to_string();
assert!(deferred_process_fields(&old, &new).is_empty());
}
#[test]
fn strict_middle_proxy_requires_a_prepared_pool() {
assert!(strict_middle_proxy_unavailable(true, false, false));
assert!(!strict_middle_proxy_unavailable(true, false, true));
assert!(!strict_middle_proxy_unavailable(true, true, false));
assert!(!strict_middle_proxy_unavailable(false, false, false));
}
}
+84 -78
View File
@@ -2,9 +2,11 @@ use std::net::IpAddr;
use std::path::Path;
use std::sync::Arc;
use arc_swap::ArcSwap;
use tokio::sync::{mpsc, watch};
use tracing::{debug, info, warn};
use tracing_subscriber::EnvFilter;
use tracing_subscriber::Registry;
use tracing_subscriber::reload;
use crate::config::hot_reload::spawn_config_watcher;
@@ -21,10 +23,11 @@ use crate::startup::{
use crate::stats::beobachten::BeobachtenStore;
use crate::stats::telemetry::TelemetryPolicy;
use crate::stats::{ReplayChecker, Stats};
use crate::tls_front::TlsFrontCache;
use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::{MePool, MeReinitTrigger};
use super::generation::RuntimeGeneration;
use super::generation::RuntimeTaskScope;
use super::helpers::write_beobachten_snapshot;
pub(crate) struct RuntimeWatches {
@@ -34,6 +37,56 @@ pub(crate) struct RuntimeWatches {
pub(crate) detected_ip_v6: Option<IpAddr>,
}
#[derive(Clone)]
pub(crate) struct RuntimeLogFilter {
handle: reload::Handle<EnvFilter, Registry>,
}
impl RuntimeLogFilter {
pub(crate) fn new(handle: reload::Handle<EnvFilter, Registry>) -> Self {
Self { handle }
}
pub(crate) fn start(
&self,
has_rust_log: bool,
effective_log_level: &LogLevel,
log_level_rx: watch::Receiver<LogLevel>,
task_scope: RuntimeTaskScope,
) {
self.apply(effective_log_level, has_rust_log);
self.spawn_watcher(log_level_rx, task_scope);
}
pub(crate) fn apply_reload(&self, level: &LogLevel) {
self.apply(level, false);
}
pub(crate) fn spawn_watcher(
&self,
mut log_level_rx: watch::Receiver<LogLevel>,
task_scope: RuntimeTaskScope,
) {
let filter = self.clone();
task_scope.spawn(async move {
loop {
if log_level_rx.changed().await.is_err() {
break;
}
let level = log_level_rx.borrow_and_update().clone();
filter.apply_reload(&level);
}
});
}
fn apply(&self, level: &LogLevel, has_rust_log: bool) {
let runtime_filter = EnvFilter::new(log_filter_spec(has_rust_log, level));
if let Err(error) = self.handle.reload(runtime_filter) {
tracing::error!(error = %error, "Failed to update runtime log filter");
}
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn spawn_runtime_tasks(
config: &Arc<ProxyConfig>,
@@ -50,14 +103,14 @@ pub(crate) async fn spawn_runtime_tasks(
rng: Arc<SecureRandom>,
ip_tracker: Arc<UserIpTracker>,
beobachten: Arc<BeobachtenStore>,
api_config_tx: watch::Sender<Arc<ProxyConfig>>,
me_pool_for_policy: Option<Arc<MePool>>,
shared_state: Arc<ProxySharedState>,
me_ready_tx: watch::Sender<u64>,
task_scope: RuntimeTaskScope,
) -> RuntimeWatches {
let um_clone = upstream_manager.clone();
let dc_overrides_for_health = config.dc_overrides.clone();
tokio::spawn(async move {
task_scope.spawn(async move {
um_clone
.run_health_checks(
prefer_ipv6,
@@ -69,19 +122,19 @@ pub(crate) async fn spawn_runtime_tasks(
});
let rc_clone = replay_checker.clone();
tokio::spawn(async move {
task_scope.spawn(async move {
rc_clone.run_periodic_cleanup().await;
});
let stats_maintenance = stats.clone();
tokio::spawn(async move {
task_scope.spawn(async move {
stats_maintenance
.run_periodic_user_stats_maintenance()
.await;
});
let ip_tracker_maintenance = ip_tracker.clone();
tokio::spawn(async move {
task_scope.spawn(async move {
ip_tracker_maintenance.run_periodic_maintenance().await;
});
@@ -104,6 +157,7 @@ pub(crate) async fn spawn_runtime_tasks(
config.clone(),
detected_ip_v4,
detected_ip_v6,
task_scope.cancellation_token(),
);
startup_tracker
.complete_component(
@@ -111,21 +165,10 @@ pub(crate) async fn spawn_runtime_tasks(
Some("config hot-reload watcher started".to_string()),
)
.await;
let mut config_rx_api_bridge = config_rx.clone();
let api_config_tx_bridge = api_config_tx.clone();
tokio::spawn(async move {
loop {
if config_rx_api_bridge.changed().await.is_err() {
break;
}
let cfg = config_rx_api_bridge.borrow_and_update().clone();
api_config_tx_bridge.send_replace(cfg);
}
});
let stats_policy = stats.clone();
let upstream_policy = upstream_manager.clone();
let mut config_rx_policy = config_rx.clone();
tokio::spawn(async move {
task_scope.spawn(async move {
loop {
if config_rx_policy.changed().await.is_err() {
break;
@@ -133,6 +176,9 @@ pub(crate) async fn spawn_runtime_tasks(
let cfg = config_rx_policy.borrow_and_update().clone();
stats_policy
.apply_telemetry_policy(TelemetryPolicy::from_config(&cfg.general.telemetry));
if let Err(error) = upstream_policy.update_dns_overrides(&cfg.network.dns_overrides) {
warn!(error = %error, "Failed to update generation DNS overrides");
}
if let Some(pool) = &me_pool_for_policy {
pool.update_runtime_transport_policy(
cfg.general.me_socks_kdf_policy,
@@ -149,7 +195,7 @@ pub(crate) async fn spawn_runtime_tasks(
let ip_tracker_policy = ip_tracker.clone();
let mut config_rx_ip_limits = config_rx.clone();
tokio::spawn(async move {
task_scope.spawn(async move {
let mut prev_limits = config_rx_ip_limits
.borrow()
.access
@@ -205,7 +251,7 @@ pub(crate) async fn spawn_runtime_tasks(
config.access.cidr_rate_limits.clone(),
);
let mut config_rx_rate_limits = config_rx.clone();
tokio::spawn(async move {
task_scope.spawn(async move {
let mut prev_user_limits = config_rx_rate_limits
.borrow()
.access
@@ -236,7 +282,7 @@ pub(crate) async fn spawn_runtime_tasks(
let shared_user_enabled = shared_state.clone();
let mut config_rx_user_enabled = config_rx.clone();
tokio::spawn(async move {
task_scope.spawn(async move {
loop {
if config_rx_user_enabled.changed().await.is_err() {
break;
@@ -257,7 +303,7 @@ pub(crate) async fn spawn_runtime_tasks(
let beobachten_writer = beobachten.clone();
let config_rx_beobachten = config_rx.clone();
tokio::spawn(async move {
task_scope.spawn(async move {
loop {
let cfg = config_rx_beobachten.borrow().clone();
let sleep_secs = cfg.general.beobachten_flush_secs.max(1);
@@ -278,7 +324,14 @@ pub(crate) async fn spawn_runtime_tasks(
});
if let Some(pool) = me_pool {
spawn_middle_proxy_runtime_tasks(config, config_rx.clone(), pool, rng, me_ready_tx);
spawn_middle_proxy_runtime_tasks(
config,
config_rx.clone(),
pool,
rng,
me_ready_tx,
task_scope,
);
}
RuntimeWatches {
@@ -295,6 +348,7 @@ pub(crate) fn spawn_middle_proxy_runtime_tasks(
pool: Arc<MePool>,
rng: Arc<SecureRandom>,
me_ready_tx: watch::Sender<u64>,
task_scope: RuntimeTaskScope,
) {
let reinit_trigger_capacity = config.general.me_reinit_trigger_channel.max(1);
let (reinit_tx, reinit_rx) = mpsc::channel::<MeReinitTrigger>(reinit_trigger_capacity);
@@ -303,7 +357,7 @@ pub(crate) fn spawn_middle_proxy_runtime_tasks(
let rng_clone_sched = rng.clone();
let config_rx_clone_sched = config_rx.clone();
let me_ready_tx_sched = me_ready_tx.clone();
tokio::spawn(async move {
task_scope.spawn(async move {
crate::transport::middle_proxy::me_reinit_scheduler(
pool_clone_sched,
rng_clone_sched,
@@ -317,7 +371,7 @@ pub(crate) fn spawn_middle_proxy_runtime_tasks(
let pool_clone = pool.clone();
let config_rx_clone = config_rx.clone();
let reinit_tx_updater = reinit_tx.clone();
tokio::spawn(async move {
task_scope.spawn(async move {
crate::transport::middle_proxy::me_config_updater(
pool_clone,
config_rx_clone,
@@ -328,37 +382,12 @@ pub(crate) fn spawn_middle_proxy_runtime_tasks(
let config_rx_clone_rot = config_rx.clone();
let reinit_tx_rotation = reinit_tx.clone();
tokio::spawn(async move {
task_scope.spawn(async move {
crate::transport::middle_proxy::me_rotation_task(config_rx_clone_rot, reinit_tx_rotation)
.await;
});
}
pub(crate) async fn apply_runtime_log_filter(
has_rust_log: bool,
effective_log_level: &LogLevel,
filter_handle: reload::Handle<EnvFilter, tracing_subscriber::Registry>,
mut log_level_rx: watch::Receiver<LogLevel>,
) {
let runtime_filter = EnvFilter::new(log_filter_spec(has_rust_log, effective_log_level));
filter_handle
.reload(runtime_filter)
.expect("Failed to switch log filter");
tokio::spawn(async move {
loop {
if log_level_rx.changed().await.is_err() {
break;
}
let level = log_level_rx.borrow_and_update().clone();
let new_filter = tracing_subscriber::EnvFilter::new(log_filter_spec(false, &level));
if let Err(e) = filter_handle.reload(new_filter) {
tracing::error!("config reload: failed to update log filter: {}", e);
}
}
});
}
pub(crate) fn log_filter_spec(has_rust_log: bool, effective_log_level: &LogLevel) -> String {
if has_rust_log {
std::env::var("RUST_LOG")
@@ -373,12 +402,7 @@ pub(crate) fn log_filter_spec(has_rust_log: bool, effective_log_level: &LogLevel
pub(crate) async fn spawn_metrics_if_configured(
config: &Arc<ProxyConfig>,
startup_tracker: &Arc<StartupTracker>,
stats: Arc<Stats>,
beobachten: Arc<BeobachtenStore>,
shared_state: Arc<ProxySharedState>,
ip_tracker: Arc<UserIpTracker>,
tls_cache: Option<Arc<TlsFrontCache>>,
config_rx: watch::Receiver<Arc<ProxyConfig>>,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
) {
// metrics_listen takes precedence; fall back to metrics_port for backward compat.
let metrics_target: Option<(u16, Option<String>)> =
@@ -408,28 +432,10 @@ pub(crate) async fn spawn_metrics_if_configured(
Some(format!("spawn metrics endpoint on {}", label)),
)
.await;
let stats = stats.clone();
let beobachten = beobachten.clone();
let shared_state = shared_state.clone();
let config_rx_metrics = config_rx.clone();
let ip_tracker_metrics = ip_tracker.clone();
let tls_cache_metrics = tls_cache.clone();
let whitelist = config.server.metrics_whitelist.clone();
let active_runtime = active_runtime.clone();
let listen_backlog = config.server.listen_backlog;
tokio::spawn(async move {
metrics::serve(
port,
listen,
listen_backlog,
stats,
beobachten,
shared_state,
ip_tracker_metrics,
tls_cache_metrics,
config_rx_metrics,
whitelist,
)
.await;
metrics::serve(port, listen, listen_backlog, active_runtime).await;
});
startup_tracker
.complete_component(
+35 -17
View File
@@ -12,17 +12,18 @@ use std::path::PathBuf;
use std::sync::Arc;
use std::time::{Duration, Instant};
use arc_swap::ArcSwap;
#[cfg(not(unix))]
use tokio::signal;
#[cfg(unix)]
use tokio::signal::unix::{SignalKind, signal};
use tracing::{info, warn};
use super::generation::RuntimeGeneration;
use super::helpers::{format_uptime, unit_label};
use super::reload_supervisor::ReloadSupervisorHandle;
use crate::stats::Stats;
use crate::synlimit_control;
use crate::transport::middle_proxy::MePool;
use super::helpers::{format_uptime, unit_label};
/// Signal that triggered shutdown.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -48,17 +49,19 @@ impl std::fmt::Display for ShutdownSignal {
/// Waits for a shutdown signal and performs graceful shutdown.
pub(crate) async fn wait_for_shutdown(
process_started_at: Instant,
me_pool: Option<Arc<MePool>>,
stats: Arc<Stats>,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
quota_state_path: PathBuf,
synlimit_controller: synlimit_control::SynlimitController,
reload_supervisor: ReloadSupervisorHandle,
) {
let signal = wait_for_shutdown_signal().await;
perform_shutdown(
signal,
process_started_at,
me_pool,
&stats,
active_runtime,
quota_state_path,
synlimit_controller,
reload_supervisor,
)
.await;
}
@@ -87,13 +90,18 @@ async fn wait_for_shutdown_signal() -> ShutdownSignal {
async fn perform_shutdown(
signal: ShutdownSignal,
process_started_at: Instant,
me_pool: Option<Arc<MePool>>,
stats: &Stats,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
quota_state_path: PathBuf,
synlimit_controller: synlimit_control::SynlimitController,
reload_supervisor: ReloadSupervisorHandle,
) {
let shutdown_started_at = Instant::now();
info!(signal = %signal, "Received shutdown signal");
reload_supervisor.quiesce().await;
let runtime = active_runtime.load_full();
let stats = runtime.stats.as_ref();
// Dump stats if SIGQUIT
if signal == ShutdownSignal::Quit {
dump_stats(stats, process_started_at);
@@ -103,12 +111,10 @@ async fn perform_shutdown(
let uptime_secs = process_started_at.elapsed().as_secs();
info!("Uptime: {}", format_uptime(uptime_secs));
if let Err(error) = synlimit_control::clear_synlimit_rules_all_backends().await {
warn!(error = %error, "Failed to clear SYN limiter rules during shutdown");
}
// Graceful ME pool shutdown
if let Some(pool) = &me_pool {
runtime.stop_sessions().await;
runtime.stop_background_tasks().await;
if let Some(pool) = runtime.current_me_pool().await {
match tokio::time::timeout(Duration::from_secs(2), pool.shutdown_send_close_conn_all())
.await
{
@@ -124,6 +130,11 @@ async fn perform_shutdown(
}
}
synlimit_controller.shutdown().await;
if let Err(error) = synlimit_control::clear_synlimit_rules_all_backends().await {
warn!(error = %error, "Failed to clear SYN limiter rules during shutdown");
}
match crate::quota_state::save_quota_state(&quota_state_path, stats).await {
Ok(()) => {
info!(
@@ -191,7 +202,10 @@ fn dump_stats(stats: &Stats, process_started_at: Instant) {
/// - SIGUSR1: Log rotation acknowledgment (for external log rotation tools)
/// - SIGUSR2: Dump runtime status to log
#[cfg(unix)]
pub(crate) fn spawn_signal_handlers(stats: Arc<Stats>, process_started_at: Instant) {
pub(crate) fn spawn_signal_handlers(
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
process_started_at: Instant,
) {
tokio::spawn(async move {
let mut sigusr1 =
signal(SignalKind::user_defined1()).expect("Failed to register SIGUSR1 handler");
@@ -204,7 +218,8 @@ pub(crate) fn spawn_signal_handlers(stats: Arc<Stats>, process_started_at: Insta
handle_sigusr1();
}
_ = sigusr2.recv() => {
handle_sigusr2(&stats, process_started_at);
let runtime = active_runtime.load_full();
handle_sigusr2(runtime.stats.as_ref(), process_started_at);
}
}
}
@@ -213,7 +228,10 @@ pub(crate) fn spawn_signal_handlers(stats: Arc<Stats>, process_started_at: Insta
/// No-op on non-Unix platforms.
#[cfg(not(unix))]
pub(crate) fn spawn_signal_handlers(_stats: Arc<Stats>, _process_started_at: Instant) {
pub(crate) fn spawn_signal_handlers(
_active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
_process_started_at: Instant,
) {
// No SIGUSR1/SIGUSR2 on non-Unix
}
+306 -146
View File
@@ -5,11 +5,86 @@ use rand::RngExt;
use tracing::warn;
use crate::config::ProxyConfig;
use crate::error::{ProxyError, Result};
use crate::startup::{COMPONENT_TLS_FRONT_BOOTSTRAP, StartupTracker};
use crate::tls_front::TlsFrontCache;
use crate::tls_front::fetcher::TlsFetchStrategy;
use crate::transport::UpstreamManager;
use super::generation::RuntimeTaskScope;
/// Readiness requirement for TLS-front cache initialization.
#[derive(Clone, Copy)]
pub(crate) enum TlsBootstrapPolicy {
BestEffort,
RequireReady,
}
#[derive(Clone)]
struct TlsFetchContext {
cache: Arc<TlsFrontCache>,
domains: Vec<String>,
mask_host: String,
primary_domain: String,
mask_unix_sock: Option<String>,
tls_fetch_scope: Option<String>,
upstream_manager: Arc<UpstreamManager>,
strategy: TlsFetchStrategy,
port: u16,
proxy_protocol: u8,
}
impl TlsFetchContext {
async fn fetch_all(&self, failure_message: &'static str) {
let mut join = tokio::task::JoinSet::new();
for domain in self.domains.clone() {
let cache = self.cache.clone();
let host = tls_fetch_host_for_domain(&self.mask_host, &self.primary_domain, &domain);
let unix_sock = self.mask_unix_sock.clone();
let scope = self.tls_fetch_scope.clone();
let upstream = self.upstream_manager.clone();
let strategy = self.strategy.clone();
let port = self.port;
let proxy_protocol = self.proxy_protocol;
join.spawn(async move {
match crate::tls_front::fetcher::fetch_real_tls_with_strategy(
&host,
port,
&domain,
&strategy,
Some(upstream),
scope.as_deref(),
proxy_protocol,
unix_sock.as_deref(),
)
.await
{
Ok(result) => cache.update_from_fetch(&domain, result).await,
Err(error) => warn!(domain = %domain, error = %error, failure_message),
}
});
}
while let Some(result) = join.join_next().await {
if let Err(error) = result {
warn!(error = %error, "TLS emulation fetch task join failed");
}
}
}
async fn fetch_all_with_budget(&self, phase: &'static str) {
if tokio::time::timeout(self.strategy.total_budget, self.fetch_all(phase))
.await
.is_err()
{
warn!(
phase,
timeout_ms = self.strategy.total_budget.as_millis(),
"TLS emulation fetch budget exhausted"
);
}
}
}
fn tls_fetch_host_for_domain(mask_host: &str, primary_tls_domain: &str, domain: &str) -> String {
if mask_host.eq_ignore_ascii_case(primary_tls_domain) {
domain.to_string()
@@ -18,12 +93,24 @@ fn tls_fetch_host_for_domain(mask_host: &str, primary_tls_domain: &str, domain:
}
}
fn readiness_error(default_domains: &[String]) -> Option<String> {
(!default_domains.is_empty()).then(|| {
format!(
"TLS-front profiles are not ready for domains: {}",
default_domains.join(", ")
)
})
}
/// Initializes the TLS-front cache and generation-owned refresh tasks.
pub(crate) async fn bootstrap_tls_front(
config: &ProxyConfig,
tls_domains: &[String],
upstream_manager: Arc<UpstreamManager>,
startup_tracker: &Arc<StartupTracker>,
) -> Option<Arc<TlsFrontCache>> {
task_scope: RuntimeTaskScope,
policy: TlsBootstrapPolicy,
) -> Result<Option<Arc<TlsFrontCache>>> {
startup_tracker
.start_component(
COMPONENT_TLS_FRONT_BOOTSTRAP,
@@ -31,26 +118,38 @@ pub(crate) async fn bootstrap_tls_front(
)
.await;
let tls_cache: Option<Arc<TlsFrontCache>> = if config.censorship.tls_emulation {
let cache = Arc::new(TlsFrontCache::new(
tls_domains,
config.censorship.fake_cert_len,
&config.censorship.tls_front_dir,
));
cache.load_from_disk().await;
if !config.censorship.tls_emulation {
startup_tracker
.skip_component(
COMPONENT_TLS_FRONT_BOOTSTRAP,
Some("censorship.tls_emulation is false".to_string()),
)
.await;
return Ok(None);
}
let port = config.censorship.mask_port;
let proxy_protocol = config.censorship.mask_proxy_protocol;
let mask_host = config
let cache = Arc::new(TlsFrontCache::new(
tls_domains,
config.censorship.fake_cert_len,
&config.censorship.tls_front_dir,
));
cache.load_from_disk().await;
let tls_fetch = config.censorship.tls_fetch.clone();
let fetch_context = TlsFetchContext {
cache: cache.clone(),
domains: tls_domains.to_vec(),
mask_host: config
.censorship
.mask_host
.clone()
.unwrap_or_else(|| config.censorship.tls_domain.clone());
let mask_unix_sock = config.censorship.mask_unix_sock.clone();
let tls_fetch_scope = (!config.censorship.tls_fetch_scope.is_empty())
.then(|| config.censorship.tls_fetch_scope.clone());
let tls_fetch = config.censorship.tls_fetch.clone();
let fetch_strategy = TlsFetchStrategy {
.unwrap_or_else(|| config.censorship.tls_domain.clone()),
primary_domain: config.censorship.tls_domain.clone(),
mask_unix_sock: config.censorship.mask_unix_sock.clone(),
tls_fetch_scope: (!config.censorship.tls_fetch_scope.is_empty())
.then(|| config.censorship.tls_fetch_scope.clone()),
upstream_manager,
strategy: TlsFetchStrategy {
profiles: tls_fetch.profiles,
strict_route: tls_fetch.strict_route,
attempt_timeout: Duration::from_millis(tls_fetch.attempt_timeout_ms.max(1)),
@@ -58,150 +157,110 @@ pub(crate) async fn bootstrap_tls_front(
grease_enabled: tls_fetch.grease_enabled,
deterministic: tls_fetch.deterministic,
profile_cache_ttl: Duration::from_secs(tls_fetch.profile_cache_ttl_secs),
};
let fetch_timeout = fetch_strategy.total_budget;
},
port: config.censorship.mask_port,
proxy_protocol: config.censorship.mask_proxy_protocol,
};
let cache_initial = cache.clone();
let domains_initial = tls_domains.to_vec();
let host_initial = mask_host.clone();
let primary_initial = config.censorship.tls_domain.clone();
let unix_sock_initial = mask_unix_sock.clone();
let scope_initial = tls_fetch_scope.clone();
let upstream_initial = upstream_manager.clone();
let strategy_initial = fetch_strategy.clone();
tokio::spawn(async move {
let mut join = tokio::task::JoinSet::new();
for domain in domains_initial {
let cache_domain = cache_initial.clone();
let host_domain =
tls_fetch_host_for_domain(&host_initial, &primary_initial, &domain);
let unix_sock_domain = unix_sock_initial.clone();
let scope_domain = scope_initial.clone();
let upstream_domain = upstream_initial.clone();
let strategy_domain = strategy_initial.clone();
join.spawn(async move {
match crate::tls_front::fetcher::fetch_real_tls_with_strategy(
&host_domain,
port,
&domain,
&strategy_domain,
Some(upstream_domain),
scope_domain.as_deref(),
proxy_protocol,
unix_sock_domain.as_deref(),
)
match policy {
TlsBootstrapPolicy::BestEffort => {
let initial_fetch = fetch_context.clone();
let fake_cert_len = config.censorship.fake_cert_len;
task_scope.spawn(async move {
initial_fetch
.fetch_all_with_budget("TLS emulation initial fetch failed")
.await;
for domain in initial_fetch
.cache
.default_profile_domains(&initial_fetch.domains)
.await
{
Ok(res) => cache_domain.update_from_fetch(&domain, res).await,
Err(e) => {
warn!(domain = %domain, error = %e, "TLS emulation initial fetch failed")
}
}
});
}
while let Some(res) = join.join_next().await {
if let Err(e) = res {
warn!(error = %e, "TLS emulation initial fetch task join failed");
}
}
});
let cache_timeout = cache.clone();
let domains_timeout = tls_domains.to_vec();
let fake_cert_len = config.censorship.fake_cert_len;
tokio::spawn(async move {
tokio::time::sleep(fetch_timeout).await;
for domain in domains_timeout {
let cached = cache_timeout.get(&domain).await;
if cached.domain == "default" {
{
warn!(
domain = %domain,
timeout_secs = fetch_timeout.as_secs(),
timeout_ms = initial_fetch.strategy.total_budget.as_millis(),
fake_cert_len,
"TLS-front fetch not ready within timeout; using cache/default fake cert fallback"
);
}
});
}
TlsBootstrapPolicy::RequireReady => {
fetch_context
.fetch_all_with_budget("TLS emulation initial fetch failed")
.await;
let default_domains = cache.default_profile_domains(tls_domains).await;
if let Some(error) = readiness_error(&default_domains) {
startup_tracker
.fail_component(COMPONENT_TLS_FRONT_BOOTSTRAP, Some(error.clone()))
.await;
return Err(ProxyError::Proxy(error));
}
});
let cache_refresh = cache.clone();
let domains_refresh = tls_domains.to_vec();
let host_refresh = mask_host.clone();
let primary_refresh = config.censorship.tls_domain.clone();
let unix_sock_refresh = mask_unix_sock.clone();
let scope_refresh = tls_fetch_scope.clone();
let upstream_refresh = upstream_manager.clone();
let strategy_refresh = fetch_strategy.clone();
tokio::spawn(async move {
loop {
let base_secs = rand::rng().random_range(4 * 3600..=6 * 3600);
let jitter_secs = rand::rng().random_range(0..=7200);
tokio::time::sleep(Duration::from_secs(base_secs + jitter_secs)).await;
let mut join = tokio::task::JoinSet::new();
for domain in domains_refresh.clone() {
let cache_domain = cache_refresh.clone();
let host_domain =
tls_fetch_host_for_domain(&host_refresh, &primary_refresh, &domain);
let unix_sock_domain = unix_sock_refresh.clone();
let scope_domain = scope_refresh.clone();
let upstream_domain = upstream_refresh.clone();
let strategy_domain = strategy_refresh.clone();
join.spawn(async move {
match crate::tls_front::fetcher::fetch_real_tls_with_strategy(
&host_domain,
port,
&domain,
&strategy_domain,
Some(upstream_domain),
scope_domain.as_deref(),
proxy_protocol,
unix_sock_domain.as_deref(),
)
.await
{
Ok(res) => cache_domain.update_from_fetch(&domain, res).await,
Err(e) => {
warn!(domain = %domain, error = %e, "TLS emulation refresh failed")
}
}
});
}
while let Some(res) = join.join_next().await {
if let Err(e) = res {
warn!(error = %e, "TLS emulation refresh task join failed");
}
}
}
});
Some(cache)
} else {
startup_tracker
.skip_component(
COMPONENT_TLS_FRONT_BOOTSTRAP,
Some("censorship.tls_emulation is false".to_string()),
)
.await;
None
};
if tls_cache.is_some() {
startup_tracker
.complete_component(
COMPONENT_TLS_FRONT_BOOTSTRAP,
Some("tls front cache is initialized".to_string()),
)
.await;
}
}
tls_cache
let refresh_context = fetch_context;
task_scope.spawn(async move {
loop {
let base_secs = rand::rng().random_range(4 * 3600..=6 * 3600);
let jitter_secs = rand::rng().random_range(0..=7200);
tokio::time::sleep(Duration::from_secs(base_secs + jitter_secs)).await;
refresh_context
.fetch_all_with_budget("TLS emulation refresh failed")
.await;
}
});
startup_tracker
.complete_component(
COMPONENT_TLS_FRONT_BOOTSTRAP,
Some("tls front cache is initialized".to_string()),
)
.await;
Ok(Some(cache))
}
#[cfg(test)]
mod tests {
use super::tls_fetch_host_for_domain;
use super::*;
use crate::startup::StartupComponentStatus;
use crate::stats::Stats;
fn test_config(cache_dir: &std::path::Path) -> ProxyConfig {
let mut config = ProxyConfig::default();
config.censorship.tls_emulation = true;
config.censorship.tls_domain = "front.example".to_string();
config.censorship.mask_host = Some("127.0.0.1".to_string());
config.censorship.mask_port = 1;
config.censorship.tls_front_dir = cache_dir.display().to_string();
config.censorship.tls_fetch.profiles.truncate(1);
config.censorship.tls_fetch.attempt_timeout_ms = 10;
config.censorship.tls_fetch.total_budget_ms = 20;
config
}
fn upstream_manager(config: &ProxyConfig) -> Arc<UpstreamManager> {
Arc::new(UpstreamManager::new(
Vec::new(),
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,
Arc::new(Stats::new()),
))
}
async fn tls_component_status(tracker: &StartupTracker) -> StartupComponentStatus {
tracker
.snapshot()
.await
.components
.into_iter()
.find(|component| component.id == COMPONENT_TLS_FRONT_BOOTSTRAP)
.unwrap()
.status
}
#[test]
fn tls_fetch_host_uses_each_domain_when_mask_host_is_primary_default() {
@@ -218,4 +277,105 @@ mod tests {
"origin.example"
);
}
#[test]
fn readiness_rejects_only_default_profiles() {
assert!(readiness_error(&[]).is_none());
assert_eq!(
readiness_error(&["front.example".to_string()]),
Some("TLS-front profiles are not ready for domains: front.example".to_string())
);
}
#[tokio::test]
async fn require_ready_rejects_default_cache_after_bounded_fetch_failure() {
let cache_dir = tempfile::tempdir().unwrap();
let config = test_config(cache_dir.path());
let domains = vec![config.censorship.tls_domain.clone()];
let tracker = Arc::new(StartupTracker::new(1));
let scope = RuntimeTaskScope::new();
let result = bootstrap_tls_front(
&config,
&domains,
upstream_manager(&config),
&tracker,
scope.clone(),
TlsBootstrapPolicy::RequireReady,
)
.await;
assert!(result.is_err());
assert_eq!(
tls_component_status(&tracker).await,
StartupComponentStatus::Failed
);
scope.stop().await;
}
#[tokio::test]
async fn require_ready_accepts_non_default_disk_cache_when_refresh_fails() {
let cache_dir = tempfile::tempdir().unwrap();
let config = test_config(cache_dir.path());
let domains = vec![config.censorship.tls_domain.clone()];
let seed = TlsFrontCache::new(&domains, config.censorship.fake_cert_len, cache_dir.path());
let mut cached = seed.default_entry().as_ref().clone();
cached.domain = domains[0].clone();
tokio::fs::write(
cache_dir.path().join("front.example.json"),
serde_json::to_vec(&cached).unwrap(),
)
.await
.unwrap();
let tracker = Arc::new(StartupTracker::new(1));
let scope = RuntimeTaskScope::new();
let cache = bootstrap_tls_front(
&config,
&domains,
upstream_manager(&config),
&tracker,
scope.clone(),
TlsBootstrapPolicy::RequireReady,
)
.await
.unwrap()
.unwrap();
assert!(cache.default_profile_domains(&domains).await.is_empty());
assert_eq!(
tls_component_status(&tracker).await,
StartupComponentStatus::Ready
);
scope.stop().await;
}
#[tokio::test]
async fn best_effort_returns_ready_and_refresh_tasks_are_scope_owned() {
let cache_dir = tempfile::tempdir().unwrap();
let config = test_config(cache_dir.path());
let domains = vec![config.censorship.tls_domain.clone()];
let tracker = Arc::new(StartupTracker::new(1));
let scope = RuntimeTaskScope::new();
let cache = bootstrap_tls_front(
&config,
&domains,
upstream_manager(&config),
&tracker,
scope.clone(),
TlsBootstrapPolicy::BestEffort,
)
.await
.unwrap();
assert!(cache.is_some());
assert_eq!(
tls_component_status(&tracker).await,
StartupComponentStatus::Ready
);
tokio::time::timeout(Duration::from_secs(1), scope.stop())
.await
.unwrap();
}
}