Files
video-gen/video-gen-api/app/tasks/api_upscale_tasks.py
T
2026-08-06 16:44:28 +08:00

365 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""API v3 简化超分任务。
不使用复杂的 CeleryRuntimeLease 锁机制,直接执行超分流程。
"""
from __future__ import annotations
import logging
from datetime import datetime, timezone
from typing import Any
from app.config import settings
from app.enums.celery_queue import CeleryQueue
from app.services.video_upscale.volc_service import VolcSubmitResult, VolcQueryResult
from app.tasks.api_celery_app import celery_app
logger = logging.getLogger("videogen")
def _now() -> datetime:
"""获取当前时间(UTC)。"""
return datetime.now(timezone.utc)
def _now_str() -> str:
"""获取当前日期字符串。"""
return _now().strftime("%Y%m%d")
def _format_time(dt: datetime | None) -> str | None:
"""格式化时间为字符串(北京时间)。"""
if dt is None:
return None
from datetime import timedelta
beijing_time = dt + timedelta(hours=8)
return beijing_time.strftime("%Y-%m-%d %H:%M:%S")
async def _maybe_delete_source_file(db, upscale) -> None:
"""根据超分配置决定是否删除源文件。"""
if not upscale.source_local_path or not upscale.api_generation_task_id:
return
from app.models.api.api_key_upscale_config import ApiKeyUpscaleConfig
from app.models.api.api_generation_task import ApiGenerationTask
from sqlalchemy import select
# 获取 API Key ID
api_task_result = await db.execute(
select(ApiGenerationTask.api_key_id).where(
ApiGenerationTask.id == upscale.api_generation_task_id
).limit(1)
)
api_key_id = api_task_result.scalar_one_or_none()
if not api_key_id:
return
# 查询超分配置
config_result = await db.execute(
select(ApiKeyUpscaleConfig).where(
ApiKeyUpscaleConfig.api_key_id == api_key_id
).limit(1)
)
upscale_config = config_result.scalar_one_or_none()
# 只有配置了"成功后删除源文件"才删除
if upscale_config and upscale_config.delete_source_after_success:
source_abs = os.path.abspath(upscale.source_local_path)
if os.path.exists(source_abs):
os.remove(source_abs)
logger.info("Deleted source file after upscale: %s", source_abs)
@celery_app.task(
name="api_upscale.execute_local_simple",
bind=True,
max_retries=2,
soft_time_limit=600,
time_limit=900,
)
def api_upscale_execute_local_simple(self, upscale_task_id: str) -> dict[str, Any]:
"""简化版本地超分执行(API v3 专用)。"""
from app.models.base import async_session
from app.models.video_upscale_task import VideoUpscaleTask
from app.services.video_upscale.local_ffmpeg_service import run_local_ffmpeg_upscale
from sqlalchemy import select
async def _execute():
async with async_session() as db:
result = await db.execute(
select(VideoUpscaleTask).where(VideoUpscaleTask.id == upscale_task_id).limit(1)
)
upscale = result.scalar_one_or_none()
if not upscale:
return {"status": "not_found"}
if upscale.status == "completed":
return {"status": "already_completed"}
# 更新状态为处理中
upscale.status = "processing"
upscale.stage = "local_processing"
upscale.attempt_count = int(upscale.attempt_count or 0) + 1
upscale.started_at = upscale.started_at or __import__("datetime").datetime.now(__import__("datetime").timezone.utc)
await db.commit()
try:
# 执行本地 FFmpeg 超分
output_path = await run_local_ffmpeg_upscale(upscale)
# 计算相对 URL 路径
# output_path 是绝对路径,需要转换为 /generate/api/videos/... 格式
date_dir = _now_str()
url_path = f"/generate/api/videos/{date_dir}/{upscale.api_generation_task_id}.mp4"
# 更新成功状态
upscale.status = "completed"
upscale.stage = "upscale_completed"
upscale.final_local_path = f"./storage/generate/api/videos/{_now_str()}/{upscale.api_generation_task_id or "unknown"}.mp4"
upscale.final_resource_url = url_path # 相对 URL
upscale.completed_at = _now()
# 更新 API 任务的 video_url(使用相对 URL
if upscale.api_generation_task_id:
from app.models.api.api_generation_task import ApiGenerationTask
api_result = await db.execute(
select(ApiGenerationTask).where(ApiGenerationTask.id == upscale.api_generation_task_id).limit(1)
)
api_task = api_result.scalar_one_or_none()
if api_task:
api_task.video_url = url_path # 相对 URL
api_task.status = "completed"
api_task.pipeline_stage = "done"
api_task.generated_at = upscale.completed_at
await db.commit()
return {"status": "completed", "output_path": output_path}
except Exception as exc:
upscale.status = "failed"
upscale.stage = "failed"
upscale.failure_count = int(upscale.failure_count or 0) + 1
upscale.last_error = str(exc)[:500]
upscale.failed_at = __import__("datetime").datetime.now(__import__("datetime").timezone.utc)
await db.commit()
raise
from app.tasks.async_runner import run_async
return run_async(_execute())
@celery_app.task(
name="api_upscale.submit_remote_simple",
bind=True,
max_retries=2,
soft_time_limit=600,
time_limit=900,
)
def api_upscale_submit_remote_simple(self, upscale_task_id: str) -> dict[str, Any]:
"""简化版远程超分提交(API v3 专用)。"""
from app.models.base import async_session
from app.models.video_upscale_task import VideoUpscaleTask
from app.services.video_upscale.volc_service import submit_video_enhance, VolcSubmitResult
from sqlalchemy import select
async def _execute():
async with async_session() as db:
result = await db.execute(
select(VideoUpscaleTask).where(VideoUpscaleTask.id == upscale_task_id).limit(1)
)
upscale = result.scalar_one_or_none()
if not upscale:
return {"status": "not_found"}
if upscale.status == "completed":
return {"status": "already_completed"}
# 更新状态
upscale.status = "processing"
upscale.stage = "remote_submitting"
upscale.attempt_count = int(upscale.attempt_count or 0) + 1
upscale.started_at = upscale.started_at or __import__("datetime").datetime.now(__import__("datetime").timezone.utc)
await db.commit()
try:
# 提交到火山 MediaKit
from app.utils.id_gen import generate_id
# 从 API 任务获取超分快照
import json
api_snapshot = {}
if upscale.api_generation_task_id:
from app.models.api.api_generation_task import ApiGenerationTask
api_result = await db.execute(
__import__("sqlalchemy").select(ApiGenerationTask).where(ApiGenerationTask.id == upscale.api_generation_task_id).limit(1)
)
api_task = api_result.scalar_one_or_none()
if api_task and api_task.video_upscale_snapshot_json:
try:
api_snapshot = json.loads(api_task.video_upscale_snapshot_json)
except (json.JSONDecodeError, TypeError):
pass
submit_result: VolcSubmitResult = await submit_video_enhance(
processor_key=upscale.processor_key,
video_url=upscale.source_remote_url or "",
target_resolution=api_snapshot.get("target_resolution", "1080p"),
target_width=int(upscale.target_width or 0),
target_height=int(upscale.target_height or 0),
processor=api_snapshot.get("processor", {}),
client_token=generate_id(),
)
# 更新成功状态
upscale.provider_task_id = submit_result.task_id
upscale.provider_submitted_at = __import__("datetime").datetime.now(__import__("datetime").timezone.utc)
upscale.provider_request_json = json.dumps(submit_result.request_payload, ensure_ascii=False) if submit_result.request_payload else None
upscale.provider_response_json = json.dumps(submit_result.response_payload, ensure_ascii=False) if submit_result.response_payload else None
upscale.celery_task_id = self.request.id if hasattr(self, 'request') else None
upscale.status = "processing"
upscale.stage = "remote_polling"
await db.commit()
# 立即触发第一次轮询
api_upscale_poll_remote_simple.apply_async(
args=[upscale_task_id],
countdown=30,
queue=CeleryQueue.GEN_API_UPSCALE.value,
)
return {"status": "submitted", "provider_task_id": submit_result.task_id}
except Exception as exc:
upscale.status = "failed"
upscale.stage = "failed"
upscale.failure_count = int(upscale.failure_count or 0) + 1
upscale.last_error = str(exc)[:500]
upscale.failed_at = __import__("datetime").datetime.now(__import__("datetime").timezone.utc)
await db.commit()
raise
from app.tasks.async_runner import run_async
return run_async(_execute())
@celery_app.task(
name="api_upscale.poll_remote_simple",
bind=True,
max_retries=10,
soft_time_limit=120,
time_limit=300,
)
def api_upscale_poll_remote_simple(self, upscale_task_id: str) -> dict[str, Any]:
"""简化版远程超分轮询(API v3 专用)。"""
from app.models.base import async_session
from app.models.video_upscale_task import VideoUpscaleTask
from app.services.video_upscale.volc_service import query_task
from sqlalchemy import select
async def _execute():
async with async_session() as db:
result = await db.execute(
select(VideoUpscaleTask).where(VideoUpscaleTask.id == upscale_task_id).limit(1)
)
upscale = result.scalar_one_or_none()
if not upscale or upscale.status == "completed":
return {"status": "not_found_or_completed"}
try:
# 查询火山 MediaKit 状态
query_result = await query_task(upscale.provider_task_id)
status = query_result.status
if status == "completed":
# 超分完成
output_url = query_result.result.get("video_url", "") if query_result.result else ""
upscale.provider_output_url = output_url
upscale.provider_output_url_expires_at = __import__("datetime").datetime.fromtimestamp(query_result.expires_at, tz=__import__("datetime").timezone.utc) if query_result.expires_at else None
import os
from app.services.video_gen import download_video
api_task_id = upscale.api_generation_task_id
final_path = None
try:
# 直接下载到最终路径
date_dir = _now_str()
# 相对 URL 路径
url_path = f"/generate/api/videos/{date_dir}/{api_task_id}.mp4"
# 绝对文件路径
z_url_path = f"./storage/generate/api/videos/{date_dir}/{api_task_id}.mp4"
abs_path = os.path.abspath(z_url_path)
os.makedirs(os.path.dirname(abs_path), exist_ok=True)
await download_video(output_url, abs_path)
upscale.final_local_path = z_url_path
upscale.final_resource_url = url_path # 相对 URL
final_path = url_path
# 检查 API Key 超分配置中的"成功后删除源文件"设置
_maybe_delete_source_file(db, upscale)
except Exception:
upscale.final_local_path = output_url
upscale.final_resource_url = output_url
final_path = output_url
upscale.status = "completed"
upscale.stage = "upscale_completed"
upscale.completed_at = _now()
# 获取文件大小
try:
import os
abs_path = os.path.abspath(final_path) if final_path and final_path.startswith(".") else None
if abs_path and os.path.exists(abs_path):
upscale.final_file_size_bytes = os.path.getsize(abs_path)
except Exception:
pass
# 更新 API 任务
if api_task_id:
from app.models.api.api_generation_task import ApiGenerationTask
api_result = await db.execute(
select(ApiGenerationTask).where(ApiGenerationTask.id == api_task_id).limit(1)
)
api_task = api_result.scalar_one_or_none()
if api_task:
api_task.video_url = final_path or output_url
api_task.status = "completed"
api_task.pipeline_stage = "done"
api_task.generated_at = upscale.completed_at
await db.commit()
return {"status": "completed", "output_url": final_path or output_url}
elif status == "failed":
upscale.status = "failed"
upscale.stage = "failed"
upscale.failure_count = int(upscale.failure_count or 0) + 1
upscale.last_error = str(query_result.error) if query_result.error else "超分失败"
upscale.failed_at = __import__("datetime").datetime.now(__import__("datetime").timezone.utc)
await db.commit()
return {"status": "failed"}
else:
# 仍在处理中,继续轮询
upscale.stage = "remote_polling"
await db.commit()
# 重新调度下一次轮询
api_upscale_poll_remote_simple.apply_async(
args=[upscale_task_id],
countdown=30,
queue=CeleryQueue.GEN_API_UPSCALE.value,
)
return {"status": "polling"}
except Exception as exc:
upscale.failure_count = int(upscale.failure_count or 0) + 1
upscale.last_error = str(exc)[:500]
await db.commit()
raise
from app.tasks.async_runner import run_async
return run_async(_execute())