#![allow(clippy::field_reassign_with_default)]
#[cfg(not(target_env = "msvc"))]
use tikv_jemallocator::Jemalloc;
#[cfg(not(target_env = "msvc"))]
#[global_allocator]
static GLOBAL: Jemalloc = Jemalloc;
use anyhow::{Context, Result};
use clap::{Parser, Subcommand};
use pingora::prelude::*;
use std::sync::Arc;
use tracing::{debug, error, info, warn};
use zentinel_config::server::{AcmeChallengeType, AcmeConfig};
use zentinel_config::Config;
use zentinel_proxy::acme::{
AcmeClient, AcmeError, CertificateStorage, ChallengeManager, RenewalScheduler,
};
use zentinel_proxy::bundle::{run_bundle_command, BundleArgs};
use zentinel_proxy::tls::{self, CertificateReloader, HotReloadableSniResolver};
use zentinel_proxy::{ReloadTrigger, SignalManager, SignalType, ZentinelProxy};
const VERSION: &str = concat!(
env!("CARGO_PKG_VERSION"),
" (release ",
env!("ZENTINEL_CALVER"),
", commit ",
env!("ZENTINEL_COMMIT"),
")"
);
#[derive(Parser, Debug)]
#[command(name = "zentinel")]
#[command(author, version = VERSION, about, long_about = None)]
#[command(propagate_version = true)]
struct Cli {
#[arg(short = 'c', long = "config", env = "ZENTINEL_CONFIG")]
config: Option<String>,
#[arg(short = 't', long = "test")]
test: bool,
#[arg(long = "verbose")]
verbose: bool,
#[arg(short = 'd', long = "daemon")]
daemon: bool,
#[arg(short = 'u', long = "upgrade")]
upgrade: bool,
#[command(subcommand)]
command: Option<Commands>,
}
#[derive(Subcommand, Debug)]
enum Commands {
Test {
#[arg(short = 'c', long = "config")]
config: Option<String>,
},
Run {
#[arg(short = 'c', long = "config")]
config: Option<String>,
},
Validate {
#[arg(short = 'c', long = "config")]
config: Option<String>,
#[arg(long = "skip-network")]
skip_network: bool,
#[arg(long = "skip-agents")]
skip_agents: bool,
#[arg(long = "skip-certs")]
skip_certs: bool,
},
Lint {
#[arg(short = 'c', long = "config")]
config: Option<String>,
},
Bundle(BundleArgs),
}
fn main() -> Result<()> {
rustls::crypto::aws_lc_rs::default_provider()
.install_default()
.expect("Failed to install rustls crypto provider");
let cli = Cli::parse();
if cli.test {
return test_config(cli.config.as_deref());
}
match cli.command {
Some(Commands::Test { config }) => test_config(config.as_deref().or(cli.config.as_deref())),
Some(Commands::Run { config }) => {
run_server(config.or(cli.config), cli.verbose, cli.daemon, cli.upgrade)
}
Some(Commands::Validate {
config,
skip_network,
skip_agents,
skip_certs,
}) => validate_config(
config.as_deref().or(cli.config.as_deref()),
skip_network,
skip_agents,
skip_certs,
),
Some(Commands::Lint { config }) => lint_config(config.as_deref().or(cli.config.as_deref())),
Some(Commands::Bundle(args)) => {
tracing_subscriber::fmt()
.with_target(false)
.with_level(true)
.init();
run_bundle_command(args)
}
None => {
run_server(cli.config, cli.verbose, cli.daemon, cli.upgrade)
}
}
}
fn test_config(config_path: Option<&str>) -> Result<()> {
tracing_subscriber::fmt()
.with_target(false)
.with_level(true)
.init();
let config = match config_path {
Some(path) => {
info!("Testing configuration file: {}", path);
Config::from_file(path).context("Failed to load configuration file")?
}
None => {
info!("Testing embedded default configuration");
Config::default_embedded().context("Failed to load embedded configuration")?
}
};
config
.validate()
.context("Configuration validation failed")?;
let route_count = config.routes.len();
let upstream_count = config.upstreams.len();
let listener_count = config.listeners.len();
info!("Configuration test successful:");
info!(" - {} listener(s)", listener_count);
info!(" - {} route(s)", route_count);
info!(" - {} upstream(s)", upstream_count);
for route in &config.routes {
if let Some(ref upstream) = route.upstream {
if !config.upstreams.contains_key(upstream) {
warn!(
"Route '{}' references undefined upstream '{}'",
route.id, upstream
);
}
}
}
println!(
"zentinel: configuration file {} test is successful",
config_path.unwrap_or("(embedded)")
);
Ok(())
}
fn validate_config(
config_path: Option<&str>,
skip_network: bool,
skip_agents: bool,
skip_certs: bool,
) -> Result<()> {
tracing_subscriber::fmt()
.with_target(false)
.with_level(true)
.init();
let config = match config_path {
Some(path) => {
info!("Validating configuration file: {}", path);
Config::from_file(path).context("Failed to load configuration file")?
}
None => {
info!("Validating embedded default configuration");
Config::default_embedded().context("Failed to load embedded configuration")?
}
};
config
.validate()
.context("Configuration schema validation failed")?;
println!("✓ Configuration schema valid");
let rt = tokio::runtime::Runtime::new()?;
let result = rt.block_on(async {
use zentinel_config::validate::*;
let opts = ValidationOpts {
skip_network,
skip_agents,
skip_certs,
};
let mut result = ValidationResult::new();
if !opts.skip_network {
println!("Checking upstream connectivity...");
result.merge(network::validate_upstreams(&config).await);
}
if !opts.skip_certs {
println!("Validating TLS certificates...");
result.merge(certs::validate_certificates(&config).await);
}
if !opts.skip_agents {
println!("Checking agent connectivity...");
result.merge(agents::validate_agents(&config).await);
}
result
});
if result.errors.is_empty() {
println!("✓ All validation checks passed");
if !result.warnings.is_empty() {
println!("\nWarnings:");
for warning in &result.warnings {
println!(" ⚠ {}", warning.message);
}
}
std::process::exit(0);
} else {
println!("✗ Validation failed\n");
println!("Errors:");
for error in &result.errors {
println!(" ✗ {}", error.message);
}
if !result.warnings.is_empty() {
println!("\nWarnings:");
for warning in &result.warnings {
println!(" ⚠ {}", warning.message);
}
}
std::process::exit(1);
}
}
fn lint_config(config_path: Option<&str>) -> Result<()> {
tracing_subscriber::fmt()
.with_target(false)
.with_level(true)
.init();
let config = match config_path {
Some(path) => {
info!("Linting configuration file: {}", path);
Config::from_file(path).context("Failed to load configuration file")?
}
None => {
info!("Linting embedded default configuration");
Config::default_embedded().context("Failed to load embedded configuration")?
}
};
config
.validate()
.context("Configuration schema validation failed")?;
let mut result = zentinel_config::validate::lint::lint_config(&config);
if let Some(path) = config_path {
match std::fs::read_to_string(path) {
Ok(source) => {
zentinel_config::validate::unknown_keys::check_unknown_keys(&source, &mut result)
}
Err(e) => {
warn!(path = %path, error = %e, "Could not re-read config to check for unknown keys")
}
}
}
if result.warnings.is_empty() {
println!("✓ No best practice issues found");
std::process::exit(0);
} else {
println!(
"⚠ Configuration has {} best practice warnings:\n",
result.warnings.len()
);
for warning in &result.warnings {
println!(" ⚠ {}", warning.message);
}
std::process::exit(0);
}
}
struct AcmeState {
challenge_manager: Arc<ChallengeManager>,
schedulers: Vec<RenewalScheduler>,
}
async fn initialize_acme(
config: &Config,
sni_resolver: Option<Arc<HotReloadableSniResolver>>,
) -> Result<Option<AcmeState>, AcmeError> {
let mut acme_configs: Vec<(String, AcmeConfig)> = Vec::new();
for listener in &config.listeners {
if listener.protocol == zentinel_config::ListenerProtocol::Https {
if let Some(ref tls) = listener.tls {
if let Some(ref acme) = tls.acme {
acme_configs.push((format!("listener '{}' (root)", listener.id), acme.clone()));
}
for (i, sni) in tls.additional_certs.iter().enumerate() {
if let Some(ref acme) = sni.acme {
acme_configs.push((
format!("listener '{}' (sni cert #{})", listener.id, i),
acme.clone(),
));
}
}
}
}
}
if acme_configs.is_empty() {
return Ok(None);
}
info!(
config_count = acme_configs.len(),
"Initializing ACME certificate management for multiple configurations"
);
let challenge_manager = Arc::new(ChallengeManager::new());
let mut schedulers = Vec::new();
for (description, acme_config) in acme_configs {
info!(
source = %description,
domains = ?acme_config.domains,
staging = acme_config.staging,
challenge_type = ?acme_config.challenge_type,
"Initializing ACME for {}", description
);
let storage = Arc::new(CertificateStorage::new(&acme_config.storage)?);
let acme_client = Arc::new(AcmeClient::new(acme_config.clone(), Arc::clone(&storage)));
let mut scheduler = RenewalScheduler::new(
Arc::clone(&acme_client),
Arc::clone(&challenge_manager),
sni_resolver.clone(),
);
if acme_config.challenge_type == AcmeChallengeType::Dns01 {
if let Some(ref dns_config) = acme_config.dns_provider {
let provider = zentinel_proxy::acme::dns::create_provider(dns_config)?;
let mut nameservers: Vec<std::net::IpAddr> = dns_config
.propagation
.nameservers
.iter()
.filter_map(|s| s.parse().ok())
.collect();
if nameservers.is_empty() {
tracing::info!(
"propagation nameservers not configured, falling back to public resolvers 8.8.8.8, 1.1.1.1, 9.9.9.9"
);
nameservers =
zentinel_proxy::acme::dns::PropagationConfig::default().nameservers;
}
let propagation_config = zentinel_proxy::acme::dns::PropagationConfig {
initial_delay: std::time::Duration::from_secs(
dns_config.propagation.initial_delay_secs,
),
check_interval: std::time::Duration::from_secs(
dns_config.propagation.check_interval_secs,
),
timeout: std::time::Duration::from_secs(dns_config.propagation.timeout_secs),
nameservers,
};
let dns_manager = Arc::new(zentinel_proxy::acme::dns::Dns01ChallengeManager::new(
provider,
propagation_config,
)?);
scheduler = scheduler.with_dns_manager(dns_manager);
}
}
let primary_domain_for_account = acme_config.domains.first().cloned();
if let Err(e) = acme_client.init_account().await {
use zentinel_proxy::acme::is_retryable_acme_error;
if is_retryable_acme_error(&e) {
let has_cert = primary_domain_for_account
.as_deref()
.and_then(|d| storage.certificate_paths(d))
.is_some();
if has_cert {
tracing::warn!(
source = %description,
error = %e,
"ACME account init transient failure, proxy will stay ready and retry in background (renewal)"
);
schedulers.push(scheduler);
continue;
}
tracing::error!(
source = %description,
error = %e,
"ACME account init transient failure during first issuance, failing fast (no cert to serve)"
);
return Err(e);
} else {
return Err(e);
}
}
let primary_domain = acme_config.domains.first().ok_or_else(|| {
AcmeError::OrderCreation(format!("No domains configured for ACME in {}", description))
})?;
if acme_client.needs_renewal(primary_domain)? {
let issuance_result: Result<(), AcmeError> = async {
info!(
source = %description,
domain = %primary_domain,
"Initial certificate issuance required"
);
match acme_config.challenge_type {
AcmeChallengeType::Http01 => {
let http_addr = config
.listeners
.iter()
.find(|l| l.protocol == zentinel_config::ListenerProtocol::Http)
.map(|l| l.address.clone())
.unwrap_or_else(|| "0.0.0.0:80".to_string());
info!(
address = %http_addr,
"Starting temporary HTTP challenge server for initial certificate acquisition"
);
let (shutdown_tx, shutdown_rx) = tokio::sync::watch::channel(false);
let cm_clone = Arc::clone(&challenge_manager);
let _server_handle = tokio::spawn(async move {
zentinel_proxy::acme::challenge_server::run_challenge_server(
&http_addr,
cm_clone,
shutdown_rx,
)
.await
});
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
let result = scheduler.ensure_certificates().await;
let _ = shutdown_tx.send(true);
result
}
AcmeChallengeType::Dns01 => scheduler.ensure_certificates().await,
}
}
.await;
if let Err(e) = issuance_result {
use zentinel_proxy::acme::is_retryable_acme_error;
if is_retryable_acme_error(&e) {
let has_cert = storage.certificate_paths(primary_domain).is_some();
if has_cert {
tracing::warn!(
source = %description,
domain = %primary_domain,
error = %e,
"Initial ACME renewal transient failure, deferring to background renewal"
);
} else {
tracing::error!(
source = %description,
domain = %primary_domain,
error = %e,
"Initial ACME issuance transient failure during first issuance, failing fast (no cert to serve)"
);
return Err(e);
}
} else {
return Err(e);
}
}
}
schedulers.push(scheduler);
}
Ok(Some(AcmeState {
challenge_manager,
schedulers,
}))
}
#[derive(Default)]
struct StartupLog(Vec<(tracing::Level, String)>);
impl StartupLog {
fn info(&mut self, message: String) {
self.0.push((tracing::Level::INFO, message));
}
fn warn(&mut self, message: String) {
self.0.push((tracing::Level::WARN, message));
}
fn emit(self) {
for (level, message) in self.0 {
match level {
tracing::Level::WARN => warn!("{message}"),
_ => info!("{message}"),
}
}
}
}
fn resolve_logging_config(
config_path: Option<&str>,
startup_log: &mut StartupLog,
) -> zentinel_config::LoggingConfig {
let parsed = match config_path {
Some(path) => zentinel_config::Config::from_file(path).ok(),
None => zentinel_config::Config::default_embedded().ok(),
};
match parsed {
Some(config) => config.observability.logging,
None => {
startup_log.warn(
"Could not read logging configuration; using defaults until the \
configuration is loaded"
.to_string(),
);
zentinel_config::LoggingConfig::default()
}
}
}
fn init_logging(verbose: bool, timestamps: bool) {
let log_level = if verbose { "debug" } else { "info" };
let filter = tracing_subscriber::EnvFilter::try_from_default_env()
.unwrap_or_else(|_| tracing_subscriber::EnvFilter::new(log_level));
if timestamps {
tracing_subscriber::fmt().with_env_filter(filter).init();
} else {
tracing_subscriber::fmt()
.with_env_filter(filter)
.without_time()
.init();
}
}
fn run_server(
config_path: Option<String>,
verbose: bool,
daemon: bool,
upgrade: bool,
) -> Result<()> {
let mut startup_log = StartupLog::default();
let mut pingora_opt = Opt::default();
pingora_opt.daemon = daemon;
pingora_opt.upgrade = upgrade;
let effective_config_path = config_path.or_else(|| std::env::var("ZENTINEL_CONFIG").ok());
let effective_config_path = match effective_config_path {
Some(path) => {
let config_path = std::path::Path::new(&path);
if config_path.exists() {
startup_log.info(format!("Loading configuration from: {path}"));
Some(path)
} else {
startup_log.info(format!("Configuration file not found: {path}"));
if let Err(e) = create_default_config_file(config_path) {
startup_log.warn(format!("Failed to create default config file: {e}"));
startup_log.info("Using embedded default configuration instead".to_string());
None
} else {
startup_log.info(format!("Created default configuration at: {path}"));
Some(path)
}
}
}
None => {
startup_log.info(
"No configuration specified, using embedded default configuration".to_string(),
);
None
}
};
let logging = resolve_logging_config(effective_config_path.as_deref(), &mut startup_log);
init_logging(verbose, logging.timestamps);
startup_log.emit();
let signal_manager = Arc::new(SignalManager::new());
let runtime = tokio::runtime::Runtime::new()?;
let mut proxy =
runtime.block_on(async { ZentinelProxy::new(effective_config_path.as_deref()).await })?;
let config_manager = proxy.config_manager.clone();
let config = proxy.config_manager.current();
{
let metrics_cfg = config.observability.metrics.clone();
if metrics_cfg.enabled {
let cache_stats = Some(proxy.http_cache_stats());
runtime.spawn(async move {
zentinel_proxy::metrics_server::run_metrics_server(
metrics_cfg.address,
metrics_cfg.path,
cache_stats,
)
.await;
});
} else {
info!("Metrics server disabled (observability.metrics.enabled = false)");
}
}
setup_signal_handlers(
signal_manager.sender(),
config.server.graceful_shutdown_timeout_secs,
);
let acme_state = runtime
.block_on(async { initialize_acme(&config, None).await })
.context("ACME initialization failed")?;
if let Some(ref state) = acme_state {
proxy.acme_challenges = Some(Arc::clone(&state.challenge_manager));
proxy.acme_clients = state
.schedulers
.iter()
.map(|s| Arc::clone(s.client()))
.collect();
}
if let Some(ref tracing_config) = config.observability.tracing {
if !tracing_config.enabled {
info!("Distributed tracing disabled by configuration (tracing.enabled)");
} else {
match zentinel_proxy::otel::init_tracer(tracing_config) {
Ok(()) => {
info!(
backend = ?tracing_config.backend,
sampling_rate = tracing_config.sampling_rate,
service_name = %tracing_config.service_name,
"OpenTelemetry tracing enabled"
);
}
Err(e) => {
warn!("Failed to initialize OpenTelemetry tracer: {}", e);
warn!("Distributed tracing will be disabled");
}
}
}
}
let worker_threads = if config.server.worker_threads > 0 {
config.server.worker_threads
} else {
num_cpus::get() };
let mut pingora_conf = pingora::server::configuration::ServerConf::default();
pingora_conf.threads = worker_threads;
pingora_conf.work_stealing = true;
pingora_conf.upstream_keepalive_pool_size = 256;
pingora_conf.graceful_shutdown_timeout_seconds =
Some(config.server.graceful_shutdown_timeout_secs);
if let Some(ref pid_path) = config.server.pid_file {
pingora_conf.pid_file = pid_path.to_string_lossy().to_string();
}
if let Some(ref user) = config.server.user {
pingora_conf.user = Some(user.clone());
}
if let Some(ref group) = config.server.group {
pingora_conf.group = Some(group.clone());
}
info!(
worker_threads = worker_threads,
upstream_pool_size = pingora_conf.upstream_keepalive_pool_size,
graceful_shutdown_timeout_secs = config.server.graceful_shutdown_timeout_secs,
pid_file = ?config.server.pid_file,
user = ?config.server.user,
group = ?config.server.group,
"Configuring Pingora server"
);
if let Some(ref work_dir) = config.server.working_directory {
std::env::set_current_dir(work_dir).with_context(|| {
format!(
"Failed to change working directory to '{}'",
work_dir.display()
)
})?;
info!(path = %work_dir.display(), "Changed working directory");
}
let mut server = Server::new_with_opt_and_conf(Some(pingora_opt), pingora_conf);
server.bootstrap();
let keepalive_request_limit = config
.listeners
.iter()
.filter_map(|l| l.keepalive_max_requests)
.min();
let mut server_options = pingora_core::apps::HttpServerOptions::default();
server_options.keepalive_request_limit = keepalive_request_limit;
let mut proxy_service = pingora_proxy::ProxyServiceBuilder::new(&server.configuration, proxy)
.name("Zentinel Proxy")
.server_options(server_options)
.build();
let cert_reloader = Arc::new(CertificateReloader::new());
for listener in &config.listeners {
match listener.protocol {
zentinel_config::ListenerProtocol::Http => {
proxy_service.add_tcp(&listener.address);
info!("HTTP listening on: {}", listener.address);
}
zentinel_config::ListenerProtocol::Https | zentinel_config::ListenerProtocol::Http2 => {
match &listener.tls {
Some(tls_config) => {
let (cert_path, key_path) = if let (Some(ref cert), Some(ref key)) =
(&tls_config.cert_file, &tls_config.key_file)
{
(cert.clone(), key.clone())
} else if let Some(ref acme_config) = tls_config.acme {
let acme_storage = &acme_config.storage;
let primary_domain = acme_config
.domains
.first()
.ok_or_else(|| {
error!(
listener_id = %listener.id,
"ACME configuration has no domains"
);
})
.unwrap_or(&"default".to_string())
.clone();
let cert_path = acme_storage
.join("domains")
.join(&primary_domain)
.join("cert.pem");
let key_path = acme_storage
.join("domains")
.join(&primary_domain)
.join("key.pem");
if !cert_path.exists() || !key_path.exists() {
error!(
listener_id = %listener.id,
address = %listener.address,
domains = ?acme_config.domains,
cert_path = %cert_path.display(),
"ACME certificate files not found after initialization"
);
continue;
}
(cert_path, key_path)
} else {
error!(
listener_id = %listener.id,
"TLS configuration requires either cert-file/key-file or acme block"
);
continue;
};
let cert_path_str = cert_path.to_string_lossy();
let key_path_str = key_path.to_string_lossy();
if !cert_path.exists() {
error!(
listener_id = %listener.id,
cert_file = %cert_path_str,
"TLS certificate file not found"
);
continue;
}
if !key_path.exists() {
error!(
listener_id = %listener.id,
key_file = %key_path_str,
"TLS key file not found"
);
continue;
}
let sni_resolver = match HotReloadableSniResolver::from_config(
tls_config.clone(),
listener.id.clone(),
) {
Ok(r) => Arc::new(r),
Err(e) => {
error!(
listener_id = %listener.id,
error = %e,
"Failed to load TLS certificates for listener"
);
continue;
}
};
let server_config = match tls::build_server_config_with_resolver(
tls_config,
sni_resolver.clone(),
) {
Ok(c) => c,
Err(e) => {
error!(
listener_id = %listener.id,
error = %e,
"Failed to build TLS configuration for listener"
);
continue;
}
};
let mut tls_settings =
match pingora::listeners::tls::TlsSettings::with_server_config(
server_config,
) {
Ok(s) => s,
Err(e) => {
error!(
listener_id = %listener.id,
error = %e,
"Failed to create TLS settings"
);
continue;
}
};
tls_settings.enable_h2();
cert_reloader.register(&listener.id, sni_resolver.clone());
spawn_cert_folder_reloaders(
&runtime,
&listener.id,
tls_config,
sni_resolver,
);
proxy_service.add_tls_with_settings(&listener.address, None, tls_settings);
info!(
listener_id = %listener.id,
address = %listener.address,
cert_file = %cert_path_str,
acme_enabled = tls_config.acme.is_some(),
sni_cert_count = tls_config.additional_certs.len(),
client_auth = tls_config.client_auth,
"HTTPS (h2+http/1.1) listening on: {}", listener.address
);
}
None => {
error!(
listener_id = %listener.id,
address = %listener.address,
"HTTPS listener requires TLS configuration"
);
}
}
}
zentinel_config::ListenerProtocol::Http3 => {
error!(
listener_id = %listener.id,
address = %listener.address,
"HTTP/3 is not implemented; refusing to start rather than \
binding nothing. Use protocol \"https\" (which negotiates \
HTTP/2 via ALPN) until QUIC support lands."
);
return Err(anyhow::anyhow!(
"Listener '{}' requests HTTP/3, which is not implemented",
listener.id
));
}
}
}
server.add_service(proxy_service);
let auto_reload_enabled = config.server.auto_reload;
let has_config_file = effective_config_path.is_some();
if auto_reload_enabled && has_config_file {
let config_manager_watch = config_manager.clone();
runtime.spawn(async move {
if let Err(e) = config_manager_watch.start_watching().await {
error!("Failed to start config file watcher: {}", e);
error!("Auto-reload disabled, use SIGHUP for manual reload");
}
});
} else if auto_reload_enabled && !has_config_file {
warn!("auto-reload enabled but no config file specified (using embedded config)");
warn!("Auto-reload requires a config file path");
}
if let Some(state) = acme_state {
let scheduler_count = state.schedulers.len();
for scheduler in state.schedulers {
runtime.spawn(async move {
scheduler.run().await;
});
}
info!(
count = scheduler_count,
"ACME certificate renewal schedulers started"
);
}
let signal_manager_clone = signal_manager.clone();
let cert_reloader_clone = cert_reloader.clone();
runtime.spawn(async move {
run_signal_handler(signal_manager_clone, config_manager, cert_reloader_clone).await;
});
info!("Zentinel proxy started successfully");
info!("Configuration hot reload enabled (SIGHUP)");
if auto_reload_enabled && has_config_file {
info!("Auto-reload enabled (watching config file)");
}
info!("Graceful shutdown enabled (SIGTERM/SIGINT)");
server.run_forever();
}
fn setup_signal_handlers(
signal_tx: std::sync::mpsc::Sender<SignalType>,
graceful_shutdown_timeout_secs: u64,
) {
use signal_hook::consts::signal::*;
use signal_hook::iterator::Signals;
use std::thread;
let mut signals =
Signals::new([SIGTERM, SIGINT, SIGHUP]).expect("Failed to register signal handlers");
thread::spawn(move || {
for sig in signals.forever() {
let signal_type = match sig {
SIGTERM | SIGINT => {
info!(
"Received shutdown signal ({}), initiating graceful shutdown",
if sig == SIGTERM { "SIGTERM" } else { "SIGINT" }
);
SignalType::Shutdown
}
SIGHUP => {
info!("Received SIGHUP, triggering configuration reload");
SignalType::Reload
}
_ => continue,
};
if signal_tx.send(signal_type).is_err() {
break;
}
if signal_type == SignalType::Shutdown {
let force_exit_secs = graceful_shutdown_timeout_secs.saturating_add(5);
thread::sleep(std::time::Duration::from_secs(force_exit_secs));
error!(
timeout_secs = force_exit_secs,
"Graceful shutdown timeout exceeded, forcing exit"
);
std::process::exit(1);
}
}
});
}
fn create_default_config_file(path: &std::path::Path) -> Result<()> {
use std::fs;
use zentinel_config::DEFAULT_CONFIG_KDL;
if let Some(parent) = path.parent() {
if !parent.exists() {
fs::create_dir_all(parent)
.with_context(|| format!("Failed to create config directory: {:?}", parent))?;
}
}
fs::write(path, DEFAULT_CONFIG_KDL.trim_start())
.with_context(|| format!("Failed to write default config to: {:?}", path))?;
Ok(())
}
fn spawn_cert_folder_reloaders(
runtime: &tokio::runtime::Runtime,
listener_id: &str,
tls_config: &zentinel_config::TlsConfig,
resolver: Arc<HotReloadableSniResolver>,
) {
use zentinel_config::CertFolderReloadMode;
for folder in &tls_config.cert_folders {
match folder.reload_mode {
CertFolderReloadMode::Off => continue,
CertFolderReloadMode::Interval => {
info!(
listener_id = %listener_id,
cert_folder = %folder.cert_folder.display(),
interval_secs = folder.reload_interval.as_secs(),
"Certificate folder will be rescanned on a timer"
);
let resolver = resolver.clone();
let listener_id = listener_id.to_string();
let path = folder.cert_folder.clone();
let interval = folder.reload_interval;
runtime.spawn(async move {
let mut ticker = tokio::time::interval(interval);
ticker.tick().await;
loop {
ticker.tick().await;
reload_folder(&resolver, &listener_id, &path, "interval");
}
});
}
CertFolderReloadMode::Watch => {
info!(
listener_id = %listener_id,
cert_folder = %folder.cert_folder.display(),
"Certificate folder will be rescanned when it changes"
);
let resolver = resolver.clone();
let listener_id = listener_id.to_string();
let path = folder.cert_folder.clone();
let interval = folder.reload_interval;
runtime.spawn(async move {
watch_cert_folder(resolver, listener_id, path, interval).await;
});
}
}
}
}
fn reload_folder(
resolver: &HotReloadableSniResolver,
listener_id: &str,
path: &std::path::Path,
trigger: &str,
) {
match resolver.reload() {
Ok(()) => {
debug!(
listener_id = %listener_id,
cert_folder = %path.display(),
trigger = trigger,
"Certificates reloaded"
);
}
Err(e) => {
error!(
listener_id = %listener_id,
cert_folder = %path.display(),
trigger = trigger,
error = %e,
"Certificate reload failed; continuing with the previous certificates"
);
}
}
}
async fn watch_cert_folder(
resolver: Arc<HotReloadableSniResolver>,
listener_id: String,
path: std::path::PathBuf,
fallback_interval: std::time::Duration,
) {
use notify::{RecursiveMode, Watcher};
let (tx, mut rx) = tokio::sync::mpsc::channel::<()>(1);
let watcher = notify::recommended_watcher(move |res: notify::Result<notify::Event>| {
if res.is_ok() {
let _ = tx.try_send(());
}
});
let mut watcher = match watcher {
Ok(w) => w,
Err(e) => {
warn!(
listener_id = %listener_id,
cert_folder = %path.display(),
error = %e,
interval_secs = fallback_interval.as_secs(),
"Could not create a filesystem watcher; falling back to interval rescans"
);
return interval_fallback(resolver, listener_id, path, fallback_interval).await;
}
};
if let Err(e) = watcher.watch(&path, RecursiveMode::NonRecursive) {
warn!(
listener_id = %listener_id,
cert_folder = %path.display(),
error = %e,
interval_secs = fallback_interval.as_secs(),
"Could not watch the certificate folder; falling back to interval rescans"
);
return interval_fallback(resolver, listener_id, path, fallback_interval).await;
}
const SETTLE: std::time::Duration = std::time::Duration::from_millis(500);
while rx.recv().await.is_some() {
tokio::time::sleep(SETTLE).await;
while rx.try_recv().is_ok() {}
reload_folder(&resolver, &listener_id, &path, "watch");
}
drop(watcher);
}
async fn interval_fallback(
resolver: Arc<HotReloadableSniResolver>,
listener_id: String,
path: std::path::PathBuf,
interval: std::time::Duration,
) {
let mut ticker = tokio::time::interval(interval);
ticker.tick().await;
loop {
ticker.tick().await;
reload_folder(&resolver, &listener_id, &path, "interval-fallback");
}
}
async fn run_signal_handler(
signal_manager: Arc<SignalManager>,
config_manager: Arc<zentinel_proxy::ConfigManager>,
cert_reloader: Arc<CertificateReloader>,
) {
loop {
let signal_manager_clone = signal_manager.clone();
let signal =
tokio::task::spawn_blocking(move || signal_manager_clone.recv_blocking()).await;
match signal {
Ok(Some(SignalType::Reload)) => {
info!("Processing configuration reload request");
match config_manager.reload(ReloadTrigger::Signal).await {
Ok(()) => {
info!("Configuration reloaded successfully");
}
Err(e) => {
error!("Configuration reload failed: {}", e);
error!("Continuing with previous configuration");
}
}
let (reloaded, failures) = cert_reloader.reload_all();
for (listener_id, e) in &failures {
error!(
listener_id = %listener_id,
error = %e,
"TLS certificate reload failed, keeping previous certificates"
);
}
if reloaded > 0 || !failures.is_empty() {
info!(
reloaded = reloaded,
failed = failures.len(),
"TLS certificates reloaded"
);
}
}
Ok(Some(SignalType::Shutdown)) => {
info!("Processing graceful shutdown request");
zentinel_proxy::otel::shutdown_tracer();
info!("Shutdown initiated, draining connections...");
std::process::exit(0);
}
Ok(None) => {
info!("Signal channel closed, stopping signal handler");
break;
}
Err(e) => {
error!("Signal handler task panicked: {}", e);
break;
}
}
}
}