1use std::any::Any;
15use std::collections::HashMap;
16use std::sync::atomic::{AtomicU64, Ordering};
17use std::sync::{Arc, Mutex, OnceLock};
18use std::time::Duration;
19
20use async_trait::async_trait;
21use camel_api::CamelError;
22
23pub trait TemplateReloadStaged: Send {
31 fn into_any(self: Box<Self>) -> Box<dyn Any>;
33}
34
35type StagedBuild = (Box<dyn TemplateReloadStaged>, u64);
37
38#[async_trait]
45pub trait TemplateReloadTarget: Send + Sync {
46 fn route_id(&self) -> &str;
48 fn reload_timeout(&self) -> Duration;
52 fn current_generation(&self) -> u64;
54 async fn build(&self) -> Result<(Box<dyn TemplateReloadStaged>, u64), CamelError>;
58 fn commit(&self, staged: Box<dyn TemplateReloadStaged>);
61}
62
63struct RegisteredTarget {
65 id: u64,
66 target: Arc<dyn TemplateReloadTarget>,
67}
68
69static NEXT_ID: AtomicU64 = AtomicU64::new(1);
71
72fn next_id() -> u64 {
73 NEXT_ID.fetch_add(1, Ordering::Relaxed)
74}
75
76pub struct TemplateReloadRegistry {
78 handlers: Mutex<Vec<RegisteredTarget>>,
79 route_locks: Mutex<HashMap<String, Arc<tokio::sync::Mutex<()>>>>,
80}
81
82impl Default for TemplateReloadRegistry {
83 fn default() -> Self {
84 Self {
85 handlers: Mutex::new(Vec::new()),
86 route_locks: Mutex::new(HashMap::new()),
87 }
88 }
89}
90
91impl TemplateReloadRegistry {
92 pub fn global() -> &'static TemplateReloadRegistry {
94 static INSTANCE: OnceLock<TemplateReloadRegistry> = OnceLock::new();
95 INSTANCE.get_or_init(TemplateReloadRegistry::default)
96 }
97
98 pub fn register(&'static self, target: Arc<dyn TemplateReloadTarget>) -> RegistrationGuard {
104 let id = next_id();
105 {
106 let mut guard = self
107 .handlers
108 .lock()
109 .expect("TemplateReloadRegistry handlers lock poisoned"); guard.push(RegisteredTarget { id, target });
111 }
112 RegistrationGuard { id, registry: self }
113 }
114
115 pub fn find_all(&self, route_id: &str) -> Vec<Arc<dyn TemplateReloadTarget>> {
118 let guard = self
119 .handlers
120 .lock()
121 .expect("TemplateReloadRegistry handlers lock poisoned"); guard
123 .iter()
124 .filter(|t| t.target.route_id() == route_id)
125 .map(|t| Arc::clone(&t.target))
126 .collect()
127 }
128
129 fn remove(&self, id: u64) {
133 let mut guard = self
134 .handlers
135 .lock()
136 .expect("TemplateReloadRegistry handlers lock poisoned"); guard.retain(|t| t.id != id);
138 }
139
140 fn route_lock(&self, route_id: &str) -> Arc<tokio::sync::Mutex<()>> {
144 let mut guard = self
145 .route_locks
146 .lock()
147 .expect("TemplateReloadRegistry route_locks lock poisoned"); guard
149 .entry(route_id.to_string())
150 .or_insert_with(|| Arc::new(tokio::sync::Mutex::new(())))
151 .clone()
152 }
153
154 pub async fn reload_route(&self, route_id: &str) -> Result<(), CamelError> {
171 let route_lock = self.route_lock(route_id);
175 let _route_guard = route_lock.lock().await;
176
177 let targets = self.find_all(route_id);
178 if targets.is_empty() {
179 return Err(CamelError::Config(format!(
180 "no template target for route '{route_id}'"
181 )));
182 }
183
184 let timeout = targets
186 .iter()
187 .map(|t| t.reload_timeout())
188 .min()
189 .unwrap_or(Duration::from_millis(5000));
190
191 tokio::time::timeout(timeout, async {
192 let built = futures::future::join_all(targets.iter().map(|t| t.build())).await;
195 let staged: Vec<StagedBuild> = built.into_iter().collect::<Result<_, _>>()?;
197
198 for (target, (_set, read_gen)) in targets.iter().zip(&staged) {
202 if *read_gen != target.current_generation() {
203 return Err(CamelError::TemplateReload("stale generation".to_string()));
204 }
205 }
206
207 for (target, (set, _)) in targets.into_iter().zip(staged) {
210 target.commit(set);
211 }
212 Ok(())
213 })
214 .await
215 .map_err(|_| CamelError::TemplateReload("reload timeout".to_string()))?
216 }
217}
218
219pub struct RegistrationGuard {
223 id: u64,
224 registry: &'static TemplateReloadRegistry,
225}
226
227impl Drop for RegistrationGuard {
228 fn drop(&mut self) {
229 self.registry.remove(self.id);
230 }
231}
232
233#[cfg(test)]
234mod tests {
235 use super::*;
236 use std::sync::Mutex as StdMutex;
237 use std::sync::atomic::{AtomicUsize, Ordering};
238
239 struct FakeStaged {
241 read_generation: u64,
242 }
243 impl TemplateReloadStaged for FakeStaged {
244 fn into_any(self: Box<Self>) -> Box<dyn Any> {
246 self
247 }
248 }
249
250 #[derive(Clone)]
252 enum BuildMode {
253 Ok,
255 Err,
257 Sleep(Duration),
259 Stale,
261 }
262
263 #[derive(Default)]
265 struct FakeState {
266 generation: AtomicU64,
267 commit_calls: AtomicUsize,
268 build_calls: AtomicUsize,
269 }
270
271 struct FakeTarget {
272 route: String,
273 timeout: Duration,
274 state: Arc<FakeState>,
275 mode: StdMutex<BuildMode>,
276 events: Option<Arc<StdMutex<Vec<&'static str>>>>,
278 }
279
280 impl FakeTarget {
281 fn new(route: &str) -> Arc<Self> {
282 Arc::new(Self {
283 route: route.to_string(),
284 timeout: Duration::from_secs(5),
285 state: Arc::new(FakeState::default()),
286 mode: StdMutex::new(BuildMode::Ok),
287 events: None,
288 })
289 }
290
291 fn set_mode(&self, mode: BuildMode) {
292 *self.mode.lock().unwrap() = mode;
293 }
294
295 fn as_dyn(self: &Arc<Self>) -> Arc<dyn TemplateReloadTarget> {
297 let concrete: Arc<Self> = Arc::clone(self);
301 concrete
302 }
303 }
304
305 #[async_trait]
306 impl TemplateReloadTarget for FakeTarget {
307 fn route_id(&self) -> &str {
308 &self.route
309 }
310 fn reload_timeout(&self) -> Duration {
311 self.timeout
312 }
313 fn current_generation(&self) -> u64 {
314 self.state.generation.load(Ordering::SeqCst)
315 }
316 async fn build(&self) -> Result<(Box<dyn TemplateReloadStaged>, u64), CamelError> {
317 self.state.build_calls.fetch_add(1, Ordering::SeqCst);
318 if let Some(ev) = &self.events {
319 ev.lock().unwrap().push("start");
320 }
321 let mode = self.mode.lock().unwrap().clone();
322 match mode {
323 BuildMode::Err => {
324 if let Some(ev) = &self.events {
325 ev.lock().unwrap().push("end");
326 }
327 return Err(CamelError::TemplateReload("fake build failed".to_string()));
328 }
329 BuildMode::Sleep(d) => {
330 tokio::time::sleep(d).await;
331 }
332 BuildMode::Ok | BuildMode::Stale => {
333 tokio::task::yield_now().await;
335 }
336 }
337 let read_gen = match mode {
338 BuildMode::Stale => self.state.generation.fetch_add(1, Ordering::SeqCst),
341 _ => self.state.generation.load(Ordering::SeqCst),
342 };
343 if let Some(ev) = &self.events {
344 ev.lock().unwrap().push("end");
345 }
346 Ok((
347 Box::new(FakeStaged {
348 read_generation: read_gen,
349 }),
350 read_gen,
351 ))
352 }
353 fn commit(&self, staged: Box<dyn TemplateReloadStaged>) {
354 let concrete = staged.into_any().downcast::<FakeStaged>().unwrap();
356 assert_eq!(
357 concrete.read_generation,
358 self.state.generation.load(Ordering::SeqCst)
359 );
360 self.state.commit_calls.fetch_add(1, Ordering::SeqCst);
361 self.state.generation.fetch_add(1, Ordering::SeqCst);
362 }
363 }
364
365 #[test]
366 fn registry_register_find_all_remove() {
367 let reg = TemplateReloadRegistry::global();
368 let route = "test-register-find-all-remove";
369 let target = FakeTarget::new(route);
370 let _guard = reg.register(target.as_dyn());
371 assert_eq!(reg.find_all(route).len(), 1);
372 drop(_guard);
373 assert_eq!(reg.find_all(route).len(), 0);
374 }
375
376 #[tokio::test]
377 async fn reload_route_all_or_nothing() {
378 let reg = TemplateReloadRegistry::global();
379 let route = "test-all-or-nothing";
380 let ok = FakeTarget::new(route);
381 let err = FakeTarget::new(route);
382 err.set_mode(BuildMode::Err);
383 let g1 = reg.register(ok.as_dyn());
384 let g2 = reg.register(err.as_dyn());
385
386 let res = reg.reload_route(route).await;
387 assert!(res.is_err(), "expected reload to fail");
388 assert_eq!(
389 ok.state.commit_calls.load(Ordering::SeqCst),
390 0,
391 "OK target must NOT be committed"
392 );
393 assert_eq!(
394 err.state.commit_calls.load(Ordering::SeqCst),
395 0,
396 "Err target must NOT be committed"
397 );
398 assert_eq!(
399 ok.state.generation.load(Ordering::SeqCst),
400 0,
401 "prior generation retained"
402 );
403 drop(g1);
404 drop(g2);
405 }
406
407 #[tokio::test]
408 async fn reload_route_commits_all_on_success() {
409 let reg = TemplateReloadRegistry::global();
410 let route = "test-commits-all-on-success";
411 let a = FakeTarget::new(route);
412 let b = FakeTarget::new(route);
413 let ga = reg.register(a.as_dyn());
414 let gb = reg.register(b.as_dyn());
415
416 let res = reg.reload_route(route).await;
417 assert!(res.is_ok(), "expected reload to succeed: {:?}", res);
418 assert_eq!(a.state.commit_calls.load(Ordering::SeqCst), 1);
419 assert_eq!(b.state.commit_calls.load(Ordering::SeqCst), 1);
420 assert_eq!(a.state.generation.load(Ordering::SeqCst), 1);
421 assert_eq!(b.state.generation.load(Ordering::SeqCst), 1);
422 drop(ga);
423 drop(gb);
424 }
425
426 #[tokio::test]
427 async fn reload_route_timeout_no_commit() {
428 let reg = TemplateReloadRegistry::global();
429 let route = "test-timeout-no-commit";
430 let slow = Arc::new(FakeTarget {
432 route: route.to_string(),
433 timeout: Duration::from_millis(40),
434 state: Arc::new(FakeState::default()),
435 mode: StdMutex::new(BuildMode::Sleep(Duration::from_millis(2_000))),
436 events: None,
437 });
438 let g = reg.register(slow.as_dyn());
439
440 let res = reg.reload_route(route).await;
441 assert!(
442 matches!(res, Err(CamelError::TemplateReload(_))),
443 "expected TemplateReload timeout error, got {res:?}"
444 );
445 assert_eq!(
446 slow.state.commit_calls.load(Ordering::SeqCst),
447 0,
448 "commit must never be called on timeout"
449 );
450 drop(g);
451 }
452
453 #[tokio::test]
454 async fn reload_route_rejects_stale_no_commit() {
455 let reg = TemplateReloadRegistry::global();
456 let route = "test-rejects-stale-no-commit";
457 let target = FakeTarget::new(route);
458 target.set_mode(BuildMode::Stale);
459 let g = reg.register(target.as_dyn());
460
461 let res = reg.reload_route(route).await;
462 assert!(
463 matches!(res, Err(CamelError::TemplateReload(_))),
464 "expected TemplateReload stale error, got {res:?}"
465 );
466 assert_eq!(
467 target.state.commit_calls.load(Ordering::SeqCst),
468 0,
469 "commit must never be called on stale rejection"
470 );
471 drop(g);
472 }
473
474 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
475 async fn reload_route_serializes_concurrent() {
476 let reg = TemplateReloadRegistry::global();
477 let route = "test-serializes-concurrent";
478 let events: Arc<StdMutex<Vec<&'static str>>> = Arc::new(StdMutex::new(Vec::new()));
479 let target = Arc::new(FakeTarget {
480 route: route.to_string(),
481 timeout: Duration::from_secs(5),
482 state: Arc::new(FakeState::default()),
483 mode: StdMutex::new(BuildMode::Ok),
484 events: Some(Arc::clone(&events)),
485 });
486 let g = reg.register(target.as_dyn());
487
488 let h1 = tokio::spawn(async move { reg.reload_route(route).await });
490 let h2 = tokio::spawn(async move { reg.reload_route(route).await });
491 let (r1, r2) = tokio::join!(h1, h2);
492 r1.unwrap().unwrap();
493 r2.unwrap().unwrap();
494
495 let evs = events.lock().unwrap().clone();
498 assert_eq!(
499 evs,
500 vec!["start", "end", "start", "end"],
501 "per-route mutex must serialize concurrent reload_route"
502 );
503 assert_eq!(target.state.commit_calls.load(Ordering::SeqCst), 2);
504 drop(g);
505 }
506}