use std::any::Any;
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use std::time::Duration;
use async_trait::async_trait;
use camel_api::CamelError;
pub trait TemplateReloadStaged: Send {
fn into_any(self: Box<Self>) -> Box<dyn Any>;
}
type StagedBuild = (Box<dyn TemplateReloadStaged>, u64);
#[async_trait]
pub trait TemplateReloadTarget: Send + Sync {
fn route_id(&self) -> &str;
fn reload_timeout(&self) -> Duration;
fn current_generation(&self) -> u64;
async fn build(&self) -> Result<(Box<dyn TemplateReloadStaged>, u64), CamelError>;
fn commit(&self, staged: Box<dyn TemplateReloadStaged>);
}
struct RegisteredTarget {
id: u64,
target: Arc<dyn TemplateReloadTarget>,
}
static NEXT_ID: AtomicU64 = AtomicU64::new(1);
fn next_id() -> u64 {
NEXT_ID.fetch_add(1, Ordering::Relaxed)
}
pub struct TemplateReloadRegistry {
handlers: Mutex<Vec<RegisteredTarget>>,
route_locks: Mutex<HashMap<String, Arc<tokio::sync::Mutex<()>>>>,
}
impl Default for TemplateReloadRegistry {
fn default() -> Self {
Self {
handlers: Mutex::new(Vec::new()),
route_locks: Mutex::new(HashMap::new()),
}
}
}
impl TemplateReloadRegistry {
pub fn global() -> &'static TemplateReloadRegistry {
static INSTANCE: OnceLock<TemplateReloadRegistry> = OnceLock::new();
INSTANCE.get_or_init(TemplateReloadRegistry::default)
}
pub fn register(&'static self, target: Arc<dyn TemplateReloadTarget>) -> RegistrationGuard {
let id = next_id();
{
let mut guard = self
.handlers
.lock()
.expect("TemplateReloadRegistry handlers lock poisoned"); guard.push(RegisteredTarget { id, target });
}
RegistrationGuard { id, registry: self }
}
pub fn find_all(&self, route_id: &str) -> Vec<Arc<dyn TemplateReloadTarget>> {
let guard = self
.handlers
.lock()
.expect("TemplateReloadRegistry handlers lock poisoned"); guard
.iter()
.filter(|t| t.target.route_id() == route_id)
.map(|t| Arc::clone(&t.target))
.collect()
}
fn remove(&self, id: u64) {
let mut guard = self
.handlers
.lock()
.expect("TemplateReloadRegistry handlers lock poisoned"); guard.retain(|t| t.id != id);
}
fn route_lock(&self, route_id: &str) -> Arc<tokio::sync::Mutex<()>> {
let mut guard = self
.route_locks
.lock()
.expect("TemplateReloadRegistry route_locks lock poisoned"); guard
.entry(route_id.to_string())
.or_insert_with(|| Arc::new(tokio::sync::Mutex::new(())))
.clone()
}
pub async fn reload_route(&self, route_id: &str) -> Result<(), CamelError> {
let route_lock = self.route_lock(route_id);
let _route_guard = route_lock.lock().await;
let targets = self.find_all(route_id);
if targets.is_empty() {
return Err(CamelError::Config(format!(
"no template target for route '{route_id}'"
)));
}
let timeout = targets
.iter()
.map(|t| t.reload_timeout())
.min()
.unwrap_or(Duration::from_millis(5000));
tokio::time::timeout(timeout, async {
let built = futures::future::join_all(targets.iter().map(|t| t.build())).await;
let staged: Vec<StagedBuild> = built.into_iter().collect::<Result<_, _>>()?;
for (target, (_set, read_gen)) in targets.iter().zip(&staged) {
if *read_gen != target.current_generation() {
return Err(CamelError::TemplateReload("stale generation".to_string()));
}
}
for (target, (set, _)) in targets.into_iter().zip(staged) {
target.commit(set);
}
Ok(())
})
.await
.map_err(|_| CamelError::TemplateReload("reload timeout".to_string()))?
}
}
pub struct RegistrationGuard {
id: u64,
registry: &'static TemplateReloadRegistry,
}
impl Drop for RegistrationGuard {
fn drop(&mut self) {
self.registry.remove(self.id);
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Mutex as StdMutex;
use std::sync::atomic::{AtomicUsize, Ordering};
struct FakeStaged {
read_generation: u64,
}
impl TemplateReloadStaged for FakeStaged {
fn into_any(self: Box<Self>) -> Box<dyn Any> {
self
}
}
#[derive(Clone)]
enum BuildMode {
Ok,
Err,
Sleep(Duration),
Stale,
}
#[derive(Default)]
struct FakeState {
generation: AtomicU64,
commit_calls: AtomicUsize,
build_calls: AtomicUsize,
}
struct FakeTarget {
route: String,
timeout: Duration,
state: Arc<FakeState>,
mode: StdMutex<BuildMode>,
events: Option<Arc<StdMutex<Vec<&'static str>>>>,
}
impl FakeTarget {
fn new(route: &str) -> Arc<Self> {
Arc::new(Self {
route: route.to_string(),
timeout: Duration::from_secs(5),
state: Arc::new(FakeState::default()),
mode: StdMutex::new(BuildMode::Ok),
events: None,
})
}
fn set_mode(&self, mode: BuildMode) {
*self.mode.lock().unwrap() = mode;
}
fn as_dyn(self: &Arc<Self>) -> Arc<dyn TemplateReloadTarget> {
let concrete: Arc<Self> = Arc::clone(self);
concrete
}
}
#[async_trait]
impl TemplateReloadTarget for FakeTarget {
fn route_id(&self) -> &str {
&self.route
}
fn reload_timeout(&self) -> Duration {
self.timeout
}
fn current_generation(&self) -> u64 {
self.state.generation.load(Ordering::SeqCst)
}
async fn build(&self) -> Result<(Box<dyn TemplateReloadStaged>, u64), CamelError> {
self.state.build_calls.fetch_add(1, Ordering::SeqCst);
if let Some(ev) = &self.events {
ev.lock().unwrap().push("start");
}
let mode = self.mode.lock().unwrap().clone();
match mode {
BuildMode::Err => {
if let Some(ev) = &self.events {
ev.lock().unwrap().push("end");
}
return Err(CamelError::TemplateReload("fake build failed".to_string()));
}
BuildMode::Sleep(d) => {
tokio::time::sleep(d).await;
}
BuildMode::Ok | BuildMode::Stale => {
tokio::task::yield_now().await;
}
}
let read_gen = match mode {
BuildMode::Stale => self.state.generation.fetch_add(1, Ordering::SeqCst),
_ => self.state.generation.load(Ordering::SeqCst),
};
if let Some(ev) = &self.events {
ev.lock().unwrap().push("end");
}
Ok((
Box::new(FakeStaged {
read_generation: read_gen,
}),
read_gen,
))
}
fn commit(&self, staged: Box<dyn TemplateReloadStaged>) {
let concrete = staged.into_any().downcast::<FakeStaged>().unwrap();
assert_eq!(
concrete.read_generation,
self.state.generation.load(Ordering::SeqCst)
);
self.state.commit_calls.fetch_add(1, Ordering::SeqCst);
self.state.generation.fetch_add(1, Ordering::SeqCst);
}
}
#[test]
fn registry_register_find_all_remove() {
let reg = TemplateReloadRegistry::global();
let route = "test-register-find-all-remove";
let target = FakeTarget::new(route);
let _guard = reg.register(target.as_dyn());
assert_eq!(reg.find_all(route).len(), 1);
drop(_guard);
assert_eq!(reg.find_all(route).len(), 0);
}
#[tokio::test]
async fn reload_route_all_or_nothing() {
let reg = TemplateReloadRegistry::global();
let route = "test-all-or-nothing";
let ok = FakeTarget::new(route);
let err = FakeTarget::new(route);
err.set_mode(BuildMode::Err);
let g1 = reg.register(ok.as_dyn());
let g2 = reg.register(err.as_dyn());
let res = reg.reload_route(route).await;
assert!(res.is_err(), "expected reload to fail");
assert_eq!(
ok.state.commit_calls.load(Ordering::SeqCst),
0,
"OK target must NOT be committed"
);
assert_eq!(
err.state.commit_calls.load(Ordering::SeqCst),
0,
"Err target must NOT be committed"
);
assert_eq!(
ok.state.generation.load(Ordering::SeqCst),
0,
"prior generation retained"
);
drop(g1);
drop(g2);
}
#[tokio::test]
async fn reload_route_commits_all_on_success() {
let reg = TemplateReloadRegistry::global();
let route = "test-commits-all-on-success";
let a = FakeTarget::new(route);
let b = FakeTarget::new(route);
let ga = reg.register(a.as_dyn());
let gb = reg.register(b.as_dyn());
let res = reg.reload_route(route).await;
assert!(res.is_ok(), "expected reload to succeed: {:?}", res);
assert_eq!(a.state.commit_calls.load(Ordering::SeqCst), 1);
assert_eq!(b.state.commit_calls.load(Ordering::SeqCst), 1);
assert_eq!(a.state.generation.load(Ordering::SeqCst), 1);
assert_eq!(b.state.generation.load(Ordering::SeqCst), 1);
drop(ga);
drop(gb);
}
#[tokio::test]
async fn reload_route_timeout_no_commit() {
let reg = TemplateReloadRegistry::global();
let route = "test-timeout-no-commit";
let slow = Arc::new(FakeTarget {
route: route.to_string(),
timeout: Duration::from_millis(40),
state: Arc::new(FakeState::default()),
mode: StdMutex::new(BuildMode::Sleep(Duration::from_millis(2_000))),
events: None,
});
let g = reg.register(slow.as_dyn());
let res = reg.reload_route(route).await;
assert!(
matches!(res, Err(CamelError::TemplateReload(_))),
"expected TemplateReload timeout error, got {res:?}"
);
assert_eq!(
slow.state.commit_calls.load(Ordering::SeqCst),
0,
"commit must never be called on timeout"
);
drop(g);
}
#[tokio::test]
async fn reload_route_rejects_stale_no_commit() {
let reg = TemplateReloadRegistry::global();
let route = "test-rejects-stale-no-commit";
let target = FakeTarget::new(route);
target.set_mode(BuildMode::Stale);
let g = reg.register(target.as_dyn());
let res = reg.reload_route(route).await;
assert!(
matches!(res, Err(CamelError::TemplateReload(_))),
"expected TemplateReload stale error, got {res:?}"
);
assert_eq!(
target.state.commit_calls.load(Ordering::SeqCst),
0,
"commit must never be called on stale rejection"
);
drop(g);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn reload_route_serializes_concurrent() {
let reg = TemplateReloadRegistry::global();
let route = "test-serializes-concurrent";
let events: Arc<StdMutex<Vec<&'static str>>> = Arc::new(StdMutex::new(Vec::new()));
let target = Arc::new(FakeTarget {
route: route.to_string(),
timeout: Duration::from_secs(5),
state: Arc::new(FakeState::default()),
mode: StdMutex::new(BuildMode::Ok),
events: Some(Arc::clone(&events)),
});
let g = reg.register(target.as_dyn());
let h1 = tokio::spawn(async move { reg.reload_route(route).await });
let h2 = tokio::spawn(async move { reg.reload_route(route).await });
let (r1, r2) = tokio::join!(h1, h2);
r1.unwrap().unwrap();
r2.unwrap().unwrap();
let evs = events.lock().unwrap().clone();
assert_eq!(
evs,
vec!["start", "end", "start", "end"],
"per-route mutex must serialize concurrent reload_route"
);
assert_eq!(target.state.commit_calls.load(Ordering::SeqCst), 2);
drop(g);
}
}