1use std::collections::HashMap;
4use std::future::Future;
5use std::panic::AssertUnwindSafe;
6use std::sync::Arc;
7
8use async_trait::async_trait;
9use chrono::{DateTime, Utc};
10use futures::{stream::FuturesUnordered, FutureExt, StreamExt};
11use serde_json::Value;
12
13#[derive(Debug, Clone, PartialEq)]
14pub struct WorkItem {
15 pub id: String,
16 pub tenant_id: String,
17 pub subject_id: String,
18 pub spec_id: String,
19 pub config: Value,
20 pub cancel_requested: bool,
21 pub lease_version: i64,
22}
23
24#[derive(Debug, Clone, PartialEq, Eq)]
25pub enum DriverOutcome {
26 Continue,
27 Reschedule { at: DateTime<Utc>, config: Value },
28 Stopped,
29}
30
31#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
32pub struct SupervisorStats {
33 pub claimed: usize,
34 pub failed: usize,
35}
36
37#[derive(Debug, Clone, PartialEq, Eq)]
38pub struct SupervisorSettings {
39 pub worker_id: String,
40 pub lease_secs: i64,
41 pub requeue_delay_secs: i64,
42 pub claim_batch: i64,
43 pub concurrency: usize,
44}
45
46impl SupervisorSettings {
47 fn concurrency(&self) -> usize {
48 self.concurrency.max(1)
49 }
50}
51
52#[async_trait]
53pub trait WorkQueue: Send + Sync {
54 async fn claim_due(
55 &self,
56 spec_ids: &[String],
57 worker_id: &str,
58 lease_secs: i64,
59 batch: i64,
60 ) -> Result<Vec<WorkItem>, String>;
61
62 async fn renew(
63 &self,
64 id: &str,
65 worker_id: &str,
66 lease_version: i64,
67 lease_secs: i64,
68 ) -> Result<(), String>;
69 async fn release(&self, id: &str, lease_version: i64, delay_secs: i64) -> Result<(), String>;
70 async fn reschedule(
71 &self,
72 id: &str,
73 lease_version: i64,
74 at: DateTime<Utc>,
75 config: Value,
76 ) -> Result<(), String>;
77 async fn complete(&self, id: &str, lease_version: i64) -> Result<(), String>;
78 async fn mark_error(&self, id: &str, lease_version: i64, error: &str) -> Result<(), String>;
79}
80
81#[async_trait]
82pub trait WorkflowDriver<Context>: Send + Sync
83where
84 Context: Send + Sync,
85{
86 fn name(&self) -> &'static str;
87 fn spec_ids(&self) -> &'static [&'static str];
88 fn validate_specs(&self) -> Result<(), String>;
89 async fn evaluate(&self, context: &Context, item: &WorkItem) -> Result<DriverOutcome, String>;
90}
91
92pub struct DriverRegistry<Context: Send + Sync> {
93 drivers: Vec<Arc<dyn WorkflowDriver<Context>>>,
94 by_spec: HashMap<&'static str, Arc<dyn WorkflowDriver<Context>>>,
95}
96
97impl<Context: Send + Sync> Default for DriverRegistry<Context> {
98 fn default() -> Self {
99 Self {
100 drivers: Vec::new(),
101 by_spec: HashMap::new(),
102 }
103 }
104}
105
106impl<Context: Send + Sync> DriverRegistry<Context> {
107 pub fn new() -> Self {
108 Self::default()
109 }
110
111 pub fn register(
114 &mut self,
115 driver: Arc<dyn WorkflowDriver<Context>>,
116 ) -> Result<&mut Self, String> {
117 for spec_id in driver.spec_ids() {
118 if let Some(existing) = self.by_spec.get(spec_id) {
119 return Err(format!(
120 "spec '{spec_id}' is claimed by both '{}' and '{}'",
121 existing.name(),
122 driver.name()
123 ));
124 }
125 self.by_spec.insert(spec_id, driver.clone());
126 }
127 self.drivers.push(driver);
128 Ok(self)
129 }
130
131 pub fn spec_ids(&self) -> Vec<String> {
132 let mut ids = self
133 .by_spec
134 .keys()
135 .map(|id| (*id).to_string())
136 .collect::<Vec<_>>();
137 ids.sort();
138 ids
139 }
140
141 pub fn for_spec(&self, spec_id: &str) -> Option<&Arc<dyn WorkflowDriver<Context>>> {
142 self.by_spec.get(spec_id)
143 }
144
145 pub fn names(&self) -> Vec<&'static str> {
146 self.drivers.iter().map(|driver| driver.name()).collect()
147 }
148
149 pub fn is_empty(&self) -> bool {
150 self.drivers.is_empty()
151 }
152
153 pub fn validate_all(&self) -> Result<(), String> {
154 for driver in &self.drivers {
155 driver
156 .validate_specs()
157 .map_err(|error| format!("driver '{}': {error}", driver.name()))?;
158 }
159 Ok(())
160 }
161}
162
163pub async fn run_due_pass<Context, Queue, BuildContext, BuildFuture>(
164 queue: &Queue,
165 registry: &DriverRegistry<Context>,
166 settings: &SupervisorSettings,
167 build_context: BuildContext,
168) -> Result<SupervisorStats, String>
169where
170 Context: Send + Sync + 'static,
171 Queue: WorkQueue,
172 BuildContext: FnOnce() -> BuildFuture,
173 BuildFuture: Future<Output = Result<Context, String>> + Send,
174{
175 let claim_limit = settings
176 .claim_batch
177 .clamp(1, i64::try_from(settings.concurrency()).unwrap_or(i64::MAX));
178 let items = queue
179 .claim_due(
180 ®istry.spec_ids(),
181 &settings.worker_id,
182 settings.lease_secs,
183 claim_limit,
184 )
185 .await
186 .map_err(|error| format!("claim: {error}"))?;
187 if items.is_empty() {
188 return Ok(SupervisorStats::default());
189 }
190
191 let context = match build_context().await {
192 Ok(context) => Arc::new(context),
193 Err(error) => {
194 for item in &items {
195 let _ = queue.mark_error(&item.id, item.lease_version, &error).await;
196 let _ = queue
197 .release(&item.id, item.lease_version, settings.requeue_delay_secs)
198 .await;
199 }
200 return Err(error);
201 }
202 };
203
204 let claimed = items.len();
205 let mut failed = 0;
206 let mut work = items.into_iter();
207 let mut tasks = FuturesUnordered::new();
208
209 let spawn_next = |tasks: &mut FuturesUnordered<_>, work: &mut std::vec::IntoIter<WorkItem>| {
210 let Some(item) = work.next() else {
211 return false;
212 };
213 let id = item.id.clone();
214 let lease_version = item.lease_version;
215 let driver = registry.for_spec(&item.spec_id).cloned();
216 let context = context.clone();
217 let renewal_period = std::time::Duration::from_millis(
218 (settings.lease_secs.clamp(1, 3600) as u64 * 1000 / 3).max(1),
219 );
220 tasks.push(async move {
221 let evaluation = async {
222 match driver {
223 Some(driver) => AssertUnwindSafe(driver.evaluate(&context, &item))
224 .catch_unwind()
225 .await
226 .map_err(|_| "driver panicked".to_string())
227 .and_then(|result| result),
228 None => Err(format!("no driver registered for spec '{}'", item.spec_id)),
229 }
230 };
231 tokio::pin!(evaluation);
232 let mut renewal = tokio::time::interval_at(
233 tokio::time::Instant::now() + renewal_period,
234 renewal_period,
235 );
236 renewal.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
237 let result = loop {
238 tokio::select! {
239 result = &mut evaluation => break result,
240 _ = renewal.tick() => {
241 if let Err(error) = queue
242 .renew(
243 &id,
244 &settings.worker_id,
245 lease_version,
246 settings.lease_secs,
247 )
248 .await
249 {
250 break Err(format!("lease renewal failed: {error}"));
251 }
252 }
253 }
254 };
255 (id, lease_version, result)
256 });
257 true
258 };
259
260 for _ in 0..settings.concurrency() {
261 if !spawn_next(&mut tasks, &mut work) {
262 break;
263 }
264 }
265
266 while let Some((id, lease_version, result)) = tasks.next().await {
267 match result {
268 Ok(outcome) => {
269 let _ = match outcome {
270 DriverOutcome::Continue => {
271 queue
272 .release(&id, lease_version, settings.requeue_delay_secs)
273 .await
274 }
275 DriverOutcome::Reschedule { at, config } => {
276 queue.reschedule(&id, lease_version, at, config).await
277 }
278 DriverOutcome::Stopped => queue.complete(&id, lease_version).await,
279 };
280 }
281 Err(error) => {
282 failed += 1;
283 let _ = queue.mark_error(&id, lease_version, &error).await;
284 let _ = queue
285 .release(&id, lease_version, settings.requeue_delay_secs)
286 .await;
287 }
288 }
289 spawn_next(&mut tasks, &mut work);
290 }
291
292 Ok(SupervisorStats { claimed, failed })
293}
294
295#[cfg(test)]
296mod tests {
297 use super::*;
298 use std::sync::Mutex;
299
300 struct Context;
301 struct Driver {
302 valid: bool,
303 }
304
305 #[async_trait]
306 impl WorkflowDriver<Context> for Driver {
307 fn name(&self) -> &'static str {
308 "driver"
309 }
310 fn spec_ids(&self) -> &'static [&'static str] {
311 &["spec"]
312 }
313 fn validate_specs(&self) -> Result<(), String> {
314 self.valid.then_some(()).ok_or_else(|| "invalid".into())
315 }
316 async fn evaluate(
317 &self,
318 _context: &Context,
319 item: &WorkItem,
320 ) -> Result<DriverOutcome, String> {
321 if let Some(milliseconds) = item.config["sleep_ms"].as_u64() {
322 tokio::time::sleep(std::time::Duration::from_millis(milliseconds)).await;
323 }
324 if item.config["fail"] == true {
325 Err("planned".into())
326 } else if item.config["reschedule"] == true {
327 Ok(DriverOutcome::Reschedule {
328 at: Utc::now(),
329 config: serde_json::json!({"next": true}),
330 })
331 } else if item.config["stop"] == true {
332 Ok(DriverOutcome::Stopped)
333 } else {
334 Ok(DriverOutcome::Continue)
335 }
336 }
337 }
338
339 #[derive(Default)]
340 struct Queue {
341 items: Mutex<Vec<WorkItem>>,
342 released: Mutex<Vec<String>>,
343 completed: Mutex<Vec<String>>,
344 rescheduled: Mutex<Vec<String>>,
345 errors: Mutex<Vec<String>>,
346 renewals: Mutex<Vec<(String, String, i64)>>,
347 renewal_error: Mutex<Option<String>>,
348 claim_limits: Mutex<Vec<i64>>,
349 }
350
351 #[async_trait]
352 impl WorkQueue for Queue {
353 async fn claim_due(
354 &self,
355 _spec_ids: &[String],
356 _worker_id: &str,
357 _lease_secs: i64,
358 batch: i64,
359 ) -> Result<Vec<WorkItem>, String> {
360 self.claim_limits.lock().unwrap().push(batch);
361 let mut items = self.items.lock().unwrap();
362 let take = items.len().min(batch as usize);
363 Ok(items.drain(..take).collect())
364 }
365
366 async fn renew(
367 &self,
368 id: &str,
369 worker_id: &str,
370 lease_version: i64,
371 _lease_secs: i64,
372 ) -> Result<(), String> {
373 self.renewals.lock().unwrap().push((
374 id.to_string(),
375 worker_id.to_string(),
376 lease_version,
377 ));
378 match self.renewal_error.lock().unwrap().clone() {
379 Some(error) => Err(error),
380 None => Ok(()),
381 }
382 }
383
384 async fn release(
385 &self,
386 id: &str,
387 _lease_version: i64,
388 _delay_secs: i64,
389 ) -> Result<(), String> {
390 self.released.lock().unwrap().push(id.into());
391 Ok(())
392 }
393
394 async fn complete(&self, id: &str, _lease_version: i64) -> Result<(), String> {
395 self.completed.lock().unwrap().push(id.into());
396 Ok(())
397 }
398
399 async fn reschedule(
400 &self,
401 id: &str,
402 _lease_version: i64,
403 _at: DateTime<Utc>,
404 _config: Value,
405 ) -> Result<(), String> {
406 self.rescheduled.lock().unwrap().push(id.into());
407 Ok(())
408 }
409
410 async fn mark_error(
411 &self,
412 id: &str,
413 _lease_version: i64,
414 _error: &str,
415 ) -> Result<(), String> {
416 self.errors.lock().unwrap().push(id.into());
417 Ok(())
418 }
419 }
420
421 fn item(id: &str, fail: bool) -> WorkItem {
422 WorkItem {
423 id: id.into(),
424 tenant_id: "tenant".into(),
425 subject_id: "subject".into(),
426 spec_id: "spec".into(),
427 config: serde_json::json!({"fail": fail}),
428 cancel_requested: false,
429 lease_version: 1,
430 }
431 }
432
433 fn stopped_item(id: &str) -> WorkItem {
434 WorkItem {
435 config: serde_json::json!({"stop": true}),
436 ..item(id, false)
437 }
438 }
439
440 fn settings() -> SupervisorSettings {
441 SupervisorSettings {
442 worker_id: "worker".into(),
443 lease_secs: 60,
444 requeue_delay_secs: 5,
445 claim_batch: 10,
446 concurrency: 3,
447 }
448 }
449
450 #[tokio::test]
451 async fn due_pass_releases_success_and_marks_failures() {
452 let queue = Queue {
453 items: Mutex::new(vec![
454 item("ok", false),
455 stopped_item("done"),
456 item("bad", true),
457 ]),
458 ..Default::default()
459 };
460 let mut registry = DriverRegistry::new();
461 registry.register(Arc::new(Driver { valid: true })).unwrap();
462 let stats = run_due_pass(&queue, ®istry, &settings(), || async { Ok(Context) })
463 .await
464 .unwrap();
465 assert_eq!(
466 stats,
467 SupervisorStats {
468 claimed: 3,
469 failed: 1
470 }
471 );
472 assert_eq!(queue.released.lock().unwrap().len(), 2);
473 assert_eq!(queue.completed.lock().unwrap().as_slice(), ["done"]);
474 assert_eq!(queue.errors.lock().unwrap().as_slice(), ["bad"]);
475 }
476
477 #[tokio::test]
478 async fn context_failure_releases_every_claim() {
479 let queue = Queue {
480 items: Mutex::new(vec![item("one", false), item("two", false)]),
481 ..Default::default()
482 };
483 let mut registry = DriverRegistry::new();
484 registry.register(Arc::new(Driver { valid: true })).unwrap();
485 let error = run_due_pass(&queue, ®istry, &settings(), || async {
486 Err::<Context, _>("context failed".into())
487 })
488 .await
489 .unwrap_err();
490 assert_eq!(error, "context failed");
491 assert_eq!(queue.released.lock().unwrap().len(), 2);
492 assert_eq!(queue.errors.lock().unwrap().len(), 2);
493 }
494
495 #[tokio::test]
496 async fn due_pass_atomically_reschedules_driver_state() {
497 let queue = Queue {
498 items: Mutex::new(vec![WorkItem {
499 config: serde_json::json!({"reschedule": true}),
500 ..item("recurring", false)
501 }]),
502 ..Default::default()
503 };
504 let mut registry = DriverRegistry::new();
505 registry.register(Arc::new(Driver { valid: true })).unwrap();
506 run_due_pass(&queue, ®istry, &settings(), || async { Ok(Context) })
507 .await
508 .unwrap();
509 assert_eq!(queue.rescheduled.lock().unwrap().as_slice(), ["recurring"]);
510 assert!(queue.released.lock().unwrap().is_empty());
511 assert!(queue.completed.lock().unwrap().is_empty());
512 }
513
514 #[tokio::test]
515 async fn due_pass_claims_only_work_that_can_start() {
516 let queue = Queue {
517 items: Mutex::new(vec![
518 item("one", false),
519 item("two", false),
520 item("queued", false),
521 ]),
522 ..Default::default()
523 };
524 let mut registry = DriverRegistry::new();
525 registry.register(Arc::new(Driver { valid: true })).unwrap();
526 let mut settings = settings();
527 settings.concurrency = 2;
528
529 let stats = run_due_pass(&queue, ®istry, &settings, || async { Ok(Context) })
530 .await
531 .unwrap();
532
533 assert_eq!(stats.claimed, 2);
534 assert_eq!(queue.claim_limits.lock().unwrap().as_slice(), [2]);
535 assert_eq!(queue.items.lock().unwrap().len(), 1);
536 }
537
538 #[tokio::test(start_paused = true)]
539 async fn long_evaluate_renews_its_lease() {
540 let queue = Queue {
541 items: Mutex::new(vec![WorkItem {
542 config: serde_json::json!({"sleep_ms": 3500}),
543 ..item("slow", false)
544 }]),
545 ..Default::default()
546 };
547 let mut registry = DriverRegistry::new();
548 registry.register(Arc::new(Driver { valid: true })).unwrap();
549 let mut settings = settings();
550 settings.lease_secs = 3;
551 settings.concurrency = 1;
552
553 let stats = run_due_pass(&queue, ®istry, &settings, || async { Ok(Context) })
554 .await
555 .unwrap();
556
557 assert_eq!(stats.failed, 0);
558 let renewals = queue.renewals.lock().unwrap();
559 assert_eq!(renewals.len(), 3);
560 assert!(renewals
561 .iter()
562 .all(|renewal| renewal == &("slow".into(), "worker".into(), 1)));
563 assert_eq!(queue.released.lock().unwrap().as_slice(), ["slow"]);
564 }
565
566 #[tokio::test(start_paused = true)]
567 async fn expired_lease_stops_evaluation() {
568 let queue = Queue {
569 items: Mutex::new(vec![WorkItem {
570 config: serde_json::json!({"sleep_ms": 5000}),
571 ..item("expired", false)
572 }]),
573 renewal_error: Mutex::new(Some("lease expired".into())),
574 ..Default::default()
575 };
576 let mut registry = DriverRegistry::new();
577 registry.register(Arc::new(Driver { valid: true })).unwrap();
578 let mut settings = settings();
579 settings.lease_secs = 3;
580 settings.concurrency = 1;
581
582 let stats = run_due_pass(&queue, ®istry, &settings, || async { Ok(Context) })
583 .await
584 .unwrap();
585
586 assert_eq!(stats.failed, 1);
587 assert_eq!(queue.renewals.lock().unwrap().len(), 1);
588 assert_eq!(queue.errors.lock().unwrap().as_slice(), ["expired"]);
589 }
590
591 #[test]
592 fn registry_fails_duplicate_ownership_and_invalid_specs() {
593 let mut registry = DriverRegistry::new();
594 assert!(registry.is_empty());
595 registry
596 .register(Arc::new(Driver { valid: false }))
597 .unwrap();
598 assert_eq!(registry.spec_ids(), ["spec"]);
599 assert_eq!(registry.names(), ["driver"]);
600 assert!(registry.for_spec("spec").is_some());
601 assert!(registry.validate_all().unwrap_err().contains("invalid"));
602 assert!(registry
603 .register(Arc::new(Driver { valid: true }))
604 .err()
605 .unwrap()
606 .contains("both"));
607 }
608}