Skip to main content

af_workflow/
supervisor.rs

1//! Database-agnostic scheduler for claimed long-running workflow instances.
2
3use 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    /// Register a driver. Duplicate spec ownership is a startup error because
112    /// dispatch must never depend on registration order.
113    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            &registry.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, &registry, &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, &registry, &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, &registry, &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, &registry, &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, &registry, &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, &registry, &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}