From ed729d22b3bb4dc998acaac0310d78ee60d7513a Mon Sep 17 00:00:00 2001 From: Simo Lin Date: Mon, 12 Jan 2026 10:47:00 -0800 Subject: [PATCH] [model-gateway] refactor workflow engine from type erasure to typed engines (#16973) --- sgl-model-gateway/src/app_context.rs | 35 +- sgl-model-gateway/src/core/job_queue.rs | 362 +++++++----------- .../src/core/steps/mcp_registration.rs | 62 ++- sgl-model-gateway/src/core/steps/mod.rs | 9 +- .../src/core/steps/tokenizer_registration.rs | 33 +- .../core/steps/wasm_module_registration.rs | 90 +++-- .../src/core/steps/wasm_module_removal.rs | 33 +- .../steps/worker/external/create_workers.rs | 21 +- .../steps/worker/external/discover_models.rs | 11 +- .../src/core/steps/worker/external/mod.rs | 10 +- .../core/steps/worker/local/create_worker.rs | 31 +- .../steps/worker/local/detect_connection.rs | 14 +- .../core/steps/worker/local/discover_dp.rs | 11 +- .../steps/worker/local/discover_metadata.rs | 22 +- .../worker/local/find_worker_to_update.rs | 16 +- .../worker/local/find_workers_to_remove.rs | 21 +- .../src/core/steps/worker/local/mod.rs | 27 +- .../steps/worker/local/register_tokenizer.rs | 15 +- .../local/remove_from_policy_registry.rs | 13 +- .../local/remove_from_worker_registry.rs | 12 +- .../local/update_policies_for_worker.rs | 18 +- .../worker/local/update_remaining_policies.rs | 14 +- .../worker/local/update_worker_properties.rs | 19 +- .../src/core/steps/worker/shared/activate.rs | 14 +- .../src/core/steps/worker/shared/register.rs | 14 +- .../steps/worker/shared/update_policies.rs | 13 +- .../src/core/steps/workflow_data.rs | 296 +++----------- .../src/core/steps/workflow_engines.rs | 167 ++++++++ sgl-model-gateway/src/server.rs | 51 +-- sgl-model-gateway/src/service_discovery.rs | 2 +- sgl-model-gateway/src/workflow/engine.rs | 59 +++ sgl-model-gateway/src/workflow/state.rs | 33 +- sgl-model-gateway/tests/common/mod.rs | 81 ++-- sgl-model-gateway/tests/common/test_app.rs | 10 +- .../tests/routing/test_pd_routing.rs | 6 +- sgl-model-gateway/tests/wasm_test.rs | 75 ++-- 36 files changed, 803 insertions(+), 917 deletions(-) create mode 100644 sgl-model-gateway/src/core/steps/workflow_engines.rs diff --git a/sgl-model-gateway/src/app_context.rs b/sgl-model-gateway/src/app_context.rs index 130b61ecc..e2e112e5b 100644 --- a/sgl-model-gateway/src/app_context.rs +++ b/sgl-model-gateway/src/app_context.rs @@ -9,8 +9,8 @@ use tracing::{debug, info}; use crate::{ config::RouterConfig, core::{ - steps::workflow_data::AnyWorkflowData, JobQueue, LoadMonitor, WorkerRegistry, - WorkerService, UNKNOWN_MODEL_ID, + steps::WorkflowEngines, JobQueue, LoadMonitor, WorkerRegistry, WorkerService, + UNKNOWN_MODEL_ID, }, data_connector::{ create_storage, ConversationItemStorage, ConversationStorage, ResponseStorage, @@ -29,12 +29,8 @@ use crate::{ }, tool_parser::ParserFactory as ToolParserFactory, wasm::{config::WasmRuntimeConfig, module_manager::WasmModuleManager}, - workflow::{InMemoryStore, WorkflowEngine}, }; -/// Type alias for the concrete workflow engine used in the application -pub type AppWorkflowEngine = WorkflowEngine>; - /// Error type for AppContext builder #[derive(Debug)] pub struct AppContextBuildError(&'static str); @@ -65,7 +61,7 @@ pub struct AppContext { pub configured_reasoning_parser: Option, pub configured_tool_parser: Option, pub worker_job_queue: Arc>>, - pub workflow_engine: Arc>>, + pub workflow_engines: Arc>, pub mcp_manager: Arc>>, pub wasm_manager: Option>, pub worker_service: Arc, @@ -95,7 +91,7 @@ pub struct AppContextBuilder { conversation_item_storage: Option>, load_monitor: Option>, worker_job_queue: Option>>>, - workflow_engine: Option>>>, + workflow_engines: Option>>, mcp_manager: Option>>>, wasm_manager: Option>, } @@ -135,7 +131,7 @@ impl AppContextBuilder { conversation_item_storage: None, load_monitor: None, worker_job_queue: None, - workflow_engine: None, + workflow_engines: None, mcp_manager: None, wasm_manager: None, } @@ -220,11 +216,8 @@ impl AppContextBuilder { self } - pub fn workflow_engine( - mut self, - workflow_engine: Arc>>, - ) -> Self { - self.workflow_engine = Some(workflow_engine); + pub fn workflow_engines(mut self, workflow_engines: Arc>) -> Self { + self.workflow_engines = Some(workflow_engines); self } @@ -286,9 +279,9 @@ impl AppContextBuilder { configured_reasoning_parser, configured_tool_parser, worker_job_queue, - workflow_engine: self - .workflow_engine - .ok_or(AppContextBuildError("workflow_engine"))?, + workflow_engines: self + .workflow_engines + .ok_or(AppContextBuildError("workflow_engines"))?, mcp_manager: self .mcp_manager .ok_or(AppContextBuildError("mcp_manager"))?, @@ -315,7 +308,7 @@ impl AppContextBuilder { .with_storage(&router_config)? .with_load_monitor(&router_config) .with_worker_job_queue() - .with_workflow_engine() + .with_workflow_engines() .with_mcp_manager(&router_config) .await? .with_wasm_manager(&router_config)? @@ -549,9 +542,9 @@ impl AppContextBuilder { self } - /// Create workflow engine OnceLock container - fn with_workflow_engine(mut self) -> Self { - self.workflow_engine = Some(Arc::new(OnceLock::new())); + /// Create workflow engines OnceLock container + fn with_workflow_engines(mut self) -> Self { + self.workflow_engines = Some(Arc::new(OnceLock::new())); self } diff --git a/sgl-model-gateway/src/core/job_queue.rs b/sgl-model-gateway/src/core/job_queue.rs index a657b3723..6412f7d6e 100644 --- a/sgl-model-gateway/src/core/job_queue.rs +++ b/sgl-model-gateway/src/core/job_queue.rs @@ -14,7 +14,7 @@ use tokio::sync::{mpsc, Semaphore}; use tracing::{debug, error, info, warn}; use crate::{ - app_context::{AppContext, AppWorkflowEngine}, + app_context::AppContext, config::{RouterConfig, RoutingMode}, core::steps::{ create_external_worker_workflow_data, create_local_worker_workflow_data, @@ -26,7 +26,7 @@ use crate::{ }, mcp::McpConfig, protocols::worker_spec::{JobStatus, WorkerConfigRequest, WorkerUpdateRequest}, - workflow::{WorkflowId, WorkflowInstanceId, WorkflowStatus}, + workflow::WorkflowId, }; /// Job types for control plane operations @@ -335,37 +335,93 @@ impl JobQueue { async fn execute_job(job: &Job, context: &Arc) -> Result { match job { Job::AddWorker { config } => { - let engine = context - .workflow_engine + let engines = context + .workflow_engines .get() - .ok_or_else(|| "Workflow engine not initialized".to_string())?; - - let instance_id = Self::start_worker_workflow(engine, config, context).await?; - - debug!( - "Started worker registration workflow for {} (instance: {})", - config.url, instance_id - ); + .ok_or_else(|| "Workflow engines not initialized".to_string())?; let timeout_duration = Duration::from_secs(context.router_config.worker_startup_timeout_secs + 30); - Self::wait_for_workflow_completion( - engine, - instance_id, - &config.url, - timeout_duration, - ) - .await + // Select workflow based on runtime field + match config.runtime.as_deref() { + Some("external") => { + let workflow_data = create_external_worker_workflow_data( + (**config).clone(), + Arc::clone(context), + ); + let instance_id = engines + .external_worker + .start_workflow( + WorkflowId::new("external_worker_registration"), + workflow_data, + ) + .await + .map_err(|e| { + format!( + "Failed to start external worker registration workflow: {:?}", + e + ) + })?; + + debug!( + "Started external worker registration workflow for {} (instance: {})", + config.url, instance_id + ); + + engines + .external_worker + .wait_for_completion(instance_id, &config.url, timeout_duration) + .await + } + _ => { + let workflow_data = create_local_worker_workflow_data( + (**config).clone(), + Arc::clone(context), + ); + let instance_id = engines + .local_worker + .start_workflow( + WorkflowId::new("local_worker_registration"), + workflow_data, + ) + .await + .map_err(|e| { + format!( + "Failed to start local worker registration workflow: {:?}", + e + ) + })?; + + debug!( + "Started local worker registration workflow for {} (instance: {})", + config.url, instance_id + ); + + engines + .local_worker + .wait_for_completion(instance_id, &config.url, timeout_duration) + .await + } + } } Job::UpdateWorker { url, update } => { - let engine = context - .workflow_engine + let engines = context + .workflow_engines .get() - .ok_or_else(|| "Workflow engine not initialized".to_string())?; + .ok_or_else(|| "Workflow engines not initialized".to_string())?; - let instance_id = - Self::start_worker_update_workflow(engine, url, update, context).await?; + let workflow_data = create_worker_update_workflow_data( + url.to_string(), + (**update).clone(), + Arc::clone(context), + ); + + let instance_id = engines + .worker_update + .start_workflow(WorkflowId::new("worker_update"), workflow_data) + .await + .map_err(|e| format!("Failed to start worker update workflow: {:?}", e))?; debug!( "Started worker update workflow for {} (instance: {})", @@ -374,15 +430,28 @@ impl JobQueue { let timeout_duration = Duration::from_secs(30); - Self::wait_for_workflow_completion(engine, instance_id, url, timeout_duration).await + engines + .worker_update + .wait_for_completion(instance_id, url, timeout_duration) + .await } Job::RemoveWorker { url } => { - let engine = context - .workflow_engine + let engines = context + .workflow_engines .get() - .ok_or_else(|| "Workflow engine not initialized".to_string())?; + .ok_or_else(|| "Workflow engines not initialized".to_string())?; - let instance_id = Self::start_worker_removal_workflow(engine, url, context).await?; + let workflow_data = create_worker_removal_workflow_data( + url.to_string(), + context.router_config.dp_aware, + Arc::clone(context), + ); + + let instance_id = engines + .worker_removal + .start_workflow(WorkflowId::new("worker_removal"), workflow_data) + .await + .map_err(|e| format!("Failed to start worker removal workflow: {:?}", e))?; debug!( "Started worker removal workflow for {} (instance: {})", @@ -391,9 +460,10 @@ impl JobQueue { let timeout_duration = Duration::from_secs(30); - let result = - Self::wait_for_workflow_completion(engine, instance_id, url, timeout_duration) - .await; + let result = engines + .worker_removal + .wait_for_completion(instance_id, url, timeout_duration) + .await; // Clean up job status when removing worker if let Some(queue) = context.worker_job_queue.get() { @@ -403,15 +473,16 @@ impl JobQueue { result } Job::AddWasmModule { config } => { - let engine = context - .workflow_engine + let engines = context + .workflow_engines .get() - .ok_or_else(|| "Workflow engine not initialized".to_string())?; + .ok_or_else(|| "Workflow engines not initialized".to_string())?; let workflow_data = create_wasm_registration_workflow_data(*config.clone(), Arc::clone(context)); - let instance_id = engine + let instance_id = engines + .wasm_registration .start_workflow(WorkflowId::new("wasm_module_registration"), workflow_data) .await .map_err(|e| { @@ -425,24 +496,22 @@ impl JobQueue { let timeout_duration = Duration::from_secs(300); // 5 minutes - Self::wait_for_workflow_completion( - engine, - instance_id, - &config.descriptor.name, - timeout_duration, - ) - .await + engines + .wasm_registration + .wait_for_completion(instance_id, &config.descriptor.name, timeout_duration) + .await } Job::RemoveWasmModule { request } => { - let engine = context - .workflow_engine + let engines = context + .workflow_engines .get() - .ok_or_else(|| "Workflow engine not initialized".to_string())?; + .ok_or_else(|| "Workflow engines not initialized".to_string())?; let workflow_data = create_wasm_removal_workflow_data(*request.clone(), Arc::clone(context)); - let instance_id = engine + let instance_id = engines + .wasm_removal .start_workflow(WorkflowId::new("wasm_module_removal"), workflow_data) .await .map_err(|e| { @@ -456,13 +525,14 @@ impl JobQueue { let timeout_duration = Duration::from_secs(60); // 1 minute - Self::wait_for_workflow_completion( - engine, - instance_id, - &request.module_uuid.to_string(), - timeout_duration, - ) - .await + engines + .wasm_removal + .wait_for_completion( + instance_id, + &request.module_uuid.to_string(), + timeout_duration, + ) + .await } Job::InitializeWorkersFromConfig { router_config } => { let api_key = router_config.api_key.clone(); @@ -640,13 +710,19 @@ impl JobQueue { Ok(format!("Submitted {} RegisterMcpServer jobs", server_count)) } Job::RegisterMcpServer { config } => { - let engine = context - .workflow_engine + let engines = context + .workflow_engines .get() - .ok_or_else(|| "Workflow engine not initialized".to_string())?; + .ok_or_else(|| "Workflow engines not initialized".to_string())?; - let instance_id = - Self::start_mcp_registration_workflow(engine, config, context).await?; + let workflow_data = + create_mcp_workflow_data((**config).clone(), Arc::clone(context)); + + let instance_id = engines + .mcp + .start_workflow(WorkflowId::new("mcp_registration"), workflow_data) + .await + .map_err(|e| format!("Failed to start MCP registration workflow: {:?}", e))?; debug!( "Started MCP registration workflow for {} (instance: {})", @@ -655,24 +731,22 @@ impl JobQueue { let timeout_duration = Duration::from_secs(7200 + 30); // 2hr + margin - Self::wait_for_workflow_completion( - engine, - instance_id, - &config.name, - timeout_duration, - ) - .await + engines + .mcp + .wait_for_completion(instance_id, &config.name, timeout_duration) + .await } Job::AddTokenizer { config } => { - let engine = context - .workflow_engine + let engines = context + .workflow_engines .get() - .ok_or_else(|| "Workflow engine not initialized".to_string())?; + .ok_or_else(|| "Workflow engines not initialized".to_string())?; let workflow_data = create_tokenizer_workflow_data(*config.clone(), Arc::clone(context)); - let instance_id = engine + let instance_id = engines + .tokenizer .start_workflow(WorkflowId::new("tokenizer_registration"), workflow_data) .await .map_err(|e| { @@ -687,13 +761,10 @@ impl JobQueue { // Allow up to 10 minutes for HuggingFace downloads let timeout_duration = Duration::from_secs(600); - Self::wait_for_workflow_completion( - engine, - instance_id, - &config.id, - timeout_duration, - ) - .await + engines + .tokenizer + .wait_for_completion(instance_id, &config.id, timeout_duration) + .await } Job::RemoveTokenizer { request } => { // Tokenizer removal is synchronous and fast @@ -710,143 +781,6 @@ impl JobQueue { } } - /// Start a workflow and return its instance ID - async fn start_worker_workflow( - engine: &Arc, - config: &WorkerConfigRequest, - context: &Arc, - ) -> Result { - // Select workflow based on runtime field - let (workflow_id, workflow_data) = match config.runtime.as_deref() { - Some("external") => ( - WorkflowId::new("external_worker_registration"), - create_external_worker_workflow_data(config.clone(), Arc::clone(context)), - ), - _ => ( - WorkflowId::new("local_worker_registration"), - create_local_worker_workflow_data(config.clone(), Arc::clone(context)), - ), - }; - - engine - .start_workflow(workflow_id, workflow_data) - .await - .map_err(|e| format!("Failed to start worker registration workflow: {:?}", e)) - } - - /// Start worker removal workflow - async fn start_worker_removal_workflow( - engine: &Arc, - url: &str, - context: &Arc, - ) -> Result { - let workflow_data = create_worker_removal_workflow_data( - url.to_string(), - context.router_config.dp_aware, - Arc::clone(context), - ); - - engine - .start_workflow(WorkflowId::new("worker_removal"), workflow_data) - .await - .map_err(|e| format!("Failed to start worker removal workflow: {:?}", e)) - } - - /// Start worker update workflow - async fn start_worker_update_workflow( - engine: &Arc, - url: &str, - update: &WorkerUpdateRequest, - context: &Arc, - ) -> Result { - let workflow_data = create_worker_update_workflow_data( - url.to_string(), - update.clone(), - Arc::clone(context), - ); - - engine - .start_workflow(WorkflowId::new("worker_update"), workflow_data) - .await - .map_err(|e| format!("Failed to start worker update workflow: {:?}", e)) - } - - /// Start MCP server registration workflow - async fn start_mcp_registration_workflow( - engine: &Arc, - config: &McpServerConfigRequest, - context: &Arc, - ) -> Result { - let workflow_data = create_mcp_workflow_data(config.clone(), Arc::clone(context)); - - engine - .start_workflow(WorkflowId::new("mcp_registration"), workflow_data) - .await - .map_err(|e| format!("Failed to start MCP registration workflow: {:?}", e)) - } - - /// Wait for workflow completion with adaptive polling - async fn wait_for_workflow_completion( - engine: &Arc, - instance_id: WorkflowInstanceId, - worker_url: &str, - timeout_duration: Duration, - ) -> Result { - let start = std::time::Instant::now(); - let mut poll_interval = Duration::from_millis(100); - let max_poll_interval = Duration::from_millis(2000); - let poll_backoff = Duration::from_millis(200); - - loop { - // Check timeout - if start.elapsed() > timeout_duration { - return Err(format!( - "Workflow timeout after {}s for worker {}", - timeout_duration.as_secs(), - worker_url - )); - } - - // Get workflow status - let state = engine - .get_status(instance_id) - .map_err(|e| format!("Failed to get workflow status: {:?}", e))?; - - let result = match state.status { - WorkflowStatus::Completed => Ok(format!( - "Worker {} registered and activated successfully via workflow", - worker_url - )), - WorkflowStatus::Failed => { - let current_step = state.current_step.as_ref(); - let step_name = current_step - .map(|s| s.to_string()) - .unwrap_or_else(|| "unknown".to_string()); - let error_msg = current_step - .and_then(|step_id| state.step_states.get(step_id)) - .and_then(|s| s.last_error.as_deref()) - .unwrap_or("Unknown error"); - Err(format!( - "Workflow failed at step {}: {}", - step_name, error_msg - )) - } - WorkflowStatus::Cancelled => { - Err(format!("Workflow cancelled for worker {}", worker_url)) - } - WorkflowStatus::Pending | WorkflowStatus::Paused | WorkflowStatus::Running => { - tokio::time::sleep(poll_interval).await; - poll_interval = (poll_interval + poll_backoff).min(max_poll_interval); - continue; - } - }; - - // Clean up terminal workflow states - engine.state_store().cleanup_if_terminal(instance_id); - return result; - } - } - /// Update job status on completion fn record_job_completion( job_type: &'static str, diff --git a/sgl-model-gateway/src/core/steps/mcp_registration.rs b/sgl-model-gateway/src/core/steps/mcp_registration.rs index a55398b1d..92a6108c1 100644 --- a/sgl-model-gateway/src/core/steps/mcp_registration.rs +++ b/sgl-model-gateway/src/core/steps/mcp_registration.rs @@ -3,7 +3,7 @@ use std::{sync::Arc, time::Duration}; use async_trait::async_trait; use tracing::{debug, error, info, warn}; -use super::workflow_data::{AnyWorkflowData, McpWorkflowData}; +use super::workflow_data::McpWorkflowData; use crate::{ app_context::AppContext, mcp::{config::McpServerConfig, manager::McpManager}, @@ -38,14 +38,14 @@ impl McpServerConfigRequest { pub struct ConnectMcpServerStep; #[async_trait] -impl StepExecutor for ConnectMcpServerStep { +impl StepExecutor for ConnectMcpServerStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_mcp()?; - let config_request = &data.config; - let app_context = data + let config_request = &context.data.config; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; @@ -76,8 +76,7 @@ impl StepExecutor for ConnectMcpServerStep { ); // Store client in typed data - let data_mut = context.data.as_mcp_mut()?; - data_mut.mcp_client = Some(Arc::new(client)); + context.data.mcp_client = Some(Arc::new(client)); Ok(StepResult::Success) } @@ -96,18 +95,19 @@ impl StepExecutor for ConnectMcpServerStep { pub struct DiscoverMcpInventoryStep; #[async_trait] -impl StepExecutor for DiscoverMcpInventoryStep { +impl StepExecutor for DiscoverMcpInventoryStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_mcp()?; - let config_request = &data.config; - let app_context = data + let config_request = &context.data.config; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; - let mcp_client = data + let mcp_client = context + .data .mcp_client .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("mcp_client".to_string()))?; @@ -149,18 +149,19 @@ impl StepExecutor for DiscoverMcpInventoryStep { pub struct RegisterMcpServerStep; #[async_trait] -impl StepExecutor for RegisterMcpServerStep { +impl StepExecutor for RegisterMcpServerStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_mcp()?; - let config_request = &data.config; - let app_context = data + let config_request = &context.data.config; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; - let mcp_client = data + let mcp_client = context + .data .mcp_client .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("mcp_client".to_string()))? @@ -203,15 +204,13 @@ impl StepExecutor for RegisterMcpServerStep { pub struct ValidateRegistrationStep; #[async_trait] -impl StepExecutor for ValidateRegistrationStep { +impl StepExecutor for ValidateRegistrationStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_mcp()?; - let config_request = &data.config; - - let client_registered = data.mcp_client.is_some(); + let config_request = &context.data.config; + let client_registered = context.data.mcp_client.is_some(); if client_registered { info!( @@ -220,8 +219,7 @@ impl StepExecutor for ValidateRegistrationStep { ); // Mark as validated - let data_mut = context.data.as_mcp_mut()?; - data_mut.validated = true; + context.data.validated = true; return Ok(StepResult::Success); } @@ -263,7 +261,7 @@ impl StepExecutor for ValidateRegistrationStep { /// - DiscoverMcpInventory: 3 retries, 10s timeout (discovery + caching) /// - RegisterMcpServer: No retry, 5s timeout (fast registration) /// - ValidateRegistration: Final validation step -pub fn create_mcp_registration_workflow() -> WorkflowDefinition { +pub fn create_mcp_registration_workflow() -> WorkflowDefinition { WorkflowDefinition::new("mcp_registration", "MCP Server Registration") .add_step( StepDefinition::new( @@ -321,11 +319,11 @@ pub fn create_mcp_registration_workflow() -> WorkflowDefinition pub fn create_mcp_workflow_data( config: McpServerConfigRequest, app_context: Arc, -) -> AnyWorkflowData { - AnyWorkflowData::Mcp(McpWorkflowData { +) -> McpWorkflowData { + McpWorkflowData { config, validated: false, app_context: Some(app_context), mcp_client: None, - }) + } } diff --git a/sgl-model-gateway/src/core/steps/mod.rs b/sgl-model-gateway/src/core/steps/mod.rs index da1b6e6fa..48a8de8a4 100644 --- a/sgl-model-gateway/src/core/steps/mod.rs +++ b/sgl-model-gateway/src/core/steps/mod.rs @@ -12,6 +12,7 @@ pub mod wasm_module_registration; pub mod wasm_module_removal; pub mod worker; pub mod workflow_data; +pub mod workflow_engines; // Worker management (registration, removal) pub use mcp_registration::{ @@ -75,8 +76,10 @@ pub use worker::{ }; // Typed workflow data structures pub use workflow_data::{ - AnyWorkflowData, ExternalWorkerWorkflowData, LocalWorkerWorkflowData, McpWorkflowData, - ProtocolUpdateRequest, TokenizerWorkflowData, WasmRegistrationWorkflowData, - WasmRemovalWorkflowData, WorkerConfigRequest, WorkerList as WorkflowWorkerList, + ExternalWorkerWorkflowData, LocalWorkerWorkflowData, McpWorkflowData, ProtocolUpdateRequest, + TokenizerWorkflowData, WasmRegistrationWorkflowData, WasmRemovalWorkflowData, + WorkerConfigRequest, WorkerList as WorkflowWorkerList, WorkerRegistrationData, WorkerRemovalWorkflowData, WorkerUpdateWorkflowData, }; +// Typed workflow engines +pub use workflow_engines::WorkflowEngines; diff --git a/sgl-model-gateway/src/core/steps/tokenizer_registration.rs b/sgl-model-gateway/src/core/steps/tokenizer_registration.rs index 7050fb1af..cf4979ab9 100644 --- a/sgl-model-gateway/src/core/steps/tokenizer_registration.rs +++ b/sgl-model-gateway/src/core/steps/tokenizer_registration.rs @@ -9,7 +9,7 @@ use async_trait::async_trait; use serde::{Deserialize, Serialize}; use tracing::{debug, error, info}; -use super::workflow_data::{AnyWorkflowData, TokenizerWorkflowData}; +use super::workflow_data::TokenizerWorkflowData; use crate::{ app_context::AppContext, tokenizer::factory, @@ -47,14 +47,14 @@ pub struct TokenizerRemovalRequest { pub struct ValidateTokenizerConfigStep; #[async_trait] -impl StepExecutor for ValidateTokenizerConfigStep { +impl StepExecutor for ValidateTokenizerConfigStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_tokenizer()?; - let config = &data.config; - let app_context = data + let config = &context.data.config; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; @@ -101,14 +101,14 @@ impl StepExecutor for ValidateTokenizerConfigStep { pub struct LoadTokenizerStep; #[async_trait] -impl StepExecutor for LoadTokenizerStep { +impl StepExecutor for LoadTokenizerStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_tokenizer()?; - let config = &data.config; - let app_context = data + let config = &context.data.config; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))? @@ -157,8 +157,7 @@ impl StepExecutor for LoadTokenizerStep { // Store vocab size in typed data if let Some(size) = vocab_size { - let data_mut = context.data.as_tokenizer_mut()?; - data_mut.vocab_size = Some(size); + context.data.vocab_size = Some(size); } Ok(StepResult::Success) @@ -191,7 +190,7 @@ impl StepExecutor for LoadTokenizerStep { /// Workflow configuration: /// - ValidateConfig: No retry, 5s timeout (fast validation) /// - LoadTokenizer: 3 retries, 5min timeout (may need to download from HuggingFace) -pub fn create_tokenizer_registration_workflow() -> WorkflowDefinition { +pub fn create_tokenizer_registration_workflow() -> WorkflowDefinition { WorkflowDefinition::new("tokenizer_registration", "Tokenizer Registration") .add_step( StepDefinition::new( @@ -222,12 +221,12 @@ pub fn create_tokenizer_registration_workflow() -> WorkflowDefinition, -) -> AnyWorkflowData { - AnyWorkflowData::Tokenizer(TokenizerWorkflowData { +) -> TokenizerWorkflowData { + TokenizerWorkflowData { config, vocab_size: None, app_context: Some(app_context), - }) + } } #[cfg(test)] diff --git a/sgl-model-gateway/src/core/steps/wasm_module_registration.rs b/sgl-model-gateway/src/core/steps/wasm_module_registration.rs index 5b6b9a9c0..b188f2ed9 100644 --- a/sgl-model-gateway/src/core/steps/wasm_module_registration.rs +++ b/sgl-model-gateway/src/core/steps/wasm_module_registration.rs @@ -10,7 +10,7 @@ use tracing::{debug, info, warn}; use uuid::Uuid; use wasmtime::{component::Component, Config, Engine}; -use super::workflow_data::{AnyWorkflowData, WasmRegistrationWorkflowData}; +use super::workflow_data::WasmRegistrationWorkflowData; use crate::{ app_context::AppContext, wasm::module::{WasmModule, WasmModuleDescriptor, WasmModuleMeta}, @@ -65,13 +65,12 @@ fn has_wasm_extension(path: &Path) -> bool { pub struct ValidateDescriptorStep; #[async_trait] -impl StepExecutor for ValidateDescriptorStep { +impl StepExecutor for ValidateDescriptorStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_wasm_registration()?; - let descriptor = &data.config.descriptor; + let descriptor = &context.data.config.descriptor; debug!("Validating WASM module descriptor: {}", descriptor.name); @@ -207,8 +206,7 @@ impl StepExecutor for ValidateDescriptorStep { let module_name = descriptor.name.clone(); // Store file size in typed data - let data_mut = context.data.as_wasm_registration_mut()?; - data_mut.file_size_bytes = Some(metadata.len()); + context.data.file_size_bytes = Some(metadata.len()); info!( "Descriptor validated successfully for module: {}", @@ -229,13 +227,12 @@ impl StepExecutor for ValidateDescriptorStep { pub struct CalculateHashStep; #[async_trait] -impl StepExecutor for CalculateHashStep { +impl StepExecutor for CalculateHashStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_wasm_registration()?; - let file_path = &data.config.descriptor.file_path; + let file_path = &context.data.config.descriptor.file_path; debug!("Calculating SHA256 hash for: {}", file_path); @@ -274,8 +271,7 @@ impl StepExecutor for CalculateHashStep { let path_for_log = file_path.clone(); // Store hash in typed data - let data_mut = context.data.as_wasm_registration_mut()?; - data_mut.sha256_hash = Some(hash); + context.data.sha256_hash = Some(hash); info!("SHA256 hash calculated for: {}", path_for_log); Ok(StepResult::Success) @@ -293,24 +289,25 @@ impl StepExecutor for CalculateHashStep { pub struct CheckDuplicateStep; #[async_trait] -impl StepExecutor for CheckDuplicateStep { +impl StepExecutor for CheckDuplicateStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_wasm_registration()?; - let app_context = data + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; - let sha256_hash = data + let sha256_hash = context + .data .sha256_hash .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("sha256_hash".to_string()))?; debug!( "Checking for duplicate SHA256 hash for module: {}", - data.config.descriptor.name + context.data.config.descriptor.name ); // Get WASM module manager from app context @@ -333,7 +330,7 @@ impl StepExecutor for CheckDuplicateStep { info!( "No duplicate found for module: {}", - data.config.descriptor.name + context.data.config.descriptor.name ); Ok(StepResult::Success) } @@ -350,13 +347,12 @@ impl StepExecutor for CheckDuplicateStep { pub struct LoadWasmBytesStep; #[async_trait] -impl StepExecutor for LoadWasmBytesStep { +impl StepExecutor for LoadWasmBytesStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_wasm_registration()?; - let file_path = &data.config.descriptor.file_path; + let file_path = &context.data.config.descriptor.file_path; debug!("Loading WASM bytes from: {}", file_path); @@ -372,8 +368,7 @@ impl StepExecutor for LoadWasmBytesStep { })?; // Store WASM bytes in typed data - let data_mut = context.data.as_wasm_registration_mut()?; - data_mut.wasm_bytes = Some(wasm_bytes); + context.data.wasm_bytes = Some(wasm_bytes); info!("WASM bytes loaded from: {}", path_for_log); Ok(StepResult::Success) @@ -391,20 +386,20 @@ impl StepExecutor for LoadWasmBytesStep { pub struct ValidateWasmComponentStep; #[async_trait] -impl StepExecutor for ValidateWasmComponentStep { +impl StepExecutor for ValidateWasmComponentStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_wasm_registration()?; - let wasm_bytes = data + let wasm_bytes = context + .data .wasm_bytes .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_bytes".to_string()))?; debug!( "Validating WASM component format for module: {}", - data.config.descriptor.name + context.data.config.descriptor.name ); // Create a temporary engine to validate the component @@ -431,7 +426,7 @@ impl StepExecutor for ValidateWasmComponentStep { info!( "WASM component validated successfully for module: {}", - data.config.descriptor.name + context.data.config.descriptor.name ); Ok(StepResult::Success) } @@ -448,29 +443,32 @@ impl StepExecutor for ValidateWasmComponentStep { pub struct RegisterModuleStep; #[async_trait] -impl StepExecutor for RegisterModuleStep { +impl StepExecutor for RegisterModuleStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_wasm_registration()?; - let app_context = data + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; - let sha256_hash = data + let sha256_hash = context + .data .sha256_hash .ok_or_else(|| WorkflowError::ContextValueNotFound("sha256_hash".to_string()))?; - let file_size_bytes = data + let file_size_bytes = context + .data .file_size_bytes .ok_or_else(|| WorkflowError::ContextValueNotFound("file_size_bytes".to_string()))?; - let wasm_bytes = data + let wasm_bytes = context + .data .wasm_bytes .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_bytes".to_string()))? .clone(); - let descriptor = &data.config.descriptor; + let descriptor = &context.data.config.descriptor; debug!("Registering WASM module in manager: {}", descriptor.name); @@ -519,8 +517,7 @@ impl StepExecutor for RegisterModuleStep { })?; // Store module UUID in typed data - let data_mut = context.data.as_wasm_registration_mut()?; - data_mut.module_uuid = Some(module_uuid); + context.data.module_uuid = Some(module_uuid); info!( "WASM module registered successfully: {} (UUID: {})", @@ -552,7 +549,8 @@ impl StepExecutor for RegisterModuleStep { /// - LoadWasmBytes: 3 retries, 60s timeout (I/O intensive) /// - ValidateWasmComponent: No retry, 30s timeout (CPU intensive validation) /// - RegisterModule: No retry, 5s timeout (fast registration) -pub fn create_wasm_module_registration_workflow() -> WorkflowDefinition { +pub fn create_wasm_module_registration_workflow() -> WorkflowDefinition +{ WorkflowDefinition::new("wasm_module_registration", "WASM Module Registration") .add_step( StepDefinition::new( @@ -627,13 +625,13 @@ pub fn create_wasm_module_registration_workflow() -> WorkflowDefinition, -) -> AnyWorkflowData { - AnyWorkflowData::WasmRegistration(WasmRegistrationWorkflowData { +) -> WasmRegistrationWorkflowData { + WasmRegistrationWorkflowData { config, wasm_bytes: None, sha256_hash: None, file_size_bytes: None, module_uuid: None, app_context: Some(app_context), - }) + } } diff --git a/sgl-model-gateway/src/core/steps/wasm_module_removal.rs b/sgl-model-gateway/src/core/steps/wasm_module_removal.rs index 83692fc5a..5a89a6245 100644 --- a/sgl-model-gateway/src/core/steps/wasm_module_removal.rs +++ b/sgl-model-gateway/src/core/steps/wasm_module_removal.rs @@ -4,7 +4,7 @@ use async_trait::async_trait; use tracing::{debug, info}; use uuid::Uuid; -use super::workflow_data::{AnyWorkflowData, WasmRemovalWorkflowData}; +use super::workflow_data::WasmRemovalWorkflowData; use crate::{ app_context::AppContext, workflow::{ @@ -37,14 +37,14 @@ impl WasmModuleRemovalRequest { pub struct FindModuleToRemoveStep; #[async_trait] -impl StepExecutor for FindModuleToRemoveStep { +impl StepExecutor for FindModuleToRemoveStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_wasm_removal()?; - let removal_request = &data.config; - let app_context = data + let removal_request = &context.data.config; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; @@ -80,8 +80,7 @@ impl StepExecutor for FindModuleToRemoveStep { let module_uuid = removal_request.module_uuid; // Store the module ID in typed data - let data_mut = context.data.as_wasm_removal_mut()?; - data_mut.module_id = Some(module_uuid.to_string()); + context.data.module_id = Some(module_uuid.to_string()); info!("Module found for removal: {}", module_uuid); Ok(StepResult::Success) @@ -98,14 +97,14 @@ impl StepExecutor for FindModuleToRemoveStep { pub struct RemoveModuleStep; #[async_trait] -impl StepExecutor for RemoveModuleStep { +impl StepExecutor for RemoveModuleStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_wasm_removal()?; - let removal_request = &data.config; - let app_context = data + let removal_request = &context.data.config; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; @@ -151,7 +150,7 @@ impl StepExecutor for RemoveModuleStep { /// Workflow configuration: /// - FindModuleToRemove: No retry, 5s timeout (fast lookup) /// - RemoveModule: No retry, 5s timeout (fast removal) -pub fn create_wasm_module_removal_workflow() -> WorkflowDefinition { +pub fn create_wasm_module_removal_workflow() -> WorkflowDefinition { WorkflowDefinition::new("wasm_module_removal", "WASM Module Removal") .add_step( StepDefinition::new( @@ -174,10 +173,10 @@ pub fn create_wasm_module_removal_workflow() -> WorkflowDefinition, -) -> AnyWorkflowData { - AnyWorkflowData::WasmRemoval(WasmRemovalWorkflowData { +) -> WasmRemovalWorkflowData { + WasmRemovalWorkflowData { config, module_id: None, app_context: Some(app_context), - }) + } } diff --git a/sgl-model-gateway/src/core/steps/worker/external/create_workers.rs b/sgl-model-gateway/src/core/steps/worker/external/create_workers.rs index ea32494a7..e23c240f4 100644 --- a/sgl-model-gateway/src/core/steps/worker/external/create_workers.rs +++ b/sgl-model-gateway/src/core/steps/worker/external/create_workers.rs @@ -8,7 +8,7 @@ use tracing::{debug, info}; use crate::{ core::{ circuit_breaker::CircuitBreakerConfig, - steps::workflow_data::{AnyWorkflowData, WorkerList}, + steps::workflow_data::{ExternalWorkerWorkflowData, WorkerList}, worker::{HealthConfig, RuntimeType, WorkerType}, BasicWorkerBuilder, ConnectionMode, Worker, }, @@ -28,18 +28,18 @@ fn normalize_external_url(url: &str) -> String { pub struct CreateExternalWorkersStep; #[async_trait] -impl StepExecutor for CreateExternalWorkersStep { +impl StepExecutor for CreateExternalWorkersStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_external_worker()?; - let config = &data.config; - let app_context = data + let config = &context.data.config; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; - let model_cards = &data.model_cards; + let model_cards = &context.data.model_cards; // Build configs from router settings let circuit_breaker_config = { @@ -150,10 +150,9 @@ impl StepExecutor for CreateExternalWorkersStep { } // Store results in workflow data - let data_mut = context.data.as_external_worker_mut()?; - data_mut.workers = Some(WorkerList::from_workers(&workers)); - data_mut.actual_workers = Some(workers); - data_mut.labels = labels; + context.data.workers = Some(WorkerList::from_workers(&workers)); + context.data.actual_workers = Some(workers); + context.data.labels = labels; Ok(StepResult::Success) } diff --git a/sgl-model-gateway/src/core/steps/worker/external/discover_models.rs b/sgl-model-gateway/src/core/steps/worker/external/discover_models.rs index eb9820c49..05edd1a7e 100644 --- a/sgl-model-gateway/src/core/steps/worker/external/discover_models.rs +++ b/sgl-model-gateway/src/core/steps/worker/external/discover_models.rs @@ -13,7 +13,7 @@ use crate::{ core::{ model_card::{ModelCard, ProviderType}, model_type::ModelType, - steps::workflow_data::AnyWorkflowData, + steps::workflow_data::ExternalWorkerWorkflowData, }, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, }; @@ -225,13 +225,12 @@ async fn fetch_models(url: &str, api_key: Option<&str>) -> Result pub struct DiscoverModelsStep; #[async_trait] -impl StepExecutor for DiscoverModelsStep { +impl StepExecutor for DiscoverModelsStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_external_worker()?; - let config = &data.config; + let config = &context.data.config; // If no API key is provided, skip model discovery and use wildcard mode. if config.api_key.as_ref().is_none_or(|k| k.is_empty()) { @@ -267,7 +266,7 @@ impl StepExecutor for DiscoverModelsStep { model_cards.iter().map(|c| &c.id).collect::>() ); - context.data.as_external_worker_mut()?.model_cards = model_cards; + context.data.model_cards = model_cards; Ok(StepResult::Success) } diff --git a/sgl-model-gateway/src/core/steps/worker/external/mod.rs b/sgl-model-gateway/src/core/steps/worker/external/mod.rs index 746efbd45..017b71db7 100644 --- a/sgl-model-gateway/src/core/steps/worker/external/mod.rs +++ b/sgl-model-gateway/src/core/steps/worker/external/mod.rs @@ -17,7 +17,7 @@ pub use discover_models::{ use super::shared::{ActivateWorkersStep, RegisterWorkersStep, UpdatePoliciesStep}; use crate::{ app_context::AppContext, - core::steps::workflow_data::{AnyWorkflowData, ExternalWorkerWorkflowData}, + core::steps::workflow_data::ExternalWorkerWorkflowData, protocols::worker_spec::WorkerConfigRequest, workflow::{BackoffStrategy, FailureAction, RetryPolicy, StepDefinition, WorkflowDefinition}, }; @@ -38,7 +38,7 @@ use crate::{ /// │ │ /// └────────────┴────────────┘ /// ``` -pub fn create_external_worker_workflow() -> WorkflowDefinition { +pub fn create_external_worker_workflow() -> WorkflowDefinition { WorkflowDefinition::new( "external_worker_registration", "External Worker Registration", @@ -110,13 +110,13 @@ pub fn create_external_worker_workflow() -> WorkflowDefinition pub fn create_external_worker_workflow_data( config: WorkerConfigRequest, app_context: Arc, -) -> AnyWorkflowData { - AnyWorkflowData::ExternalWorker(ExternalWorkerWorkflowData { +) -> ExternalWorkerWorkflowData { + ExternalWorkerWorkflowData { config, model_cards: Vec::new(), workers: None, labels: std::collections::HashMap::new(), app_context: Some(app_context), actual_workers: None, - }) + } } diff --git a/sgl-model-gateway/src/core/steps/worker/local/create_worker.rs b/sgl-model-gateway/src/core/steps/worker/local/create_worker.rs index d7f53d445..1fc20fc87 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/create_worker.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/create_worker.rs @@ -10,7 +10,7 @@ use crate::{ core::{ circuit_breaker::CircuitBreakerConfig, model_card::ModelCard, - steps::workflow_data::{AnyWorkflowData, LocalWorkerWorkflowData}, + steps::workflow_data::LocalWorkerWorkflowData, worker::{HealthConfig, RuntimeType, WorkerType}, BasicWorkerBuilder, ConnectionMode, DPAwareWorkerBuilder, Worker, UNKNOWN_MODEL_ID, }, @@ -29,22 +29,22 @@ use crate::{ pub struct CreateLocalWorkerStep; #[async_trait] -impl StepExecutor for CreateLocalWorkerStep { +impl StepExecutor for CreateLocalWorkerStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_local_worker()?; - let config = &data.config; - let app_context = data + let config = &context.data.config; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; - let connection_mode = data - .connection_mode - .as_ref() - .ok_or_else(|| WorkflowError::ContextValueNotFound("connection_mode".to_string()))?; - let discovered_labels = &data.discovered_labels; + let connection_mode = + context.data.connection_mode.as_ref().ok_or_else(|| { + WorkflowError::ContextValueNotFound("connection_mode".to_string()) + })?; + let discovered_labels = &context.data.discovered_labels; // Check if worker already exists if app_context @@ -100,7 +100,7 @@ impl StepExecutor for CreateLocalWorkerStep { let worker_type = parse_worker_type(config); // Get runtime type (for gRPC workers) - let runtime_type = determine_runtime_type(connection_mode, data, config); + let runtime_type = determine_runtime_type(connection_mode, &context.data, config); // Build circuit breaker config let circuit_breaker_config = build_circuit_breaker_config(app_context); @@ -121,7 +121,7 @@ impl StepExecutor for CreateLocalWorkerStep { // Create workers - always output as Vec for unified downstream handling let workers = if config.dp_aware { create_dp_aware_workers( - data, + &context.data, &normalized_url, model_card, worker_type, @@ -147,9 +147,8 @@ impl StepExecutor for CreateLocalWorkerStep { }; // Update workflow data - let data_mut = context.data.as_local_worker_mut()?; - data_mut.actual_workers = Some(workers); - data_mut.final_labels = final_labels; + context.data.actual_workers = Some(workers); + context.data.final_labels = final_labels; Ok(StepResult::Success) } diff --git a/sgl-model-gateway/src/core/steps/worker/local/detect_connection.rs b/sgl-model-gateway/src/core/steps/worker/local/detect_connection.rs index 58d31fdd6..9908de6a6 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/detect_connection.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/detect_connection.rs @@ -8,7 +8,7 @@ use tracing::debug; use super::strip_protocol; use crate::{ - core::{steps::workflow_data::AnyWorkflowData, ConnectionMode}, + core::{steps::workflow_data::LocalWorkerWorkflowData, ConnectionMode}, routers::grpc::client::GrpcClient, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, }; @@ -86,14 +86,14 @@ async fn try_grpc_health_check( pub struct DetectConnectionModeStep; #[async_trait] -impl StepExecutor for DetectConnectionModeStep { +impl StepExecutor for DetectConnectionModeStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_local_worker()?; - let config = &data.config; - let app_context = data + let config = &context.data.config; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; @@ -134,7 +134,7 @@ impl StepExecutor for DetectConnectionModeStep { } }; - context.data.as_local_worker_mut()?.connection_mode = Some(connection_mode); + context.data.connection_mode = Some(connection_mode); Ok(StepResult::Success) } diff --git a/sgl-model-gateway/src/core/steps/worker/local/discover_dp.rs b/sgl-model-gateway/src/core/steps/worker/local/discover_dp.rs index e5c3b7f9e..c2a52b32f 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/discover_dp.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/discover_dp.rs @@ -5,7 +5,7 @@ use tracing::debug; use super::discover_metadata::get_server_info; use crate::{ - core::{steps::workflow_data::AnyWorkflowData, UNKNOWN_MODEL_ID}, + core::{steps::workflow_data::LocalWorkerWorkflowData, UNKNOWN_MODEL_ID}, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, }; @@ -41,13 +41,12 @@ pub async fn get_dp_info(url: &str, api_key: Option<&str>) -> Result for DiscoverDPInfoStep { +impl StepExecutor for DiscoverDPInfoStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_local_worker()?; - let config = &data.config; + let config = &context.data.config; if !config.dp_aware { debug!( @@ -71,7 +70,7 @@ impl StepExecutor for DiscoverDPInfoStep { dp_info.dp_size, config.url, dp_info.model_id ); - context.data.as_local_worker_mut()?.dp_info = Some(dp_info); + context.data.dp_info = Some(dp_info); Ok(StepResult::Success) } diff --git a/sgl-model-gateway/src/core/steps/worker/local/discover_metadata.rs b/sgl-model-gateway/src/core/steps/worker/local/discover_metadata.rs index 501b28b11..6524d4295 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/discover_metadata.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/discover_metadata.rs @@ -11,7 +11,7 @@ use tracing::{debug, warn}; use super::strip_protocol; use crate::{ - core::{steps::workflow_data::AnyWorkflowData, ConnectionMode}, + core::{steps::workflow_data::LocalWorkerWorkflowData, ConnectionMode}, routers::grpc::client::GrpcClient, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, }; @@ -218,17 +218,16 @@ async fn fetch_grpc_metadata( pub struct DiscoverMetadataStep; #[async_trait] -impl StepExecutor for DiscoverMetadataStep { +impl StepExecutor for DiscoverMetadataStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_local_worker()?; - let config = &data.config; - let connection_mode = data - .connection_mode - .as_ref() - .ok_or_else(|| WorkflowError::ContextValueNotFound("connection_mode".to_string()))?; + let config = &context.data.config; + let connection_mode = + context.data.connection_mode.as_ref().ok_or_else(|| { + WorkflowError::ContextValueNotFound("connection_mode".to_string()) + })?; debug!( "Discovering metadata for {} ({:?})", @@ -301,11 +300,10 @@ impl StepExecutor for DiscoverMetadataStep { ); // Update workflow data - let data_mut = context.data.as_local_worker_mut()?; - data_mut.discovered_labels = discovered_labels; + context.data.discovered_labels = discovered_labels; if let Some(runtime) = detected_runtime { debug!("Detected runtime type: {}", runtime); - data_mut.detected_runtime_type = Some(runtime); + context.data.detected_runtime_type = Some(runtime); } Ok(StepResult::Success) diff --git a/sgl-model-gateway/src/core/steps/worker/local/find_worker_to_update.rs b/sgl-model-gateway/src/core/steps/worker/local/find_worker_to_update.rs index 2a811a944..1dc52f32c 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/find_worker_to_update.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/find_worker_to_update.rs @@ -5,7 +5,7 @@ use tracing::debug; use super::find_workers_by_url; use crate::{ - core::steps::workflow_data::AnyWorkflowData, + core::steps::workflow_data::WorkerUpdateWorkflowData, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, }; @@ -16,15 +16,15 @@ use crate::{ pub struct FindWorkerToUpdateStep; #[async_trait] -impl StepExecutor for FindWorkerToUpdateStep { +impl StepExecutor for FindWorkerToUpdateStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_worker_update()?; - let worker_url = &data.worker_url; - let dp_aware = data.dp_aware; - let app_context = data + let worker_url = &context.data.worker_url; + let dp_aware = context.data.dp_aware; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; @@ -50,7 +50,7 @@ impl StepExecutor for FindWorkerToUpdateStep { worker_url ); - context.data.as_worker_update_mut()?.workers_to_update = Some(workers_to_update); + context.data.workers_to_update = Some(workers_to_update); Ok(StepResult::Success) } diff --git a/sgl-model-gateway/src/core/steps/worker/local/find_workers_to_remove.rs b/sgl-model-gateway/src/core/steps/worker/local/find_workers_to_remove.rs index d627539d3..22635a1ea 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/find_workers_to_remove.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/find_workers_to_remove.rs @@ -7,7 +7,7 @@ use tracing::debug; use super::find_workers_by_url; use crate::{ - core::steps::workflow_data::{AnyWorkflowData, WorkerList}, + core::steps::workflow_data::{WorkerList, WorkerRemovalWorkflowData}, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, }; @@ -25,14 +25,14 @@ pub struct WorkerRemovalRequest { pub struct FindWorkersToRemoveStep; #[async_trait] -impl StepExecutor for FindWorkersToRemoveStep { +impl StepExecutor for FindWorkersToRemoveStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_worker_removal()?; - let request = &data.config; - let app_context = data + let request = &context.data.config; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; @@ -70,11 +70,10 @@ impl StepExecutor for FindWorkersToRemoveStep { .collect(); // Update workflow data - let data_mut = context.data.as_worker_removal_mut()?; - data_mut.workers_to_remove = Some(WorkerList::from_workers(&workers_to_remove)); - data_mut.actual_workers_to_remove = Some(workers_to_remove); - data_mut.worker_urls = worker_urls; - data_mut.affected_models = affected_models; + context.data.workers_to_remove = Some(WorkerList::from_workers(&workers_to_remove)); + context.data.actual_workers_to_remove = Some(workers_to_remove); + context.data.worker_urls = worker_urls; + context.data.affected_models = affected_models; Ok(StepResult::Success) } diff --git a/sgl-model-gateway/src/core/steps/worker/local/mod.rs b/sgl-model-gateway/src/core/steps/worker/local/mod.rs index 411252024..49ada6bdb 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/mod.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/mod.rs @@ -40,8 +40,7 @@ use crate::{ config::RouterConfig, core::{ steps::workflow_data::{ - AnyWorkflowData, LocalWorkerWorkflowData, WorkerRemovalWorkflowData, - WorkerUpdateWorkflowData, + LocalWorkerWorkflowData, WorkerRemovalWorkflowData, WorkerUpdateWorkflowData, }, Worker, WorkerRegistry, }, @@ -76,7 +75,7 @@ pub(crate) fn find_workers_by_url( pub fn create_local_worker_workflow( router_config: &RouterConfig, -) -> WorkflowDefinition { +) -> WorkflowDefinition { let detect_timeout = Duration::from_secs(router_config.worker_startup_timeout_secs); // Calculate max_attempts based on timeout @@ -208,7 +207,7 @@ pub fn create_local_worker_workflow( /// │ /// update_remaining_policies /// ``` -pub fn create_worker_removal_workflow() -> WorkflowDefinition { +pub fn create_worker_removal_workflow() -> WorkflowDefinition { WorkflowDefinition::new("worker_removal", "Remove worker from router") .add_step( StepDefinition::new( @@ -273,7 +272,7 @@ pub fn create_worker_removal_workflow() -> WorkflowDefinition { /// │ /// update_policies_for_worker /// ``` -pub fn create_worker_update_workflow() -> WorkflowDefinition { +pub fn create_worker_update_workflow() -> WorkflowDefinition { WorkflowDefinition::new("worker_update", "Update worker properties") .add_step( StepDefinition::new( @@ -319,8 +318,8 @@ pub fn create_worker_update_workflow() -> WorkflowDefinition { pub fn create_local_worker_workflow_data( config: WorkerConfigRequest, app_context: Arc, -) -> AnyWorkflowData { - AnyWorkflowData::LocalWorker(LocalWorkerWorkflowData { +) -> LocalWorkerWorkflowData { + LocalWorkerWorkflowData { config, connection_mode: None, discovered_labels: std::collections::HashMap::new(), @@ -330,7 +329,7 @@ pub fn create_local_worker_workflow_data( detected_runtime_type: None, app_context: Some(app_context), actual_workers: None, - }) + } } /// Helper to create initial workflow data for worker removal @@ -338,15 +337,15 @@ pub fn create_worker_removal_workflow_data( url: String, dp_aware: bool, app_context: Arc, -) -> AnyWorkflowData { - AnyWorkflowData::WorkerRemoval(WorkerRemovalWorkflowData { +) -> WorkerRemovalWorkflowData { + WorkerRemovalWorkflowData { config: WorkerRemovalRequest { url, dp_aware }, workers_to_remove: None, worker_urls: Vec::new(), affected_models: std::collections::HashSet::new(), app_context: Some(app_context), actual_workers_to_remove: None, - }) + } } /// Helper to create initial workflow data for worker update @@ -354,15 +353,15 @@ pub fn create_worker_update_workflow_data( worker_url: String, update_config: WorkerUpdateRequest, app_context: Arc, -) -> AnyWorkflowData { +) -> WorkerUpdateWorkflowData { // Determine if this is a DP-aware update based on URL pattern let dp_aware = worker_url.contains('@'); - AnyWorkflowData::WorkerUpdate(WorkerUpdateWorkflowData { + WorkerUpdateWorkflowData { config: update_config, worker_url, dp_aware, app_context: Some(app_context), workers_to_update: None, updated_workers: None, - }) + } } diff --git a/sgl-model-gateway/src/core/steps/worker/local/register_tokenizer.rs b/sgl-model-gateway/src/core/steps/worker/local/register_tokenizer.rs index 6b630cc64..a6c401e02 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/register_tokenizer.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/register_tokenizer.rs @@ -4,7 +4,7 @@ use async_trait::async_trait; use tracing::{debug, warn}; use crate::{ - core::steps::workflow_data::AnyWorkflowData, + core::steps::workflow_data::LocalWorkerWorkflowData, tokenizer::{factory, TokenizerRegistry}, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, }; @@ -13,18 +13,19 @@ use crate::{ pub struct RegisterTokenizerStep; #[async_trait] -impl StepExecutor for RegisterTokenizerStep { +impl StepExecutor for RegisterTokenizerStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_local_worker()?; - let labels = &data.final_labels; - let app_context = data + let labels = &context.data.final_labels; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; - let workers = data + let workers = context + .data .actual_workers .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("workers".to_string()))?; diff --git a/sgl-model-gateway/src/core/steps/worker/local/remove_from_policy_registry.rs b/sgl-model-gateway/src/core/steps/worker/local/remove_from_policy_registry.rs index 216cd5bdd..a68ab1758 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/remove_from_policy_registry.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/remove_from_policy_registry.rs @@ -4,7 +4,7 @@ use async_trait::async_trait; use tracing::debug; use crate::{ - core::steps::workflow_data::AnyWorkflowData, + core::steps::workflow_data::WorkerRemovalWorkflowData, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, }; @@ -15,17 +15,18 @@ use crate::{ pub struct RemoveFromPolicyRegistryStep; #[async_trait] -impl StepExecutor for RemoveFromPolicyRegistryStep { +impl StepExecutor for RemoveFromPolicyRegistryStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_worker_removal()?; - let app_context = data + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; - let workers_to_remove = data + let workers_to_remove = context + .data .actual_workers_to_remove .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("workers_to_remove".to_string()))?; diff --git a/sgl-model-gateway/src/core/steps/worker/local/remove_from_worker_registry.rs b/sgl-model-gateway/src/core/steps/worker/local/remove_from_worker_registry.rs index 57f04953f..808d995ed 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/remove_from_worker_registry.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/remove_from_worker_registry.rs @@ -6,7 +6,7 @@ use async_trait::async_trait; use tracing::{debug, warn}; use crate::{ - core::steps::workflow_data::AnyWorkflowData, + core::steps::workflow_data::WorkerRemovalWorkflowData, observability::metrics::Metrics, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, }; @@ -17,17 +17,17 @@ use crate::{ pub struct RemoveFromWorkerRegistryStep; #[async_trait] -impl StepExecutor for RemoveFromWorkerRegistryStep { +impl StepExecutor for RemoveFromWorkerRegistryStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_worker_removal()?; - let app_context = data + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; - let worker_urls = &data.worker_urls; + let worker_urls = &context.data.worker_urls; debug!( "Removing {} worker(s) from worker registry", diff --git a/sgl-model-gateway/src/core/steps/worker/local/update_policies_for_worker.rs b/sgl-model-gateway/src/core/steps/worker/local/update_policies_for_worker.rs index 511933d22..1ff926389 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/update_policies_for_worker.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/update_policies_for_worker.rs @@ -6,7 +6,7 @@ use async_trait::async_trait; use tracing::debug; use crate::{ - core::steps::workflow_data::AnyWorkflowData, + core::steps::workflow_data::WorkerUpdateWorkflowData, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, }; @@ -17,20 +17,20 @@ use crate::{ pub struct UpdatePoliciesForWorkerStep; #[async_trait] -impl StepExecutor for UpdatePoliciesForWorkerStep { +impl StepExecutor for UpdatePoliciesForWorkerStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_worker_update()?; - let app_context = data + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; - let updated_workers = data - .updated_workers - .as_ref() - .ok_or_else(|| WorkflowError::ContextValueNotFound("updated_workers".to_string()))?; + let updated_workers = + context.data.updated_workers.as_ref().ok_or_else(|| { + WorkflowError::ContextValueNotFound("updated_workers".to_string()) + })?; // Collect affected models let affected_models: HashSet = updated_workers diff --git a/sgl-model-gateway/src/core/steps/worker/local/update_remaining_policies.rs b/sgl-model-gateway/src/core/steps/worker/local/update_remaining_policies.rs index 3d95c7b21..0bdb10471 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/update_remaining_policies.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/update_remaining_policies.rs @@ -4,7 +4,7 @@ use async_trait::async_trait; use tracing::{debug, info}; use crate::{ - core::steps::workflow_data::AnyWorkflowData, + core::steps::workflow_data::WorkerRemovalWorkflowData, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, }; @@ -15,18 +15,18 @@ use crate::{ pub struct UpdateRemainingPoliciesStep; #[async_trait] -impl StepExecutor for UpdateRemainingPoliciesStep { +impl StepExecutor for UpdateRemainingPoliciesStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_worker_removal()?; - let app_context = data + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; - let affected_models = &data.affected_models; - let worker_urls = &data.worker_urls; + let affected_models = &context.data.affected_models; + let worker_urls = &context.data.worker_urls; debug!( "Updating cache-aware policies for {} affected model(s)", diff --git a/sgl-model-gateway/src/core/steps/worker/local/update_worker_properties.rs b/sgl-model-gateway/src/core/steps/worker/local/update_worker_properties.rs index b9ea8c123..d5b40bf9d 100644 --- a/sgl-model-gateway/src/core/steps/worker/local/update_worker_properties.rs +++ b/sgl-model-gateway/src/core/steps/worker/local/update_worker_properties.rs @@ -6,7 +6,9 @@ use async_trait::async_trait; use tracing::{debug, info}; use crate::{ - core::{steps::workflow_data::AnyWorkflowData, BasicWorkerBuilder, HealthConfig, Worker}, + core::{ + steps::workflow_data::WorkerUpdateWorkflowData, BasicWorkerBuilder, HealthConfig, Worker, + }, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, }; @@ -17,19 +19,20 @@ use crate::{ pub struct UpdateWorkerPropertiesStep; #[async_trait] -impl StepExecutor for UpdateWorkerPropertiesStep { +impl StepExecutor for UpdateWorkerPropertiesStep { async fn execute( &self, - context: &mut WorkflowContext, + context: &mut WorkflowContext, ) -> WorkflowResult { - let data = context.data.as_worker_update()?; - let request = &data.config; - let app_context = data + let request = &context.data.config; + let app_context = context + .data .app_context .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))? .clone(); - let workers_to_update = data + let workers_to_update = context + .data .workers_to_update .as_ref() .ok_or_else(|| WorkflowError::ContextValueNotFound("workers_to_update".to_string()))? @@ -137,7 +140,7 @@ impl StepExecutor for UpdateWorkerPropertiesStep { } // Store updated workers for subsequent steps - context.data.as_worker_update_mut()?.updated_workers = Some(updated_workers); + context.data.updated_workers = Some(updated_workers); Ok(StepResult::Success) } diff --git a/sgl-model-gateway/src/core/steps/worker/shared/activate.rs b/sgl-model-gateway/src/core/steps/worker/shared/activate.rs index a89c1287b..a56f355a2 100644 --- a/sgl-model-gateway/src/core/steps/worker/shared/activate.rs +++ b/sgl-model-gateway/src/core/steps/worker/shared/activate.rs @@ -4,21 +4,21 @@ use async_trait::async_trait; use tracing::info; use crate::{ - core::steps::workflow_data::AnyWorkflowData, - workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, + core::steps::workflow_data::WorkerRegistrationData, + workflow::{ + StepExecutor, StepResult, WorkflowContext, WorkflowData, WorkflowError, WorkflowResult, + }, }; /// Unified step to activate workers by marking them as healthy. /// /// This is the final step in any worker registration workflow. +/// Works with any workflow data type that implements `WorkerRegistrationData`. pub struct ActivateWorkersStep; #[async_trait] -impl StepExecutor for ActivateWorkersStep { - async fn execute( - &self, - context: &mut WorkflowContext, - ) -> WorkflowResult { +impl StepExecutor for ActivateWorkersStep { + async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult { let workers = context .data .get_actual_workers() diff --git a/sgl-model-gateway/src/core/steps/worker/shared/register.rs b/sgl-model-gateway/src/core/steps/worker/shared/register.rs index 7fa67ec6b..fe239bbb2 100644 --- a/sgl-model-gateway/src/core/steps/worker/shared/register.rs +++ b/sgl-model-gateway/src/core/steps/worker/shared/register.rs @@ -6,23 +6,23 @@ use async_trait::async_trait; use tracing::debug; use crate::{ - core::steps::workflow_data::AnyWorkflowData, + core::steps::workflow_data::WorkerRegistrationData, observability::metrics::Metrics, - workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, + workflow::{ + StepExecutor, StepResult, WorkflowContext, WorkflowData, WorkflowError, WorkflowResult, + }, }; /// Unified step to register workers in the registry. /// /// Works with both single workers and batches. Always expects `workers` key /// in context containing `Vec>`. +/// Works with any workflow data type that implements `WorkerRegistrationData`. pub struct RegisterWorkersStep; #[async_trait] -impl StepExecutor for RegisterWorkersStep { - async fn execute( - &self, - context: &mut WorkflowContext, - ) -> WorkflowResult { +impl StepExecutor for RegisterWorkersStep { + async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult { let app_context = context .data .get_app_context() diff --git a/sgl-model-gateway/src/core/steps/worker/shared/update_policies.rs b/sgl-model-gateway/src/core/steps/worker/shared/update_policies.rs index 4722d0866..c9527d73b 100644 --- a/sgl-model-gateway/src/core/steps/worker/shared/update_policies.rs +++ b/sgl-model-gateway/src/core/steps/worker/shared/update_policies.rs @@ -6,8 +6,10 @@ use async_trait::async_trait; use tracing::{debug, warn}; use crate::{ - core::{steps::workflow_data::AnyWorkflowData, Worker}, - workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, + core::{steps::workflow_data::WorkerRegistrationData, Worker}, + workflow::{ + StepExecutor, StepResult, WorkflowContext, WorkflowData, WorkflowError, WorkflowResult, + }, }; /// Unified step to update policy registry for registered workers. @@ -81,11 +83,8 @@ impl UpdatePoliciesStep { } #[async_trait] -impl StepExecutor for UpdatePoliciesStep { - async fn execute( - &self, - context: &mut WorkflowContext, - ) -> WorkflowResult { +impl StepExecutor for UpdatePoliciesStep { + async fn execute(&self, context: &mut WorkflowContext) -> WorkflowResult { let app_context = context .data .get_app_context() diff --git a/sgl-model-gateway/src/core/steps/workflow_data.rs b/sgl-model-gateway/src/core/steps/workflow_data.rs index 9874176a0..570a3904b 100644 --- a/sgl-model-gateway/src/core/steps/workflow_data.rs +++ b/sgl-model-gateway/src/core/steps/workflow_data.rs @@ -1,7 +1,14 @@ //! Typed workflow data structures //! //! This module defines the typed data structures for all workflows, enabling -//! compile-time type safety and state persistence. +//! compile-time type safety and state persistence. Each workflow has its own +//! strongly-typed data structure, and steps are typed to their specific workflow. +//! +//! # Shared Step Trait +//! +//! For steps that are shared between local and external worker workflows, +//! we use the `WorkerRegistrationData` trait. This trait provides a common +//! interface for accessing worker data while maintaining full type safety. use std::{collections::HashMap, sync::Arc}; @@ -26,6 +33,26 @@ use crate::{ workflow::{WorkflowData, WorkflowError}, }; +// ============================================================================ +// Shared trait for worker registration workflows +// ============================================================================ + +/// Trait for workflow data that supports worker registration operations. +/// +/// This trait is implemented by both `LocalWorkerWorkflowData` and +/// `ExternalWorkerWorkflowData`, allowing shared steps to work with either +/// workflow type while maintaining full type safety. +pub trait WorkerRegistrationData: WorkflowData { + /// Get the application context (transient, not serialized). + fn get_app_context(&self) -> Option<&Arc>; + + /// Get the actual worker objects (transient, not serialized). + fn get_actual_workers(&self) -> Option<&Vec>>; + + /// Get the labels for policy registration. + fn get_labels(&self) -> Option<&HashMap>; +} + /// Wrapper for worker list that can be serialized #[derive(Debug, Clone, Default, Serialize, Deserialize)] pub struct WorkerList { @@ -119,6 +146,20 @@ impl LocalWorkerWorkflowData { } } +impl WorkerRegistrationData for LocalWorkerWorkflowData { + fn get_app_context(&self) -> Option<&Arc> { + self.app_context.as_ref() + } + + fn get_actual_workers(&self) -> Option<&Vec>> { + self.actual_workers.as_ref() + } + + fn get_labels(&self) -> Option<&HashMap> { + Some(&self.final_labels) + } +} + /// Data for external worker registration workflow #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ExternalWorkerWorkflowData { @@ -154,6 +195,20 @@ impl ExternalWorkerWorkflowData { } } +impl WorkerRegistrationData for ExternalWorkerWorkflowData { + fn get_app_context(&self) -> Option<&Arc> { + self.app_context.as_ref() + } + + fn get_actual_workers(&self) -> Option<&Vec>> { + self.actual_workers.as_ref() + } + + fn get_labels(&self) -> Option<&HashMap> { + Some(&self.labels) + } +} + /// Data for worker removal workflow #[derive(Debug, Clone, Serialize, Deserialize)] pub struct WorkerRemovalWorkflowData { @@ -318,242 +373,3 @@ impl WasmRemovalWorkflowData { Ok(()) } } - -// ============================================================================ -// Unified enum for all workflow types -// ============================================================================ - -/// Macro to generate type-safe accessor methods for AnyWorkflowData variants. -/// -/// This reduces boilerplate and ensures consistent error handling across all accessors. -macro_rules! impl_workflow_accessor { - ($fn_name:ident, $fn_name_mut:ident, $variant:ident, $ty:ty, $type_name:expr) => { - /// Extract the inner data, returning an error if this is a different variant. - #[must_use = "this returns the result of the operation, without modifying the original"] - pub fn $fn_name(&self) -> Result<&$ty, WorkflowError> { - match self { - AnyWorkflowData::$variant(data) => Ok(data), - _ => Err(WorkflowError::TypeMismatch { - expected: $type_name, - actual: self.concrete_type(), - }), - } - } - - /// Extract the inner data mutably, returning an error if this is a different variant. - pub fn $fn_name_mut(&mut self) -> Result<&mut $ty, WorkflowError> { - // Store the type name before the mutable borrow - let actual = self.concrete_type(); - match self { - AnyWorkflowData::$variant(data) => Ok(data), - _ => Err(WorkflowError::TypeMismatch { - expected: $type_name, - actual, - }), - } - } - }; -} - -/// Macro to generate From implementations for AnyWorkflowData variants. -macro_rules! impl_from_workflow_data { - ($variant:ident, $ty:ty) => { - impl From<$ty> for AnyWorkflowData { - fn from(data: $ty) -> Self { - AnyWorkflowData::$variant(data) - } - } - }; -} - -/// Unified workflow data enum covering all workflow types. -/// -/// This allows a single `WorkflowEngine` to handle all workflows -/// while maintaining type safety at the step level. -/// -/// # Type Erasure -/// -/// `AnyWorkflowData` implements `WorkflowData` with `workflow_type()` returning `"any"`. -/// This is intentional: the static method cannot know the runtime variant. Use -/// [`concrete_type()`](Self::concrete_type) to get the actual workflow type at runtime. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub enum AnyWorkflowData { - Tokenizer(TokenizerWorkflowData), - LocalWorker(LocalWorkerWorkflowData), - ExternalWorker(ExternalWorkerWorkflowData), - WorkerRemoval(WorkerRemovalWorkflowData), - WorkerUpdate(WorkerUpdateWorkflowData), - Mcp(McpWorkflowData), - WasmRegistration(WasmRegistrationWorkflowData), - WasmRemoval(WasmRemovalWorkflowData), -} - -impl WorkflowData for AnyWorkflowData { - /// Returns `"any"` as this is a type-erased container. - /// - /// Use [`concrete_type()`](Self::concrete_type) to get the actual workflow type at runtime. - fn workflow_type() -> &'static str { - "any" - } -} - -// Generate From implementations for ergonomic construction -impl_from_workflow_data!(Tokenizer, TokenizerWorkflowData); -impl_from_workflow_data!(LocalWorker, LocalWorkerWorkflowData); -impl_from_workflow_data!(ExternalWorker, ExternalWorkerWorkflowData); -impl_from_workflow_data!(WorkerRemoval, WorkerRemovalWorkflowData); -impl_from_workflow_data!(WorkerUpdate, WorkerUpdateWorkflowData); -impl_from_workflow_data!(Mcp, McpWorkflowData); -impl_from_workflow_data!(WasmRegistration, WasmRegistrationWorkflowData); -impl_from_workflow_data!(WasmRemoval, WasmRemovalWorkflowData); - -impl AnyWorkflowData { - /// Get the concrete workflow type name at runtime. - /// - /// Unlike the static `workflow_type()` method, this returns the actual - /// type of the contained workflow data. - #[must_use] - pub fn concrete_type(&self) -> &'static str { - match self { - AnyWorkflowData::Tokenizer(_) => TokenizerWorkflowData::workflow_type(), - AnyWorkflowData::LocalWorker(_) => LocalWorkerWorkflowData::workflow_type(), - AnyWorkflowData::ExternalWorker(_) => ExternalWorkerWorkflowData::workflow_type(), - AnyWorkflowData::WorkerRemoval(_) => WorkerRemovalWorkflowData::workflow_type(), - AnyWorkflowData::WorkerUpdate(_) => WorkerUpdateWorkflowData::workflow_type(), - AnyWorkflowData::Mcp(_) => McpWorkflowData::workflow_type(), - AnyWorkflowData::WasmRegistration(_) => WasmRegistrationWorkflowData::workflow_type(), - AnyWorkflowData::WasmRemoval(_) => WasmRemovalWorkflowData::workflow_type(), - } - } - - // Generate all accessor methods using the macro - impl_workflow_accessor!( - as_tokenizer, - as_tokenizer_mut, - Tokenizer, - TokenizerWorkflowData, - "tokenizer_registration" - ); - impl_workflow_accessor!( - as_local_worker, - as_local_worker_mut, - LocalWorker, - LocalWorkerWorkflowData, - "local_worker_registration" - ); - impl_workflow_accessor!( - as_external_worker, - as_external_worker_mut, - ExternalWorker, - ExternalWorkerWorkflowData, - "external_worker_registration" - ); - impl_workflow_accessor!( - as_worker_removal, - as_worker_removal_mut, - WorkerRemoval, - WorkerRemovalWorkflowData, - "worker_removal" - ); - impl_workflow_accessor!( - as_worker_update, - as_worker_update_mut, - WorkerUpdate, - WorkerUpdateWorkflowData, - "worker_update" - ); - impl_workflow_accessor!(as_mcp, as_mcp_mut, Mcp, McpWorkflowData, "mcp_registration"); - impl_workflow_accessor!( - as_wasm_registration, - as_wasm_registration_mut, - WasmRegistration, - WasmRegistrationWorkflowData, - "wasm_module_registration" - ); - impl_workflow_accessor!( - as_wasm_removal, - as_wasm_removal_mut, - WasmRemoval, - WasmRemovalWorkflowData, - "wasm_module_removal" - ); - - // ======================================================================== - // Helper methods for shared worker steps - // ======================================================================== - - /// Get app_context from any workflow data type that has it. - #[must_use] - pub fn get_app_context(&self) -> Option<&Arc> { - match self { - AnyWorkflowData::Tokenizer(d) => d.app_context.as_ref(), - AnyWorkflowData::LocalWorker(d) => d.app_context.as_ref(), - AnyWorkflowData::ExternalWorker(d) => d.app_context.as_ref(), - AnyWorkflowData::WorkerRemoval(d) => d.app_context.as_ref(), - AnyWorkflowData::WorkerUpdate(d) => d.app_context.as_ref(), - AnyWorkflowData::Mcp(d) => d.app_context.as_ref(), - AnyWorkflowData::WasmRegistration(d) => d.app_context.as_ref(), - AnyWorkflowData::WasmRemoval(d) => d.app_context.as_ref(), - } - } - - /// Get actual workers from local or external worker workflows. - #[must_use] - pub fn get_actual_workers(&self) -> Option<&Vec>> { - match self { - AnyWorkflowData::LocalWorker(d) => d.actual_workers.as_ref(), - AnyWorkflowData::ExternalWorker(d) => d.actual_workers.as_ref(), - _ => None, - } - } - - /// Set actual workers for local or external worker workflows. - pub fn set_actual_workers( - &mut self, - workers: Vec>, - ) -> Result<(), WorkflowError> { - match self { - AnyWorkflowData::LocalWorker(d) => { - d.workers = Some(WorkerList::from_workers(&workers)); - d.actual_workers = Some(workers); - Ok(()) - } - AnyWorkflowData::ExternalWorker(d) => { - d.workers = Some(WorkerList::from_workers(&workers)); - d.actual_workers = Some(workers); - Ok(()) - } - _ => Err(WorkflowError::TypeMismatch { - expected: "LocalWorker or ExternalWorker", - actual: self.concrete_type(), - }), - } - } - - /// Get labels for policy configuration (from local or external worker workflows). - #[must_use] - pub fn get_labels(&self) -> Option<&HashMap> { - match self { - AnyWorkflowData::LocalWorker(d) => Some(&d.final_labels), - AnyWorkflowData::ExternalWorker(d) => Some(&d.labels), - _ => None, - } - } - - /// Validate that all transient fields are properly initialized. - /// - /// Call this after deserializing workflow state to ensure runtime fields - /// have been repopulated. - pub fn validate_initialized(&self) -> Result<(), WorkflowError> { - match self { - AnyWorkflowData::Tokenizer(d) => d.validate_initialized(), - AnyWorkflowData::LocalWorker(d) => d.validate_initialized(), - AnyWorkflowData::ExternalWorker(d) => d.validate_initialized(), - AnyWorkflowData::WorkerRemoval(d) => d.validate_initialized(), - AnyWorkflowData::WorkerUpdate(d) => d.validate_initialized(), - AnyWorkflowData::Mcp(d) => d.validate_initialized(), - AnyWorkflowData::WasmRegistration(d) => d.validate_initialized(), - AnyWorkflowData::WasmRemoval(d) => d.validate_initialized(), - } - } -} diff --git a/sgl-model-gateway/src/core/steps/workflow_engines.rs b/sgl-model-gateway/src/core/steps/workflow_engines.rs new file mode 100644 index 000000000..3c4088819 --- /dev/null +++ b/sgl-model-gateway/src/core/steps/workflow_engines.rs @@ -0,0 +1,167 @@ +//! Typed workflow engines collection +//! +//! This module provides a collection of typed workflow engines for different workflow types. +//! Each workflow type has its own engine with compile-time type safety. + +use std::sync::Arc; + +use super::{ + create_external_worker_workflow, create_local_worker_workflow, + create_mcp_registration_workflow, create_tokenizer_registration_workflow, + create_wasm_module_registration_workflow, create_wasm_module_removal_workflow, + create_worker_removal_workflow, create_worker_update_workflow, ExternalWorkerWorkflowData, + LocalWorkerWorkflowData, McpWorkflowData, TokenizerWorkflowData, WasmRegistrationWorkflowData, + WasmRemovalWorkflowData, WorkerRemovalWorkflowData, WorkerUpdateWorkflowData, +}; +use crate::{ + config::RouterConfig, + workflow::{EventSubscriber, InMemoryStore, WorkflowEngine}, +}; + +/// Type alias for local worker workflow engine +pub type LocalWorkerEngine = + WorkflowEngine>; + +/// Type alias for external worker workflow engine +pub type ExternalWorkerEngine = + WorkflowEngine>; + +/// Type alias for worker removal workflow engine +pub type WorkerRemovalEngine = + WorkflowEngine>; + +/// Type alias for worker update workflow engine +pub type WorkerUpdateEngine = + WorkflowEngine>; + +/// Type alias for MCP registration workflow engine +pub type McpEngine = WorkflowEngine>; + +/// Type alias for tokenizer registration workflow engine +pub type TokenizerEngine = + WorkflowEngine>; + +/// Type alias for WASM registration workflow engine +pub type WasmRegistrationEngine = + WorkflowEngine>; + +/// Type alias for WASM removal workflow engine +pub type WasmRemovalEngine = + WorkflowEngine>; + +/// Collection of typed workflow engines +/// +/// Each workflow type has its own engine with compile-time type safety. +/// This replaces the old `WorkflowEngine` approach. +#[derive(Clone, Debug)] +pub struct WorkflowEngines { + /// Engine for local worker registration workflows + pub local_worker: Arc, + /// Engine for external worker registration workflows + pub external_worker: Arc, + /// Engine for worker removal workflows + pub worker_removal: Arc, + /// Engine for worker update workflows + pub worker_update: Arc, + /// Engine for MCP server registration workflows + pub mcp: Arc, + /// Engine for tokenizer registration workflows + pub tokenizer: Arc, + /// Engine for WASM module registration workflows + pub wasm_registration: Arc, + /// Engine for WASM module removal workflows + pub wasm_removal: Arc, +} + +impl WorkflowEngines { + /// Create and initialize all workflow engines with their workflow definitions + pub fn new(router_config: &RouterConfig) -> Self { + // Create local worker engine + let local_worker = WorkflowEngine::new(); + local_worker + .register_workflow(create_local_worker_workflow(router_config)) + .expect("local_worker_registration workflow should be valid"); + + // Create external worker engine + let external_worker = WorkflowEngine::new(); + external_worker + .register_workflow(create_external_worker_workflow()) + .expect("external_worker_registration workflow should be valid"); + + // Create worker removal engine + let worker_removal = WorkflowEngine::new(); + worker_removal + .register_workflow(create_worker_removal_workflow()) + .expect("worker_removal workflow should be valid"); + + // Create worker update engine + let worker_update = WorkflowEngine::new(); + worker_update + .register_workflow(create_worker_update_workflow()) + .expect("worker_update workflow should be valid"); + + // Create MCP engine + let mcp = WorkflowEngine::new(); + mcp.register_workflow(create_mcp_registration_workflow()) + .expect("mcp_registration workflow should be valid"); + + // Create tokenizer engine + let tokenizer = WorkflowEngine::new(); + tokenizer + .register_workflow(create_tokenizer_registration_workflow()) + .expect("tokenizer_registration workflow should be valid"); + + // Create WASM registration engine + let wasm_registration = WorkflowEngine::new(); + wasm_registration + .register_workflow(create_wasm_module_registration_workflow()) + .expect("wasm_module_registration workflow should be valid"); + + // Create WASM removal engine + let wasm_removal = WorkflowEngine::new(); + wasm_removal + .register_workflow(create_wasm_module_removal_workflow()) + .expect("wasm_module_removal workflow should be valid"); + + Self { + local_worker: Arc::new(local_worker), + external_worker: Arc::new(external_worker), + worker_removal: Arc::new(worker_removal), + worker_update: Arc::new(worker_update), + mcp: Arc::new(mcp), + tokenizer: Arc::new(tokenizer), + wasm_registration: Arc::new(wasm_registration), + wasm_removal: Arc::new(wasm_removal), + } + } + + /// Subscribe an event subscriber to all workflow engines + pub async fn subscribe_all(&self, subscriber: Arc) { + self.local_worker + .event_bus() + .subscribe(subscriber.clone()) + .await; + self.external_worker + .event_bus() + .subscribe(subscriber.clone()) + .await; + self.worker_removal + .event_bus() + .subscribe(subscriber.clone()) + .await; + self.worker_update + .event_bus() + .subscribe(subscriber.clone()) + .await; + self.mcp.event_bus().subscribe(subscriber.clone()).await; + self.tokenizer + .event_bus() + .subscribe(subscriber.clone()) + .await; + self.wasm_registration + .event_bus() + .subscribe(subscriber.clone()) + .await; + self.wasm_removal.event_bus().subscribe(subscriber).await; + } +} diff --git a/sgl-model-gateway/src/server.rs b/sgl-model-gateway/src/server.rs index 5b855996e..fdc893486 100644 --- a/sgl-model-gateway/src/server.rs +++ b/sgl-model-gateway/src/server.rs @@ -20,16 +20,11 @@ use tokio::{signal, spawn}; use tracing::{debug, error, info, warn, Level}; use crate::{ - app_context::{AppContext, AppWorkflowEngine}, + app_context::AppContext, config::{RouterConfig, RoutingMode}, core::{ job_queue::{JobQueue, JobQueueConfig}, - steps::{ - create_external_worker_workflow, create_local_worker_workflow, - create_mcp_registration_workflow, create_tokenizer_registration_workflow, - create_wasm_module_registration_workflow, create_wasm_module_removal_workflow, - create_worker_removal_workflow, create_worker_update_workflow, - }, + steps::WorkflowEngines, worker::WorkerType, worker_manager::WorkerManager, Job, @@ -730,44 +725,18 @@ pub async fn startup(config: ServerConfig) -> Result<(), Box + 'static> WorkflowEngine { self.state_store.load(instance_id) } + /// Wait for a workflow to complete with adaptive polling + /// + /// Returns Ok with success message on completion, Err on failure/timeout/cancellation. + /// Automatically cleans up terminal workflow states. + pub async fn wait_for_completion( + &self, + instance_id: WorkflowInstanceId, + label: &str, + timeout_duration: Duration, + ) -> Result { + let start = std::time::Instant::now(); + let mut poll_interval = Duration::from_millis(100); + let max_poll_interval = Duration::from_millis(2000); + let poll_backoff = Duration::from_millis(200); + + loop { + if start.elapsed() > timeout_duration { + return Err(format!( + "Workflow timeout after {}s for {}", + timeout_duration.as_secs(), + label + )); + } + + let state = self + .get_status(instance_id) + .map_err(|e| format!("Failed to get workflow status: {:?}", e))?; + + let result = match state.status { + WorkflowStatus::Completed => { + Ok(format!("{} completed successfully via workflow", label)) + } + WorkflowStatus::Failed => { + let current_step = state.current_step.as_ref(); + let step_name = current_step + .map(|s| s.to_string()) + .unwrap_or_else(|| "unknown".to_string()); + let error_msg = current_step + .and_then(|step_id| state.step_states.get(step_id)) + .and_then(|s| s.last_error.as_deref()) + .unwrap_or("Unknown error"); + Err(format!( + "Workflow failed at step {}: {}", + step_name, error_msg + )) + } + WorkflowStatus::Cancelled => Err(format!("Workflow cancelled for {}", label)), + WorkflowStatus::Pending | WorkflowStatus::Paused | WorkflowStatus::Running => { + tokio::time::sleep(poll_interval).await; + poll_interval = (poll_interval + poll_backoff).min(max_poll_interval); + continue; + } + }; + + self.state_store.cleanup_if_terminal(instance_id); + return result; + } + } + /// Clone engine for async execution fn clone_for_execution(&self) -> Self { Self { diff --git a/sgl-model-gateway/src/workflow/state.rs b/sgl-model-gateway/src/workflow/state.rs index de548bb1e..5a2d898a0 100644 --- a/sgl-model-gateway/src/workflow/state.rs +++ b/sgl-model-gateway/src/workflow/state.rs @@ -42,6 +42,10 @@ pub trait StateStore: Send + Sync + Clone { /// Get just the workflow context without cloning the entire state fn get_context(&self, instance_id: WorkflowInstanceId) -> WorkflowResult>; + + /// Clean up a specific workflow immediately if it's in a terminal state + /// Returns true if the workflow was removed, false otherwise + fn cleanup_if_terminal(&self, instance_id: WorkflowInstanceId) -> bool; } /// In-memory state storage for workflow instances @@ -72,21 +76,6 @@ impl InMemoryStore { pub fn count(&self) -> usize { self.states.read().len() } - - /// Clean up a specific completed workflow immediately - pub fn cleanup_if_terminal(&self, instance_id: WorkflowInstanceId) -> bool { - let mut states = self.states.write(); - if let Some(state) = states.get(&instance_id) { - if matches!( - state.status, - WorkflowStatus::Completed | WorkflowStatus::Failed | WorkflowStatus::Cancelled - ) { - states.remove(&instance_id); - return true; - } - } - false - } } impl Default for InMemoryStore { @@ -189,4 +178,18 @@ impl StateStore for InMemoryStore { } removed_count } + + fn cleanup_if_terminal(&self, instance_id: WorkflowInstanceId) -> bool { + let mut states = self.states.write(); + if let Some(state) = states.get(&instance_id) { + if matches!( + state.status, + WorkflowStatus::Completed | WorkflowStatus::Failed | WorkflowStatus::Cancelled + ) { + states.remove(&instance_id); + return true; + } + } + false + } } diff --git a/sgl-model-gateway/tests/common/mod.rs b/sgl-model-gateway/tests/common/mod.rs index 09f7df33a..0c56c4435 100644 --- a/sgl-model-gateway/tests/common/mod.rs +++ b/sgl-model-gateway/tests/common/mod.rs @@ -322,9 +322,9 @@ pub async fn create_test_context(config: RouterConfig) -> Arc { config.worker_startup_check_interval_secs, ))); - // Create empty OnceLock for worker job queue, workflow engine, and mcp manager + // Create empty OnceLock for worker job queue, workflow engines, and mcp manager let worker_job_queue = Arc::new(OnceLock::new()); - let workflow_engine = Arc::new(OnceLock::new()); + let workflow_engines = Arc::new(OnceLock::new()); let mcp_manager_lock = Arc::new(OnceLock::new()); let app_context = Arc::new( @@ -342,7 +342,7 @@ pub async fn create_test_context(config: RouterConfig) -> Arc { .conversation_item_storage(conversation_item_storage) .load_monitor(load_monitor) .worker_job_queue(worker_job_queue) - .workflow_engine(workflow_engine) + .workflow_engines(workflow_engines) .mcp_manager(mcp_manager_lock) .build() .unwrap(), @@ -356,22 +356,13 @@ pub async fn create_test_context(config: RouterConfig) -> Arc { .set(job_queue) .expect("JobQueue should only be initialized once"); - // Initialize WorkflowEngine and register workflows - use smg::{ - core::steps::{create_local_worker_workflow, create_worker_removal_workflow}, - workflow::WorkflowEngine, - }; - let engine = Arc::new(WorkflowEngine::new()); - engine - .register_workflow(create_local_worker_workflow(&config)) - .expect("worker_registration workflow should be valid"); - engine - .register_workflow(create_worker_removal_workflow()) - .expect("worker_removal workflow should be valid"); + // Initialize typed workflow engines + use smg::core::steps::WorkflowEngines; + let engines = WorkflowEngines::new(&config); app_context - .workflow_engine - .set(engine) - .expect("WorkflowEngine should only be initialized once"); + .workflow_engines + .set(engines) + .expect("WorkflowEngines should only be initialized once"); // Register external workers for OpenAI mode if let RoutingMode::OpenAI { worker_urls, .. } = &config.mode { @@ -451,9 +442,9 @@ pub async fn create_test_context_with_parsers(config: RouterConfig) -> Arc Arc Arc Arc { // Initialize empty OnceLocks let worker_job_queue = Arc::new(OnceLock::new()); - let workflow_engine = Arc::new(OnceLock::new()); + let workflow_engines = Arc::new(OnceLock::new()); // Initialize MCP manager with empty config let mcp_manager_lock = Arc::new(OnceLock::new()); @@ -208,7 +208,7 @@ pub async fn create_test_app_context() -> Arc { .conversation_item_storage(conversation_item_storage) .load_monitor(None) .worker_job_queue(worker_job_queue) - .workflow_engine(workflow_engine) + .workflow_engines(workflow_engines) .mcp_manager(mcp_manager_lock) .build() .unwrap(), diff --git a/sgl-model-gateway/tests/routing/test_pd_routing.rs b/sgl-model-gateway/tests/routing/test_pd_routing.rs index 85fc29aab..b5a01f675 100644 --- a/sgl-model-gateway/tests/routing/test_pd_routing.rs +++ b/sgl-model-gateway/tests/routing/test_pd_routing.rs @@ -247,9 +247,9 @@ mod pd_routing_unit_tests { config.worker_startup_check_interval_secs, ))); - // Create empty OnceLock for worker job queue, workflow engine, and mcp manager + // Create empty OnceLock for worker job queue, workflow engines, and mcp manager let worker_job_queue = Arc::new(OnceLock::new()); - let workflow_engine = Arc::new(OnceLock::new()); + let workflow_engines = Arc::new(OnceLock::new()); let mcp_manager = Arc::new(OnceLock::new()); Arc::new( @@ -267,7 +267,7 @@ mod pd_routing_unit_tests { .conversation_item_storage(conversation_item_storage) .load_monitor(load_monitor) .worker_job_queue(worker_job_queue) - .workflow_engine(workflow_engine) + .workflow_engines(workflow_engines) .mcp_manager(mcp_manager) .build() .unwrap(), diff --git a/sgl-model-gateway/tests/wasm_test.rs b/sgl-model-gateway/tests/wasm_test.rs index 21d15d2c3..09e1d6fda 100644 --- a/sgl-model-gateway/tests/wasm_test.rs +++ b/sgl-model-gateway/tests/wasm_test.rs @@ -18,10 +18,7 @@ use axum::{ use smg::{ app_context::AppContext, config::RouterConfig, - core::{ - steps::{create_wasm_module_registration_workflow, create_wasm_module_removal_workflow}, - LoadMonitor, WorkerRegistry, - }, + core::{LoadMonitor, WorkerRegistry}, data_connector::{ MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage, }, @@ -71,10 +68,10 @@ async fn create_test_context_with_wasm() -> Arc { config.worker_startup_check_interval_secs, ))); - // Create empty OnceLock for worker job queue, workflow engine, and mcp manager + // Create empty OnceLock for worker job queue, workflow engines, and mcp manager use std::sync::OnceLock; let worker_job_queue = Arc::new(OnceLock::new()); - let workflow_engine = Arc::new(OnceLock::new()); + let workflow_engines = Arc::new(OnceLock::new()); let mcp_manager_lock = Arc::new(OnceLock::new()); let app_context = Arc::new( @@ -92,7 +89,7 @@ async fn create_test_context_with_wasm() -> Arc { .conversation_item_storage(conversation_item_storage) .load_monitor(load_monitor) .worker_job_queue(worker_job_queue) - .workflow_engine(workflow_engine) + .workflow_engines(workflow_engines) .mcp_manager(mcp_manager_lock) .wasm_manager(Some(wasm_manager)) .build() @@ -107,28 +104,13 @@ async fn create_test_context_with_wasm() -> Arc { .set(job_queue) .expect("JobQueue should only be initialized once"); - // Initialize WorkflowEngine and register workflows - use smg::{ - core::steps::{create_local_worker_workflow, create_worker_removal_workflow}, - workflow::WorkflowEngine, - }; - let engine = Arc::new(WorkflowEngine::new()); - engine - .register_workflow(create_local_worker_workflow(&config)) - .expect("worker_registration workflow should be valid"); - engine - .register_workflow(create_worker_removal_workflow()) - .expect("worker_removal workflow should be valid"); - engine - .register_workflow(create_wasm_module_registration_workflow()) - .expect("wasm_module_registration workflow should be valid"); - engine - .register_workflow(create_wasm_module_removal_workflow()) - .expect("wasm_module_removal workflow should be valid"); + // Initialize WorkflowEngines + use smg::core::steps::WorkflowEngines; + let engines = WorkflowEngines::new(&config); app_context - .workflow_engine - .set(engine) - .expect("WorkflowEngine should only be initialized once"); + .workflow_engines + .set(engines) + .expect("WorkflowEngines should only be initialized once"); // Initialize MCP manager with empty config use smg::mcp::{McpConfig, McpManager}; @@ -678,17 +660,14 @@ async fn test_wasm_module_execution() { .as_ref() .expect("WASM manager should be initialized"); - let engine = app_context - .workflow_engine + let engines = app_context + .workflow_engines .get() - .expect("Workflow engine should be initialized"); + .expect("Workflow engines should be initialized"); // Create workflow context for registration use smg::{ - core::steps::{ - workflow_data::{AnyWorkflowData, WasmRegistrationWorkflowData}, - WasmModuleConfigRequest, - }, + core::steps::{WasmModuleConfigRequest, WasmRegistrationWorkflowData}, workflow::WorkflowId, }; @@ -703,17 +682,18 @@ async fn test_wasm_module_execution() { }; let config_request = WasmModuleConfigRequest { descriptor }; - let workflow_data = AnyWorkflowData::WasmRegistration(WasmRegistrationWorkflowData { + let workflow_data = WasmRegistrationWorkflowData { config: config_request, wasm_bytes: None, sha256_hash: None, file_size_bytes: None, module_uuid: None, app_context: Some(app_context.clone()), - }); + }; // Start workflow - let instance_id = engine + let instance_id = engines + .wasm_registration .start_workflow(WorkflowId::new("wasm_module_registration"), workflow_data) .await .expect("Failed to start workflow"); @@ -721,24 +701,25 @@ async fn test_wasm_module_execution() { // Wait for workflow to complete let timeout = Duration::from_secs(30); let start = std::time::Instant::now(); - let mut module_uuid: Option = None; - loop { + let module_uuid = loop { if start.elapsed() > timeout { panic!("Workflow timeout"); } - let state = engine + let state = engines + .wasm_registration .get_status(instance_id) .expect("Failed to get workflow status"); match state.status { smg::workflow::WorkflowStatus::Completed => { // Extract module UUID from typed workflow data - if let AnyWorkflowData::WasmRegistration(ref data) = state.context.data { - module_uuid = data.module_uuid; - } - break; + break state + .context + .data + .module_uuid + .expect("Module UUID should be in context"); } smg::workflow::WorkflowStatus::Failed => { panic!("Workflow failed: {:?}", state); @@ -747,9 +728,7 @@ async fn test_wasm_module_execution() { tokio::time::sleep(Duration::from_millis(100)).await; } } - } - - let module_uuid = module_uuid.expect("Module UUID should be in context"); + }; // Verify module is registered let module = wasm_manager