use std::collections::HashMap; use std::sync::{Arc, Mutex}; use std::time::Duration; use anyhow::Result; use tokio::time::sleep; use tracing::{info, debug, error}; use crate::data::models::hedra_lipsync_record::HedraLipSyncStatus; use crate::data::repositories::hedra_lipsync_repository::HedraLipSyncRepository; use crate::infrastructure::bowong_text_video_agent_service::BowongTextVideoAgentService; use crate::infrastructure::event_bus::EventBusManager; use crate::data::models::bowong_text_video_agent::HedraTaskStatusParams; /// Hedra 任务轮询器 /// 负责在后台轮询 Hedra 任务状态并发送事件通知 pub struct HedraTaskPoller { /// 正在轮询的任务 polling_tasks: Arc>>, /// BowongTextVideoAgent 服务 service: Arc, /// 事件总线管理器 event_bus: Arc, /// 数据库仓库 repository: Arc, /// 是否正在运行 is_running: Arc>, } /// 轮询任务信息 #[derive(Debug, Clone)] struct PollingTask { task_id: String, record_id: String, start_time: std::time::Instant, poll_count: u32, max_polls: u32, } impl HedraTaskPoller { /// 创建新的 Hedra 任务轮询器 pub fn new( service: Arc, event_bus: Arc, repository: Arc, ) -> Self { Self { polling_tasks: Arc::new(Mutex::new(HashMap::new())), service, event_bus, repository, is_running: Arc::new(Mutex::new(false)), } } /// 添加任务到轮询队列 pub fn add_task(&self, task_id: String, record_id: String) -> Result<()> { let mut tasks = self.polling_tasks.lock().unwrap(); let polling_task = PollingTask { task_id: task_id.clone(), record_id: record_id.clone(), start_time: std::time::Instant::now(), poll_count: 0, max_polls: 360, // 30分钟 (5秒间隔 * 360次) }; tasks.insert(task_id.clone(), polling_task); info!("添加 Hedra 任务到轮询队列: task_id={}, record_id={}", task_id, record_id); Ok(()) } /// 移除轮询任务 pub fn remove_task(&self, task_id: &str) { let mut tasks = self.polling_tasks.lock().unwrap(); if tasks.remove(task_id).is_some() { debug!("从轮询队列移除任务: {}", task_id); } } /// 启动轮询器 pub async fn start(&self) -> Result<()> { { let mut running = self.is_running.lock().unwrap(); if *running { return Ok(()); } *running = true; } info!("启动 Hedra 任务轮询器"); // 启动轮询循环 let polling_tasks = Arc::clone(&self.polling_tasks); let service = Arc::clone(&self.service); let event_bus = Arc::clone(&self.event_bus); let repository = Arc::clone(&self.repository); let is_running = Arc::clone(&self.is_running); tokio::spawn(async move { while *is_running.lock().unwrap() { // 获取当前需要轮询的任务 let tasks_to_poll: Vec = { let tasks = polling_tasks.lock().unwrap(); tasks.values().cloned().collect() }; // 轮询每个任务 for mut task in tasks_to_poll { if let Err(e) = Self::poll_single_task( &mut task, &service, &event_bus, &repository, &polling_tasks, ).await { error!("轮询任务 {} 时出错: {}", task.task_id, e); } } // 等待 5 秒后继续下一轮轮询 sleep(Duration::from_secs(5)).await; } info!("Hedra 任务轮询器已停止"); }); Ok(()) } /// 停止轮询器 pub fn stop(&self) { let mut running = self.is_running.lock().unwrap(); *running = false; info!("停止 Hedra 任务轮询器"); } /// 轮询单个任务 async fn poll_single_task( task: &mut PollingTask, service: &BowongTextVideoAgentService, event_bus: &EventBusManager, repository: &HedraLipSyncRepository, polling_tasks: &Arc>>, ) -> Result<()> { task.poll_count += 1; debug!("轮询 Hedra 任务: {} (第 {} 次)", task.task_id, task.poll_count); // 检查是否超过最大轮询次数 if task.poll_count > task.max_polls { error!("任务 {} 轮询超时,移除任务", task.task_id); // 发送超时事件 let _ = event_bus.publish_hedra_task_progress( task.task_id.clone(), task.record_id.clone(), "failed".to_string(), None, None, Some("任务轮询超时".to_string()), ).await; // 更新数据库状态 let _ = repository.update_status_and_error( &task.record_id, HedraLipSyncStatus::Failed, Some("任务轮询超时".to_string()), ); // 从轮询队列移除 polling_tasks.lock().unwrap().remove(&task.task_id); return Ok(()); } // 查询任务状态 let status_params = HedraTaskStatusParams { task_id: task.task_id.clone(), }; match service.hedra_query_task_status(&status_params).await { Ok(status_response) => { debug!("任务 {} 状态: {}", task.task_id, status_response.status); // 映射状态 let mapped_status = match status_response.status.as_str() { "success" => "completed", "running" => "processing", "failed" => "failed", "cancelled" => "cancelled", "pending" | _ => "pending", }; // 提取视频URL let video_url = if status_response.status == "success" { status_response.result .as_ref() .and_then(|result| { result.get("video_url") .or_else(|| result.get("url")) .or_else(|| result.get("video")) .and_then(|v| v.as_str()) .map(|s| s.to_string()) }) } else { None }; // 发送进度事件 let _ = event_bus.publish_hedra_task_progress( task.task_id.clone(), task.record_id.clone(), mapped_status.to_string(), status_response.progress.map(|p| p as f64), video_url.clone(), status_response.error.clone(), ).await; // 更新数据库 let db_status = match mapped_status { "completed" => HedraLipSyncStatus::Completed, "processing" => HedraLipSyncStatus::Processing, "failed" => HedraLipSyncStatus::Failed, "cancelled" => HedraLipSyncStatus::Cancelled, _ => HedraLipSyncStatus::Pending, }; let _ = repository.update_progress_and_result( &task.record_id, db_status, status_response.progress.unwrap_or(0.0) as i32, video_url, status_response.error, ); // 如果任务完成或失败,从轮询队列移除 if matches!(mapped_status, "completed" | "failed" | "cancelled") { info!("任务 {} 已完成,状态: {}", task.task_id, mapped_status); polling_tasks.lock().unwrap().remove(&task.task_id); } } Err(e) => { error!("查询任务 {} 状态失败: {}", task.task_id, e); // 如果连续查询失败次数过多,标记为失败 if task.poll_count > 10 { let _ = event_bus.publish_hedra_task_progress( task.task_id.clone(), task.record_id.clone(), "failed".to_string(), None, None, Some(format!("状态查询失败: {}", e)), ).await; polling_tasks.lock().unwrap().remove(&task.task_id); } } } Ok(()) } /// 获取当前轮询任务数量 pub fn get_polling_task_count(&self) -> usize { self.polling_tasks.lock().unwrap().len() } }