[model-gateway] Fix workflow engine race conditions and add graceful shutdown (#16963)

This commit is contained in:
Simo Lin
2026-01-12 06:37:39 -08:00
committed by GitHub
parent 76f69b7753
commit d0092decb1
3 changed files with 340 additions and 30 deletions

View File

@@ -6,14 +6,20 @@
use std::{
collections::{HashMap, HashSet, VecDeque},
sync::Arc,
sync::{
atomic::{AtomicUsize, Ordering},
Arc,
},
time::Duration,
};
use backoff::{backoff::Backoff, ExponentialBackoffBuilder};
use chrono::Utc;
use parking_lot::RwLock;
use tokio::{sync::mpsc, time::timeout};
use tokio::{
sync::{mpsc, watch},
time::timeout,
};
use super::{
definition::{StepDefinition, WorkflowDefinition},
@@ -94,18 +100,143 @@ impl Backoff for LinearBackoff {
}
/// Main workflow execution engine
///
/// # Graceful Shutdown
///
/// The engine supports graceful shutdown via [`shutdown()`](Self::shutdown):
///
/// ```ignore
/// // Trigger shutdown - stops accepting new workflows
/// engine.shutdown();
///
/// // Wait for all running workflows to complete (with timeout)
/// if !engine.wait_for_shutdown(Duration::from_secs(30)).await {
/// // Force cancel remaining workflows
/// engine.force_cancel_all().await;
/// }
/// ```
pub struct WorkflowEngine {
definitions: Arc<RwLock<HashMap<WorkflowId, Arc<WorkflowDefinition>>>>,
state_store: WorkflowStateStore,
event_bus: Arc<EventBus>,
/// Shutdown signal sender - when true, engine is shutting down
shutdown_tx: Arc<watch::Sender<bool>>,
/// Shutdown signal receiver for cloning to tasks
shutdown_rx: watch::Receiver<bool>,
/// Count of active workflow executions
active_workflows: Arc<AtomicUsize>,
}
impl WorkflowEngine {
pub fn new() -> Self {
let (shutdown_tx, shutdown_rx) = watch::channel(false);
Self {
definitions: Arc::new(RwLock::new(HashMap::new())),
state_store: WorkflowStateStore::new(),
event_bus: Arc::new(EventBus::new()),
shutdown_tx: Arc::new(shutdown_tx),
shutdown_rx,
active_workflows: Arc::new(AtomicUsize::new(0)),
}
}
/// Check if the engine is shutting down
pub fn is_shutting_down(&self) -> bool {
*self.shutdown_rx.borrow()
}
/// Initiate graceful shutdown
///
/// This will:
/// - Stop accepting new workflows (start_workflow will return an error)
/// - Stop the cleanup task
/// - Allow running workflows to complete
///
/// Use [`wait_for_shutdown`](Self::wait_for_shutdown) to wait for completion.
pub fn shutdown(&self) {
let _ = self.shutdown_tx.send(true);
tracing::info!("Workflow engine shutdown initiated");
}
/// Wait for all active workflows to complete
///
/// Returns `true` if all workflows completed within the timeout,
/// `false` if the timeout was reached with workflows still running.
///
/// Uses simple polling - appropriate for shutdown which happens once per process.
pub async fn wait_for_shutdown(&self, timeout_duration: Duration) -> bool {
let start = tokio::time::Instant::now();
loop {
let active = self.active_workflows.load(Ordering::Acquire);
if active == 0 {
tracing::info!("All workflows completed, shutdown complete");
return true;
}
if start.elapsed() >= timeout_duration {
tracing::warn!(
remaining_workflows = active,
"Shutdown timeout reached with workflows still running"
);
return false;
}
tracing::debug!(
active_workflows = active,
"Waiting for workflows to complete"
);
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
/// Force cancel all running workflows
///
/// This should be called after `wait_for_shutdown` times out if you need
/// to ensure all workflows are stopped. Note that this cancels workflows
/// at the state level; running steps may still complete.
pub async fn force_cancel_all(&self) -> usize {
let active_states = match self.state_store.list_active() {
Ok(states) => states,
Err(e) => {
tracing::error!(error = ?e, "Failed to list active workflows for force cancel");
return 0;
}
};
let mut cancelled = 0;
for state in active_states {
if let Err(e) = self.cancel_workflow(state.instance_id).await {
tracing::warn!(
instance_id = %state.instance_id,
error = ?e,
"Failed to cancel workflow during force shutdown"
);
} else {
cancelled += 1;
}
}
tracing::info!(cancelled_count = cancelled, "Force cancelled workflows");
cancelled
}
/// Get the number of currently active workflow executions
pub fn active_workflow_count(&self) -> usize {
self.active_workflows.load(Ordering::Acquire)
}
/// Decrement active workflow count
fn workflow_finished(&self) {
self.active_workflows.fetch_sub(1, Ordering::Release);
}
/// Create a guard that decrements active_workflows on drop.
/// This ensures the count is decremented even if a task panics.
fn active_workflow_guard(&self) -> ActiveWorkflowGuard {
ActiveWorkflowGuard {
active_workflows: Arc::clone(&self.active_workflows),
}
}
@@ -114,6 +245,8 @@ impl WorkflowEngine {
/// This prevents unbounded memory growth by removing completed/failed workflows
/// that are older than the specified TTL.
///
/// The task will automatically stop when [`shutdown()`](Self::shutdown) is called.
///
/// # Arguments
///
/// * `ttl` - Time-to-live for terminal workflows (default: 1 hour)
@@ -130,14 +263,22 @@ impl WorkflowEngine {
let state_store = self.state_store.clone();
let ttl = ttl.unwrap_or(Duration::from_secs(3600)); // 1 hour default
let interval = interval.unwrap_or(Duration::from_secs(300)); // 5 minutes default
let mut shutdown_rx = self.shutdown_rx.clone();
tokio::spawn(async move {
let mut ticker = tokio::time::interval(interval);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
ticker.tick().await;
state_store.cleanup_old_workflows(ttl);
tokio::select! {
_ = ticker.tick() => {
state_store.cleanup_old_workflows(ttl);
}
_ = shutdown_rx.changed() => {
tracing::info!("Cleanup task stopping due to shutdown");
break;
}
}
}
})
}
@@ -163,37 +304,41 @@ impl WorkflowEngine {
}
/// Start a new workflow instance
///
/// Returns `Err(WorkflowError::ShuttingDown)` if the engine is shutting down.
pub async fn start_workflow(
&self,
definition_id: WorkflowId,
context: WorkflowContext,
) -> WorkflowResult<WorkflowInstanceId> {
// Get workflow definition
let definition = {
let definitions = self.definitions.read();
definitions
.get(&definition_id)
.cloned()
.ok_or_else(|| WorkflowError::DefinitionNotFound(definition_id.clone()))?
};
// Guard increments counter and decrements on drop unless committed.
// This handles all error paths automatically.
let guard = StartGuard::new(self);
if self.is_shutting_down() {
return Err(WorkflowError::ShuttingDown);
}
let definition = self
.definitions
.read()
.get(&definition_id)
.cloned()
.ok_or_else(|| WorkflowError::DefinitionNotFound(definition_id.clone()))?;
// Create new workflow instance
let instance_id = context.instance_id;
let mut state = WorkflowState::new(instance_id, definition_id.clone());
state.status = WorkflowStatus::Running;
state.context = context;
// Initialize step states
for step in &definition.steps {
state
.step_states
.insert(step.id.clone(), StepState::default());
}
// Save initial state
self.state_store.save(state)?;
// Emit workflow started event
self.event_bus
.publish(WorkflowEvent::WorkflowStarted {
instance_id,
@@ -201,11 +346,15 @@ impl WorkflowEngine {
})
.await;
// Execute workflow in background
// Commit the guard - from here the spawned task takes ownership of the count
guard.commit();
let engine = self.clone_for_execution();
let def = Arc::clone(&definition);
tokio::spawn(async move {
if let Err(e) = engine.execute_workflow(instance_id, def).await {
let _guard = engine.active_workflow_guard();
let result = engine.execute_workflow(instance_id, def).await;
if let Err(e) = result {
tracing::error!(instance_id = %instance_id, error = ?e, "Workflow execution failed");
}
});
@@ -290,9 +439,11 @@ impl WorkflowEngine {
}
// Launch ready steps in parallel
let mut tasks_launched = 0;
for step_idx in ready_step_indices {
let step = &definition.steps[step_idx];
tracker.write().running.insert(step.id.clone());
tasks_launched += 1;
let engine = self.clone_for_execution();
let def = Arc::clone(&definition);
@@ -306,6 +457,11 @@ impl WorkflowEngine {
.execute_step_with_retry(instance_id, step, &def)
.await;
let signal = match &result {
Ok(r) => *r,
Err(_) => StepResult::Failure,
};
{
let mut t = tracker.write();
t.running.remove(&step_id);
@@ -337,18 +493,30 @@ impl WorkflowEngine {
}
},
}
}
let signal = match result {
Ok(r) => r,
Err(_) => StepResult::Failure,
};
let _ = tx.send((step_id, signal)).await;
if let Err(e) = tx.try_send((step_id.clone(), signal)) {
use mpsc::error::TrySendError;
match e {
TrySendError::Full(_) => {
tracing::error!(
step_id = %step_id,
"Channel full when sending step completion - this is a bug"
);
}
TrySendError::Closed(_) => {
tracing::debug!(
step_id = %step_id,
"Channel closed, workflow likely cancelled"
);
}
}
}
}
});
}
// Wait for at least one step to complete (if any running)
if !tracker.read().running.is_empty() {
let should_wait = tasks_launched > 0 || !tracker.read().running.is_empty();
if should_wait {
if let Some((completed_step_id, result)) = rx.recv().await {
tracing::debug!(
step_id = %completed_step_id,
@@ -605,10 +773,60 @@ impl WorkflowEngine {
definitions: Arc::clone(&self.definitions),
state_store: self.state_store.clone(),
event_bus: Arc::clone(&self.event_bus),
shutdown_tx: Arc::clone(&self.shutdown_tx),
shutdown_rx: self.shutdown_rx.clone(),
active_workflows: Arc::clone(&self.active_workflows),
}
}
}
/// RAII guard that decrements active_workflows count on drop.
/// Ensures proper cleanup even if a workflow task panics.
struct ActiveWorkflowGuard {
active_workflows: Arc<AtomicUsize>,
}
impl Drop for ActiveWorkflowGuard {
fn drop(&mut self) {
self.active_workflows.fetch_sub(1, Ordering::Release);
}
}
/// RAII guard for start_workflow that increments on creation and decrements on drop
/// unless commit() is called. Handles all error paths automatically.
struct StartGuard<'a> {
engine: &'a WorkflowEngine,
committed: bool,
}
impl<'a> StartGuard<'a> {
fn new(engine: &'a WorkflowEngine) -> Self {
engine.active_workflows.fetch_add(1, Ordering::AcqRel);
Self {
engine,
committed: false,
}
}
fn commit(mut self) {
self.committed = true;
}
}
impl Drop for StartGuard<'_> {
fn drop(&mut self) {
if !self.committed {
self.engine.workflow_finished();
}
}
}
impl Clone for WorkflowEngine {
fn clone(&self) -> Self {
self.clone_for_execution()
}
}
impl Default for WorkflowEngine {
fn default() -> Self {
Self::new()

View File

@@ -8,6 +8,9 @@ use tracing::{error, info, warn};
use super::types::{StepId, WorkflowId, WorkflowInstanceId};
/// Default timeout for subscriber event handlers
const DEFAULT_SUBSCRIBER_TIMEOUT: Duration = Duration::from_secs(30);
/// Events emitted by the workflow engine
#[derive(Debug, Clone)]
pub enum WorkflowEvent {
@@ -58,14 +61,33 @@ pub trait EventSubscriber: Send + Sync {
}
/// Event bus for publishing and subscribing to workflow events
///
/// # Subscriber Isolation
///
/// Each subscriber is notified in a separate spawned task with a timeout.
/// This ensures that:
/// - A slow subscriber doesn't block other subscribers
/// - A panicking subscriber doesn't affect other subscribers
/// - Event publishing returns quickly regardless of subscriber behavior
pub struct EventBus {
subscribers: Arc<RwLock<Vec<Arc<dyn EventSubscriber>>>>,
/// Timeout for each subscriber's event handler
subscriber_timeout: Duration,
}
impl EventBus {
pub fn new() -> Self {
Self {
subscribers: Arc::new(RwLock::new(Vec::new())),
subscriber_timeout: DEFAULT_SUBSCRIBER_TIMEOUT,
}
}
/// Create an EventBus with a custom subscriber timeout
pub fn with_timeout(timeout: Duration) -> Self {
Self {
subscribers: Arc::new(RwLock::new(Vec::new())),
subscriber_timeout: timeout,
}
}
@@ -74,11 +96,78 @@ impl EventBus {
self.subscribers.write().await.push(subscriber);
}
/// Unsubscribe from workflow events
///
/// Removes the subscriber by Arc pointer equality.
/// Returns true if the subscriber was found and removed.
pub async fn unsubscribe(&self, subscriber: &Arc<dyn EventSubscriber>) -> bool {
let mut subs = self.subscribers.write().await;
let len_before = subs.len();
subs.retain(|s| !Arc::ptr_eq(s, subscriber));
subs.len() < len_before
}
/// Publish an event to all subscribers concurrently
///
/// Each subscriber is notified in a separate spawned task with a timeout.
/// This method returns after spawning all notification tasks, without
/// waiting for subscribers to complete (fire-and-forget).
///
/// Subscriber failures (timeout or panic) are logged but don't affect
/// other subscribers or the caller.
pub async fn publish(&self, event: WorkflowEvent) {
let subscribers = self.subscribers.read().await;
let futures: Vec<_> = subscribers.iter().map(|s| s.on_event(&event)).collect();
futures::future::join_all(futures).await;
let subscribers: Vec<_> = self.subscribers.read().await.iter().cloned().collect();
let timeout = self.subscriber_timeout;
for (idx, subscriber) in subscribers.into_iter().enumerate() {
let event = event.clone();
tokio::spawn(async move {
let result = tokio::time::timeout(timeout, subscriber.on_event(&event)).await;
match result {
Ok(()) => {}
Err(_) => {
warn!(
subscriber_index = idx,
timeout_secs = timeout.as_secs(),
"Event subscriber timed out"
);
}
}
});
}
}
/// Publish an event and wait for all subscribers to complete
///
/// Unlike `publish`, this method waits for all subscribers to finish
/// (or timeout). Use this when you need to ensure all subscribers
/// have processed the event before continuing.
pub async fn publish_and_wait(&self, event: WorkflowEvent) {
let subscribers: Vec<_> = self.subscribers.read().await.iter().cloned().collect();
let timeout = self.subscriber_timeout;
let handles: Vec<_> = subscribers
.into_iter()
.enumerate()
.map(|(idx, subscriber)| {
let event = event.clone();
tokio::spawn(async move {
let result = tokio::time::timeout(timeout, subscriber.on_event(&event)).await;
if result.is_err() {
warn!(
subscriber_index = idx,
timeout_secs = timeout.as_secs(),
"Event subscriber timed out"
);
}
})
})
.collect();
// Wait for all spawned tasks, ignoring individual failures (panics)
for handle in handles {
let _ = handle.await;
}
}
}

View File

@@ -244,7 +244,7 @@ pub enum StepResult {
}
/// Error kinds for workflow operations
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[derive(Debug, Clone, thiserror::Error)]
pub enum WorkflowError {
#[error("Workflow not found: {0}")]
NotFound(WorkflowInstanceId),
@@ -269,6 +269,9 @@ pub enum WorkflowError {
#[error("Context value not found: {0}")]
ContextValueNotFound(String),
#[error("Engine is shutting down, not accepting new workflows")]
ShuttingDown,
}
pub type WorkflowResult<T> = Result<T, WorkflowError>;