use std::collections::HashMap;
use std::future::Future;
use std::panic::AssertUnwindSafe;
use std::sync::Arc;
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use futures::{stream::FuturesUnordered, FutureExt, StreamExt};
use serde_json::Value;
#[derive(Debug, Clone, PartialEq)]
pub struct WorkItem {
pub id: String,
pub tenant_id: String,
pub subject_id: String,
pub spec_id: String,
pub config: Value,
pub cancel_requested: bool,
pub lease_version: i64,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum DriverOutcome {
Continue,
Reschedule { at: DateTime<Utc>, config: Value },
Stopped,
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct SupervisorStats {
pub claimed: usize,
pub failed: usize,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SupervisorSettings {
pub worker_id: String,
pub lease_secs: i64,
pub requeue_delay_secs: i64,
pub claim_batch: i64,
pub concurrency: usize,
}
impl SupervisorSettings {
fn concurrency(&self) -> usize {
self.concurrency.max(1)
}
}
#[async_trait]
pub trait WorkQueue: Send + Sync {
async fn claim_due(
&self,
spec_ids: &[String],
worker_id: &str,
lease_secs: i64,
batch: i64,
) -> Result<Vec<WorkItem>, String>;
async fn renew(
&self,
id: &str,
worker_id: &str,
lease_version: i64,
lease_secs: i64,
) -> Result<(), String>;
async fn release(&self, id: &str, lease_version: i64, delay_secs: i64) -> Result<(), String>;
async fn reschedule(
&self,
id: &str,
lease_version: i64,
at: DateTime<Utc>,
config: Value,
) -> Result<(), String>;
async fn complete(&self, id: &str, lease_version: i64) -> Result<(), String>;
async fn mark_error(&self, id: &str, lease_version: i64, error: &str) -> Result<(), String>;
}
#[async_trait]
pub trait WorkflowDriver<Context>: Send + Sync
where
Context: Send + Sync,
{
fn name(&self) -> &'static str;
fn spec_ids(&self) -> &'static [&'static str];
fn validate_specs(&self) -> Result<(), String>;
async fn evaluate(&self, context: &Context, item: &WorkItem) -> Result<DriverOutcome, String>;
}
pub struct DriverRegistry<Context: Send + Sync> {
drivers: Vec<Arc<dyn WorkflowDriver<Context>>>,
by_spec: HashMap<&'static str, Arc<dyn WorkflowDriver<Context>>>,
}
impl<Context: Send + Sync> Default for DriverRegistry<Context> {
fn default() -> Self {
Self {
drivers: Vec::new(),
by_spec: HashMap::new(),
}
}
}
impl<Context: Send + Sync> DriverRegistry<Context> {
pub fn new() -> Self {
Self::default()
}
pub fn register(
&mut self,
driver: Arc<dyn WorkflowDriver<Context>>,
) -> Result<&mut Self, String> {
for spec_id in driver.spec_ids() {
if let Some(existing) = self.by_spec.get(spec_id) {
return Err(format!(
"spec '{spec_id}' is claimed by both '{}' and '{}'",
existing.name(),
driver.name()
));
}
self.by_spec.insert(spec_id, driver.clone());
}
self.drivers.push(driver);
Ok(self)
}
pub fn spec_ids(&self) -> Vec<String> {
let mut ids = self
.by_spec
.keys()
.map(|id| (*id).to_string())
.collect::<Vec<_>>();
ids.sort();
ids
}
pub fn for_spec(&self, spec_id: &str) -> Option<&Arc<dyn WorkflowDriver<Context>>> {
self.by_spec.get(spec_id)
}
pub fn names(&self) -> Vec<&'static str> {
self.drivers.iter().map(|driver| driver.name()).collect()
}
pub fn is_empty(&self) -> bool {
self.drivers.is_empty()
}
pub fn validate_all(&self) -> Result<(), String> {
for driver in &self.drivers {
driver
.validate_specs()
.map_err(|error| format!("driver '{}': {error}", driver.name()))?;
}
Ok(())
}
}
pub async fn run_due_pass<Context, Queue, BuildContext, BuildFuture>(
queue: &Queue,
registry: &DriverRegistry<Context>,
settings: &SupervisorSettings,
build_context: BuildContext,
) -> Result<SupervisorStats, String>
where
Context: Send + Sync + 'static,
Queue: WorkQueue,
BuildContext: FnOnce() -> BuildFuture,
BuildFuture: Future<Output = Result<Context, String>> + Send,
{
let claim_limit = settings
.claim_batch
.clamp(1, i64::try_from(settings.concurrency()).unwrap_or(i64::MAX));
let items = queue
.claim_due(
®istry.spec_ids(),
&settings.worker_id,
settings.lease_secs,
claim_limit,
)
.await
.map_err(|error| format!("claim: {error}"))?;
if items.is_empty() {
return Ok(SupervisorStats::default());
}
let context = match build_context().await {
Ok(context) => Arc::new(context),
Err(error) => {
for item in &items {
let _ = queue.mark_error(&item.id, item.lease_version, &error).await;
let _ = queue
.release(&item.id, item.lease_version, settings.requeue_delay_secs)
.await;
}
return Err(error);
}
};
let claimed = items.len();
let mut failed = 0;
let mut work = items.into_iter();
let mut tasks = FuturesUnordered::new();
let spawn_next = |tasks: &mut FuturesUnordered<_>, work: &mut std::vec::IntoIter<WorkItem>| {
let Some(item) = work.next() else {
return false;
};
let id = item.id.clone();
let lease_version = item.lease_version;
let driver = registry.for_spec(&item.spec_id).cloned();
let context = context.clone();
let renewal_period = std::time::Duration::from_millis(
(settings.lease_secs.clamp(1, 3600) as u64 * 1000 / 3).max(1),
);
tasks.push(async move {
let evaluation = async {
match driver {
Some(driver) => AssertUnwindSafe(driver.evaluate(&context, &item))
.catch_unwind()
.await
.map_err(|_| "driver panicked".to_string())
.and_then(|result| result),
None => Err(format!("no driver registered for spec '{}'", item.spec_id)),
}
};
tokio::pin!(evaluation);
let mut renewal = tokio::time::interval_at(
tokio::time::Instant::now() + renewal_period,
renewal_period,
);
renewal.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
let result = loop {
tokio::select! {
result = &mut evaluation => break result,
_ = renewal.tick() => {
if let Err(error) = queue
.renew(
&id,
&settings.worker_id,
lease_version,
settings.lease_secs,
)
.await
{
break Err(format!("lease renewal failed: {error}"));
}
}
}
};
(id, lease_version, result)
});
true
};
for _ in 0..settings.concurrency() {
if !spawn_next(&mut tasks, &mut work) {
break;
}
}
while let Some((id, lease_version, result)) = tasks.next().await {
match result {
Ok(outcome) => {
let _ = match outcome {
DriverOutcome::Continue => {
queue
.release(&id, lease_version, settings.requeue_delay_secs)
.await
}
DriverOutcome::Reschedule { at, config } => {
queue.reschedule(&id, lease_version, at, config).await
}
DriverOutcome::Stopped => queue.complete(&id, lease_version).await,
};
}
Err(error) => {
failed += 1;
let _ = queue.mark_error(&id, lease_version, &error).await;
let _ = queue
.release(&id, lease_version, settings.requeue_delay_secs)
.await;
}
}
spawn_next(&mut tasks, &mut work);
}
Ok(SupervisorStats { claimed, failed })
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Mutex;
struct Context;
struct Driver {
valid: bool,
}
#[async_trait]
impl WorkflowDriver<Context> for Driver {
fn name(&self) -> &'static str {
"driver"
}
fn spec_ids(&self) -> &'static [&'static str] {
&["spec"]
}
fn validate_specs(&self) -> Result<(), String> {
self.valid.then_some(()).ok_or_else(|| "invalid".into())
}
async fn evaluate(
&self,
_context: &Context,
item: &WorkItem,
) -> Result<DriverOutcome, String> {
if let Some(milliseconds) = item.config["sleep_ms"].as_u64() {
tokio::time::sleep(std::time::Duration::from_millis(milliseconds)).await;
}
if item.config["fail"] == true {
Err("planned".into())
} else if item.config["reschedule"] == true {
Ok(DriverOutcome::Reschedule {
at: Utc::now(),
config: serde_json::json!({"next": true}),
})
} else if item.config["stop"] == true {
Ok(DriverOutcome::Stopped)
} else {
Ok(DriverOutcome::Continue)
}
}
}
#[derive(Default)]
struct Queue {
items: Mutex<Vec<WorkItem>>,
released: Mutex<Vec<String>>,
completed: Mutex<Vec<String>>,
rescheduled: Mutex<Vec<String>>,
errors: Mutex<Vec<String>>,
renewals: Mutex<Vec<(String, String, i64)>>,
renewal_error: Mutex<Option<String>>,
claim_limits: Mutex<Vec<i64>>,
}
#[async_trait]
impl WorkQueue for Queue {
async fn claim_due(
&self,
_spec_ids: &[String],
_worker_id: &str,
_lease_secs: i64,
batch: i64,
) -> Result<Vec<WorkItem>, String> {
self.claim_limits.lock().unwrap().push(batch);
let mut items = self.items.lock().unwrap();
let take = items.len().min(batch as usize);
Ok(items.drain(..take).collect())
}
async fn renew(
&self,
id: &str,
worker_id: &str,
lease_version: i64,
_lease_secs: i64,
) -> Result<(), String> {
self.renewals.lock().unwrap().push((
id.to_string(),
worker_id.to_string(),
lease_version,
));
match self.renewal_error.lock().unwrap().clone() {
Some(error) => Err(error),
None => Ok(()),
}
}
async fn release(
&self,
id: &str,
_lease_version: i64,
_delay_secs: i64,
) -> Result<(), String> {
self.released.lock().unwrap().push(id.into());
Ok(())
}
async fn complete(&self, id: &str, _lease_version: i64) -> Result<(), String> {
self.completed.lock().unwrap().push(id.into());
Ok(())
}
async fn reschedule(
&self,
id: &str,
_lease_version: i64,
_at: DateTime<Utc>,
_config: Value,
) -> Result<(), String> {
self.rescheduled.lock().unwrap().push(id.into());
Ok(())
}
async fn mark_error(
&self,
id: &str,
_lease_version: i64,
_error: &str,
) -> Result<(), String> {
self.errors.lock().unwrap().push(id.into());
Ok(())
}
}
fn item(id: &str, fail: bool) -> WorkItem {
WorkItem {
id: id.into(),
tenant_id: "tenant".into(),
subject_id: "subject".into(),
spec_id: "spec".into(),
config: serde_json::json!({"fail": fail}),
cancel_requested: false,
lease_version: 1,
}
}
fn stopped_item(id: &str) -> WorkItem {
WorkItem {
config: serde_json::json!({"stop": true}),
..item(id, false)
}
}
fn settings() -> SupervisorSettings {
SupervisorSettings {
worker_id: "worker".into(),
lease_secs: 60,
requeue_delay_secs: 5,
claim_batch: 10,
concurrency: 3,
}
}
#[tokio::test]
async fn due_pass_releases_success_and_marks_failures() {
let queue = Queue {
items: Mutex::new(vec![
item("ok", false),
stopped_item("done"),
item("bad", true),
]),
..Default::default()
};
let mut registry = DriverRegistry::new();
registry.register(Arc::new(Driver { valid: true })).unwrap();
let stats = run_due_pass(&queue, ®istry, &settings(), || async { Ok(Context) })
.await
.unwrap();
assert_eq!(
stats,
SupervisorStats {
claimed: 3,
failed: 1
}
);
assert_eq!(queue.released.lock().unwrap().len(), 2);
assert_eq!(queue.completed.lock().unwrap().as_slice(), ["done"]);
assert_eq!(queue.errors.lock().unwrap().as_slice(), ["bad"]);
}
#[tokio::test]
async fn context_failure_releases_every_claim() {
let queue = Queue {
items: Mutex::new(vec![item("one", false), item("two", false)]),
..Default::default()
};
let mut registry = DriverRegistry::new();
registry.register(Arc::new(Driver { valid: true })).unwrap();
let error = run_due_pass(&queue, ®istry, &settings(), || async {
Err::<Context, _>("context failed".into())
})
.await
.unwrap_err();
assert_eq!(error, "context failed");
assert_eq!(queue.released.lock().unwrap().len(), 2);
assert_eq!(queue.errors.lock().unwrap().len(), 2);
}
#[tokio::test]
async fn due_pass_atomically_reschedules_driver_state() {
let queue = Queue {
items: Mutex::new(vec![WorkItem {
config: serde_json::json!({"reschedule": true}),
..item("recurring", false)
}]),
..Default::default()
};
let mut registry = DriverRegistry::new();
registry.register(Arc::new(Driver { valid: true })).unwrap();
run_due_pass(&queue, ®istry, &settings(), || async { Ok(Context) })
.await
.unwrap();
assert_eq!(queue.rescheduled.lock().unwrap().as_slice(), ["recurring"]);
assert!(queue.released.lock().unwrap().is_empty());
assert!(queue.completed.lock().unwrap().is_empty());
}
#[tokio::test]
async fn due_pass_claims_only_work_that_can_start() {
let queue = Queue {
items: Mutex::new(vec![
item("one", false),
item("two", false),
item("queued", false),
]),
..Default::default()
};
let mut registry = DriverRegistry::new();
registry.register(Arc::new(Driver { valid: true })).unwrap();
let mut settings = settings();
settings.concurrency = 2;
let stats = run_due_pass(&queue, ®istry, &settings, || async { Ok(Context) })
.await
.unwrap();
assert_eq!(stats.claimed, 2);
assert_eq!(queue.claim_limits.lock().unwrap().as_slice(), [2]);
assert_eq!(queue.items.lock().unwrap().len(), 1);
}
#[tokio::test(start_paused = true)]
async fn long_evaluate_renews_its_lease() {
let queue = Queue {
items: Mutex::new(vec![WorkItem {
config: serde_json::json!({"sleep_ms": 3500}),
..item("slow", false)
}]),
..Default::default()
};
let mut registry = DriverRegistry::new();
registry.register(Arc::new(Driver { valid: true })).unwrap();
let mut settings = settings();
settings.lease_secs = 3;
settings.concurrency = 1;
let stats = run_due_pass(&queue, ®istry, &settings, || async { Ok(Context) })
.await
.unwrap();
assert_eq!(stats.failed, 0);
let renewals = queue.renewals.lock().unwrap();
assert_eq!(renewals.len(), 3);
assert!(renewals
.iter()
.all(|renewal| renewal == &("slow".into(), "worker".into(), 1)));
assert_eq!(queue.released.lock().unwrap().as_slice(), ["slow"]);
}
#[tokio::test(start_paused = true)]
async fn expired_lease_stops_evaluation() {
let queue = Queue {
items: Mutex::new(vec![WorkItem {
config: serde_json::json!({"sleep_ms": 5000}),
..item("expired", false)
}]),
renewal_error: Mutex::new(Some("lease expired".into())),
..Default::default()
};
let mut registry = DriverRegistry::new();
registry.register(Arc::new(Driver { valid: true })).unwrap();
let mut settings = settings();
settings.lease_secs = 3;
settings.concurrency = 1;
let stats = run_due_pass(&queue, ®istry, &settings, || async { Ok(Context) })
.await
.unwrap();
assert_eq!(stats.failed, 1);
assert_eq!(queue.renewals.lock().unwrap().len(), 1);
assert_eq!(queue.errors.lock().unwrap().as_slice(), ["expired"]);
}
#[test]
fn registry_fails_duplicate_ownership_and_invalid_specs() {
let mut registry = DriverRegistry::new();
assert!(registry.is_empty());
registry
.register(Arc::new(Driver { valid: false }))
.unwrap();
assert_eq!(registry.spec_ids(), ["spec"]);
assert_eq!(registry.names(), ["driver"]);
assert!(registry.for_spec("spec").is_some());
assert!(registry.validate_all().unwrap_err().contains("invalid"));
assert!(registry
.register(Arc::new(Driver { valid: true }))
.err()
.unwrap()
.contains("both"));
}
}