use std::sync::Arc;
use std::time::Duration;
use sqlx::PgPool;
use tokio::sync::Notify;
use crate::runtime::config::UdbConfig;
use crate::runtime::metrics::MetricsRecorder;
use super::resources::{ordered_resource_types, resource_type_to_db};
use super::{sourcing, store};
pub const ENV_RELOAD_INTERVAL_MS: &str = "UDB_CONTROL_RELOAD_INTERVAL_MS";
const DEFAULT_RELOAD_INTERVAL_MS: u64 = 1_000;
pub fn reload_interval() -> Duration {
let ms = std::env::var(ENV_RELOAD_INTERVAL_MS)
.ok()
.and_then(|raw| raw.trim().parse::<u64>().ok())
.unwrap_or(DEFAULT_RELOAD_INTERVAL_MS)
.clamp(100, 60_000);
Duration::from_millis(ms)
}
pub type AuthzVersionFn = Arc<dyn Fn() -> String + Send + Sync>;
#[derive(Clone)]
pub struct SubscriberHandle {
pool: PgPool,
config: Arc<UdbConfig>,
reload_notify: Arc<Notify>,
metrics: Option<Arc<dyn MetricsRecorder>>,
authz_version: Option<AuthzVersionFn>,
}
impl SubscriberHandle {
pub fn new(pool: PgPool, config: Arc<UdbConfig>) -> Self {
Self {
pool,
config,
reload_notify: Arc::new(Notify::new()),
metrics: None,
authz_version: None,
}
}
pub fn with_metrics(mut self, metrics: Arc<dyn MetricsRecorder>) -> Self {
self.metrics = Some(metrics);
self
}
pub fn with_authz_version(mut self, probe: AuthzVersionFn) -> Self {
self.authz_version = Some(probe);
self
}
pub fn reload_notify(&self) -> Arc<Notify> {
self.reload_notify.clone()
}
}
#[derive(Default, Clone)]
pub(crate) struct LastSeen {
fingerprint: String,
authz_version: String,
per_type: std::collections::BTreeMap<&'static str, String>,
}
pub fn spawn_control_plane_subscriber(handle: SubscriberHandle) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
let interval = reload_interval();
let mut ticker = tokio::time::interval(interval);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
let mut last = LastSeen::default();
let mut seeded = false;
loop {
ticker.tick().await;
match run_once(&handle, &last, seeded).await {
Ok(next) => {
last = next;
seeded = true;
}
Err(err) => {
tracing::warn!(
error = %err,
"control-plane reload subscriber tick failed; will retry"
);
}
}
}
})
}
pub(crate) async fn run_once(
handle: &SubscriberHandle,
last: &LastSeen,
seeded: bool,
) -> Result<LastSeen, String> {
let reload_started = std::time::Instant::now();
sourcing::resync(&handle.pool, &handle.config)
.await
.map_err(|status| format!("resync failed: {status}"))?;
let mut per_type = std::collections::BTreeMap::new();
for rt in ordered_resource_types() {
let version = store::world_version(&handle.pool, *rt, None, &[])
.await
.map_err(|status| format!("world_version failed: {status}"))?;
per_type.insert(resource_type_to_db(*rt), version);
}
let fingerprint = store::fleet_world_fingerprint(&handle.pool)
.await
.map_err(|status| format!("fleet fingerprint failed: {status}"))?;
let authz_version = handle
.authz_version
.as_ref()
.map(|probe| probe())
.unwrap_or_default();
let next = LastSeen {
fingerprint: fingerprint.clone(),
authz_version: authz_version.clone(),
per_type: per_type.clone(),
};
if !seeded {
return Ok(next);
}
let registry_changed = fingerprint != last.fingerprint;
let authz_changed = authz_version != last.authz_version;
if registry_changed || authz_changed {
if let Some(metrics) = handle.metrics.as_ref() {
metrics.observe_policy_reload_seconds(reload_started.elapsed().as_secs_f64());
if registry_changed {
if let Ok(Some(emit_unix)) = store::latest_updated_at_unix(&handle.pool).await {
let now_unix = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs_f64())
.unwrap_or(0.0);
let lag = (now_unix - emit_unix as f64).max(0.0);
metrics.observe_policy_invalidation_lag_seconds(lag);
}
}
for (rt_db, version) in &per_type {
if last.per_type.get(rt_db) != Some(version) {
metrics.inc_control_reload_applied(rt_db);
}
}
if authz_changed && !registry_changed {
metrics.inc_control_reload_applied(resource_type_to_db(
crate::proto::udb::core::control::entity::v1::ResourceType::MethodSecurityPolicy,
));
}
}
handle.reload_notify.notify_waiters();
}
Ok(next)
}
#[cfg(test)]
mod tests {
use super::*;
fn last_with(fp: &str, authz: &str, pairs: &[(&'static str, &str)]) -> LastSeen {
LastSeen {
fingerprint: fp.to_string(),
authz_version: authz.to_string(),
per_type: pairs.iter().map(|(k, v)| (*k, v.to_string())).collect(),
}
}
#[test]
fn reload_interval_clamps_env() {
unsafe {
std::env::remove_var(ENV_RELOAD_INTERVAL_MS);
}
assert_eq!(reload_interval(), Duration::from_millis(1_000));
}
#[test]
fn change_detection_fires_only_on_real_change() {
let prev = last_with("fp1", "az1", &[("RESOURCE_TYPE_ROUTING_POLICY", "r1")]);
let same = (prev.fingerprint.clone(), prev.authz_version.clone());
let changed = same.0 != prev.fingerprint || same.1 != prev.authz_version;
assert!(!changed, "identical fingerprint+authz must not reload");
let reg = ("fp2".to_string(), prev.authz_version.clone());
assert!(reg.0 != prev.fingerprint || reg.1 != prev.authz_version);
let az = (prev.fingerprint.clone(), "az2".to_string());
assert!(az.0 != prev.fingerprint || az.1 != prev.authz_version);
}
#[test]
fn per_type_reload_counts_only_changed_types() {
let last = last_with(
"fp1",
"az1",
&[
("RESOURCE_TYPE_BACKEND_TARGET_DEFINITION", "b1"),
("RESOURCE_TYPE_ROUTING_POLICY", "r1"),
],
);
let next: std::collections::BTreeMap<&'static str, String> = [
("RESOURCE_TYPE_BACKEND_TARGET_DEFINITION", "b1".to_string()), ("RESOURCE_TYPE_ROUTING_POLICY", "r2".to_string()), ]
.into_iter()
.collect();
let changed: Vec<&'static str> = next
.iter()
.filter(|(k, v)| last.per_type.get(*k) != Some(*v))
.map(|(k, _)| *k)
.collect();
assert_eq!(changed, vec!["RESOURCE_TYPE_ROUTING_POLICY"]);
}
}