[model-gateway] refactor workflow engine from type erasure to typed engines (#16973)
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
// 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,8 +460,9 @@ impl JobQueue {
|
||||
|
||||
let timeout_duration = Duration::from_secs(30);
|
||||
|
||||
let result =
|
||||
Self::wait_for_workflow_completion(engine, instance_id, url, timeout_duration)
|
||||
let result = engines
|
||||
.worker_removal
|
||||
.wait_for_completion(instance_id, url, timeout_duration)
|
||||
.await;
|
||||
|
||||
// Clean up job status when removing worker
|
||||
@@ -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,
|
||||
)
|
||||
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,8 +525,9 @@ impl JobQueue {
|
||||
|
||||
let timeout_duration = Duration::from_secs(60); // 1 minute
|
||||
|
||||
Self::wait_for_workflow_completion(
|
||||
engine,
|
||||
engines
|
||||
.wasm_removal
|
||||
.wait_for_completion(
|
||||
instance_id,
|
||||
&request.module_uuid.to_string(),
|
||||
timeout_duration,
|
||||
@@ -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,
|
||||
)
|
||||
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,12 +761,9 @@ 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,
|
||||
)
|
||||
engines
|
||||
.tokenizer
|
||||
.wait_for_completion(instance_id, &config.id, timeout_duration)
|
||||
.await
|
||||
}
|
||||
Job::RemoveTokenizer { request } => {
|
||||
@@ -710,143 +781,6 @@ impl JobQueue {
|
||||
}
|
||||
}
|
||||
|
||||
/// Start a workflow and return its instance ID
|
||||
async fn start_worker_workflow(
|
||||
engine: &Arc<AppWorkflowEngine>,
|
||||
config: &WorkerConfigRequest,
|
||||
context: &Arc<AppContext>,
|
||||
) -> Result<WorkflowInstanceId, String> {
|
||||
// Select workflow based on runtime field
|
||||
let (workflow_id, workflow_data) = match config.runtime.as_deref() {
|
||||
Some("external") => (
|
||||
WorkflowId::new("external_worker_registration"),
|
||||
create_external_worker_workflow_data(config.clone(), Arc::clone(context)),
|
||||
),
|
||||
_ => (
|
||||
WorkflowId::new("local_worker_registration"),
|
||||
create_local_worker_workflow_data(config.clone(), Arc::clone(context)),
|
||||
),
|
||||
};
|
||||
|
||||
engine
|
||||
.start_workflow(workflow_id, workflow_data)
|
||||
.await
|
||||
.map_err(|e| format!("Failed to start worker registration workflow: {:?}", e))
|
||||
}
|
||||
|
||||
/// Start worker removal workflow
|
||||
async fn start_worker_removal_workflow(
|
||||
engine: &Arc<AppWorkflowEngine>,
|
||||
url: &str,
|
||||
context: &Arc<AppContext>,
|
||||
) -> Result<WorkflowInstanceId, String> {
|
||||
let workflow_data = create_worker_removal_workflow_data(
|
||||
url.to_string(),
|
||||
context.router_config.dp_aware,
|
||||
Arc::clone(context),
|
||||
);
|
||||
|
||||
engine
|
||||
.start_workflow(WorkflowId::new("worker_removal"), workflow_data)
|
||||
.await
|
||||
.map_err(|e| format!("Failed to start worker removal workflow: {:?}", e))
|
||||
}
|
||||
|
||||
/// Start worker update workflow
|
||||
async fn start_worker_update_workflow(
|
||||
engine: &Arc<AppWorkflowEngine>,
|
||||
url: &str,
|
||||
update: &WorkerUpdateRequest,
|
||||
context: &Arc<AppContext>,
|
||||
) -> Result<WorkflowInstanceId, String> {
|
||||
let workflow_data = create_worker_update_workflow_data(
|
||||
url.to_string(),
|
||||
update.clone(),
|
||||
Arc::clone(context),
|
||||
);
|
||||
|
||||
engine
|
||||
.start_workflow(WorkflowId::new("worker_update"), workflow_data)
|
||||
.await
|
||||
.map_err(|e| format!("Failed to start worker update workflow: {:?}", e))
|
||||
}
|
||||
|
||||
/// Start MCP server registration workflow
|
||||
async fn start_mcp_registration_workflow(
|
||||
engine: &Arc<AppWorkflowEngine>,
|
||||
config: &McpServerConfigRequest,
|
||||
context: &Arc<AppContext>,
|
||||
) -> Result<WorkflowInstanceId, String> {
|
||||
let workflow_data = create_mcp_workflow_data(config.clone(), Arc::clone(context));
|
||||
|
||||
engine
|
||||
.start_workflow(WorkflowId::new("mcp_registration"), workflow_data)
|
||||
.await
|
||||
.map_err(|e| format!("Failed to start MCP registration workflow: {:?}", e))
|
||||
}
|
||||
|
||||
/// Wait for workflow completion with adaptive polling
|
||||
async fn wait_for_workflow_completion(
|
||||
engine: &Arc<AppWorkflowEngine>,
|
||||
instance_id: WorkflowInstanceId,
|
||||
worker_url: &str,
|
||||
timeout_duration: Duration,
|
||||
) -> Result<String, String> {
|
||||
let start = std::time::Instant::now();
|
||||
let mut poll_interval = Duration::from_millis(100);
|
||||
let max_poll_interval = Duration::from_millis(2000);
|
||||
let poll_backoff = Duration::from_millis(200);
|
||||
|
||||
loop {
|
||||
// Check timeout
|
||||
if start.elapsed() > timeout_duration {
|
||||
return Err(format!(
|
||||
"Workflow timeout after {}s for worker {}",
|
||||
timeout_duration.as_secs(),
|
||||
worker_url
|
||||
));
|
||||
}
|
||||
|
||||
// Get workflow status
|
||||
let state = engine
|
||||
.get_status(instance_id)
|
||||
.map_err(|e| format!("Failed to get workflow status: {:?}", e))?;
|
||||
|
||||
let result = match state.status {
|
||||
WorkflowStatus::Completed => Ok(format!(
|
||||
"Worker {} registered and activated successfully via workflow",
|
||||
worker_url
|
||||
)),
|
||||
WorkflowStatus::Failed => {
|
||||
let current_step = state.current_step.as_ref();
|
||||
let step_name = current_step
|
||||
.map(|s| s.to_string())
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
let error_msg = current_step
|
||||
.and_then(|step_id| state.step_states.get(step_id))
|
||||
.and_then(|s| s.last_error.as_deref())
|
||||
.unwrap_or("Unknown error");
|
||||
Err(format!(
|
||||
"Workflow failed at step {}: {}",
|
||||
step_name, error_msg
|
||||
))
|
||||
}
|
||||
WorkflowStatus::Cancelled => {
|
||||
Err(format!("Workflow cancelled for worker {}", worker_url))
|
||||
}
|
||||
WorkflowStatus::Pending | WorkflowStatus::Paused | WorkflowStatus::Running => {
|
||||
tokio::time::sleep(poll_interval).await;
|
||||
poll_interval = (poll_interval + poll_backoff).min(max_poll_interval);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
// Clean up terminal workflow states
|
||||
engine.state_store().cleanup_if_terminal(instance_id);
|
||||
return result;
|
||||
}
|
||||
}
|
||||
|
||||
/// Update job status on completion
|
||||
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,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
);
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user