use async_trait::async_trait;
use scirs2_core::random::Rng;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::{Duration, SystemTime};
use thiserror::Error;
use tokio::sync::{RwLock, Semaphore};
use tracing::{debug, error, info, warn};
use voirs_sdk::{AudioBuffer, VoirsError};
use crate::caching::CacheConfig;
use crate::pronunciation::PronunciationEvaluatorImpl;
use crate::quality::QualityEvaluator;
use crate::traits::{
PronunciationEvaluator as PronunciationEvaluatorTrait, PronunciationScore,
QualityEvaluationConfig, QualityEvaluator as QualityEvaluatorTrait, QualityScore,
};
#[derive(Error, Debug)]
pub enum WorkflowError {
#[error("Stage '{stage}' execution failed: {message}")]
StageExecutionError {
stage: String,
message: String,
#[source]
source: Option<Box<dyn std::error::Error + Send + Sync>>,
},
#[error("Workflow validation failed: {message}")]
ValidationError {
message: String,
},
#[error("Workflow configuration error: {message}")]
ConfigurationError {
message: String,
},
#[error("Dependency error: {message}")]
DependencyError {
message: String,
},
#[error("Workflow timed out after {duration:?}")]
TimeoutError {
duration: Duration,
},
#[error("Condition evaluation failed: {message}")]
ConditionError {
message: String,
},
#[error("VoiRS error: {0}")]
VoirsError(#[from] VoirsError),
#[error("Serialization error: {0}")]
SerializationError(#[from] serde_json::Error),
#[error("IO error: {0}")]
IoError(#[from] std::io::Error),
#[error("Evaluation error: {0}")]
EvaluationError(#[from] crate::EvaluationError),
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum StageType {
QualityEvaluation,
PronunciationEvaluation,
Preprocessing,
FeatureExtraction,
StatisticalAnalysis,
ExportResults,
Custom(String),
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub enum StageCondition {
Always,
OnSuccess,
OnFailure,
MetricThreshold {
metric: String,
min_value: f64,
max_value: Option<f64>,
},
Custom(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StageConfig {
pub name: String,
pub stage_type: StageType,
pub condition: StageCondition,
pub max_retries: usize,
pub timeout_seconds: Option<u64>,
pub enable_cache: bool,
pub parameters: HashMap<String, serde_json::Value>,
pub dependencies: Vec<String>,
}
impl Default for StageConfig {
fn default() -> Self {
Self {
name: String::new(),
stage_type: StageType::Custom("default".to_string()),
condition: StageCondition::Always,
max_retries: 3,
timeout_seconds: Some(300), enable_cache: true,
parameters: HashMap::new(),
dependencies: Vec::new(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StageResult {
pub stage_name: String,
pub status: StageStatus,
pub duration_ms: u64,
pub quality_score: Option<QualityScore>,
pub custom_results: HashMap<String, serde_json::Value>,
pub error_message: Option<String>,
pub retry_count: usize,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum StageStatus {
Pending,
Running,
Success,
Failed,
Skipped,
Timeout,
}
#[async_trait]
pub trait WorkflowStageExecutor: Send + Sync {
fn config(&self) -> &StageConfig;
async fn execute(
&self,
audio: &AudioBuffer,
reference: Option<&AudioBuffer>,
context: &WorkflowContext,
) -> Result<StageResult, WorkflowError>;
fn validate(&self) -> Result<(), WorkflowError> {
Ok(())
}
fn dependencies(&self) -> Vec<String> {
self.config().dependencies.clone()
}
}
pub struct QualityEvaluationStage {
config: StageConfig,
evaluator: Arc<RwLock<QualityEvaluator>>,
}
impl QualityEvaluationStage {
pub async fn new(config: StageConfig) -> Result<Self, WorkflowError> {
let evaluator =
QualityEvaluator::new()
.await
.map_err(|e| WorkflowError::ConfigurationError {
message: format!("Failed to create quality evaluator: {}", e),
})?;
Ok(Self {
config,
evaluator: Arc::new(RwLock::new(evaluator)),
})
}
}
#[async_trait]
impl WorkflowStageExecutor for QualityEvaluationStage {
fn config(&self) -> &StageConfig {
&self.config
}
async fn execute(
&self,
audio: &AudioBuffer,
reference: Option<&AudioBuffer>,
_context: &WorkflowContext,
) -> Result<StageResult, WorkflowError> {
let start = SystemTime::now();
let evaluator = self.evaluator.read().await;
let eval_config = QualityEvaluationConfig::default();
let quality_score = evaluator
.evaluate_quality(audio, reference, Some(&eval_config))
.await?;
let duration_ms = SystemTime::now()
.duration_since(start)
.unwrap_or(Duration::ZERO)
.as_millis() as u64;
Ok(StageResult {
stage_name: self.config.name.clone(),
status: StageStatus::Success,
duration_ms,
quality_score: Some(quality_score),
custom_results: HashMap::new(),
error_message: None,
retry_count: 0,
})
}
}
pub struct PronunciationEvaluationStage {
config: StageConfig,
evaluator: Arc<RwLock<PronunciationEvaluatorImpl>>,
}
impl PronunciationEvaluationStage {
pub async fn new(config: StageConfig) -> Result<Self, WorkflowError> {
let evaluator = PronunciationEvaluatorImpl::new().await.map_err(|e| {
WorkflowError::ConfigurationError {
message: format!("Failed to create pronunciation evaluator: {}", e),
}
})?;
Ok(Self {
config,
evaluator: Arc::new(RwLock::new(evaluator)),
})
}
}
#[async_trait]
impl WorkflowStageExecutor for PronunciationEvaluationStage {
fn config(&self) -> &StageConfig {
&self.config
}
async fn execute(
&self,
audio: &AudioBuffer,
_reference: Option<&AudioBuffer>,
context: &WorkflowContext,
) -> Result<StageResult, WorkflowError> {
let start = SystemTime::now();
let expected_text = "Hello world";
let _language = "en-US";
let evaluator = self.evaluator.read().await;
let pronunciation_score = evaluator
.evaluate_pronunciation(audio, expected_text, None)
.await?;
let duration_ms = SystemTime::now()
.duration_since(start)
.unwrap_or(Duration::ZERO)
.as_millis() as u64;
let mut custom_results = HashMap::new();
custom_results.insert(
"pronunciation_score".to_string(),
serde_json::json!({
"overall_score": pronunciation_score.overall_score,
"fluency_score": pronunciation_score.fluency_score,
"rhythm_score": pronunciation_score.rhythm_score,
}),
);
Ok(StageResult {
stage_name: self.config.name.clone(),
status: StageStatus::Success,
duration_ms,
quality_score: None,
custom_results,
error_message: None,
retry_count: 0,
})
}
}
pub struct ExportResultsStage {
config: StageConfig,
}
impl ExportResultsStage {
pub fn new(config: StageConfig) -> Self {
Self { config }
}
}
#[async_trait]
impl WorkflowStageExecutor for ExportResultsStage {
fn config(&self) -> &StageConfig {
&self.config
}
async fn execute(
&self,
_audio: &AudioBuffer,
_reference: Option<&AudioBuffer>,
context: &WorkflowContext,
) -> Result<StageResult, WorkflowError> {
let start = SystemTime::now();
let output_path = self
.config
.parameters
.get("output_path")
.and_then(|v| v.as_str())
.ok_or_else(|| WorkflowError::ConfigurationError {
message: "Missing 'output_path' parameter for export stage".to_string(),
})?;
let workflow_results = context.get_all_results();
let output_data = serde_json::to_string_pretty(&workflow_results)?;
tokio::fs::write(output_path, output_data).await?;
let duration_ms = SystemTime::now()
.duration_since(start)
.unwrap_or(Duration::ZERO)
.as_millis() as u64;
info!("Exported workflow results to: {}", output_path);
Ok(StageResult {
stage_name: self.config.name.clone(),
status: StageStatus::Success,
duration_ms,
quality_score: None,
custom_results: HashMap::new(),
error_message: None,
retry_count: 0,
})
}
}
#[derive(Clone)]
pub struct WorkflowContext {
parameters: Arc<RwLock<HashMap<String, serde_json::Value>>>,
results: Arc<RwLock<HashMap<String, StageResult>>>,
cache_config: Option<CacheConfig>,
}
impl WorkflowContext {
pub fn new() -> Self {
Self {
parameters: Arc::new(RwLock::new(HashMap::new())),
results: Arc::new(RwLock::new(HashMap::new())),
cache_config: None,
}
}
pub async fn set_parameter(&self, key: String, value: serde_json::Value) {
let mut params = self.parameters.write().await;
params.insert(key, value);
}
pub fn get_parameter(&self, key: &str) -> Option<serde_json::Value> {
None
}
pub async fn set_result(&self, stage_name: String, result: StageResult) {
let mut results = self.results.write().await;
results.insert(stage_name, result);
}
pub async fn get_result(&self, stage_name: &str) -> Option<StageResult> {
let results = self.results.read().await;
results.get(stage_name).cloned()
}
pub fn get_all_results(&self) -> HashMap<String, StageResult> {
HashMap::new()
}
pub fn with_cache(mut self, cache_config: CacheConfig) -> Self {
self.cache_config = Some(cache_config);
self
}
}
impl Default for WorkflowContext {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkflowConfig {
pub name: String,
pub description: Option<String>,
pub max_parallel_stages: usize,
pub global_timeout_seconds: Option<u64>,
pub enable_cache: bool,
pub cache_config: Option<CacheConfig>,
pub retry_policy: RetryPolicy,
}
impl Default for WorkflowConfig {
fn default() -> Self {
Self {
name: "default_workflow".to_string(),
description: None,
max_parallel_stages: 4,
global_timeout_seconds: Some(1800), enable_cache: true,
cache_config: None,
retry_policy: RetryPolicy::default(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RetryPolicy {
pub max_attempts: usize,
pub initial_delay_ms: u64,
pub backoff_multiplier: f64,
pub max_delay_ms: u64,
}
impl Default for RetryPolicy {
fn default() -> Self {
Self {
max_attempts: 3,
initial_delay_ms: 1000,
backoff_multiplier: 2.0,
max_delay_ms: 30000,
}
}
}
pub struct Workflow {
config: WorkflowConfig,
stages: Vec<Arc<dyn WorkflowStageExecutor>>,
context: WorkflowContext,
semaphore: Arc<Semaphore>,
}
impl Workflow {
pub fn new(config: WorkflowConfig, stages: Vec<Arc<dyn WorkflowStageExecutor>>) -> Self {
let semaphore = Arc::new(Semaphore::new(config.max_parallel_stages));
let mut context = WorkflowContext::new();
if config.enable_cache {
let cache_config = config.cache_config.clone().unwrap_or_default();
context = context.with_cache(cache_config);
}
Self {
config,
stages,
context,
semaphore,
}
}
pub async fn execute(
&self,
audio: &AudioBuffer,
reference: Option<&AudioBuffer>,
) -> Result<WorkflowResult, WorkflowError> {
let workflow_start = SystemTime::now();
info!("Starting workflow: {}", self.config.name);
let mut stage_results = Vec::new();
for stage in &self.stages {
let _permit =
self.semaphore
.acquire()
.await
.map_err(|e| WorkflowError::StageExecutionError {
stage: stage.config().name.clone(),
message: format!("Failed to acquire semaphore: {}", e),
source: None,
})?;
if !self
.should_execute_stage(stage.config(), &stage_results)
.await?
{
info!("Skipping stage '{}' due to condition", stage.config().name);
stage_results.push(StageResult {
stage_name: stage.config().name.clone(),
status: StageStatus::Skipped,
duration_ms: 0,
quality_score: None,
custom_results: HashMap::new(),
error_message: None,
retry_count: 0,
});
continue;
}
let result = self
.execute_stage_with_retry(stage.as_ref(), audio, reference)
.await?;
self.context
.set_result(stage.config().name.clone(), result.clone())
.await;
stage_results.push(result);
}
let total_duration_ms = SystemTime::now()
.duration_since(workflow_start)
.unwrap_or(Duration::ZERO)
.as_millis() as u64;
info!(
"Workflow '{}' completed in {}ms",
self.config.name, total_duration_ms
);
Ok(WorkflowResult {
workflow_name: self.config.name.clone(),
stage_results,
total_duration_ms,
status: WorkflowStatus::Success,
error_message: None,
})
}
async fn should_execute_stage(
&self,
config: &StageConfig,
previous_results: &[StageResult],
) -> Result<bool, WorkflowError> {
match &config.condition {
StageCondition::Always => Ok(true),
StageCondition::OnSuccess => {
if let Some(last_result) = previous_results.last() {
Ok(last_result.status == StageStatus::Success)
} else {
Ok(true)
}
}
StageCondition::OnFailure => {
if let Some(last_result) = previous_results.last() {
Ok(last_result.status == StageStatus::Failed)
} else {
Ok(false)
}
}
StageCondition::MetricThreshold {
metric,
min_value,
max_value,
} => {
for result in previous_results.iter().rev() {
if let Some(quality) = &result.quality_score {
let value = match metric.as_str() {
"overall_score" => quality.overall_score as f64,
_ => continue, };
let meets_min = value >= *min_value;
let meets_max = max_value.is_none_or(|max| value <= max);
return Ok(meets_min && meets_max);
}
}
Ok(false)
}
StageCondition::Custom(_) => {
warn!("Custom conditions not yet implemented, defaulting to true");
Ok(true)
}
}
}
async fn execute_stage_with_retry(
&self,
stage: &dyn WorkflowStageExecutor,
audio: &AudioBuffer,
reference: Option<&AudioBuffer>,
) -> Result<StageResult, WorkflowError> {
let config = stage.config();
let max_retries = config.max_retries;
let mut retry_count = 0;
let mut last_error = None;
while retry_count <= max_retries {
match stage.execute(audio, reference, &self.context).await {
Ok(mut result) => {
result.retry_count = retry_count;
return Ok(result);
}
Err(e) => {
error!(
"Stage '{}' failed (attempt {}/{}): {}",
config.name,
retry_count + 1,
max_retries + 1,
e
);
last_error = Some(e);
retry_count += 1;
if retry_count <= max_retries {
let delay_ms = self.calculate_backoff_delay(retry_count);
tokio::time::sleep(Duration::from_millis(delay_ms)).await;
}
}
}
}
Ok(StageResult {
stage_name: config.name.clone(),
status: StageStatus::Failed,
duration_ms: 0,
quality_score: None,
custom_results: HashMap::new(),
error_message: Some(
last_error
.map(|e| e.to_string())
.unwrap_or_else(|| "Unknown error".to_string()),
),
retry_count,
})
}
fn calculate_backoff_delay(&self, retry_count: usize) -> u64 {
let policy = &self.config.retry_policy;
let delay =
policy.initial_delay_ms as f64 * policy.backoff_multiplier.powi(retry_count as i32 - 1);
delay.min(policy.max_delay_ms as f64) as u64
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkflowResult {
pub workflow_name: String,
pub stage_results: Vec<StageResult>,
pub total_duration_ms: u64,
pub status: WorkflowStatus,
pub error_message: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum WorkflowStatus {
Success,
PartialSuccess,
Failed,
Cancelled,
}
pub struct WorkflowBuilder {
config: WorkflowConfig,
stages: Vec<Arc<dyn WorkflowStageExecutor>>,
}
impl WorkflowBuilder {
pub fn new(name: impl Into<String>) -> Self {
Self {
config: WorkflowConfig {
name: name.into(),
..Default::default()
},
stages: Vec::new(),
}
}
pub fn description(mut self, description: impl Into<String>) -> Self {
self.config.description = Some(description.into());
self
}
pub fn max_parallel_stages(mut self, max: usize) -> Self {
self.config.max_parallel_stages = max;
self
}
pub fn global_timeout(mut self, seconds: u64) -> Self {
self.config.global_timeout_seconds = Some(seconds);
self
}
pub fn enable_cache(mut self, enable: bool) -> Self {
self.config.enable_cache = enable;
self
}
pub fn cache_config(mut self, config: CacheConfig) -> Self {
self.config.cache_config = Some(config);
self
}
pub fn add_stage(mut self, stage: Arc<dyn WorkflowStageExecutor>) -> Self {
self.stages.push(stage);
self
}
pub fn build(self) -> Result<Workflow, WorkflowError> {
if self.stages.is_empty() {
return Err(WorkflowError::ValidationError {
message: "Workflow must have at least one stage".to_string(),
});
}
self.validate_dependencies()?;
Ok(Workflow::new(self.config, self.stages))
}
fn validate_dependencies(&self) -> Result<(), WorkflowError> {
let stage_names: Vec<String> = self
.stages
.iter()
.map(|s| s.config().name.clone())
.collect();
for stage in &self.stages {
for dep in stage.dependencies() {
if !stage_names.contains(&dep) {
return Err(WorkflowError::DependencyError {
message: format!(
"Stage '{}' depends on '{}' which is not in the workflow",
stage.config().name,
dep
),
});
}
}
}
Ok(())
}
}
pub struct WorkflowStage;
impl WorkflowStage {
pub fn quality_evaluation(config: StageConfig) -> Arc<dyn WorkflowStageExecutor> {
Arc::new(ExportResultsStage::new(config))
}
pub fn pronunciation_evaluation(config: StageConfig) -> Arc<dyn WorkflowStageExecutor> {
Arc::new(ExportResultsStage::new(config))
}
pub fn export_results(config: StageConfig) -> Arc<dyn WorkflowStageExecutor> {
Arc::new(ExportResultsStage::new(config))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_stage_config_default() {
let config = StageConfig::default();
assert_eq!(config.max_retries, 3);
assert!(config.enable_cache);
assert_eq!(config.condition, StageCondition::Always);
}
#[test]
fn test_workflow_config_default() {
let config = WorkflowConfig::default();
assert_eq!(config.max_parallel_stages, 4);
assert!(config.enable_cache);
}
#[test]
fn test_retry_policy_default() {
let policy = RetryPolicy::default();
assert_eq!(policy.max_attempts, 3);
assert_eq!(policy.initial_delay_ms, 1000);
assert!((policy.backoff_multiplier - 2.0).abs() < f64::EPSILON);
}
#[test]
fn test_workflow_builder() {
let builder = WorkflowBuilder::new("test_workflow")
.description("Test workflow")
.max_parallel_stages(2)
.enable_cache(false);
assert_eq!(builder.config.name, "test_workflow");
assert_eq!(builder.config.max_parallel_stages, 2);
assert!(!builder.config.enable_cache);
}
#[test]
fn test_workflow_context() {
let context = WorkflowContext::new();
assert!(context.cache_config.is_none());
}
#[test]
fn test_stage_status() {
let status = StageStatus::Success;
assert_eq!(status, StageStatus::Success);
assert_ne!(status, StageStatus::Failed);
}
#[test]
fn test_workflow_error_display() {
let error = WorkflowError::ValidationError {
message: "Test error".to_string(),
};
assert!(error.to_string().contains("Test error"));
}
#[tokio::test]
async fn test_export_results_stage_creation() {
let config = StageConfig {
name: "export".to_string(),
stage_type: StageType::ExportResults,
..Default::default()
};
let stage = ExportResultsStage::new(config);
assert_eq!(stage.config().name, "export");
}
#[tokio::test]
async fn test_workflow_context_async_operations() {
let context = WorkflowContext::new();
context
.set_parameter("test_key".to_string(), serde_json::json!("test_value"))
.await;
let result = StageResult {
stage_name: "test_stage".to_string(),
status: StageStatus::Success,
duration_ms: 100,
quality_score: None,
custom_results: HashMap::new(),
error_message: None,
retry_count: 0,
};
context
.set_result("test_stage".to_string(), result.clone())
.await;
let retrieved = context.get_result("test_stage").await;
assert!(retrieved.is_some());
assert_eq!(retrieved.unwrap().stage_name, "test_stage");
}
}