* Copyright(c) 2024-2026 China Telecom Cloud Technologies Co., Ltd. All rights
* reserved. ctscat is licensed under Mulan PSL v2. You can use this software
* according to the terms and conditions of the Mulan PSL V2. You may obtain a
* copy of Mulan PSL v2 at: http://license.coscl.org.cn/MulanPSL2.
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY
* KIND, EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT,
* MERCHANTABILITY OR FIT FOR A PARTICULAR PURPOSE. See the Mulan PSL v2 for
* more details.
*/
use anyhow::{Context, Result};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::path::Path;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkflowConfig {
pub name: String,
#[serde(default)]
pub description: String,
#[serde(default = "default_version")]
pub version: String,
pub tasks: Vec<TaskConfig>,
#[serde(default)]
pub variables: HashMap<String, String>,
#[serde(default = "default_execution_policy")]
pub execution_policy: ExecutionPolicy,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TaskConfig {
pub name: String,
pub task_type: String,
#[serde(default)]
pub parameters: HashMap<String, serde_yaml::Value>,
#[serde(default)]
pub depends_on: Vec<String>,
#[serde(default)]
pub retry: RetryConfig,
#[serde(default = "default_timeout")]
pub timeout: u64,
#[serde(default)]
pub parallel: bool,
#[serde(default)]
pub condition: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct RetryConfig {
#[serde(default = "default_max_retries")]
pub max_attempts: u32,
#[serde(default = "default_retry_interval")]
pub interval: u64,
#[serde(default = "default_backoff_multiplier")]
pub backoff_multiplier: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(rename_all = "snake_case")]
pub enum ExecutionPolicy {
Sequential,
Parallel,
DAG,
}
impl WorkflowConfig {
pub fn from_file<P: AsRef<Path>>(path: P) -> Result<Self> {
let content = std::fs::read_to_string(&path).context("读取工作流文件失败")?;
Self::from_yaml(&content)
}
pub fn from_yaml(yaml: &str) -> Result<Self> {
serde_yaml::from_str(yaml).context("解析工作流配置失败")
}
pub fn validate(&self) -> Result<()> {
let mut task_names = std::collections::HashSet::new();
for task in &self.tasks {
if !task_names.insert(&task.name) {
anyhow::bail!("任务名称重复: {}", task.name);
}
}
for task in &self.tasks {
for dep in &task.depends_on {
if !task_names.contains(dep) {
anyhow::bail!("任务 {} 依赖不存在的任务: {}", task.name, dep);
}
}
}
if self.execution_policy == ExecutionPolicy::Sequential {
for task in &self.tasks {
if task.parallel {
anyhow::bail!("顺序执行模式下任务 {} 不能设置为并行", task.name);
}
}
}
Ok(())
}
pub fn topological_sort(&self) -> Result<Vec<&TaskConfig>> {
let mut sorted = Vec::new();
let mut visited = std::collections::HashSet::new();
let mut temp_visited = std::collections::HashSet::new();
fn visit<'a>(
task_name: &str,
tasks: &'a [TaskConfig],
sorted: &mut Vec<&'a TaskConfig>,
visited: &mut std::collections::HashSet<String>,
temp_visited: &mut std::collections::HashSet<String>,
) -> Result<()> {
if visited.contains(task_name) {
return Ok(());
}
if temp_visited.contains(task_name) {
anyhow::bail!("任务依赖中存在循环: {}", task_name);
}
temp_visited.insert(task_name.to_string());
if let Some(task) = tasks.iter().find(|t| t.name == task_name) {
for dep in &task.depends_on {
visit(dep, tasks, sorted, visited, temp_visited)?;
}
sorted.push(task);
}
temp_visited.remove(task_name);
visited.insert(task_name.to_string());
Ok(())
}
for task in &self.tasks {
visit(
&task.name,
&self.tasks,
&mut sorted,
&mut visited,
&mut temp_visited,
)?;
}
Ok(sorted)
}
pub fn substitute_variables(&mut self, vars: &HashMap<String, String>) {
for (key, value) in vars {
self.variables.insert(key.clone(), value.clone());
}
}
}
fn default_version() -> String {
"1.0".to_string()
}
fn default_execution_policy() -> ExecutionPolicy {
ExecutionPolicy::Sequential
}
fn default_timeout() -> u64 {
3600
}
fn default_max_retries() -> u32 {
0
}
fn default_retry_interval() -> u64 {
5
}
fn default_backoff_multiplier() -> f32 {
2.0
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_workflow_from_yaml() {
let yaml = r#"
name: test_workflow
description: Test workflow
version: 1.0
tasks:
- name: task1
task_type: sync
parameters:
tracking_id: 1
depends_on: []
- name: task2
task_type: classify
parameters:
limit: 100
depends_on:
- task1
"#;
let workflow = WorkflowConfig::from_yaml(yaml);
assert!(workflow.is_ok());
let w = workflow.unwrap();
assert_eq!(w.name, "test_workflow");
assert_eq!(w.tasks.len(), 2);
}
#[test]
fn test_workflow_validation() {
let yaml = r#"
name: test_workflow
tasks:
- name: task1
task_type: sync
parameters: {}
depends_on: []
- name: task2
task_type: classify
parameters: {}
depends_on:
- task1
"#;
let workflow = WorkflowConfig::from_yaml(yaml).unwrap();
assert!(workflow.validate().is_ok());
}
#[test]
fn test_topological_sort() {
let yaml = r#"
name: test_workflow
tasks:
- name: task1
task_type: sync
parameters: {}
depends_on: []
- name: task2
task_type: classify
parameters: {}
depends_on:
- task1
- name: task3
task_type: compare
parameters: {}
depends_on:
- task2
"#;
let workflow = WorkflowConfig::from_yaml(yaml).unwrap();
let sorted = workflow.topological_sort().unwrap();
assert_eq!(sorted.len(), 3);
assert_eq!(sorted[0].name, "task1");
assert_eq!(sorted[1].name, "task2");
assert_eq!(sorted[2].name, "task3");
}
}