from datetime import datetime, timezone from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.models import GenerationTaskModel from packages.domain import GenerationTask from packages.domain.generation_task import GenerationTaskStatus def _to_domain(model: GenerationTaskModel) -> GenerationTask: """Convert ORM model to domain entity.""" return GenerationTask( id=model.id, project_id=model.project_id, strategy_id=model.strategy_id, asset_library_id=model.asset_library_id, voice_library_id=model.voice_library_id, template_id=model.template_id, asset_ids=list(model.asset_ids or []), title_ids=list(model.title_ids or []), voice_ids=list(model.voice_ids or []), status=GenerationTaskStatus(model.status) if model.status else GenerationTaskStatus.PENDING, progress=model.progress, result_count=int(model.result_count or 0), error_message=model.error_message, error_info=dict(model.error_info) if model.error_info else {}, retry_count=model.retry_count or 0, auto_retry_enabled=bool(model.auto_retry_enabled), auto_retry_max=model.auto_retry_max or 0, started_at=model.started_at, completed_at=model.completed_at, created_by_user_id=model.created_by_user_id, source_edit_plan_id=model.source_edit_plan_id or "", asset_select_mode=model.asset_select_mode or "", batch_id=model.batch_id or "", video_title=getattr(model, "video_title", "") or "", resolution=getattr(model, "resolution", "") or "", bgm_config=dict(getattr(model, "bgm_config", {}) or {}), is_preview=bool(getattr(model, "is_preview", False)), source_task_id=getattr(model, "source_task_id", "") or "", celery_task_id=getattr(model, "celery_task_id", "") or "", output_width=getattr(model, "output_width", 1280) or 1280, output_height=getattr(model, "output_height", 720) or 720, cover_url=getattr(model, "cover_url", "") or "", title_config=dict(getattr(model, "title_config", {}) or {}), logs=model.logs or "[]", created_at=model.created_at, updated_at=model.updated_at, ) class SQLAlchemyGenerationTaskRepository: def __init__(self, session: Session): self.session = session def create(self, task: GenerationTask) -> GenerationTask: model = GenerationTaskModel( id=task.id, project_id=task.project_id, strategy_id=task.strategy_id, asset_library_id=task.asset_library_id, voice_library_id=task.voice_library_id, template_id=task.template_id, asset_ids=task.asset_ids, title_ids=task.title_ids, voice_ids=task.voice_ids, status=task.status, progress=task.progress, result_count=task.result_count, error_message=task.error_message, error_info=task.error_info or None, retry_count=task.retry_count or 0, auto_retry_enabled=task.auto_retry_enabled, auto_retry_max=task.auto_retry_max or 0, started_at=task.started_at, completed_at=task.completed_at, created_by_user_id=task.created_by_user_id, source_edit_plan_id=task.source_edit_plan_id or None, asset_select_mode=task.asset_select_mode or "", batch_id=task.batch_id or "", video_title=task.video_title or "", resolution=task.resolution or "", bgm_config=task.bgm_config or {}, is_preview=task.is_preview or False, source_task_id=task.source_task_id or "", celery_task_id=getattr(task, "celery_task_id", "") or "", output_width=task.output_width, output_height=task.output_height, cover_url=task.cover_url or "", title_config=dict(task.title_config) if task.title_config else {}, logs=task.logs, created_at=task.created_at, updated_at=task.updated_at, ) self.session.add(model) self.session.commit() return task def get(self, task_id: str) -> GenerationTask | None: model = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.id == task_id).first() if model is None: return None return _to_domain(model) def list_by_project(self, project_id: str) -> list[GenerationTask]: models = ( self.session.query(GenerationTaskModel) .filter(GenerationTaskModel.project_id == project_id) .order_by(GenerationTaskModel.created_at.desc()) .all() ) return [_to_domain(m) for m in models] def list_by_user(self, user_id: str) -> list[GenerationTask]: models = ( self.session.query(GenerationTaskModel) .filter(GenerationTaskModel.created_by_user_id == user_id) .order_by(GenerationTaskModel.created_at.desc()) .all() ) return [_to_domain(m) for m in models] def count_by_user(self, user_id: str) -> int: return self.session.query(GenerationTaskModel).filter(GenerationTaskModel.created_by_user_id == user_id).count() def count_pending_by_user(self, user_id: str) -> int: return ( self.session.query(GenerationTaskModel) .filter( GenerationTaskModel.created_by_user_id == user_id, GenerationTaskModel.status == GenerationTaskStatus.PENDING.value, ) .count() ) def count_pending_total(self) -> int: return ( self.session.query(GenerationTaskModel) .filter(GenerationTaskModel.status == GenerationTaskStatus.PENDING.value) .count() ) def count_running_by_user(self, user_id: str) -> int: """统计指定用户处于 running 状态的任务数(用于限流提示展示)。""" return ( self.session.query(GenerationTaskModel) .filter( GenerationTaskModel.created_by_user_id == user_id, GenerationTaskModel.status == GenerationTaskStatus.RUNNING.value, ) .count() ) def count_running_total(self) -> int: """统计全局处于 running 状态的任务数(worker 实际在执行的任务数)。""" return ( self.session.query(GenerationTaskModel) .filter(GenerationTaskModel.status == GenerationTaskStatus.RUNNING.value) .count() ) def estimate_avg_duration_seconds(self, limit: int = 20, default_seconds: float = 120.0) -> float: """估算最近完成任务的平均耗时(秒),用于 429 限流提示的等待预估。 取最近 N 条 completed 任务的 (completed_at - started_at) 平均值; 无足够历史数据时返回 default_seconds。 用 Python 侧计算差值,避免 SQLite/PostgreSQL 方言差异。 """ rows = ( self.session.query(GenerationTaskModel.started_at, GenerationTaskModel.completed_at) .filter( GenerationTaskModel.status == GenerationTaskStatus.COMPLETED.value, GenerationTaskModel.started_at.isnot(None), GenerationTaskModel.completed_at.isnot(None), ) .order_by(GenerationTaskModel.completed_at.desc()) .limit(limit) .all() ) durations = [ (completed - started).total_seconds() for started, completed in rows if completed and started and (completed - started).total_seconds() > 0 ] if not durations: return default_seconds return sum(durations) / len(durations) def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]: models = ( self.session.query(GenerationTaskModel) .filter(GenerationTaskModel.created_by_user_id == user_id) .order_by(GenerationTaskModel.created_at.desc()) .limit(limit) .all() ) return [_to_domain(m) for m in models] def list_latest_completed_preview(self, user_id: str, template_id: str, limit: int = 1) -> list[GenerationTask]: """按用户+模板查找最近已完成的预览任务。""" models = ( self.session.query(GenerationTaskModel) .filter( GenerationTaskModel.created_by_user_id == user_id, GenerationTaskModel.template_id == template_id, GenerationTaskModel.is_preview, GenerationTaskModel.status == "completed", ) .order_by(GenerationTaskModel.created_at.desc()) .limit(limit) .all() ) return [_to_domain(m) for m in models] def list_by_source_edit_plan(self, plan_id: str) -> list[GenerationTask]: models = ( self.session.query(GenerationTaskModel) .filter(GenerationTaskModel.source_edit_plan_id == plan_id) .order_by(GenerationTaskModel.created_at.desc()) .all() ) return [_to_domain(m) for m in models] def list_by_user_filtered( self, user_id: str, *, status: str | None = None, limit: int | None = None, offset: int = 0, ) -> list[GenerationTask]: """按用户+状态筛选任务列表。""" query = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.created_by_user_id == user_id) if status: query = query.filter(GenerationTaskModel.status == status) query = query.order_by(GenerationTaskModel.created_at.desc()) if offset: query = query.offset(offset) if limit: query = query.limit(limit) return [_to_domain(m) for m in query.all()] def count_by_user_filtered( self, user_id: str, *, status: str | None = None, ) -> int: """按用户+状态筛选计数。""" query = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.created_by_user_id == user_id) if status: query = query.filter(GenerationTaskModel.status == status) return query.count() def list_by_project_filtered( self, project_id: str, *, status: str | None = None, limit: int | None = None, offset: int = 0, ) -> list[GenerationTask]: """按项目+状态筛选任务列表。""" query = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.project_id == project_id) if status: query = query.filter(GenerationTaskModel.status == status) query = query.order_by(GenerationTaskModel.created_at.desc()) if offset: query = query.offset(offset) if limit: query = query.limit(limit) return [_to_domain(m) for m in query.all()] def count_by_project_filtered( self, project_id: str, *, status: str | None = None, ) -> int: """按项目+状态筛选计数。""" query = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.project_id == project_id) if status: query = query.filter(GenerationTaskModel.status == status) return query.count() def update(self, task: GenerationTask) -> GenerationTask: model = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.id == task.id).first() if model is None: raise ValueError(f"GenerationTask {task.id} not found") model.project_id = task.project_id model.asset_library_id = task.asset_library_id model.strategy_id = task.strategy_id model.voice_library_id = task.voice_library_id model.template_id = task.template_id model.asset_ids = task.asset_ids model.title_ids = task.title_ids model.voice_ids = task.voice_ids model.status = task.status model.progress = task.progress model.result_count = task.result_count model.error_message = task.error_message model.error_info = task.error_info or None model.retry_count = task.retry_count or 0 model.auto_retry_enabled = task.auto_retry_enabled model.auto_retry_max = task.auto_retry_max or 0 model.started_at = task.started_at model.completed_at = task.completed_at model.source_edit_plan_id = task.source_edit_plan_id or None model.asset_select_mode = task.asset_select_mode or "" model.batch_id = task.batch_id or "" if hasattr(model, "video_title"): model.video_title = task.video_title or "" if hasattr(model, "resolution"): model.resolution = task.resolution or "" if hasattr(model, "bgm_config"): model.bgm_config = task.bgm_config or {} if hasattr(model, "is_preview"): model.is_preview = task.is_preview or False model.source_task_id = task.source_task_id or "" model.celery_task_id = getattr(task, "celery_task_id", "") or model.celery_task_id or "" model.output_width = task.output_width model.output_height = task.output_height model.cover_url = task.cover_url or "" model.title_config = dict(task.title_config) if task.title_config else {} model.logs = task.logs self.session.commit() return task def cleanup_stale_running(self, timeout_minutes: int = 10) -> int: """清理超时未更新的 running 任务(孤儿任务)。 Returns: 清理的任务数量(仅计数,保持旧签名兼容) """ items = self.cleanup_stale_running_with_ids(timeout_minutes) return len(items) def cleanup_stale_running_with_ids(self, timeout_minutes: int = 10) -> list[tuple[str, str]]: """同 cleanup_stale_running,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。""" from datetime import timedelta cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes) models = ( self.session.query(GenerationTaskModel) .filter( GenerationTaskModel.status == GenerationTaskStatus.RUNNING.value, GenerationTaskModel.updated_at < cutoff, ) .all() ) if not models: return [] result: list[tuple[str, str]] = [] for model in models: result.append((model.id, getattr(model, "celery_task_id", "") or "")) model.status = GenerationTaskStatus.FAILED.value model.error_message = "任务执行中断(worker重启/超时)" model.error_info = { "error_type": "WorkerInterrupted", "message": "任务在运行中中断,可能因 worker 重启或超时", "failed_at": datetime.now(timezone.utc).isoformat(), } model.completed_at = datetime.now(timezone.utc) self.session.commit() return result def cleanup_stale_pending(self, timeout_minutes: int = 30) -> int: """清理超时的 pending 任务(未被 Worker 拉取的任务)。 Returns: 清理的任务数量(仅计数,保持旧签名兼容) """ items = self.cleanup_stale_pending_with_ids(timeout_minutes) return len(items) def cleanup_stale_pending_with_ids(self, timeout_minutes: int = 30) -> list[tuple[str, str]]: """同 cleanup_stale_pending,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。""" from datetime import timedelta cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes) models = ( self.session.query(GenerationTaskModel) .filter( GenerationTaskModel.status == GenerationTaskStatus.PENDING.value, GenerationTaskModel.created_at < cutoff, ) .all() ) if not models: return [] error_info = { "error_type": "PendingTimeout", "message": f"任务在 pending 状态停留超过 {timeout_minutes} 分钟,自动清理", "failed_at": datetime.now(timezone.utc).isoformat(), } result: list[tuple[str, str]] = [] for model in models: result.append((model.id, getattr(model, "celery_task_id", "") or "")) model.status = GenerationTaskStatus.FAILED.value model.error_message = "pending timeout: auto cleanup" model.error_info = error_info model.completed_at = datetime.now(timezone.utc) self.session.commit() return result