189 lines
9.5 KiB
Python
189 lines
9.5 KiB
Python
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()
|