From 526fd0082f367a7136a053522f163a7d6b21a066 Mon Sep 17 00:00:00 2001 From: Simo Lin Date: Fri, 12 Dec 2025 09:03:28 -0800 Subject: [PATCH] [model-gateway] refactor: workflow engine cleanup and minor optimization (#15001) --- sgl-model-gateway/src/workflow/engine.rs | 28 ++++++------------------ sgl-model-gateway/src/workflow/event.rs | 6 +++++ sgl-model-gateway/src/workflow/mod.rs | 1 - 3 files changed, 13 insertions(+), 22 deletions(-) diff --git a/sgl-model-gateway/src/workflow/engine.rs b/sgl-model-gateway/src/workflow/engine.rs index 56540ef95..c5d8e6998 100644 --- a/sgl-model-gateway/src/workflow/engine.rs +++ b/sgl-model-gateway/src/workflow/engine.rs @@ -22,8 +22,6 @@ use super::{ types::*, }; -/// Consolidated step execution tracking to minimize lock contention. -/// Single lock instead of 4 separate locks. #[derive(Default)] struct StepTracker { completed: HashSet, @@ -213,24 +211,17 @@ impl WorkflowEngine { let start_time = std::time::Instant::now(); let step_count = definition.steps.len(); - // Consolidated tracking state - single lock instead of 4 let tracker: Arc> = Arc::new(RwLock::new(StepTracker::default())); - - // Channel for step completion notifications - // Capacity equals step count to prevent blocking on send let (tx, mut rx) = mpsc::channel::<(StepId, StepResult)>(step_count.max(1)); loop { - // Check if workflow was cancelled - let state = self.state_store.load(instance_id)?; - if state.status == WorkflowStatus::Cancelled { + if self.state_store.is_cancelled(instance_id)? { self.event_bus .publish(WorkflowEvent::WorkflowCancelled { instance_id }) .await; return Ok(()); } - // Single lock acquisition for all tracking state reads let (ready_step_indices, total_processed, running_count, blocked_by_failure) = { let t = tracker.read(); @@ -305,7 +296,6 @@ impl WorkflowEngine { .execute_step_with_retry(instance_id, step, &def) .await; - // Single lock acquisition to update tracking state { let mut t = tracker.write(); t.running.remove(&step_id); @@ -339,7 +329,6 @@ impl WorkflowEngine { } } - // Send completion signal (result determines the message) let signal = match result { Ok(r) => r, Err(_) => StepResult::Failure, @@ -360,19 +349,19 @@ impl WorkflowEngine { } } - // Check final status - single lock read - let has_failures = !tracker.read().failed.is_empty(); + let failed_step = { + let t = tracker.read(); + t.failed.iter().next().cloned() + }; - if has_failures { + if let Some(ref step) = failed_step { self.state_store.update(instance_id, |s| { s.status = WorkflowStatus::Failed; })?; - - let failed_step = tracker.read().failed.iter().next().cloned(); self.event_bus .publish(WorkflowEvent::WorkflowFailed { instance_id, - failed_step: failed_step.unwrap_or_else(|| StepId::new("unknown")), + failed_step: step.clone(), error: "One or more steps failed".to_string(), }) .await; @@ -413,7 +402,6 @@ impl WorkflowEngine { let mut backoff = Self::create_backoff(&retry_policy.backoff); loop { - // Check for cancellation before starting/retrying step (optimized: no full state load) if self.state_store.is_cancelled(instance_id)? { return Err(WorkflowError::Cancelled(instance_id)); } @@ -441,7 +429,6 @@ impl WorkflowEngine { }) .await; - // Get current context (optimized: only clone context, not full state) let mut context = self.state_store.get_context(instance_id)?; // Execute step with timeout @@ -450,7 +437,6 @@ impl WorkflowEngine { let step_duration = step_start.elapsed(); - // Save updated context (use std::mem::replace to avoid extra clone) self.state_store.update(instance_id, |s| { s.context = std::mem::replace(&mut context, WorkflowContext::new(instance_id)); })?; diff --git a/sgl-model-gateway/src/workflow/event.rs b/sgl-model-gateway/src/workflow/event.rs index 1fd5202e8..a21aaad75 100644 --- a/sgl-model-gateway/src/workflow/event.rs +++ b/sgl-model-gateway/src/workflow/event.rs @@ -89,6 +89,12 @@ impl Default for EventBus { } } +impl std::fmt::Debug for EventBus { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("EventBus").finish_non_exhaustive() + } +} + /// Logging subscriber that logs events using tracing pub struct LoggingSubscriber; diff --git a/sgl-model-gateway/src/workflow/mod.rs b/sgl-model-gateway/src/workflow/mod.rs index 133bdde60..aa8e625db 100644 --- a/sgl-model-gateway/src/workflow/mod.rs +++ b/sgl-model-gateway/src/workflow/mod.rs @@ -7,7 +7,6 @@ mod executor; mod state; pub mod types; -// Re-export main types pub use definition::{StepDefinition, WorkflowDefinition}; pub use engine::WorkflowEngine; pub use event::{EventBus, EventSubscriber, LoggingSubscriber, WorkflowEvent};