Team Ai
Apppublic

PCNUSMSE/transcript_service

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes
task_manager.py462 linesDownload Raw Back to core
1"""任务管理模块2 3提供任务状态跟踪、进度管理和任务队列功能。4"""5 6import asyncio7import time8import uuid9from dataclasses import dataclass, field10from datetime import datetime, timedelta11from enum import Enum12from pathlib import Path13from typing import Dict, List, Optional, Callable, Any14from concurrent.futures import ThreadPoolExecutor15 16from ..core.config import get_config17from ..utils.logger import get_task_logger18from ..services.file_validator import get_file_validator19from ..services.oss_service import get_oss_service20from ..services.paraformer_service import get_paraformer_service21 22 23class TaskStatus(Enum):24    """任务状态"""25    PENDING = "pending"26    VALIDATING = "validating"27    UPLOADING = "uploading"28    TRANSCRIBING = "transcribing"29    COMPLETED = "completed"30    FAILED = "failed"31    CANCELLED = "cancelled"32 33 34class TaskPriority(Enum):35    """任务优先级"""36    LOW = 137    NORMAL = 238    HIGH = 339    URGENT = 440 41 42@dataclass43class TaskProgress:44    """任务进度信息"""45    stage: str = ""46    current: int = 047    total: int = 10048    message: str = ""49    percentage: float = 0.050    51    def update(self, current: int = None, total: int = None, message: str = None):52        """更新进度信息"""53        if current is not None:54            self.current = current55        if total is not None:56            self.total = total57        if message is not None:58            self.message = message59        60        if self.total > 0:61            self.percentage = min(100.0, (self.current / self.total) * 100)62 63 64@dataclass65class TaskResult:66    """任务结果"""67    success: bool = False68    data: Optional[Dict] = None69    error_message: Optional[str] = None70    processed_files: List[str] = field(default_factory=list)71    failed_files: List[str] = field(default_factory=list)72    transcription_results: Optional[Dict] = None73    duration: float = 0.074    75    def to_dict(self) -> Dict:76        """转换为字典格式"""77        return {78            'success': self.success,79            'data': self.data,80            'error_message': self.error_message,81            'processed_files': self.processed_files,82            'failed_files': self.failed_files,83            'transcription_results': self.transcription_results,84            'duration': self.duration85        }86 87 88@dataclass89class Task:90    """任务信息"""91    id: str = field(default_factory=lambda: str(uuid.uuid4())[:8])92    status: TaskStatus = TaskStatus.PENDING93    priority: TaskPriority = TaskPriority.NORMAL94    file_paths: List[Path] = field(default_factory=list)95    progress: TaskProgress = field(default_factory=TaskProgress)96    result: TaskResult = field(default_factory=TaskResult)97    created_at: datetime = field(default_factory=datetime.now)98    started_at: Optional[datetime] = None99    completed_at: Optional[datetime] = None100    callback: Optional[Callable] = None101    metadata: Dict[str, Any] = field(default_factory=dict)102    103    def to_dict(self) -> Dict:104        """转换为字典格式"""105        return {106            'id': self.id,107            'status': self.status.value,108            'priority': self.priority.value,109            'file_count': len(self.file_paths),110            'file_names': [fp.name for fp in self.file_paths],111            'progress': {112                'stage': self.progress.stage,113                'current': self.progress.current,114                'total': self.progress.total,115                'percentage': self.progress.percentage,116                'message': self.progress.message117            },118            'result': self.result.to_dict(),119            'created_at': self.created_at.isoformat() if self.created_at else None,120            'started_at': self.started_at.isoformat() if self.started_at else None,121            'completed_at': self.completed_at.isoformat() if self.completed_at else None,122            'metadata': self.metadata123        }124 125 126class TaskManager:127    """任务管理器"""128    129    def __init__(self):130        """初始化任务管理器"""131        self.config = get_config()132        self.logger = get_task_logger(logger_name="transcript_service.task")133        134        # 任务存储135        self.tasks: Dict[str, Task] = {}136        self.task_queue: asyncio.Queue = asyncio.Queue(maxsize=self.config.task.queue_size)137        138        # 服务实例139        self.file_validator = get_file_validator()140        self.oss_service = get_oss_service()141        self.paraformer_service = get_paraformer_service()142        143        # 工作线程池144        self.executor = ThreadPoolExecutor(max_workers=self.config.app.concurrent_tasks)145        146        # 状态回调147        self.status_callbacks: List[Callable] = []148        149        # 任务处理器状态150        self._processor_started = False151        152        # 启动任务处理器153        self._start_task_processor()154    155    def add_status_callback(self, callback: Callable):156        """添加状态变化回调函数157        158        Args:159            callback: 回调函数160        """161        self.status_callbacks.append(callback)162    163    def _notify_status_change(self, task: Task):164        """通知状态变化"""165        for callback in self.status_callbacks:166            try:167                callback(task)168            except Exception as e:169                self.logger.error(f"回调函数执行失败: {str(e)}")170    171    async def create_task(self, file_paths: List[Path], priority: TaskPriority = TaskPriority.NORMAL, metadata = None) -> str:172        """创建新任务173        174        Args:175            file_paths: 文件路径列表176            priority: 任务优先级177            metadata: 任务元数据178            179        Returns:180            任务ID181        """182        # 确保任务处理器已启动183        if not self._processor_started:184            self._ensure_processor_started()185            186        task = Task(187            file_paths=file_paths,188            priority=priority,189            metadata=metadata or {}190        )191        192        self.tasks[task.id] = task193        194        # 添加到队列195        await self.task_queue.put(task.id)196        197        self.logger.info(f"创建任务: {task.id}, 文件数量: {len(file_paths)}")198        return task.id199    200    def get_task(self, task_id: str) -> Optional[Task]:201        """获取任务信息202        203        Args:204            task_id: 任务ID205            206        Returns:207            任务对象208        """209        return self.tasks.get(task_id)210    211    def get_all_tasks(self) -> List[Task]:212        """获取所有任务"""213        return list(self.tasks.values())214    215    def get_tasks_by_status(self, status: TaskStatus) -> List[Task]:216        """根据状态获取任务"""217        return [task for task in self.tasks.values() if task.status == status]218    219    async def cancel_task(self, task_id: str) -> bool:220        """取消任务221        222        Args:223            task_id: 任务ID224            225        Returns:226            是否成功取消227        """228        task = self.get_task(task_id)229        if not task:230            return False231        232        if task.status in [TaskStatus.COMPLETED, TaskStatus.FAILED, TaskStatus.CANCELLED]:233            return False234        235        task.status = TaskStatus.CANCELLED236        task.completed_at = datetime.now()237        task.progress.message = "任务已取消"238        239        self._notify_status_change(task)240        self.logger.info(f"任务已取消: {task_id}")241        return True242    243    def _start_task_processor(self):244        """启动任务处理器"""245        try:246            # 只有在有运行的事件循环时才启动任务处理器247            loop = asyncio.get_running_loop()248            asyncio.create_task(self._process_tasks())249        except RuntimeError:250            # 没有运行的事件循环,延迟启动251            self.logger.debug("没有运行的事件循环,任务处理器将在需要时启动")252            self._processor_started = False253        else:254            self._processor_started = True255    256    def _ensure_processor_started(self):257        """确保任务处理器已启动"""258        if not self._processor_started:259            try:260                loop = asyncio.get_running_loop()261                asyncio.create_task(self._process_tasks())262                self._processor_started = True263            except RuntimeError:264                self.logger.warning("无法启动任务处理器:没有运行的事件循环")265    266    async def _process_tasks(self):267        """处理任务队列"""268        while True:269            try:270                # 从队列获取任务271                task_id = await self.task_queue.get()272                task = self.get_task(task_id)273                274                if not task or task.status == TaskStatus.CANCELLED:275                    self.task_queue.task_done()276                    continue277                278                # 处理任务279                await self._execute_task(task)280                self.task_queue.task_done()281                282            except Exception as e:283                self.logger.exception(f"处理任务队列时发生错误: {str(e)}")284                await asyncio.sleep(1)285    286    async def _execute_task(self, task: Task):287        """执行任务288        289        Args:290            task: 任务对象291        """292        try:293            # 设置任务日志上下文294            self.logger.set_task_id(task.id)295            296            task.status = TaskStatus.VALIDATING297            task.started_at = datetime.now()298            task.progress.stage = "文件验证"299            task.progress.update(0, 100, "开始验证文件")300            self._notify_status_change(task)301            302            # 1. 文件验证303            valid_files, invalid_files = await self._validate_files(task)304            if not valid_files:305                task.status = TaskStatus.FAILED306                task.result.error_message = "没有有效的文件"307                task.result.failed_files = [str(f[0]) for f in invalid_files]308                task.completed_at = datetime.now()309                self._notify_status_change(task)310                return311            312            # 2. 文件上传313            task.status = TaskStatus.UPLOADING314            task.progress.stage = "文件上传"315            task.progress.update(0, len(valid_files), "开始上传文件到OSS")316            self._notify_status_change(task)317            318            upload_results = await self._upload_files(task, valid_files)319            successful_uploads = [r for r in upload_results if r[1]]320            321            if not successful_uploads:322                task.status = TaskStatus.FAILED323                task.result.error_message = "文件上传失败"324                task.completed_at = datetime.now()325                self._notify_status_change(task)326                return327            328            # 3. 转录处理329            task.status = TaskStatus.TRANSCRIBING330            task.progress.stage = "语音转录"331            task.progress.update(0, 100, "开始语音转录")332            self._notify_status_change(task)333            334            file_urls = [r[2] for r in successful_uploads]335            success, transcription_result, error = await self._transcribe_audio(task, file_urls)336            337            # 4. 完成任务338            task.completed_at = datetime.now()339            task.result.duration = (task.completed_at - task.started_at).total_seconds()340            341            if success:342                task.status = TaskStatus.COMPLETED343                task.result.success = True344                task.result.transcription_results = transcription_result345                task.result.processed_files = [r[0] for r in successful_uploads]346                task.progress.update(100, 100, "转录完成")347            else:348                task.status = TaskStatus.FAILED349                task.result.error_message = error350            351            self._notify_status_change(task)352            353        except Exception as e:354            task.status = TaskStatus.FAILED355            task.result.error_message = f"任务执行失败: {str(e)}"356            task.completed_at = datetime.now()357            self.logger.exception(f"执行任务时发生错误: {task.id}")358            self._notify_status_change(task)359        finally:360            self.logger.clear_task_id()361    362    async def _validate_files(self, task: Task) -> tuple:363        """验证文件"""364        self.logger.info(f"开始验证 {len(task.file_paths)} 个文件")365        366        valid_files, invalid_files = self.file_validator.validate_multiple_files(task.file_paths)367        368        task.progress.update(100, 100, f"验证完成: {len(valid_files)} 个有效文件")369        self.logger.info(f"文件验证完成: {len(valid_files)} 个有效文件, {len(invalid_files)} 个无效文件")370        371        return valid_files, invalid_files372    373    async def _upload_files(self, task: Task, file_paths: List[Path]) -> List[tuple]:374        """上传文件"""375        self.logger.info(f"开始上传 {len(file_paths)} 个文件")376        377        results = []378        for i, file_path in enumerate(file_paths):379            if task.status == TaskStatus.CANCELLED:380                break381            382            success, url_or_error, object_key = await self.oss_service.upload_file(file_path, task.id)383            results.append((file_path.name, success, url_or_error, object_key))384            385            # 更新进度386            task.progress.update(i + 1, len(file_paths), f"已上传 {i + 1}/{len(file_paths)} 个文件")387            self._notify_status_change(task)388        389        self.logger.info(f"文件上传完成: {len([r for r in results if r[1]])} 个成功")390        return results391    392    async def _transcribe_audio(self, task: Task, file_urls: List[str]) -> tuple:393        """转录音频"""394        self.logger.info(f"开始转录 {len(file_urls)} 个音频文件")395        396        # 提取Paraformer参数397        paraformer_params = None398        if 'paraformer_params' in task.metadata:399            paraformer_params = task.metadata['paraformer_params']400            self.logger.info(f"使用自定义Paraformer参数: {paraformer_params}")401        402        success, results, error = await self.paraformer_service.batch_process_with_retry(403            file_urls, task.id, paraformer_params404        )405        406        if success:407            task.progress.update(100, 100, "转录完成")408            self.logger.info(f"转录完成: {len(file_urls)} 个文件")409        else:410            self.logger.error(f"转录失败: {error}")411        412        return success, results, error413    414    def cleanup_completed_tasks(self, hours: int = 24):415        """清理已完成的任务416        417        Args:418            hours: 保留时间(小时)419        """420        cutoff_time = datetime.now() - timedelta(hours=hours)421        to_remove = []422        423        for task_id, task in self.tasks.items():424            if (task.status in [TaskStatus.COMPLETED, TaskStatus.FAILED, TaskStatus.CANCELLED] and425                task.completed_at and task.completed_at < cutoff_time):426                to_remove.append(task_id)427        428        for task_id in to_remove:429            del self.tasks[task_id]430        431        self.logger.info(f"清理了 {len(to_remove)} 个过期任务")432    433    def get_statistics(self) -> Dict:434        """获取任务统计信息"""435        stats = {436            'total_tasks': len(self.tasks),437            'pending': len(self.get_tasks_by_status(TaskStatus.PENDING)),438            'validating': len(self.get_tasks_by_status(TaskStatus.VALIDATING)),439            'uploading': len(self.get_tasks_by_status(TaskStatus.UPLOADING)),440            'transcribing': len(self.get_tasks_by_status(TaskStatus.TRANSCRIBING)),441            'completed': len(self.get_tasks_by_status(TaskStatus.COMPLETED)),442            'failed': len(self.get_tasks_by_status(TaskStatus.FAILED)),443            'cancelled': len(self.get_tasks_by_status(TaskStatus.CANCELLED)),444            'queue_size': self.task_queue.qsize()445        }446        return stats447 448 449# 全局任务管理器实例450task_manager = None451 452 453def get_task_manager() -> TaskManager:454    """获取任务管理器实例455    456    Returns:457        任务管理器实例458    """459    global task_manager460    if task_manager is None:461        task_manager = TaskManager()462    return task_manager