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