[model-gateway] refactor workflow engine from type erasure to typed engines (#16973)

This commit is contained in:
Simo Lin
2026-01-12 10:47:00 -08:00
committed by GitHub
parent fa51b85466
commit ed729d22b3
36 changed files with 803 additions and 917 deletions
+14 -21
View File
@@ -9,8 +9,8 @@ use tracing::{debug, info};
use crate::{ use crate::{
config::RouterConfig, config::RouterConfig,
core::{ core::{
steps::workflow_data::AnyWorkflowData, JobQueue, LoadMonitor, WorkerRegistry, steps::WorkflowEngines, JobQueue, LoadMonitor, WorkerRegistry, WorkerService,
WorkerService, UNKNOWN_MODEL_ID, UNKNOWN_MODEL_ID,
}, },
data_connector::{ data_connector::{
create_storage, ConversationItemStorage, ConversationStorage, ResponseStorage, create_storage, ConversationItemStorage, ConversationStorage, ResponseStorage,
@@ -29,12 +29,8 @@ use crate::{
}, },
tool_parser::ParserFactory as ToolParserFactory, tool_parser::ParserFactory as ToolParserFactory,
wasm::{config::WasmRuntimeConfig, module_manager::WasmModuleManager}, wasm::{config::WasmRuntimeConfig, module_manager::WasmModuleManager},
workflow::{InMemoryStore, WorkflowEngine},
}; };
/// Type alias for the concrete workflow engine used in the application
pub type AppWorkflowEngine = WorkflowEngine<AnyWorkflowData, InMemoryStore<AnyWorkflowData>>;
/// Error type for AppContext builder /// Error type for AppContext builder
#[derive(Debug)] #[derive(Debug)]
pub struct AppContextBuildError(&'static str); pub struct AppContextBuildError(&'static str);
@@ -65,7 +61,7 @@ pub struct AppContext {
pub configured_reasoning_parser: Option<String>, pub configured_reasoning_parser: Option<String>,
pub configured_tool_parser: Option<String>, pub configured_tool_parser: Option<String>,
pub worker_job_queue: Arc<OnceLock<Arc<JobQueue>>>, pub worker_job_queue: Arc<OnceLock<Arc<JobQueue>>>,
pub workflow_engine: Arc<OnceLock<Arc<AppWorkflowEngine>>>, pub workflow_engines: Arc<OnceLock<WorkflowEngines>>,
pub mcp_manager: Arc<OnceLock<Arc<McpManager>>>, pub mcp_manager: Arc<OnceLock<Arc<McpManager>>>,
pub wasm_manager: Option<Arc<WasmModuleManager>>, pub wasm_manager: Option<Arc<WasmModuleManager>>,
pub worker_service: Arc<WorkerService>, pub worker_service: Arc<WorkerService>,
@@ -95,7 +91,7 @@ pub struct AppContextBuilder {
conversation_item_storage: Option<Arc<dyn ConversationItemStorage>>, conversation_item_storage: Option<Arc<dyn ConversationItemStorage>>,
load_monitor: Option<Arc<LoadMonitor>>, load_monitor: Option<Arc<LoadMonitor>>,
worker_job_queue: Option<Arc<OnceLock<Arc<JobQueue>>>>, worker_job_queue: Option<Arc<OnceLock<Arc<JobQueue>>>>,
workflow_engine: Option<Arc<OnceLock<Arc<AppWorkflowEngine>>>>, workflow_engines: Option<Arc<OnceLock<WorkflowEngines>>>,
mcp_manager: Option<Arc<OnceLock<Arc<McpManager>>>>, mcp_manager: Option<Arc<OnceLock<Arc<McpManager>>>>,
wasm_manager: Option<Arc<WasmModuleManager>>, wasm_manager: Option<Arc<WasmModuleManager>>,
} }
@@ -135,7 +131,7 @@ impl AppContextBuilder {
conversation_item_storage: None, conversation_item_storage: None,
load_monitor: None, load_monitor: None,
worker_job_queue: None, worker_job_queue: None,
workflow_engine: None, workflow_engines: None,
mcp_manager: None, mcp_manager: None,
wasm_manager: None, wasm_manager: None,
} }
@@ -220,11 +216,8 @@ impl AppContextBuilder {
self self
} }
pub fn workflow_engine( pub fn workflow_engines(mut self, workflow_engines: Arc<OnceLock<WorkflowEngines>>) -> Self {
mut self, self.workflow_engines = Some(workflow_engines);
workflow_engine: Arc<OnceLock<Arc<AppWorkflowEngine>>>,
) -> Self {
self.workflow_engine = Some(workflow_engine);
self self
} }
@@ -286,9 +279,9 @@ impl AppContextBuilder {
configured_reasoning_parser, configured_reasoning_parser,
configured_tool_parser, configured_tool_parser,
worker_job_queue, worker_job_queue,
workflow_engine: self workflow_engines: self
.workflow_engine .workflow_engines
.ok_or(AppContextBuildError("workflow_engine"))?, .ok_or(AppContextBuildError("workflow_engines"))?,
mcp_manager: self mcp_manager: self
.mcp_manager .mcp_manager
.ok_or(AppContextBuildError("mcp_manager"))?, .ok_or(AppContextBuildError("mcp_manager"))?,
@@ -315,7 +308,7 @@ impl AppContextBuilder {
.with_storage(&router_config)? .with_storage(&router_config)?
.with_load_monitor(&router_config) .with_load_monitor(&router_config)
.with_worker_job_queue() .with_worker_job_queue()
.with_workflow_engine() .with_workflow_engines()
.with_mcp_manager(&router_config) .with_mcp_manager(&router_config)
.await? .await?
.with_wasm_manager(&router_config)? .with_wasm_manager(&router_config)?
@@ -549,9 +542,9 @@ impl AppContextBuilder {
self self
} }
/// Create workflow engine OnceLock container /// Create workflow engines OnceLock container
fn with_workflow_engine(mut self) -> Self { fn with_workflow_engines(mut self) -> Self {
self.workflow_engine = Some(Arc::new(OnceLock::new())); self.workflow_engines = Some(Arc::new(OnceLock::new()));
self self
} }
+137 -203
View File
@@ -14,7 +14,7 @@ use tokio::sync::{mpsc, Semaphore};
use tracing::{debug, error, info, warn}; use tracing::{debug, error, info, warn};
use crate::{ use crate::{
app_context::{AppContext, AppWorkflowEngine}, app_context::AppContext,
config::{RouterConfig, RoutingMode}, config::{RouterConfig, RoutingMode},
core::steps::{ core::steps::{
create_external_worker_workflow_data, create_local_worker_workflow_data, create_external_worker_workflow_data, create_local_worker_workflow_data,
@@ -26,7 +26,7 @@ use crate::{
}, },
mcp::McpConfig, mcp::McpConfig,
protocols::worker_spec::{JobStatus, WorkerConfigRequest, WorkerUpdateRequest}, protocols::worker_spec::{JobStatus, WorkerConfigRequest, WorkerUpdateRequest},
workflow::{WorkflowId, WorkflowInstanceId, WorkflowStatus}, workflow::WorkflowId,
}; };
/// Job types for control plane operations /// Job types for control plane operations
@@ -335,37 +335,93 @@ impl JobQueue {
async fn execute_job(job: &Job, context: &Arc<AppContext>) -> Result<String, String> { async fn execute_job(job: &Job, context: &Arc<AppContext>) -> Result<String, String> {
match job { match job {
Job::AddWorker { config } => { Job::AddWorker { config } => {
let engine = context let engines = context
.workflow_engine .workflow_engines
.get() .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_workflow(engine, config, context).await?;
debug!(
"Started worker registration workflow for {} (instance: {})",
config.url, instance_id
);
let timeout_duration = let timeout_duration =
Duration::from_secs(context.router_config.worker_startup_timeout_secs + 30); Duration::from_secs(context.router_config.worker_startup_timeout_secs + 30);
Self::wait_for_workflow_completion( // Select workflow based on runtime field
engine, match config.runtime.as_deref() {
instance_id, Some("external") => {
&config.url, let workflow_data = create_external_worker_workflow_data(
timeout_duration, (**config).clone(),
Arc::clone(context),
);
let instance_id = engines
.external_worker
.start_workflow(
WorkflowId::new("external_worker_registration"),
workflow_data,
) )
.await .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 } => { Job::UpdateWorker { url, update } => {
let engine = context let engines = context
.workflow_engine .workflow_engines
.get() .get()
.ok_or_else(|| "Workflow engine not initialized".to_string())?; .ok_or_else(|| "Workflow engines not initialized".to_string())?;
let instance_id = let workflow_data = create_worker_update_workflow_data(
Self::start_worker_update_workflow(engine, url, update, context).await?; 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!( debug!(
"Started worker update workflow for {} (instance: {})", "Started worker update workflow for {} (instance: {})",
@@ -374,15 +430,28 @@ impl JobQueue {
let timeout_duration = Duration::from_secs(30); 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 } => { Job::RemoveWorker { url } => {
let engine = context let engines = context
.workflow_engine .workflow_engines
.get() .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!( debug!(
"Started worker removal workflow for {} (instance: {})", "Started worker removal workflow for {} (instance: {})",
@@ -391,8 +460,9 @@ impl JobQueue {
let timeout_duration = Duration::from_secs(30); let timeout_duration = Duration::from_secs(30);
let result = let result = engines
Self::wait_for_workflow_completion(engine, instance_id, url, timeout_duration) .worker_removal
.wait_for_completion(instance_id, url, timeout_duration)
.await; .await;
// Clean up job status when removing worker // Clean up job status when removing worker
@@ -403,15 +473,16 @@ impl JobQueue {
result result
} }
Job::AddWasmModule { config } => { Job::AddWasmModule { config } => {
let engine = context let engines = context
.workflow_engine .workflow_engines
.get() .get()
.ok_or_else(|| "Workflow engine not initialized".to_string())?; .ok_or_else(|| "Workflow engines not initialized".to_string())?;
let workflow_data = let workflow_data =
create_wasm_registration_workflow_data(*config.clone(), Arc::clone(context)); 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) .start_workflow(WorkflowId::new("wasm_module_registration"), workflow_data)
.await .await
.map_err(|e| { .map_err(|e| {
@@ -425,24 +496,22 @@ impl JobQueue {
let timeout_duration = Duration::from_secs(300); // 5 minutes let timeout_duration = Duration::from_secs(300); // 5 minutes
Self::wait_for_workflow_completion( engines
engine, .wasm_registration
instance_id, .wait_for_completion(instance_id, &config.descriptor.name, timeout_duration)
&config.descriptor.name,
timeout_duration,
)
.await .await
} }
Job::RemoveWasmModule { request } => { Job::RemoveWasmModule { request } => {
let engine = context let engines = context
.workflow_engine .workflow_engines
.get() .get()
.ok_or_else(|| "Workflow engine not initialized".to_string())?; .ok_or_else(|| "Workflow engines not initialized".to_string())?;
let workflow_data = let workflow_data =
create_wasm_removal_workflow_data(*request.clone(), Arc::clone(context)); 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) .start_workflow(WorkflowId::new("wasm_module_removal"), workflow_data)
.await .await
.map_err(|e| { .map_err(|e| {
@@ -456,8 +525,9 @@ impl JobQueue {
let timeout_duration = Duration::from_secs(60); // 1 minute let timeout_duration = Duration::from_secs(60); // 1 minute
Self::wait_for_workflow_completion( engines
engine, .wasm_removal
.wait_for_completion(
instance_id, instance_id,
&request.module_uuid.to_string(), &request.module_uuid.to_string(),
timeout_duration, timeout_duration,
@@ -640,13 +710,19 @@ impl JobQueue {
Ok(format!("Submitted {} RegisterMcpServer jobs", server_count)) Ok(format!("Submitted {} RegisterMcpServer jobs", server_count))
} }
Job::RegisterMcpServer { config } => { Job::RegisterMcpServer { config } => {
let engine = context let engines = context
.workflow_engine .workflow_engines
.get() .get()
.ok_or_else(|| "Workflow engine not initialized".to_string())?; .ok_or_else(|| "Workflow engines not initialized".to_string())?;
let instance_id = let workflow_data =
Self::start_mcp_registration_workflow(engine, config, context).await?; 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!( debug!(
"Started MCP registration workflow for {} (instance: {})", "Started MCP registration workflow for {} (instance: {})",
@@ -655,24 +731,22 @@ impl JobQueue {
let timeout_duration = Duration::from_secs(7200 + 30); // 2hr + margin let timeout_duration = Duration::from_secs(7200 + 30); // 2hr + margin
Self::wait_for_workflow_completion( engines
engine, .mcp
instance_id, .wait_for_completion(instance_id, &config.name, timeout_duration)
&config.name,
timeout_duration,
)
.await .await
} }
Job::AddTokenizer { config } => { Job::AddTokenizer { config } => {
let engine = context let engines = context
.workflow_engine .workflow_engines
.get() .get()
.ok_or_else(|| "Workflow engine not initialized".to_string())?; .ok_or_else(|| "Workflow engines not initialized".to_string())?;
let workflow_data = let workflow_data =
create_tokenizer_workflow_data(*config.clone(), Arc::clone(context)); 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) .start_workflow(WorkflowId::new("tokenizer_registration"), workflow_data)
.await .await
.map_err(|e| { .map_err(|e| {
@@ -687,12 +761,9 @@ impl JobQueue {
// Allow up to 10 minutes for HuggingFace downloads // Allow up to 10 minutes for HuggingFace downloads
let timeout_duration = Duration::from_secs(600); let timeout_duration = Duration::from_secs(600);
Self::wait_for_workflow_completion( engines
engine, .tokenizer
instance_id, .wait_for_completion(instance_id, &config.id, timeout_duration)
&config.id,
timeout_duration,
)
.await .await
} }
Job::RemoveTokenizer { request } => { Job::RemoveTokenizer { request } => {
@@ -710,143 +781,6 @@ impl JobQueue {
} }
} }
/// Start a workflow and return its instance ID
async fn start_worker_workflow(
engine: &Arc<AppWorkflowEngine>,
config: &WorkerConfigRequest,
context: &Arc<AppContext>,
) -> Result<WorkflowInstanceId, String> {
// 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<AppWorkflowEngine>,
url: &str,
context: &Arc<AppContext>,
) -> Result<WorkflowInstanceId, String> {
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<AppWorkflowEngine>,
url: &str,
update: &WorkerUpdateRequest,
context: &Arc<AppContext>,
) -> Result<WorkflowInstanceId, String> {
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<AppWorkflowEngine>,
config: &McpServerConfigRequest,
context: &Arc<AppContext>,
) -> Result<WorkflowInstanceId, String> {
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<AppWorkflowEngine>,
instance_id: WorkflowInstanceId,
worker_url: &str,
timeout_duration: Duration,
) -> Result<String, String> {
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 /// Update job status on completion
fn record_job_completion( fn record_job_completion(
job_type: &'static str, job_type: &'static str,
@@ -3,7 +3,7 @@ use std::{sync::Arc, time::Duration};
use async_trait::async_trait; use async_trait::async_trait;
use tracing::{debug, error, info, warn}; use tracing::{debug, error, info, warn};
use super::workflow_data::{AnyWorkflowData, McpWorkflowData}; use super::workflow_data::McpWorkflowData;
use crate::{ use crate::{
app_context::AppContext, app_context::AppContext,
mcp::{config::McpServerConfig, manager::McpManager}, mcp::{config::McpServerConfig, manager::McpManager},
@@ -38,14 +38,14 @@ impl McpServerConfigRequest {
pub struct ConnectMcpServerStep; pub struct ConnectMcpServerStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for ConnectMcpServerStep { impl StepExecutor<McpWorkflowData> for ConnectMcpServerStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<McpWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_mcp()?; let config_request = &context.data.config;
let config_request = &data.config; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
@@ -76,8 +76,7 @@ impl StepExecutor<AnyWorkflowData> for ConnectMcpServerStep {
); );
// Store client in typed data // Store client in typed data
let data_mut = context.data.as_mcp_mut()?; context.data.mcp_client = Some(Arc::new(client));
data_mut.mcp_client = Some(Arc::new(client));
Ok(StepResult::Success) Ok(StepResult::Success)
} }
@@ -96,18 +95,19 @@ impl StepExecutor<AnyWorkflowData> for ConnectMcpServerStep {
pub struct DiscoverMcpInventoryStep; pub struct DiscoverMcpInventoryStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for DiscoverMcpInventoryStep { impl StepExecutor<McpWorkflowData> for DiscoverMcpInventoryStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<McpWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_mcp()?; let config_request = &context.data.config;
let config_request = &data.config; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
let mcp_client = data let mcp_client = context
.data
.mcp_client .mcp_client
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("mcp_client".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("mcp_client".to_string()))?;
@@ -149,18 +149,19 @@ impl StepExecutor<AnyWorkflowData> for DiscoverMcpInventoryStep {
pub struct RegisterMcpServerStep; pub struct RegisterMcpServerStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for RegisterMcpServerStep { impl StepExecutor<McpWorkflowData> for RegisterMcpServerStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<McpWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_mcp()?; let config_request = &context.data.config;
let config_request = &data.config; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
let mcp_client = data let mcp_client = context
.data
.mcp_client .mcp_client
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("mcp_client".to_string()))? .ok_or_else(|| WorkflowError::ContextValueNotFound("mcp_client".to_string()))?
@@ -203,15 +204,13 @@ impl StepExecutor<AnyWorkflowData> for RegisterMcpServerStep {
pub struct ValidateRegistrationStep; pub struct ValidateRegistrationStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for ValidateRegistrationStep { impl StepExecutor<McpWorkflowData> for ValidateRegistrationStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<McpWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_mcp()?; let config_request = &context.data.config;
let config_request = &data.config; let client_registered = context.data.mcp_client.is_some();
let client_registered = data.mcp_client.is_some();
if client_registered { if client_registered {
info!( info!(
@@ -220,8 +219,7 @@ impl StepExecutor<AnyWorkflowData> for ValidateRegistrationStep {
); );
// Mark as validated // Mark as validated
let data_mut = context.data.as_mcp_mut()?; context.data.validated = true;
data_mut.validated = true;
return Ok(StepResult::Success); return Ok(StepResult::Success);
} }
@@ -263,7 +261,7 @@ impl StepExecutor<AnyWorkflowData> for ValidateRegistrationStep {
/// - DiscoverMcpInventory: 3 retries, 10s timeout (discovery + caching) /// - DiscoverMcpInventory: 3 retries, 10s timeout (discovery + caching)
/// - RegisterMcpServer: No retry, 5s timeout (fast registration) /// - RegisterMcpServer: No retry, 5s timeout (fast registration)
/// - ValidateRegistration: Final validation step /// - ValidateRegistration: Final validation step
pub fn create_mcp_registration_workflow() -> WorkflowDefinition<AnyWorkflowData> { pub fn create_mcp_registration_workflow() -> WorkflowDefinition<McpWorkflowData> {
WorkflowDefinition::new("mcp_registration", "MCP Server Registration") WorkflowDefinition::new("mcp_registration", "MCP Server Registration")
.add_step( .add_step(
StepDefinition::new( StepDefinition::new(
@@ -321,11 +319,11 @@ pub fn create_mcp_registration_workflow() -> WorkflowDefinition<AnyWorkflowData>
pub fn create_mcp_workflow_data( pub fn create_mcp_workflow_data(
config: McpServerConfigRequest, config: McpServerConfigRequest,
app_context: Arc<AppContext>, app_context: Arc<AppContext>,
) -> AnyWorkflowData { ) -> McpWorkflowData {
AnyWorkflowData::Mcp(McpWorkflowData { McpWorkflowData {
config, config,
validated: false, validated: false,
app_context: Some(app_context), app_context: Some(app_context),
mcp_client: None, mcp_client: None,
}) }
} }
+6 -3
View File
@@ -12,6 +12,7 @@ pub mod wasm_module_registration;
pub mod wasm_module_removal; pub mod wasm_module_removal;
pub mod worker; pub mod worker;
pub mod workflow_data; pub mod workflow_data;
pub mod workflow_engines;
// Worker management (registration, removal) // Worker management (registration, removal)
pub use mcp_registration::{ pub use mcp_registration::{
@@ -75,8 +76,10 @@ pub use worker::{
}; };
// Typed workflow data structures // Typed workflow data structures
pub use workflow_data::{ pub use workflow_data::{
AnyWorkflowData, ExternalWorkerWorkflowData, LocalWorkerWorkflowData, McpWorkflowData, ExternalWorkerWorkflowData, LocalWorkerWorkflowData, McpWorkflowData, ProtocolUpdateRequest,
ProtocolUpdateRequest, TokenizerWorkflowData, WasmRegistrationWorkflowData, TokenizerWorkflowData, WasmRegistrationWorkflowData, WasmRemovalWorkflowData,
WasmRemovalWorkflowData, WorkerConfigRequest, WorkerList as WorkflowWorkerList, WorkerConfigRequest, WorkerList as WorkflowWorkerList, WorkerRegistrationData,
WorkerRemovalWorkflowData, WorkerUpdateWorkflowData, WorkerRemovalWorkflowData, WorkerUpdateWorkflowData,
}; };
// Typed workflow engines
pub use workflow_engines::WorkflowEngines;
@@ -9,7 +9,7 @@ use async_trait::async_trait;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use tracing::{debug, error, info}; use tracing::{debug, error, info};
use super::workflow_data::{AnyWorkflowData, TokenizerWorkflowData}; use super::workflow_data::TokenizerWorkflowData;
use crate::{ use crate::{
app_context::AppContext, app_context::AppContext,
tokenizer::factory, tokenizer::factory,
@@ -47,14 +47,14 @@ pub struct TokenizerRemovalRequest {
pub struct ValidateTokenizerConfigStep; pub struct ValidateTokenizerConfigStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for ValidateTokenizerConfigStep { impl StepExecutor<TokenizerWorkflowData> for ValidateTokenizerConfigStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<TokenizerWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_tokenizer()?; let config = &context.data.config;
let config = &data.config; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
@@ -101,14 +101,14 @@ impl StepExecutor<AnyWorkflowData> for ValidateTokenizerConfigStep {
pub struct LoadTokenizerStep; pub struct LoadTokenizerStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for LoadTokenizerStep { impl StepExecutor<TokenizerWorkflowData> for LoadTokenizerStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<TokenizerWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_tokenizer()?; let config = &context.data.config;
let config = &data.config; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))? .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?
@@ -157,8 +157,7 @@ impl StepExecutor<AnyWorkflowData> for LoadTokenizerStep {
// Store vocab size in typed data // Store vocab size in typed data
if let Some(size) = vocab_size { if let Some(size) = vocab_size {
let data_mut = context.data.as_tokenizer_mut()?; context.data.vocab_size = Some(size);
data_mut.vocab_size = Some(size);
} }
Ok(StepResult::Success) Ok(StepResult::Success)
@@ -191,7 +190,7 @@ impl StepExecutor<AnyWorkflowData> for LoadTokenizerStep {
/// Workflow configuration: /// Workflow configuration:
/// - ValidateConfig: No retry, 5s timeout (fast validation) /// - ValidateConfig: No retry, 5s timeout (fast validation)
/// - LoadTokenizer: 3 retries, 5min timeout (may need to download from HuggingFace) /// - LoadTokenizer: 3 retries, 5min timeout (may need to download from HuggingFace)
pub fn create_tokenizer_registration_workflow() -> WorkflowDefinition<AnyWorkflowData> { pub fn create_tokenizer_registration_workflow() -> WorkflowDefinition<TokenizerWorkflowData> {
WorkflowDefinition::new("tokenizer_registration", "Tokenizer Registration") WorkflowDefinition::new("tokenizer_registration", "Tokenizer Registration")
.add_step( .add_step(
StepDefinition::new( StepDefinition::new(
@@ -222,12 +221,12 @@ pub fn create_tokenizer_registration_workflow() -> WorkflowDefinition<AnyWorkflo
pub fn create_tokenizer_workflow_data( pub fn create_tokenizer_workflow_data(
config: TokenizerConfigRequest, config: TokenizerConfigRequest,
app_context: Arc<AppContext>, app_context: Arc<AppContext>,
) -> AnyWorkflowData { ) -> TokenizerWorkflowData {
AnyWorkflowData::Tokenizer(TokenizerWorkflowData { TokenizerWorkflowData {
config, config,
vocab_size: None, vocab_size: None,
app_context: Some(app_context), app_context: Some(app_context),
}) }
} }
#[cfg(test)] #[cfg(test)]
@@ -10,7 +10,7 @@ use tracing::{debug, info, warn};
use uuid::Uuid; use uuid::Uuid;
use wasmtime::{component::Component, Config, Engine}; use wasmtime::{component::Component, Config, Engine};
use super::workflow_data::{AnyWorkflowData, WasmRegistrationWorkflowData}; use super::workflow_data::WasmRegistrationWorkflowData;
use crate::{ use crate::{
app_context::AppContext, app_context::AppContext,
wasm::module::{WasmModule, WasmModuleDescriptor, WasmModuleMeta}, wasm::module::{WasmModule, WasmModuleDescriptor, WasmModuleMeta},
@@ -65,13 +65,12 @@ fn has_wasm_extension(path: &Path) -> bool {
pub struct ValidateDescriptorStep; pub struct ValidateDescriptorStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for ValidateDescriptorStep { impl StepExecutor<WasmRegistrationWorkflowData> for ValidateDescriptorStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WasmRegistrationWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_wasm_registration()?; let descriptor = &context.data.config.descriptor;
let descriptor = &data.config.descriptor;
debug!("Validating WASM module descriptor: {}", descriptor.name); debug!("Validating WASM module descriptor: {}", descriptor.name);
@@ -207,8 +206,7 @@ impl StepExecutor<AnyWorkflowData> for ValidateDescriptorStep {
let module_name = descriptor.name.clone(); let module_name = descriptor.name.clone();
// Store file size in typed data // Store file size in typed data
let data_mut = context.data.as_wasm_registration_mut()?; context.data.file_size_bytes = Some(metadata.len());
data_mut.file_size_bytes = Some(metadata.len());
info!( info!(
"Descriptor validated successfully for module: {}", "Descriptor validated successfully for module: {}",
@@ -229,13 +227,12 @@ impl StepExecutor<AnyWorkflowData> for ValidateDescriptorStep {
pub struct CalculateHashStep; pub struct CalculateHashStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for CalculateHashStep { impl StepExecutor<WasmRegistrationWorkflowData> for CalculateHashStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WasmRegistrationWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_wasm_registration()?; let file_path = &context.data.config.descriptor.file_path;
let file_path = &data.config.descriptor.file_path;
debug!("Calculating SHA256 hash for: {}", file_path); debug!("Calculating SHA256 hash for: {}", file_path);
@@ -274,8 +271,7 @@ impl StepExecutor<AnyWorkflowData> for CalculateHashStep {
let path_for_log = file_path.clone(); let path_for_log = file_path.clone();
// Store hash in typed data // Store hash in typed data
let data_mut = context.data.as_wasm_registration_mut()?; context.data.sha256_hash = Some(hash);
data_mut.sha256_hash = Some(hash);
info!("SHA256 hash calculated for: {}", path_for_log); info!("SHA256 hash calculated for: {}", path_for_log);
Ok(StepResult::Success) Ok(StepResult::Success)
@@ -293,24 +289,25 @@ impl StepExecutor<AnyWorkflowData> for CalculateHashStep {
pub struct CheckDuplicateStep; pub struct CheckDuplicateStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for CheckDuplicateStep { impl StepExecutor<WasmRegistrationWorkflowData> for CheckDuplicateStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WasmRegistrationWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_wasm_registration()?; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
let sha256_hash = data let sha256_hash = context
.data
.sha256_hash .sha256_hash
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("sha256_hash".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("sha256_hash".to_string()))?;
debug!( debug!(
"Checking for duplicate SHA256 hash for module: {}", "Checking for duplicate SHA256 hash for module: {}",
data.config.descriptor.name context.data.config.descriptor.name
); );
// Get WASM module manager from app context // Get WASM module manager from app context
@@ -333,7 +330,7 @@ impl StepExecutor<AnyWorkflowData> for CheckDuplicateStep {
info!( info!(
"No duplicate found for module: {}", "No duplicate found for module: {}",
data.config.descriptor.name context.data.config.descriptor.name
); );
Ok(StepResult::Success) Ok(StepResult::Success)
} }
@@ -350,13 +347,12 @@ impl StepExecutor<AnyWorkflowData> for CheckDuplicateStep {
pub struct LoadWasmBytesStep; pub struct LoadWasmBytesStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for LoadWasmBytesStep { impl StepExecutor<WasmRegistrationWorkflowData> for LoadWasmBytesStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WasmRegistrationWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_wasm_registration()?; let file_path = &context.data.config.descriptor.file_path;
let file_path = &data.config.descriptor.file_path;
debug!("Loading WASM bytes from: {}", file_path); debug!("Loading WASM bytes from: {}", file_path);
@@ -372,8 +368,7 @@ impl StepExecutor<AnyWorkflowData> for LoadWasmBytesStep {
})?; })?;
// Store WASM bytes in typed data // Store WASM bytes in typed data
let data_mut = context.data.as_wasm_registration_mut()?; context.data.wasm_bytes = Some(wasm_bytes);
data_mut.wasm_bytes = Some(wasm_bytes);
info!("WASM bytes loaded from: {}", path_for_log); info!("WASM bytes loaded from: {}", path_for_log);
Ok(StepResult::Success) Ok(StepResult::Success)
@@ -391,20 +386,20 @@ impl StepExecutor<AnyWorkflowData> for LoadWasmBytesStep {
pub struct ValidateWasmComponentStep; pub struct ValidateWasmComponentStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for ValidateWasmComponentStep { impl StepExecutor<WasmRegistrationWorkflowData> for ValidateWasmComponentStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WasmRegistrationWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_wasm_registration()?; let wasm_bytes = context
let wasm_bytes = data .data
.wasm_bytes .wasm_bytes
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_bytes".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_bytes".to_string()))?;
debug!( debug!(
"Validating WASM component format for module: {}", "Validating WASM component format for module: {}",
data.config.descriptor.name context.data.config.descriptor.name
); );
// Create a temporary engine to validate the component // Create a temporary engine to validate the component
@@ -431,7 +426,7 @@ impl StepExecutor<AnyWorkflowData> for ValidateWasmComponentStep {
info!( info!(
"WASM component validated successfully for module: {}", "WASM component validated successfully for module: {}",
data.config.descriptor.name context.data.config.descriptor.name
); );
Ok(StepResult::Success) Ok(StepResult::Success)
} }
@@ -448,29 +443,32 @@ impl StepExecutor<AnyWorkflowData> for ValidateWasmComponentStep {
pub struct RegisterModuleStep; pub struct RegisterModuleStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for RegisterModuleStep { impl StepExecutor<WasmRegistrationWorkflowData> for RegisterModuleStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WasmRegistrationWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_wasm_registration()?; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
let sha256_hash = data let sha256_hash = context
.data
.sha256_hash .sha256_hash
.ok_or_else(|| WorkflowError::ContextValueNotFound("sha256_hash".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("sha256_hash".to_string()))?;
let file_size_bytes = data let file_size_bytes = context
.data
.file_size_bytes .file_size_bytes
.ok_or_else(|| WorkflowError::ContextValueNotFound("file_size_bytes".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("file_size_bytes".to_string()))?;
let wasm_bytes = data let wasm_bytes = context
.data
.wasm_bytes .wasm_bytes
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_bytes".to_string()))? .ok_or_else(|| WorkflowError::ContextValueNotFound("wasm_bytes".to_string()))?
.clone(); .clone();
let descriptor = &data.config.descriptor; let descriptor = &context.data.config.descriptor;
debug!("Registering WASM module in manager: {}", descriptor.name); debug!("Registering WASM module in manager: {}", descriptor.name);
@@ -519,8 +517,7 @@ impl StepExecutor<AnyWorkflowData> for RegisterModuleStep {
})?; })?;
// Store module UUID in typed data // Store module UUID in typed data
let data_mut = context.data.as_wasm_registration_mut()?; context.data.module_uuid = Some(module_uuid);
data_mut.module_uuid = Some(module_uuid);
info!( info!(
"WASM module registered successfully: {} (UUID: {})", "WASM module registered successfully: {} (UUID: {})",
@@ -552,7 +549,8 @@ impl StepExecutor<AnyWorkflowData> for RegisterModuleStep {
/// - LoadWasmBytes: 3 retries, 60s timeout (I/O intensive) /// - LoadWasmBytes: 3 retries, 60s timeout (I/O intensive)
/// - ValidateWasmComponent: No retry, 30s timeout (CPU intensive validation) /// - ValidateWasmComponent: No retry, 30s timeout (CPU intensive validation)
/// - RegisterModule: No retry, 5s timeout (fast registration) /// - RegisterModule: No retry, 5s timeout (fast registration)
pub fn create_wasm_module_registration_workflow() -> WorkflowDefinition<AnyWorkflowData> { pub fn create_wasm_module_registration_workflow() -> WorkflowDefinition<WasmRegistrationWorkflowData>
{
WorkflowDefinition::new("wasm_module_registration", "WASM Module Registration") WorkflowDefinition::new("wasm_module_registration", "WASM Module Registration")
.add_step( .add_step(
StepDefinition::new( StepDefinition::new(
@@ -627,13 +625,13 @@ pub fn create_wasm_module_registration_workflow() -> WorkflowDefinition<AnyWorkf
pub fn create_wasm_registration_workflow_data( pub fn create_wasm_registration_workflow_data(
config: WasmModuleConfigRequest, config: WasmModuleConfigRequest,
app_context: Arc<AppContext>, app_context: Arc<AppContext>,
) -> AnyWorkflowData { ) -> WasmRegistrationWorkflowData {
AnyWorkflowData::WasmRegistration(WasmRegistrationWorkflowData { WasmRegistrationWorkflowData {
config, config,
wasm_bytes: None, wasm_bytes: None,
sha256_hash: None, sha256_hash: None,
file_size_bytes: None, file_size_bytes: None,
module_uuid: None, module_uuid: None,
app_context: Some(app_context), app_context: Some(app_context),
}) }
} }
@@ -4,7 +4,7 @@ use async_trait::async_trait;
use tracing::{debug, info}; use tracing::{debug, info};
use uuid::Uuid; use uuid::Uuid;
use super::workflow_data::{AnyWorkflowData, WasmRemovalWorkflowData}; use super::workflow_data::WasmRemovalWorkflowData;
use crate::{ use crate::{
app_context::AppContext, app_context::AppContext,
workflow::{ workflow::{
@@ -37,14 +37,14 @@ impl WasmModuleRemovalRequest {
pub struct FindModuleToRemoveStep; pub struct FindModuleToRemoveStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for FindModuleToRemoveStep { impl StepExecutor<WasmRemovalWorkflowData> for FindModuleToRemoveStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WasmRemovalWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_wasm_removal()?; let removal_request = &context.data.config;
let removal_request = &data.config; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
@@ -80,8 +80,7 @@ impl StepExecutor<AnyWorkflowData> for FindModuleToRemoveStep {
let module_uuid = removal_request.module_uuid; let module_uuid = removal_request.module_uuid;
// Store the module ID in typed data // Store the module ID in typed data
let data_mut = context.data.as_wasm_removal_mut()?; context.data.module_id = Some(module_uuid.to_string());
data_mut.module_id = Some(module_uuid.to_string());
info!("Module found for removal: {}", module_uuid); info!("Module found for removal: {}", module_uuid);
Ok(StepResult::Success) Ok(StepResult::Success)
@@ -98,14 +97,14 @@ impl StepExecutor<AnyWorkflowData> for FindModuleToRemoveStep {
pub struct RemoveModuleStep; pub struct RemoveModuleStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for RemoveModuleStep { impl StepExecutor<WasmRemovalWorkflowData> for RemoveModuleStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WasmRemovalWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_wasm_removal()?; let removal_request = &context.data.config;
let removal_request = &data.config; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
@@ -151,7 +150,7 @@ impl StepExecutor<AnyWorkflowData> for RemoveModuleStep {
/// Workflow configuration: /// Workflow configuration:
/// - FindModuleToRemove: No retry, 5s timeout (fast lookup) /// - FindModuleToRemove: No retry, 5s timeout (fast lookup)
/// - RemoveModule: No retry, 5s timeout (fast removal) /// - RemoveModule: No retry, 5s timeout (fast removal)
pub fn create_wasm_module_removal_workflow() -> WorkflowDefinition<AnyWorkflowData> { pub fn create_wasm_module_removal_workflow() -> WorkflowDefinition<WasmRemovalWorkflowData> {
WorkflowDefinition::new("wasm_module_removal", "WASM Module Removal") WorkflowDefinition::new("wasm_module_removal", "WASM Module Removal")
.add_step( .add_step(
StepDefinition::new( StepDefinition::new(
@@ -174,10 +173,10 @@ pub fn create_wasm_module_removal_workflow() -> WorkflowDefinition<AnyWorkflowDa
pub fn create_wasm_removal_workflow_data( pub fn create_wasm_removal_workflow_data(
config: WasmModuleRemovalRequest, config: WasmModuleRemovalRequest,
app_context: Arc<AppContext>, app_context: Arc<AppContext>,
) -> AnyWorkflowData { ) -> WasmRemovalWorkflowData {
AnyWorkflowData::WasmRemoval(WasmRemovalWorkflowData { WasmRemovalWorkflowData {
config, config,
module_id: None, module_id: None,
app_context: Some(app_context), app_context: Some(app_context),
}) }
} }
@@ -8,7 +8,7 @@ use tracing::{debug, info};
use crate::{ use crate::{
core::{ core::{
circuit_breaker::CircuitBreakerConfig, circuit_breaker::CircuitBreakerConfig,
steps::workflow_data::{AnyWorkflowData, WorkerList}, steps::workflow_data::{ExternalWorkerWorkflowData, WorkerList},
worker::{HealthConfig, RuntimeType, WorkerType}, worker::{HealthConfig, RuntimeType, WorkerType},
BasicWorkerBuilder, ConnectionMode, Worker, BasicWorkerBuilder, ConnectionMode, Worker,
}, },
@@ -28,18 +28,18 @@ fn normalize_external_url(url: &str) -> String {
pub struct CreateExternalWorkersStep; pub struct CreateExternalWorkersStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for CreateExternalWorkersStep { impl StepExecutor<ExternalWorkerWorkflowData> for CreateExternalWorkersStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<ExternalWorkerWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_external_worker()?; let config = &context.data.config;
let config = &data.config; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .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 // Build configs from router settings
let circuit_breaker_config = { let circuit_breaker_config = {
@@ -150,10 +150,9 @@ impl StepExecutor<AnyWorkflowData> for CreateExternalWorkersStep {
} }
// Store results in workflow data // Store results in workflow data
let data_mut = context.data.as_external_worker_mut()?; context.data.workers = Some(WorkerList::from_workers(&workers));
data_mut.workers = Some(WorkerList::from_workers(&workers)); context.data.actual_workers = Some(workers);
data_mut.actual_workers = Some(workers); context.data.labels = labels;
data_mut.labels = labels;
Ok(StepResult::Success) Ok(StepResult::Success)
} }
@@ -13,7 +13,7 @@ use crate::{
core::{ core::{
model_card::{ModelCard, ProviderType}, model_card::{ModelCard, ProviderType},
model_type::ModelType, model_type::ModelType,
steps::workflow_data::AnyWorkflowData, steps::workflow_data::ExternalWorkerWorkflowData,
}, },
workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -225,13 +225,12 @@ async fn fetch_models(url: &str, api_key: Option<&str>) -> Result<Vec<ModelCard>
pub struct DiscoverModelsStep; pub struct DiscoverModelsStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for DiscoverModelsStep { impl StepExecutor<ExternalWorkerWorkflowData> for DiscoverModelsStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<ExternalWorkerWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_external_worker()?; let config = &context.data.config;
let config = &data.config;
// If no API key is provided, skip model discovery and use wildcard mode. // 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()) { if config.api_key.as_ref().is_none_or(|k| k.is_empty()) {
@@ -267,7 +266,7 @@ impl StepExecutor<AnyWorkflowData> for DiscoverModelsStep {
model_cards.iter().map(|c| &c.id).collect::<Vec<_>>() model_cards.iter().map(|c| &c.id).collect::<Vec<_>>()
); );
context.data.as_external_worker_mut()?.model_cards = model_cards; context.data.model_cards = model_cards;
Ok(StepResult::Success) Ok(StepResult::Success)
} }
+5 -5
View File
@@ -17,7 +17,7 @@ pub use discover_models::{
use super::shared::{ActivateWorkersStep, RegisterWorkersStep, UpdatePoliciesStep}; use super::shared::{ActivateWorkersStep, RegisterWorkersStep, UpdatePoliciesStep};
use crate::{ use crate::{
app_context::AppContext, app_context::AppContext,
core::steps::workflow_data::{AnyWorkflowData, ExternalWorkerWorkflowData}, core::steps::workflow_data::ExternalWorkerWorkflowData,
protocols::worker_spec::WorkerConfigRequest, protocols::worker_spec::WorkerConfigRequest,
workflow::{BackoffStrategy, FailureAction, RetryPolicy, StepDefinition, WorkflowDefinition}, workflow::{BackoffStrategy, FailureAction, RetryPolicy, StepDefinition, WorkflowDefinition},
}; };
@@ -38,7 +38,7 @@ use crate::{
/// │ │ /// │ │
/// └────────────┴────────────┘ /// └────────────┴────────────┘
/// ``` /// ```
pub fn create_external_worker_workflow() -> WorkflowDefinition<AnyWorkflowData> { pub fn create_external_worker_workflow() -> WorkflowDefinition<ExternalWorkerWorkflowData> {
WorkflowDefinition::new( WorkflowDefinition::new(
"external_worker_registration", "external_worker_registration",
"External Worker Registration", "External Worker Registration",
@@ -110,13 +110,13 @@ pub fn create_external_worker_workflow() -> WorkflowDefinition<AnyWorkflowData>
pub fn create_external_worker_workflow_data( pub fn create_external_worker_workflow_data(
config: WorkerConfigRequest, config: WorkerConfigRequest,
app_context: Arc<AppContext>, app_context: Arc<AppContext>,
) -> AnyWorkflowData { ) -> ExternalWorkerWorkflowData {
AnyWorkflowData::ExternalWorker(ExternalWorkerWorkflowData { ExternalWorkerWorkflowData {
config, config,
model_cards: Vec::new(), model_cards: Vec::new(),
workers: None, workers: None,
labels: std::collections::HashMap::new(), labels: std::collections::HashMap::new(),
app_context: Some(app_context), app_context: Some(app_context),
actual_workers: None, actual_workers: None,
}) }
} }
@@ -10,7 +10,7 @@ use crate::{
core::{ core::{
circuit_breaker::CircuitBreakerConfig, circuit_breaker::CircuitBreakerConfig,
model_card::ModelCard, model_card::ModelCard,
steps::workflow_data::{AnyWorkflowData, LocalWorkerWorkflowData}, steps::workflow_data::LocalWorkerWorkflowData,
worker::{HealthConfig, RuntimeType, WorkerType}, worker::{HealthConfig, RuntimeType, WorkerType},
BasicWorkerBuilder, ConnectionMode, DPAwareWorkerBuilder, Worker, UNKNOWN_MODEL_ID, BasicWorkerBuilder, ConnectionMode, DPAwareWorkerBuilder, Worker, UNKNOWN_MODEL_ID,
}, },
@@ -29,22 +29,22 @@ use crate::{
pub struct CreateLocalWorkerStep; pub struct CreateLocalWorkerStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for CreateLocalWorkerStep { impl StepExecutor<LocalWorkerWorkflowData> for CreateLocalWorkerStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<LocalWorkerWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_local_worker()?; let config = &context.data.config;
let config = &data.config; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
let connection_mode = data let connection_mode =
.connection_mode context.data.connection_mode.as_ref().ok_or_else(|| {
.as_ref() WorkflowError::ContextValueNotFound("connection_mode".to_string())
.ok_or_else(|| WorkflowError::ContextValueNotFound("connection_mode".to_string()))?; })?;
let discovered_labels = &data.discovered_labels; let discovered_labels = &context.data.discovered_labels;
// Check if worker already exists // Check if worker already exists
if app_context if app_context
@@ -100,7 +100,7 @@ impl StepExecutor<AnyWorkflowData> for CreateLocalWorkerStep {
let worker_type = parse_worker_type(config); let worker_type = parse_worker_type(config);
// Get runtime type (for gRPC workers) // 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 // Build circuit breaker config
let circuit_breaker_config = build_circuit_breaker_config(app_context); let circuit_breaker_config = build_circuit_breaker_config(app_context);
@@ -121,7 +121,7 @@ impl StepExecutor<AnyWorkflowData> for CreateLocalWorkerStep {
// Create workers - always output as Vec for unified downstream handling // Create workers - always output as Vec for unified downstream handling
let workers = if config.dp_aware { let workers = if config.dp_aware {
create_dp_aware_workers( create_dp_aware_workers(
data, &context.data,
&normalized_url, &normalized_url,
model_card, model_card,
worker_type, worker_type,
@@ -147,9 +147,8 @@ impl StepExecutor<AnyWorkflowData> for CreateLocalWorkerStep {
}; };
// Update workflow data // Update workflow data
let data_mut = context.data.as_local_worker_mut()?; context.data.actual_workers = Some(workers);
data_mut.actual_workers = Some(workers); context.data.final_labels = final_labels;
data_mut.final_labels = final_labels;
Ok(StepResult::Success) Ok(StepResult::Success)
} }
@@ -8,7 +8,7 @@ use tracing::debug;
use super::strip_protocol; use super::strip_protocol;
use crate::{ use crate::{
core::{steps::workflow_data::AnyWorkflowData, ConnectionMode}, core::{steps::workflow_data::LocalWorkerWorkflowData, ConnectionMode},
routers::grpc::client::GrpcClient, routers::grpc::client::GrpcClient,
workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -86,14 +86,14 @@ async fn try_grpc_health_check(
pub struct DetectConnectionModeStep; pub struct DetectConnectionModeStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for DetectConnectionModeStep { impl StepExecutor<LocalWorkerWorkflowData> for DetectConnectionModeStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<LocalWorkerWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_local_worker()?; let config = &context.data.config;
let config = &data.config; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
@@ -134,7 +134,7 @@ impl StepExecutor<AnyWorkflowData> for DetectConnectionModeStep {
} }
}; };
context.data.as_local_worker_mut()?.connection_mode = Some(connection_mode); context.data.connection_mode = Some(connection_mode);
Ok(StepResult::Success) Ok(StepResult::Success)
} }
@@ -5,7 +5,7 @@ use tracing::debug;
use super::discover_metadata::get_server_info; use super::discover_metadata::get_server_info;
use crate::{ 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}, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -41,13 +41,12 @@ pub async fn get_dp_info(url: &str, api_key: Option<&str>) -> Result<DpInfo, Str
pub struct DiscoverDPInfoStep; pub struct DiscoverDPInfoStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for DiscoverDPInfoStep { impl StepExecutor<LocalWorkerWorkflowData> for DiscoverDPInfoStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<LocalWorkerWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_local_worker()?; let config = &context.data.config;
let config = &data.config;
if !config.dp_aware { if !config.dp_aware {
debug!( debug!(
@@ -71,7 +70,7 @@ impl StepExecutor<AnyWorkflowData> for DiscoverDPInfoStep {
dp_info.dp_size, config.url, dp_info.model_id 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) Ok(StepResult::Success)
} }
@@ -11,7 +11,7 @@ use tracing::{debug, warn};
use super::strip_protocol; use super::strip_protocol;
use crate::{ use crate::{
core::{steps::workflow_data::AnyWorkflowData, ConnectionMode}, core::{steps::workflow_data::LocalWorkerWorkflowData, ConnectionMode},
routers::grpc::client::GrpcClient, routers::grpc::client::GrpcClient,
workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -218,17 +218,16 @@ async fn fetch_grpc_metadata(
pub struct DiscoverMetadataStep; pub struct DiscoverMetadataStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for DiscoverMetadataStep { impl StepExecutor<LocalWorkerWorkflowData> for DiscoverMetadataStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<LocalWorkerWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_local_worker()?; let config = &context.data.config;
let config = &data.config; let connection_mode =
let connection_mode = data context.data.connection_mode.as_ref().ok_or_else(|| {
.connection_mode WorkflowError::ContextValueNotFound("connection_mode".to_string())
.as_ref() })?;
.ok_or_else(|| WorkflowError::ContextValueNotFound("connection_mode".to_string()))?;
debug!( debug!(
"Discovering metadata for {} ({:?})", "Discovering metadata for {} ({:?})",
@@ -301,11 +300,10 @@ impl StepExecutor<AnyWorkflowData> for DiscoverMetadataStep {
); );
// Update workflow data // Update workflow data
let data_mut = context.data.as_local_worker_mut()?; context.data.discovered_labels = discovered_labels;
data_mut.discovered_labels = discovered_labels;
if let Some(runtime) = detected_runtime { if let Some(runtime) = detected_runtime {
debug!("Detected runtime type: {}", runtime); debug!("Detected runtime type: {}", runtime);
data_mut.detected_runtime_type = Some(runtime); context.data.detected_runtime_type = Some(runtime);
} }
Ok(StepResult::Success) Ok(StepResult::Success)
@@ -5,7 +5,7 @@ use tracing::debug;
use super::find_workers_by_url; use super::find_workers_by_url;
use crate::{ use crate::{
core::steps::workflow_data::AnyWorkflowData, core::steps::workflow_data::WorkerUpdateWorkflowData,
workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -16,15 +16,15 @@ use crate::{
pub struct FindWorkerToUpdateStep; pub struct FindWorkerToUpdateStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for FindWorkerToUpdateStep { impl StepExecutor<WorkerUpdateWorkflowData> for FindWorkerToUpdateStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WorkerUpdateWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_worker_update()?; let worker_url = &context.data.worker_url;
let worker_url = &data.worker_url; let dp_aware = context.data.dp_aware;
let dp_aware = data.dp_aware; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
@@ -50,7 +50,7 @@ impl StepExecutor<AnyWorkflowData> for FindWorkerToUpdateStep {
worker_url 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) Ok(StepResult::Success)
} }
@@ -7,7 +7,7 @@ use tracing::debug;
use super::find_workers_by_url; use super::find_workers_by_url;
use crate::{ use crate::{
core::steps::workflow_data::{AnyWorkflowData, WorkerList}, core::steps::workflow_data::{WorkerList, WorkerRemovalWorkflowData},
workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepId, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -25,14 +25,14 @@ pub struct WorkerRemovalRequest {
pub struct FindWorkersToRemoveStep; pub struct FindWorkersToRemoveStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for FindWorkersToRemoveStep { impl StepExecutor<WorkerRemovalWorkflowData> for FindWorkersToRemoveStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WorkerRemovalWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_worker_removal()?; let request = &context.data.config;
let request = &data.config; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
@@ -70,11 +70,10 @@ impl StepExecutor<AnyWorkflowData> for FindWorkersToRemoveStep {
.collect(); .collect();
// Update workflow data // Update workflow data
let data_mut = context.data.as_worker_removal_mut()?; context.data.workers_to_remove = Some(WorkerList::from_workers(&workers_to_remove));
data_mut.workers_to_remove = Some(WorkerList::from_workers(&workers_to_remove)); context.data.actual_workers_to_remove = Some(workers_to_remove);
data_mut.actual_workers_to_remove = Some(workers_to_remove); context.data.worker_urls = worker_urls;
data_mut.worker_urls = worker_urls; context.data.affected_models = affected_models;
data_mut.affected_models = affected_models;
Ok(StepResult::Success) Ok(StepResult::Success)
} }
@@ -40,8 +40,7 @@ use crate::{
config::RouterConfig, config::RouterConfig,
core::{ core::{
steps::workflow_data::{ steps::workflow_data::{
AnyWorkflowData, LocalWorkerWorkflowData, WorkerRemovalWorkflowData, LocalWorkerWorkflowData, WorkerRemovalWorkflowData, WorkerUpdateWorkflowData,
WorkerUpdateWorkflowData,
}, },
Worker, WorkerRegistry, Worker, WorkerRegistry,
}, },
@@ -76,7 +75,7 @@ pub(crate) fn find_workers_by_url(
pub fn create_local_worker_workflow( pub fn create_local_worker_workflow(
router_config: &RouterConfig, router_config: &RouterConfig,
) -> WorkflowDefinition<AnyWorkflowData> { ) -> WorkflowDefinition<LocalWorkerWorkflowData> {
let detect_timeout = Duration::from_secs(router_config.worker_startup_timeout_secs); let detect_timeout = Duration::from_secs(router_config.worker_startup_timeout_secs);
// Calculate max_attempts based on timeout // Calculate max_attempts based on timeout
@@ -208,7 +207,7 @@ pub fn create_local_worker_workflow(
/// │ /// │
/// update_remaining_policies /// update_remaining_policies
/// ``` /// ```
pub fn create_worker_removal_workflow() -> WorkflowDefinition<AnyWorkflowData> { pub fn create_worker_removal_workflow() -> WorkflowDefinition<WorkerRemovalWorkflowData> {
WorkflowDefinition::new("worker_removal", "Remove worker from router") WorkflowDefinition::new("worker_removal", "Remove worker from router")
.add_step( .add_step(
StepDefinition::new( StepDefinition::new(
@@ -273,7 +272,7 @@ pub fn create_worker_removal_workflow() -> WorkflowDefinition<AnyWorkflowData> {
/// │ /// │
/// update_policies_for_worker /// update_policies_for_worker
/// ``` /// ```
pub fn create_worker_update_workflow() -> WorkflowDefinition<AnyWorkflowData> { pub fn create_worker_update_workflow() -> WorkflowDefinition<WorkerUpdateWorkflowData> {
WorkflowDefinition::new("worker_update", "Update worker properties") WorkflowDefinition::new("worker_update", "Update worker properties")
.add_step( .add_step(
StepDefinition::new( StepDefinition::new(
@@ -319,8 +318,8 @@ pub fn create_worker_update_workflow() -> WorkflowDefinition<AnyWorkflowData> {
pub fn create_local_worker_workflow_data( pub fn create_local_worker_workflow_data(
config: WorkerConfigRequest, config: WorkerConfigRequest,
app_context: Arc<AppContext>, app_context: Arc<AppContext>,
) -> AnyWorkflowData { ) -> LocalWorkerWorkflowData {
AnyWorkflowData::LocalWorker(LocalWorkerWorkflowData { LocalWorkerWorkflowData {
config, config,
connection_mode: None, connection_mode: None,
discovered_labels: std::collections::HashMap::new(), discovered_labels: std::collections::HashMap::new(),
@@ -330,7 +329,7 @@ pub fn create_local_worker_workflow_data(
detected_runtime_type: None, detected_runtime_type: None,
app_context: Some(app_context), app_context: Some(app_context),
actual_workers: None, actual_workers: None,
}) }
} }
/// Helper to create initial workflow data for worker removal /// Helper to create initial workflow data for worker removal
@@ -338,15 +337,15 @@ pub fn create_worker_removal_workflow_data(
url: String, url: String,
dp_aware: bool, dp_aware: bool,
app_context: Arc<AppContext>, app_context: Arc<AppContext>,
) -> AnyWorkflowData { ) -> WorkerRemovalWorkflowData {
AnyWorkflowData::WorkerRemoval(WorkerRemovalWorkflowData { WorkerRemovalWorkflowData {
config: WorkerRemovalRequest { url, dp_aware }, config: WorkerRemovalRequest { url, dp_aware },
workers_to_remove: None, workers_to_remove: None,
worker_urls: Vec::new(), worker_urls: Vec::new(),
affected_models: std::collections::HashSet::new(), affected_models: std::collections::HashSet::new(),
app_context: Some(app_context), app_context: Some(app_context),
actual_workers_to_remove: None, actual_workers_to_remove: None,
}) }
} }
/// Helper to create initial workflow data for worker update /// Helper to create initial workflow data for worker update
@@ -354,15 +353,15 @@ pub fn create_worker_update_workflow_data(
worker_url: String, worker_url: String,
update_config: WorkerUpdateRequest, update_config: WorkerUpdateRequest,
app_context: Arc<AppContext>, app_context: Arc<AppContext>,
) -> AnyWorkflowData { ) -> WorkerUpdateWorkflowData {
// Determine if this is a DP-aware update based on URL pattern // Determine if this is a DP-aware update based on URL pattern
let dp_aware = worker_url.contains('@'); let dp_aware = worker_url.contains('@');
AnyWorkflowData::WorkerUpdate(WorkerUpdateWorkflowData { WorkerUpdateWorkflowData {
config: update_config, config: update_config,
worker_url, worker_url,
dp_aware, dp_aware,
app_context: Some(app_context), app_context: Some(app_context),
workers_to_update: None, workers_to_update: None,
updated_workers: None, updated_workers: None,
}) }
} }
@@ -4,7 +4,7 @@ use async_trait::async_trait;
use tracing::{debug, warn}; use tracing::{debug, warn};
use crate::{ use crate::{
core::steps::workflow_data::AnyWorkflowData, core::steps::workflow_data::LocalWorkerWorkflowData,
tokenizer::{factory, TokenizerRegistry}, tokenizer::{factory, TokenizerRegistry},
workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -13,18 +13,19 @@ use crate::{
pub struct RegisterTokenizerStep; pub struct RegisterTokenizerStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for RegisterTokenizerStep { impl StepExecutor<LocalWorkerWorkflowData> for RegisterTokenizerStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<LocalWorkerWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_local_worker()?; let labels = &context.data.final_labels;
let labels = &data.final_labels; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
let workers = data let workers = context
.data
.actual_workers .actual_workers
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("workers".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("workers".to_string()))?;
@@ -4,7 +4,7 @@ use async_trait::async_trait;
use tracing::debug; use tracing::debug;
use crate::{ use crate::{
core::steps::workflow_data::AnyWorkflowData, core::steps::workflow_data::WorkerRemovalWorkflowData,
workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -15,17 +15,18 @@ use crate::{
pub struct RemoveFromPolicyRegistryStep; pub struct RemoveFromPolicyRegistryStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for RemoveFromPolicyRegistryStep { impl StepExecutor<WorkerRemovalWorkflowData> for RemoveFromPolicyRegistryStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WorkerRemovalWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_worker_removal()?; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
let workers_to_remove = data let workers_to_remove = context
.data
.actual_workers_to_remove .actual_workers_to_remove
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("workers_to_remove".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("workers_to_remove".to_string()))?;
@@ -6,7 +6,7 @@ use async_trait::async_trait;
use tracing::{debug, warn}; use tracing::{debug, warn};
use crate::{ use crate::{
core::steps::workflow_data::AnyWorkflowData, core::steps::workflow_data::WorkerRemovalWorkflowData,
observability::metrics::Metrics, observability::metrics::Metrics,
workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -17,17 +17,17 @@ use crate::{
pub struct RemoveFromWorkerRegistryStep; pub struct RemoveFromWorkerRegistryStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for RemoveFromWorkerRegistryStep { impl StepExecutor<WorkerRemovalWorkflowData> for RemoveFromWorkerRegistryStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WorkerRemovalWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_worker_removal()?; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
let worker_urls = &data.worker_urls; let worker_urls = &context.data.worker_urls;
debug!( debug!(
"Removing {} worker(s) from worker registry", "Removing {} worker(s) from worker registry",
@@ -6,7 +6,7 @@ use async_trait::async_trait;
use tracing::debug; use tracing::debug;
use crate::{ use crate::{
core::steps::workflow_data::AnyWorkflowData, core::steps::workflow_data::WorkerUpdateWorkflowData,
workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -17,20 +17,20 @@ use crate::{
pub struct UpdatePoliciesForWorkerStep; pub struct UpdatePoliciesForWorkerStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for UpdatePoliciesForWorkerStep { impl StepExecutor<WorkerUpdateWorkflowData> for UpdatePoliciesForWorkerStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WorkerUpdateWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_worker_update()?; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
let updated_workers = data let updated_workers =
.updated_workers context.data.updated_workers.as_ref().ok_or_else(|| {
.as_ref() WorkflowError::ContextValueNotFound("updated_workers".to_string())
.ok_or_else(|| WorkflowError::ContextValueNotFound("updated_workers".to_string()))?; })?;
// Collect affected models // Collect affected models
let affected_models: HashSet<String> = updated_workers let affected_models: HashSet<String> = updated_workers
@@ -4,7 +4,7 @@ use async_trait::async_trait;
use tracing::{debug, info}; use tracing::{debug, info};
use crate::{ use crate::{
core::steps::workflow_data::AnyWorkflowData, core::steps::workflow_data::WorkerRemovalWorkflowData,
workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -15,18 +15,18 @@ use crate::{
pub struct UpdateRemainingPoliciesStep; pub struct UpdateRemainingPoliciesStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for UpdateRemainingPoliciesStep { impl StepExecutor<WorkerRemovalWorkflowData> for UpdateRemainingPoliciesStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WorkerRemovalWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_worker_removal()?; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?; .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?;
let affected_models = &data.affected_models; let affected_models = &context.data.affected_models;
let worker_urls = &data.worker_urls; let worker_urls = &context.data.worker_urls;
debug!( debug!(
"Updating cache-aware policies for {} affected model(s)", "Updating cache-aware policies for {} affected model(s)",
@@ -6,7 +6,9 @@ use async_trait::async_trait;
use tracing::{debug, info}; use tracing::{debug, info};
use crate::{ use crate::{
core::{steps::workflow_data::AnyWorkflowData, BasicWorkerBuilder, HealthConfig, Worker}, core::{
steps::workflow_data::WorkerUpdateWorkflowData, BasicWorkerBuilder, HealthConfig, Worker,
},
workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult},
}; };
@@ -17,19 +19,20 @@ use crate::{
pub struct UpdateWorkerPropertiesStep; pub struct UpdateWorkerPropertiesStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for UpdateWorkerPropertiesStep { impl StepExecutor<WorkerUpdateWorkflowData> for UpdateWorkerPropertiesStep {
async fn execute( async fn execute(
&self, &self,
context: &mut WorkflowContext<AnyWorkflowData>, context: &mut WorkflowContext<WorkerUpdateWorkflowData>,
) -> WorkflowResult<StepResult> { ) -> WorkflowResult<StepResult> {
let data = context.data.as_worker_update()?; let request = &context.data.config;
let request = &data.config; let app_context = context
let app_context = data .data
.app_context .app_context
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))? .ok_or_else(|| WorkflowError::ContextValueNotFound("app_context".to_string()))?
.clone(); .clone();
let workers_to_update = data let workers_to_update = context
.data
.workers_to_update .workers_to_update
.as_ref() .as_ref()
.ok_or_else(|| WorkflowError::ContextValueNotFound("workers_to_update".to_string()))? .ok_or_else(|| WorkflowError::ContextValueNotFound("workers_to_update".to_string()))?
@@ -137,7 +140,7 @@ impl StepExecutor<AnyWorkflowData> for UpdateWorkerPropertiesStep {
} }
// Store updated workers for subsequent steps // 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) Ok(StepResult::Success)
} }
@@ -4,21 +4,21 @@ use async_trait::async_trait;
use tracing::info; use tracing::info;
use crate::{ use crate::{
core::steps::workflow_data::AnyWorkflowData, core::steps::workflow_data::WorkerRegistrationData,
workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{
StepExecutor, StepResult, WorkflowContext, WorkflowData, WorkflowError, WorkflowResult,
},
}; };
/// Unified step to activate workers by marking them as healthy. /// Unified step to activate workers by marking them as healthy.
/// ///
/// This is the final step in any worker registration workflow. /// This is the final step in any worker registration workflow.
/// Works with any workflow data type that implements `WorkerRegistrationData`.
pub struct ActivateWorkersStep; pub struct ActivateWorkersStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for ActivateWorkersStep { impl<D: WorkerRegistrationData + WorkflowData> StepExecutor<D> for ActivateWorkersStep {
async fn execute( async fn execute(&self, context: &mut WorkflowContext<D>) -> WorkflowResult<StepResult> {
&self,
context: &mut WorkflowContext<AnyWorkflowData>,
) -> WorkflowResult<StepResult> {
let workers = context let workers = context
.data .data
.get_actual_workers() .get_actual_workers()
@@ -6,23 +6,23 @@ use async_trait::async_trait;
use tracing::debug; use tracing::debug;
use crate::{ use crate::{
core::steps::workflow_data::AnyWorkflowData, core::steps::workflow_data::WorkerRegistrationData,
observability::metrics::Metrics, observability::metrics::Metrics,
workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{
StepExecutor, StepResult, WorkflowContext, WorkflowData, WorkflowError, WorkflowResult,
},
}; };
/// Unified step to register workers in the registry. /// Unified step to register workers in the registry.
/// ///
/// Works with both single workers and batches. Always expects `workers` key /// Works with both single workers and batches. Always expects `workers` key
/// in context containing `Vec<Arc<dyn Worker>>`. /// in context containing `Vec<Arc<dyn Worker>>`.
/// Works with any workflow data type that implements `WorkerRegistrationData`.
pub struct RegisterWorkersStep; pub struct RegisterWorkersStep;
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for RegisterWorkersStep { impl<D: WorkerRegistrationData + WorkflowData> StepExecutor<D> for RegisterWorkersStep {
async fn execute( async fn execute(&self, context: &mut WorkflowContext<D>) -> WorkflowResult<StepResult> {
&self,
context: &mut WorkflowContext<AnyWorkflowData>,
) -> WorkflowResult<StepResult> {
let app_context = context let app_context = context
.data .data
.get_app_context() .get_app_context()
@@ -6,8 +6,10 @@ use async_trait::async_trait;
use tracing::{debug, warn}; use tracing::{debug, warn};
use crate::{ use crate::{
core::{steps::workflow_data::AnyWorkflowData, Worker}, core::{steps::workflow_data::WorkerRegistrationData, Worker},
workflow::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}, workflow::{
StepExecutor, StepResult, WorkflowContext, WorkflowData, WorkflowError, WorkflowResult,
},
}; };
/// Unified step to update policy registry for registered workers. /// Unified step to update policy registry for registered workers.
@@ -81,11 +83,8 @@ impl UpdatePoliciesStep {
} }
#[async_trait] #[async_trait]
impl StepExecutor<AnyWorkflowData> for UpdatePoliciesStep { impl<D: WorkerRegistrationData + WorkflowData> StepExecutor<D> for UpdatePoliciesStep {
async fn execute( async fn execute(&self, context: &mut WorkflowContext<D>) -> WorkflowResult<StepResult> {
&self,
context: &mut WorkflowContext<AnyWorkflowData>,
) -> WorkflowResult<StepResult> {
let app_context = context let app_context = context
.data .data
.get_app_context() .get_app_context()
+56 -240
View File
@@ -1,7 +1,14 @@
//! Typed workflow data structures //! Typed workflow data structures
//! //!
//! This module defines the typed data structures for all workflows, enabling //! 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}; use std::{collections::HashMap, sync::Arc};
@@ -26,6 +33,26 @@ use crate::{
workflow::{WorkflowData, WorkflowError}, 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<AppContext>>;
/// Get the actual worker objects (transient, not serialized).
fn get_actual_workers(&self) -> Option<&Vec<Arc<dyn Worker>>>;
/// Get the labels for policy registration.
fn get_labels(&self) -> Option<&HashMap<String, String>>;
}
/// Wrapper for worker list that can be serialized /// Wrapper for worker list that can be serialized
#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct WorkerList { pub struct WorkerList {
@@ -119,6 +146,20 @@ impl LocalWorkerWorkflowData {
} }
} }
impl WorkerRegistrationData for LocalWorkerWorkflowData {
fn get_app_context(&self) -> Option<&Arc<AppContext>> {
self.app_context.as_ref()
}
fn get_actual_workers(&self) -> Option<&Vec<Arc<dyn Worker>>> {
self.actual_workers.as_ref()
}
fn get_labels(&self) -> Option<&HashMap<String, String>> {
Some(&self.final_labels)
}
}
/// Data for external worker registration workflow /// Data for external worker registration workflow
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExternalWorkerWorkflowData { pub struct ExternalWorkerWorkflowData {
@@ -154,6 +195,20 @@ impl ExternalWorkerWorkflowData {
} }
} }
impl WorkerRegistrationData for ExternalWorkerWorkflowData {
fn get_app_context(&self) -> Option<&Arc<AppContext>> {
self.app_context.as_ref()
}
fn get_actual_workers(&self) -> Option<&Vec<Arc<dyn Worker>>> {
self.actual_workers.as_ref()
}
fn get_labels(&self) -> Option<&HashMap<String, String>> {
Some(&self.labels)
}
}
/// Data for worker removal workflow /// Data for worker removal workflow
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkerRemovalWorkflowData { pub struct WorkerRemovalWorkflowData {
@@ -318,242 +373,3 @@ impl WasmRemovalWorkflowData {
Ok(()) 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<AnyWorkflowData>` 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<AppContext>> {
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<Arc<dyn Worker>>> {
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<Arc<dyn Worker>>,
) -> 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<String, String>> {
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(),
}
}
}
@@ -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<LocalWorkerWorkflowData, InMemoryStore<LocalWorkerWorkflowData>>;
/// Type alias for external worker workflow engine
pub type ExternalWorkerEngine =
WorkflowEngine<ExternalWorkerWorkflowData, InMemoryStore<ExternalWorkerWorkflowData>>;
/// Type alias for worker removal workflow engine
pub type WorkerRemovalEngine =
WorkflowEngine<WorkerRemovalWorkflowData, InMemoryStore<WorkerRemovalWorkflowData>>;
/// Type alias for worker update workflow engine
pub type WorkerUpdateEngine =
WorkflowEngine<WorkerUpdateWorkflowData, InMemoryStore<WorkerUpdateWorkflowData>>;
/// Type alias for MCP registration workflow engine
pub type McpEngine = WorkflowEngine<McpWorkflowData, InMemoryStore<McpWorkflowData>>;
/// Type alias for tokenizer registration workflow engine
pub type TokenizerEngine =
WorkflowEngine<TokenizerWorkflowData, InMemoryStore<TokenizerWorkflowData>>;
/// Type alias for WASM registration workflow engine
pub type WasmRegistrationEngine =
WorkflowEngine<WasmRegistrationWorkflowData, InMemoryStore<WasmRegistrationWorkflowData>>;
/// Type alias for WASM removal workflow engine
pub type WasmRemovalEngine =
WorkflowEngine<WasmRemovalWorkflowData, InMemoryStore<WasmRemovalWorkflowData>>;
/// Collection of typed workflow engines
///
/// Each workflow type has its own engine with compile-time type safety.
/// This replaces the old `WorkflowEngine<AnyWorkflowData, ...>` approach.
#[derive(Clone, Debug)]
pub struct WorkflowEngines {
/// Engine for local worker registration workflows
pub local_worker: Arc<LocalWorkerEngine>,
/// Engine for external worker registration workflows
pub external_worker: Arc<ExternalWorkerEngine>,
/// Engine for worker removal workflows
pub worker_removal: Arc<WorkerRemovalEngine>,
/// Engine for worker update workflows
pub worker_update: Arc<WorkerUpdateEngine>,
/// Engine for MCP server registration workflows
pub mcp: Arc<McpEngine>,
/// Engine for tokenizer registration workflows
pub tokenizer: Arc<TokenizerEngine>,
/// Engine for WASM module registration workflows
pub wasm_registration: Arc<WasmRegistrationEngine>,
/// Engine for WASM module removal workflows
pub wasm_removal: Arc<WasmRemovalEngine>,
}
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<S: EventSubscriber + 'static>(&self, subscriber: Arc<S>) {
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;
}
}
+10 -41
View File
@@ -20,16 +20,11 @@ use tokio::{signal, spawn};
use tracing::{debug, error, info, warn, Level}; use tracing::{debug, error, info, warn, Level};
use crate::{ use crate::{
app_context::{AppContext, AppWorkflowEngine}, app_context::AppContext,
config::{RouterConfig, RoutingMode}, config::{RouterConfig, RoutingMode},
core::{ core::{
job_queue::{JobQueue, JobQueueConfig}, job_queue::{JobQueue, JobQueueConfig},
steps::{ steps::WorkflowEngines,
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,
},
worker::WorkerType, worker::WorkerType,
worker_manager::WorkerManager, worker_manager::WorkerManager,
Job, Job,
@@ -730,44 +725,18 @@ pub async fn startup(config: ServerConfig) -> Result<(), Box<dyn std::error::Err
.set(worker_job_queue) .set(worker_job_queue)
.expect("JobQueue should only be initialized once"); .expect("JobQueue should only be initialized once");
// Initialize workflow engine and register workflows // Initialize typed workflow engines
let engine = Arc::new(AppWorkflowEngine::new()); let engines = WorkflowEngines::new(&config.router_config);
engine // Subscribe logging to all workflow engines
.event_bus() engines.subscribe_all(Arc::new(LoggingSubscriber)).await;
.subscribe(Arc::new(LoggingSubscriber))
.await;
engine
.register_workflow(create_local_worker_workflow(&config.router_config))
.expect("local_worker_registration workflow should be valid");
engine
.register_workflow(create_external_worker_workflow())
.expect("external_worker_registration workflow should be valid");
engine
.register_workflow(create_worker_removal_workflow())
.expect("worker_removal workflow should be valid");
engine
.register_workflow(create_worker_update_workflow())
.expect("worker_update workflow should be valid");
engine
.register_workflow(create_mcp_registration_workflow())
.expect("mcp_registration 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");
engine
.register_workflow(create_tokenizer_registration_workflow())
.expect("tokenizer_registration workflow should be valid");
app_context app_context
.workflow_engine .workflow_engines
.set(engine) .set(engines)
.expect("WorkflowEngine should only be initialized once"); .expect("WorkflowEngines should only be initialized once");
debug!( debug!(
"Workflow engine initialized with worker and MCP registration workflows (health check timeout: {}s)", "Workflow engines initialized (health check timeout: {}s)",
config.router_config.health_check.timeout_secs config.router_config.health_check.timeout_secs
); );
+1 -1
View File
@@ -643,7 +643,7 @@ mod tests {
configured_reasoning_parser: None, configured_reasoning_parser: None,
configured_tool_parser: None, configured_tool_parser: None,
worker_job_queue: worker_job_queue.clone(), worker_job_queue: worker_job_queue.clone(),
workflow_engine: Arc::new(std::sync::OnceLock::new()), workflow_engines: Arc::new(std::sync::OnceLock::new()),
mcp_manager: Arc::new(std::sync::OnceLock::new()), mcp_manager: Arc::new(std::sync::OnceLock::new()),
tokenizer_registry: Arc::new(crate::tokenizer::registry::TokenizerRegistry::new()), tokenizer_registry: Arc::new(crate::tokenizer::registry::TokenizerRegistry::new()),
wasm_manager: None, wasm_manager: None,
+59
View File
@@ -781,6 +781,65 @@ impl<D: WorkflowData, S: StateStore<D> + 'static> WorkflowEngine<D, S> {
self.state_store.load(instance_id) 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<String, String> {
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 /// Clone engine for async execution
fn clone_for_execution(&self) -> Self { fn clone_for_execution(&self) -> Self {
Self { Self {
+18 -15
View File
@@ -42,6 +42,10 @@ pub trait StateStore<D: WorkflowData>: Send + Sync + Clone {
/// Get just the workflow context without cloning the entire state /// Get just the workflow context without cloning the entire state
fn get_context(&self, instance_id: WorkflowInstanceId) -> WorkflowResult<WorkflowContext<D>>; fn get_context(&self, instance_id: WorkflowInstanceId) -> WorkflowResult<WorkflowContext<D>>;
/// 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 /// In-memory state storage for workflow instances
@@ -72,21 +76,6 @@ impl<D: WorkflowData> InMemoryStore<D> {
pub fn count(&self) -> usize { pub fn count(&self) -> usize {
self.states.read().len() 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<D: WorkflowData> Default for InMemoryStore<D> { impl<D: WorkflowData> Default for InMemoryStore<D> {
@@ -189,4 +178,18 @@ impl<D: WorkflowData> StateStore<D> for InMemoryStore<D> {
} }
removed_count 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
}
} }
+27 -54
View File
@@ -322,9 +322,9 @@ pub async fn create_test_context(config: RouterConfig) -> Arc<AppContext> {
config.worker_startup_check_interval_secs, 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 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 mcp_manager_lock = Arc::new(OnceLock::new());
let app_context = Arc::new( let app_context = Arc::new(
@@ -342,7 +342,7 @@ pub async fn create_test_context(config: RouterConfig) -> Arc<AppContext> {
.conversation_item_storage(conversation_item_storage) .conversation_item_storage(conversation_item_storage)
.load_monitor(load_monitor) .load_monitor(load_monitor)
.worker_job_queue(worker_job_queue) .worker_job_queue(worker_job_queue)
.workflow_engine(workflow_engine) .workflow_engines(workflow_engines)
.mcp_manager(mcp_manager_lock) .mcp_manager(mcp_manager_lock)
.build() .build()
.unwrap(), .unwrap(),
@@ -356,22 +356,13 @@ pub async fn create_test_context(config: RouterConfig) -> Arc<AppContext> {
.set(job_queue) .set(job_queue)
.expect("JobQueue should only be initialized once"); .expect("JobQueue should only be initialized once");
// Initialize WorkflowEngine and register workflows // Initialize typed workflow engines
use smg::{ use smg::core::steps::WorkflowEngines;
core::steps::{create_local_worker_workflow, create_worker_removal_workflow}, let engines = WorkflowEngines::new(&config);
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");
app_context app_context
.workflow_engine .workflow_engines
.set(engine) .set(engines)
.expect("WorkflowEngine should only be initialized once"); .expect("WorkflowEngines should only be initialized once");
// Register external workers for OpenAI mode // Register external workers for OpenAI mode
if let RoutingMode::OpenAI { worker_urls, .. } = &config.mode { if let RoutingMode::OpenAI { worker_urls, .. } = &config.mode {
@@ -451,9 +442,9 @@ pub async fn create_test_context_with_parsers(config: RouterConfig) -> Arc<AppCo
config.worker_startup_check_interval_secs, 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 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 mcp_manager_lock = Arc::new(OnceLock::new());
// Initialize parser factories // Initialize parser factories
@@ -475,7 +466,7 @@ pub async fn create_test_context_with_parsers(config: RouterConfig) -> Arc<AppCo
.conversation_item_storage(conversation_item_storage) .conversation_item_storage(conversation_item_storage)
.load_monitor(load_monitor) .load_monitor(load_monitor)
.worker_job_queue(worker_job_queue) .worker_job_queue(worker_job_queue)
.workflow_engine(workflow_engine) .workflow_engines(workflow_engines)
.mcp_manager(mcp_manager_lock) .mcp_manager(mcp_manager_lock)
.build() .build()
.unwrap(), .unwrap(),
@@ -489,22 +480,13 @@ pub async fn create_test_context_with_parsers(config: RouterConfig) -> Arc<AppCo
.set(job_queue) .set(job_queue)
.expect("JobQueue should only be initialized once"); .expect("JobQueue should only be initialized once");
// Initialize WorkflowEngine and register workflows // Initialize typed workflow engines
use smg::{ use smg::core::steps::WorkflowEngines;
core::steps::{create_local_worker_workflow, create_worker_removal_workflow}, let engines = WorkflowEngines::new(&config);
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");
app_context app_context
.workflow_engine .workflow_engines
.set(engine) .set(engines)
.expect("WorkflowEngine should only be initialized once"); .expect("WorkflowEngines should only be initialized once");
// Register external workers for OpenAI mode // Register external workers for OpenAI mode
if let RoutingMode::OpenAI { worker_urls, .. } = &config.mode { if let RoutingMode::OpenAI { worker_urls, .. } = &config.mode {
@@ -588,9 +570,9 @@ pub async fn create_test_context_with_mcp_config(
config.worker_startup_check_interval_secs, 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 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 mcp_manager_lock = Arc::new(OnceLock::new());
let app_context = Arc::new( let app_context = Arc::new(
@@ -608,7 +590,7 @@ pub async fn create_test_context_with_mcp_config(
.conversation_item_storage(conversation_item_storage) .conversation_item_storage(conversation_item_storage)
.load_monitor(load_monitor) .load_monitor(load_monitor)
.worker_job_queue(worker_job_queue) .worker_job_queue(worker_job_queue)
.workflow_engine(workflow_engine) .workflow_engines(workflow_engines)
.mcp_manager(mcp_manager_lock) .mcp_manager(mcp_manager_lock)
.build() .build()
.unwrap(), .unwrap(),
@@ -622,22 +604,13 @@ pub async fn create_test_context_with_mcp_config(
.set(job_queue) .set(job_queue)
.expect("JobQueue should only be initialized once"); .expect("JobQueue should only be initialized once");
// Initialize WorkflowEngine and register workflows // Initialize typed workflow engines
use smg::{ use smg::core::steps::WorkflowEngines;
core::steps::{create_local_worker_workflow, create_worker_removal_workflow}, let engines = WorkflowEngines::new(&config);
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");
app_context app_context
.workflow_engine .workflow_engines
.set(engine) .set(engines)
.expect("WorkflowEngine should only be initialized once"); .expect("WorkflowEngines should only be initialized once");
// Register external workers for OpenAI mode // Register external workers for OpenAI mode
if let RoutingMode::OpenAI { worker_urls, .. } = &config.mode { if let RoutingMode::OpenAI { worker_urls, .. } = &config.mode {
+5 -5
View File
@@ -58,9 +58,9 @@ pub fn create_test_app(
router_config.worker_startup_check_interval_secs, router_config.worker_startup_check_interval_secs,
))); )));
// Create empty OnceLock for worker job queue and workflow engine // Create empty OnceLock for worker job queue and workflow engines
let worker_job_queue = Arc::new(OnceLock::new()); let worker_job_queue = Arc::new(OnceLock::new());
let workflow_engine = Arc::new(OnceLock::new()); let workflow_engines = Arc::new(OnceLock::new());
// Create AppContext using builder pattern // Create AppContext using builder pattern
let app_context = Arc::new( let app_context = Arc::new(
@@ -78,7 +78,7 @@ pub fn create_test_app(
.conversation_item_storage(conversation_item_storage) .conversation_item_storage(conversation_item_storage)
.load_monitor(load_monitor) .load_monitor(load_monitor)
.worker_job_queue(worker_job_queue) .worker_job_queue(worker_job_queue)
.workflow_engine(workflow_engine) .workflow_engines(workflow_engines)
.build() .build()
.unwrap(), .unwrap(),
); );
@@ -168,7 +168,7 @@ pub async fn create_test_app_context() -> Arc<AppContext> {
// Initialize empty OnceLocks // Initialize empty OnceLocks
let worker_job_queue = Arc::new(OnceLock::new()); 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 // Initialize MCP manager with empty config
let mcp_manager_lock = Arc::new(OnceLock::new()); let mcp_manager_lock = Arc::new(OnceLock::new());
@@ -208,7 +208,7 @@ pub async fn create_test_app_context() -> Arc<AppContext> {
.conversation_item_storage(conversation_item_storage) .conversation_item_storage(conversation_item_storage)
.load_monitor(None) .load_monitor(None)
.worker_job_queue(worker_job_queue) .worker_job_queue(worker_job_queue)
.workflow_engine(workflow_engine) .workflow_engines(workflow_engines)
.mcp_manager(mcp_manager_lock) .mcp_manager(mcp_manager_lock)
.build() .build()
.unwrap(), .unwrap(),
@@ -247,9 +247,9 @@ mod pd_routing_unit_tests {
config.worker_startup_check_interval_secs, 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 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()); let mcp_manager = Arc::new(OnceLock::new());
Arc::new( Arc::new(
@@ -267,7 +267,7 @@ mod pd_routing_unit_tests {
.conversation_item_storage(conversation_item_storage) .conversation_item_storage(conversation_item_storage)
.load_monitor(load_monitor) .load_monitor(load_monitor)
.worker_job_queue(worker_job_queue) .worker_job_queue(worker_job_queue)
.workflow_engine(workflow_engine) .workflow_engines(workflow_engines)
.mcp_manager(mcp_manager) .mcp_manager(mcp_manager)
.build() .build()
.unwrap(), .unwrap(),
+27 -48
View File
@@ -18,10 +18,7 @@ use axum::{
use smg::{ use smg::{
app_context::AppContext, app_context::AppContext,
config::RouterConfig, config::RouterConfig,
core::{ core::{LoadMonitor, WorkerRegistry},
steps::{create_wasm_module_registration_workflow, create_wasm_module_removal_workflow},
LoadMonitor, WorkerRegistry,
},
data_connector::{ data_connector::{
MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage, MemoryConversationItemStorage, MemoryConversationStorage, MemoryResponseStorage,
}, },
@@ -71,10 +68,10 @@ async fn create_test_context_with_wasm() -> Arc<AppContext> {
config.worker_startup_check_interval_secs, 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; use std::sync::OnceLock;
let worker_job_queue = Arc::new(OnceLock::new()); 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 mcp_manager_lock = Arc::new(OnceLock::new());
let app_context = Arc::new( let app_context = Arc::new(
@@ -92,7 +89,7 @@ async fn create_test_context_with_wasm() -> Arc<AppContext> {
.conversation_item_storage(conversation_item_storage) .conversation_item_storage(conversation_item_storage)
.load_monitor(load_monitor) .load_monitor(load_monitor)
.worker_job_queue(worker_job_queue) .worker_job_queue(worker_job_queue)
.workflow_engine(workflow_engine) .workflow_engines(workflow_engines)
.mcp_manager(mcp_manager_lock) .mcp_manager(mcp_manager_lock)
.wasm_manager(Some(wasm_manager)) .wasm_manager(Some(wasm_manager))
.build() .build()
@@ -107,28 +104,13 @@ async fn create_test_context_with_wasm() -> Arc<AppContext> {
.set(job_queue) .set(job_queue)
.expect("JobQueue should only be initialized once"); .expect("JobQueue should only be initialized once");
// Initialize WorkflowEngine and register workflows // Initialize WorkflowEngines
use smg::{ use smg::core::steps::WorkflowEngines;
core::steps::{create_local_worker_workflow, create_worker_removal_workflow}, let engines = WorkflowEngines::new(&config);
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");
app_context app_context
.workflow_engine .workflow_engines
.set(engine) .set(engines)
.expect("WorkflowEngine should only be initialized once"); .expect("WorkflowEngines should only be initialized once");
// Initialize MCP manager with empty config // Initialize MCP manager with empty config
use smg::mcp::{McpConfig, McpManager}; use smg::mcp::{McpConfig, McpManager};
@@ -678,17 +660,14 @@ async fn test_wasm_module_execution() {
.as_ref() .as_ref()
.expect("WASM manager should be initialized"); .expect("WASM manager should be initialized");
let engine = app_context let engines = app_context
.workflow_engine .workflow_engines
.get() .get()
.expect("Workflow engine should be initialized"); .expect("Workflow engines should be initialized");
// Create workflow context for registration // Create workflow context for registration
use smg::{ use smg::{
core::steps::{ core::steps::{WasmModuleConfigRequest, WasmRegistrationWorkflowData},
workflow_data::{AnyWorkflowData, WasmRegistrationWorkflowData},
WasmModuleConfigRequest,
},
workflow::WorkflowId, workflow::WorkflowId,
}; };
@@ -703,17 +682,18 @@ async fn test_wasm_module_execution() {
}; };
let config_request = WasmModuleConfigRequest { descriptor }; let config_request = WasmModuleConfigRequest { descriptor };
let workflow_data = AnyWorkflowData::WasmRegistration(WasmRegistrationWorkflowData { let workflow_data = WasmRegistrationWorkflowData {
config: config_request, config: config_request,
wasm_bytes: None, wasm_bytes: None,
sha256_hash: None, sha256_hash: None,
file_size_bytes: None, file_size_bytes: None,
module_uuid: None, module_uuid: None,
app_context: Some(app_context.clone()), app_context: Some(app_context.clone()),
}); };
// Start workflow // Start workflow
let instance_id = engine let instance_id = engines
.wasm_registration
.start_workflow(WorkflowId::new("wasm_module_registration"), workflow_data) .start_workflow(WorkflowId::new("wasm_module_registration"), workflow_data)
.await .await
.expect("Failed to start workflow"); .expect("Failed to start workflow");
@@ -721,24 +701,25 @@ async fn test_wasm_module_execution() {
// Wait for workflow to complete // Wait for workflow to complete
let timeout = Duration::from_secs(30); let timeout = Duration::from_secs(30);
let start = std::time::Instant::now(); let start = std::time::Instant::now();
let mut module_uuid: Option<Uuid> = None;
loop { let module_uuid = loop {
if start.elapsed() > timeout { if start.elapsed() > timeout {
panic!("Workflow timeout"); panic!("Workflow timeout");
} }
let state = engine let state = engines
.wasm_registration
.get_status(instance_id) .get_status(instance_id)
.expect("Failed to get workflow status"); .expect("Failed to get workflow status");
match state.status { match state.status {
smg::workflow::WorkflowStatus::Completed => { smg::workflow::WorkflowStatus::Completed => {
// Extract module UUID from typed workflow data // Extract module UUID from typed workflow data
if let AnyWorkflowData::WasmRegistration(ref data) = state.context.data { break state
module_uuid = data.module_uuid; .context
} .data
break; .module_uuid
.expect("Module UUID should be in context");
} }
smg::workflow::WorkflowStatus::Failed => { smg::workflow::WorkflowStatus::Failed => {
panic!("Workflow failed: {:?}", state); panic!("Workflow failed: {:?}", state);
@@ -747,9 +728,7 @@ async fn test_wasm_module_execution() {
tokio::time::sleep(Duration::from_millis(100)).await; tokio::time::sleep(Duration::from_millis(100)).await;
} }
} }
} };
let module_uuid = module_uuid.expect("Module UUID should be in context");
// Verify module is registered // Verify module is registered
let module = wasm_manager let module = wasm_manager