use crate::workflow::node::NodeType;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkflowDefinition {
pub metadata: WorkflowMetadata,
#[serde(default)]
pub config: WorkflowConfig,
#[serde(default)]
pub nodes: Vec<NodeDefinition>,
#[serde(default)]
pub edges: Vec<EdgeDefinition>,
#[serde(default)]
pub agents: HashMap<String, LlmAgentConfig>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkflowMetadata {
pub id: String,
pub name: String,
#[serde(default)]
pub description: String,
#[serde(default)]
pub version: Option<String>,
#[serde(default)]
pub author: Option<String>,
#[serde(default)]
pub tags: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkflowConfig {
#[serde(default = "default_max_parallel")]
pub max_parallel: usize,
#[serde(default = "default_timeout")]
pub default_timeout_ms: u64,
#[serde(default)]
pub enable_checkpoints: bool,
#[serde(default)]
pub retry_policy: Option<RetryPolicy>,
}
impl Default for WorkflowConfig {
fn default() -> Self {
Self {
max_parallel: default_max_parallel(),
default_timeout_ms: default_timeout(),
enable_checkpoints: false,
retry_policy: None,
}
}
}
fn default_max_parallel() -> usize {
10
}
fn default_timeout() -> u64 {
60000 }
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum NodeDefinition {
Start {
id: String,
#[serde(default)]
name: Option<String>,
},
End {
id: String,
#[serde(default)]
name: Option<String>,
},
Task {
id: String,
name: String,
#[serde(flatten)]
executor: TaskExecutorDef,
#[serde(default)]
config: NodeConfigDef,
},
LlmAgent {
id: String,
name: String,
agent: AgentRef,
#[serde(default)]
prompt_template: Option<String>,
#[serde(default)]
config: NodeConfigDef,
},
Condition {
id: String,
name: String,
condition: ConditionDef,
#[serde(default)]
config: NodeConfigDef,
},
Parallel {
id: String,
name: String,
#[serde(default)]
config: NodeConfigDef,
},
Join {
id: String,
name: String,
#[serde(default)]
wait_for: Vec<String>,
#[serde(default)]
config: NodeConfigDef,
},
Loop {
id: String,
name: String,
#[serde(flatten)]
body: TaskExecutorDef,
condition: LoopConditionDef,
#[serde(default)]
max_iterations: u32,
#[serde(default)]
config: NodeConfigDef,
},
Transform {
id: String,
name: String,
#[serde(flatten)]
transform: TransformDef,
#[serde(default)]
config: NodeConfigDef,
},
SubWorkflow {
id: String,
name: String,
workflow_id: String,
#[serde(default)]
config: NodeConfigDef,
},
Wait {
id: String,
name: String,
event_type: String,
#[serde(default)]
config: NodeConfigDef,
},
}
impl NodeDefinition {
pub fn id(&self) -> &str {
match self {
NodeDefinition::Start { id, .. } => id,
NodeDefinition::End { id, .. } => id,
NodeDefinition::Task { id, .. } => id,
NodeDefinition::LlmAgent { id, .. } => id,
NodeDefinition::Condition { id, .. } => id,
NodeDefinition::Parallel { id, .. } => id,
NodeDefinition::Join { id, .. } => id,
NodeDefinition::Loop { id, .. } => id,
NodeDefinition::Transform { id, .. } => id,
NodeDefinition::SubWorkflow { id, .. } => id,
NodeDefinition::Wait { id, .. } => id,
}
}
pub fn node_type(&self) -> NodeType {
match self {
NodeDefinition::Start { .. } => NodeType::Start,
NodeDefinition::End { .. } => NodeType::End,
NodeDefinition::Task { .. } => NodeType::Task,
NodeDefinition::LlmAgent { .. } => NodeType::Agent,
NodeDefinition::Condition { .. } => NodeType::Condition,
NodeDefinition::Parallel { .. } => NodeType::Parallel,
NodeDefinition::Join { .. } => NodeType::Join,
NodeDefinition::Loop { .. } => NodeType::Loop,
NodeDefinition::Transform { .. } => NodeType::Transform,
NodeDefinition::SubWorkflow { .. } => NodeType::SubWorkflow,
NodeDefinition::Wait { .. } => NodeType::Wait,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EdgeDefinition {
pub from: String,
pub to: String,
#[serde(default)]
pub condition: Option<String>,
#[serde(default)]
pub label: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LlmAgentConfig {
pub model: String,
#[serde(default)]
pub system_prompt: Option<String>,
#[serde(default)]
pub temperature: Option<f32>,
#[serde(default)]
pub max_tokens: Option<u32>,
#[serde(default)]
pub context_window_size: Option<usize>,
#[serde(default)]
pub user_id: Option<String>,
#[serde(default)]
pub tenant_id: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum AgentRef {
Registry { agent_id: String },
Inline(Box<LlmAgentConfig>),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "executor_type", rename_all = "snake_case")]
pub enum TaskExecutorDef {
Function { function: String },
Http {
url: String,
#[serde(default)]
method: Option<String>,
},
Script { script: String },
None,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "condition_type", rename_all = "snake_case")]
pub enum ConditionDef {
Expression { expr: String },
Value {
field: String,
operator: String,
value: serde_json::Value,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "condition_type", rename_all = "snake_case")]
pub enum LoopConditionDef {
While { expr: String },
Until { expr: String },
Count { max: u32 },
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "transform_type", rename_all = "snake_case")]
pub enum TransformDef {
Template { template: String },
Expression { expr: String },
MapReduce {
#[serde(default)]
map: Option<String>,
#[serde(default)]
reduce: Option<String>,
},
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct NodeConfigDef {
#[serde(default)]
pub retry_policy: Option<RetryPolicy>,
#[serde(default)]
pub timeout_ms: Option<u64>,
#[serde(default)]
pub metadata: HashMap<String, String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RetryPolicy {
#[serde(default = "default_max_retries")]
pub max_retries: u32,
#[serde(default = "default_retry_delay")]
pub retry_delay_ms: u64,
#[serde(default = "default_exponential_backoff")]
pub exponential_backoff: bool,
#[serde(default = "default_max_delay")]
pub max_delay_ms: u64,
}
impl Default for RetryPolicy {
fn default() -> Self {
Self {
max_retries: default_max_retries(),
retry_delay_ms: default_retry_delay(),
exponential_backoff: default_exponential_backoff(),
max_delay_ms: default_max_delay(),
}
}
}
fn default_max_retries() -> u32 {
3
}
fn default_retry_delay() -> u64 {
1000
}
fn default_exponential_backoff() -> bool {
true
}
fn default_max_delay() -> u64 {
30000
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TimeoutConfig {
pub execution_timeout_ms: u64,
#[serde(default = "default_cancel_on_timeout")]
pub cancel_on_timeout: bool,
}
impl Default for TimeoutConfig {
fn default() -> Self {
Self {
execution_timeout_ms: 60000,
cancel_on_timeout: true,
}
}
}
fn default_cancel_on_timeout() -> bool {
true
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_workflow_yaml() {
let yaml = r#"
metadata:
id: test_workflow
name: Test Workflow
description: A test workflow
nodes:
- type: start
id: start
- type: llm_agent
id: agent1
name: First Agent
agent:
agent_id: my_agent
- type: end
id: end
edges:
- from: start
to: agent1
- from: agent1
to: end
"#;
let def: WorkflowDefinition = serde_yaml::from_str(yaml).unwrap();
assert_eq!(def.metadata.id, "test_workflow");
assert_eq!(def.nodes.len(), 3);
assert_eq!(def.edges.len(), 2);
}
#[test]
fn test_parse_agent_config() {
let yaml = r#"
metadata:
id: agent_config_test
name: Agent Config Test
nodes:
- type: start
id: start
- type: end
id: end
agents:
my_agent:
model: gpt-4
system_prompt: "You are helpful"
temperature: 0.7
max_tokens: 2000
"#;
let def: WorkflowDefinition = serde_yaml::from_str(yaml).unwrap();
assert_eq!(def.agents.len(), 1);
let agent = def.agents.get("my_agent").unwrap();
assert_eq!(agent.model, "gpt-4");
assert_eq!(agent.temperature, Some(0.7));
}
#[test]
fn test_parse_toml() {
let toml = r#"
[metadata]
id = "test_workflow"
name = "Test Workflow"
[[nodes]]
type = "start"
id = "start"
[[nodes]]
type = "end"
id = "end"
[[edges]]
from = "start"
to = "end"
"#;
let def: WorkflowDefinition = toml::from_str(toml).unwrap();
assert_eq!(def.metadata.id, "test_workflow");
assert_eq!(def.nodes.len(), 2);
}
}