from __future__ import annotations import argparse import asyncio import json from datetime import datetime, timezone from sqlalchemy import or_, select from app.config import settings from app.enums.video_upscale import LOCAL_PROCESSOR_KEYS, VideoUpscaleTaskStatus from app.models.base import async_session from app.models.chat_generation_task import ChatGenerationTask from app.models.generation_record import GenerationRecord from app.models.module_generation_step import ModuleGenerationStep from app.models.shot_replicate_segment import ShotReplicateSegment from app.models.video_upscale_task import VideoUpscaleTask from app.services.video_upscale.media_service import is_valid_file from app.services.video_upscale.task_service import reset_failed_upscale_task_for_manual_retry def _parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description="人工恢复视频超分任务") parser.add_argument("--task-id", action="append", default=[], help="ChatGenerationTask.id,可重复传入") parser.add_argument("--task-ids", default="", help="逗号分隔的 ChatGenerationTask.id") parser.add_argument("--generation-record-id", action="append", default=[], help="GenerationRecord.id,可重复传入") parser.add_argument("--generation-record-ids", default="", help="逗号分隔的 GenerationRecord.id") parser.add_argument("--project-id", action="append", default=[], help="Project.id,可重复传入") parser.add_argument("--generation-mode", default="", help="按 ChatGenerationTask.generation_mode 筛选") parser.add_argument("--module-owner-id", default="", help="ModuleGenerationProject.id") parser.add_argument("--shot-task-set-id", default="", help="ShotReplicateTaskSet.id") parser.add_argument("--shot-segment-id", action="append", default=[], help="ShotReplicateSegment.id,可重复传入") parser.add_argument("--failed-only", action=argparse.BooleanOptionalAction, default=True) parser.add_argument("--limit", type=int, default=100) parser.add_argument("--dry-run", action="store_true") parser.add_argument("--enqueue", action=argparse.BooleanOptionalAction, default=True) parser.add_argument("--force-resubmit", action="store_true", help="远程任务清空 provider task/result 后从 source.mp4 重新提交") return parser async def _collect_chat_task_ids(db, args: argparse.Namespace) -> list[str]: ids = [str(item).strip() for item in args.task_id if str(item).strip()] ids.extend(item.strip() for item in str(args.task_ids or "").split(",") if item.strip()) project_ids: list[str] = [str(args.module_owner_id).strip()] if args.module_owner_id else [] segment_ids = [str(item).strip() for item in args.shot_segment_id if str(item).strip()] if args.shot_task_set_id or segment_ids: query = select(ShotReplicateSegment.module_project_id).where( ShotReplicateSegment.deleted_at.is_(None), ShotReplicateSegment.module_project_id.is_not(None), ) if args.shot_task_set_id: query = query.where(ShotReplicateSegment.task_set_id == str(args.shot_task_set_id).strip()) if segment_ids: query = query.where(ShotReplicateSegment.id.in_(segment_ids)) result = await db.execute(query) project_ids.extend(str(value) for value in result.scalars().all() if value) project_ids = list(dict.fromkeys(item for item in project_ids if item)) if project_ids: result = await db.execute( select(ModuleGenerationStep.chat_task_id).where( ModuleGenerationStep.project_id.in_(project_ids), ModuleGenerationStep.chat_task_id.isnot(None), ModuleGenerationStep.deleted_at.is_(None), ) ) ids.extend(str(value) for value in result.scalars().all() if value) return list(dict.fromkeys(ids)) async def _collect_generation_record_ids(db, args: argparse.Namespace) -> list[str]: ids = [str(item).strip() for item in args.generation_record_id if str(item).strip()] ids.extend(item.strip() for item in str(args.generation_record_ids or "").split(",") if item.strip()) project_ids = [str(item).strip() for item in args.project_id if str(item).strip()] if project_ids: result = await db.execute( select(GenerationRecord.id).where( GenerationRecord.project_id.in_(project_ids), GenerationRecord.deleted_at.is_(None), ) ) ids.extend(str(value) for value in result.scalars().all() if value) return list(dict.fromkeys(ids)) def _owner_preview(upscale: VideoUpscaleTask, chat: ChatGenerationTask | None, record: GenerationRecord | None) -> dict: owner = chat or record return { "upscale_task_id": upscale.id, "owner_type": "chat_generation_task" if chat else "generation_record", "owner_id": owner.id if owner else None, "chat_task_id": chat.id if chat else None, "generation_record_id": record.id if record else None, "project_id": record.project_id if record else None, "generation_mode": chat.generation_mode if chat else None, "status": upscale.status, "stage": upscale.stage, "processor_key": upscale.processor_key, "source_local_path": upscale.source_local_path, "provider_task_id": upscale.provider_task_id, "provider_output_url_expires_at": upscale.provider_output_url_expires_at, } async def _run(args: argparse.Namespace) -> dict: async with async_session() as db: chat_ids = await _collect_chat_task_ids(db, args) record_ids = await _collect_generation_record_ids(db, args) query = ( select(VideoUpscaleTask, ChatGenerationTask, GenerationRecord) .outerjoin(ChatGenerationTask, ChatGenerationTask.id == VideoUpscaleTask.chat_generation_task_id) .outerjoin(GenerationRecord, GenerationRecord.id == VideoUpscaleTask.generation_record_id) .where( or_( (ChatGenerationTask.id.isnot(None) & ChatGenerationTask.deleted_at.is_(None)), (GenerationRecord.id.isnot(None) & GenerationRecord.deleted_at.is_(None)), ) ) .order_by(VideoUpscaleTask.updated_at.asc()) .limit(max(1, min(int(args.limit or 100), 1000))) ) owner_filters = [] if chat_ids: owner_filters.append(VideoUpscaleTask.chat_generation_task_id.in_(chat_ids)) if record_ids: owner_filters.append(VideoUpscaleTask.generation_record_id.in_(record_ids)) if owner_filters: query = query.where(or_(*owner_filters)) if args.generation_mode: query = query.where(ChatGenerationTask.generation_mode == args.generation_mode) if args.failed_only: query = query.where(VideoUpscaleTask.status == VideoUpscaleTaskStatus.FAILED.value) result = await db.execute(query) rows = result.all() preview = [_owner_preview(upscale, chat, record) for upscale, chat, record in rows] if args.dry_run or not args.enqueue: return {"dry_run": True, "matched": len(preview), "items": preview} from app.tasks.video_upscale_tasks import download_remote_result, execute_local, finalize, poll_remote, submit_remote enqueued = [] for upscale, chat, record in rows: reset = await reset_failed_upscale_task_for_manual_retry( db, upscale_task_id=upscale.id, force_resubmit=bool(args.force_resubmit), ) if is_valid_file(reset.final_local_path) and not args.force_resubmit: action = "finalize" finalize.apply_async(args=[reset.id], queue=settings.VIDEO_UPSCALE_LOCAL_QUEUE) elif reset.processor_key in LOCAL_PROCESSOR_KEYS: action = "local" execute_local.apply_async(args=[reset.id], queue=settings.VIDEO_UPSCALE_LOCAL_QUEUE) else: expires_at = reset.provider_output_url_expires_at if expires_at and expires_at.tzinfo is None: expires_at = expires_at.replace(tzinfo=timezone.utc) remaining = (expires_at - datetime.now(timezone.utc)).total_seconds() if expires_at else None if reset.provider_output_url and remaining is not None and remaining >= 2 * 3600 and not args.force_resubmit: action = "download" download_remote_result.apply_async(args=[reset.id], queue=settings.VIDEO_UPSCALE_REMOTE_QUEUE) elif reset.provider_task_id and not args.force_resubmit: action = "poll" poll_remote.apply_async(args=[reset.id], queue=settings.VIDEO_UPSCALE_REMOTE_QUEUE) else: action = "submit" submit_remote.apply_async(args=[reset.id], queue=settings.VIDEO_UPSCALE_REMOTE_QUEUE) owner = chat or record enqueued.append( { "upscale_task_id": reset.id, "owner_type": "chat_generation_task" if chat else "generation_record", "owner_id": owner.id if owner else None, "action": action, } ) return {"dry_run": False, "matched": len(preview), "enqueued": enqueued} def main() -> None: args = _parser().parse_args() result = asyncio.run(_run(args)) print(json.dumps(result, ensure_ascii=False, indent=2, default=str)) if __name__ == "__main__": main()