Files
video-gen/video-gen-api/app/services/video_queue.py
T
2026-07-11 12:48:29 +08:00

377 lines
15 KiB
Python

import asyncio
import json
import logging
import os
from datetime import datetime
from types import SimpleNamespace
from sqlalchemy import select
from app.models.base import async_session
from app.models.generation_record import GenerationRecord
from app.models.image_engine import ImageEngine
from app.models.video_engine import VideoEngine
from app.enums.model_pricing import PricingSnapshotStage, ProviderCostStatus
from app.services.video_gen import poll_task_status, download_video, _log_video_response
from app.services.image_gen import download_image, is_sync_image_provider_result_uncertain
from app.services.media_token_usage_snapshot_service import (
mark_media_provider_cost_status,
sync_generation_record_media_token_snapshot,
)
from app.services.resource_accounting_service import (
record_generation_record_generated_resource,
safe_file_size,
)
from app.services.video_cover_service import create_video_cover_for_local_video
from app.config import settings
from app.services.generation_refund_service import mark_generation_record_failed_and_refund_once
logger = logging.getLogger("videogen")
POLL_INTERVAL = 30 # seconds between polls
MAX_POLLS = 60 # max 30 minutes total
def _loads(data: str | None) -> dict:
if not data:
return {}
try:
value = json.loads(data)
return value if isinstance(value, dict) else {}
except Exception:
return {}
async def _get_runtime_engine(db, record: GenerationRecord):
"""使用生成开始时冻结的引擎快照,数据库行只读取当前密钥。"""
if not record.engine_id:
raise ValueError("生成记录缺少锁定的 engine_id")
snapshot = _loads(record.engine_snapshot_json)
model = ImageEngine if record.gen_type == "image" else VideoEngine
engine = (await db.execute(select(model).where(model.id == record.engine_id).limit(1))).scalar_one_or_none()
if not engine:
raise ValueError("锁定的生成引擎不存在")
return SimpleNamespace(
id=record.engine_id,
name=snapshot.get("name") or engine.name,
provider=snapshot.get("provider") or engine.provider,
api_base=snapshot.get("api_base") or engine.api_base,
api_key=engine.api_key,
model_name=snapshot.get("model_name") or engine.model_name,
generate_url=snapshot.get("generate_url") or getattr(engine, "generate_url", ""),
query_url=snapshot.get("query_url") or getattr(engine, "query_url", ""),
default_size=snapshot.get("default_size") or getattr(engine, "default_size", "2K"),
)
class TaskQueue:
def __init__(self):
self.queue: asyncio.Queue[str] = asyncio.Queue()
self.running = False
self._active: dict[str, int] = {} # record_id -> poll count
async def enqueue(self, record_id: str):
"""Add a record to the polling queue."""
await self.queue.put(record_id)
async def recover(self):
"""Recover in-progress tasks from DB on startup."""
async with async_session() as db:
result = await db.execute(
select(GenerationRecord).where(
GenerationRecord.status == "generating",
GenerationRecord.seedance_task_id.isnot(None),
GenerationRecord.deleted_at.is_(None),
)
)
records = result.scalars().all()
for record in records:
await self.queue.put(record.id)
logger.info(f"Recovered task: {record.id} (seedance: {record.seedance_task_id})")
async def run(self):
"""Main polling loop."""
self.running = True
logger.info("Video queue started")
while self.running:
try:
record_id = await asyncio.wait_for(self.queue.get(), timeout=5.0)
except asyncio.TimeoutError:
continue
try:
await self._process(record_id)
except Exception as e:
logger.error(f"Error processing {record_id}: {e}")
finally:
self.queue.task_done()
logger.info("Video queue stopped")
async def _process(self, record_id: str):
"""Process a single record: poll status and update DB."""
async with async_session() as db:
result = await db.execute(
select(GenerationRecord).where(
GenerationRecord.id == record_id,
GenerationRecord.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
record = result.scalar_one_or_none()
if not record or record.status != "generating":
return
if record.gen_type == "video":
await self._process_video(db, record)
else:
await self._process_image(db, record)
async def _process_video(self, db, record):
"""Process video generation task."""
record_id = record.id
if not record.seedance_task_id:
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message="缺少外部任务ID",
)
await db.commit()
return
try:
engine = await _get_runtime_engine(db, record)
poll_result = await poll_task_status(engine, record.seedance_task_id)
except Exception as e:
logger.error(f"Poll error for {record_id}: {e}")
count = self._active.get(record_id, 0) + 1
self._active[record_id] = count
if count >= MAX_POLLS:
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=f"轮询超时: {e}",
)
await db.commit()
del self._active[record_id]
else:
await asyncio.sleep(POLL_INTERVAL)
await self.queue.put(record_id)
return
status = poll_result["status"]
try:
resp_data = json.loads(poll_result.get("response_data", "{}"))
except (json.JSONDecodeError, TypeError):
resp_data = {}
_log_video_response(record_id, resp_data, poll_result.get("error"))
if status == "succeeded":
file_url = poll_result.get("video_url", "")
resp_data["video_url"] = file_url
resp_data["task_id"] = record.seedance_task_id
record.provider_response_json = json.dumps(resp_data, ensure_ascii=False, default=str)
record.video_tokens_used = poll_result.get("video_tokens", 0)
# 视频 Provider 已完成时核算供应商成本;下载失败不影响已发生的供应商费用。
await sync_generation_record_media_token_snapshot(
db,
record,
provider_response=resp_data,
stage=PricingSnapshotStage.PROVIDER_ASYNC_COMPLETED.value,
)
storage_path = None
file_size_bytes = 0
if settings.STORAGE_TYPE == "local" and file_url:
try:
date_dir = datetime.now().strftime("%Y/%m/%d")
dest_dir = os.path.join(settings.STORAGE_LOCAL_PATH, date_dir)
os.makedirs(dest_dir, exist_ok=True)
dest = os.path.join(dest_dir, f"{record_id}.mp4")
await download_video(file_url, dest)
record.video_url = f"/generate/videos/{date_dir}/{record_id}.mp4"
cover_url, _cover_storage_path = create_video_cover_for_local_video(
record_id=record_id,
video_path=dest,
date_dir=date_dir,
log_prefix=f"GenerationRecord视频封面生成 record_id={record_id}",
)
record.video_cover_url = cover_url
storage_path = dest
file_size_bytes = safe_file_size(dest)
except Exception as e:
logger.warning(f"Download failed, using remote URL: {e}")
record.video_url = file_url
else:
record.video_url = file_url
record.status = "completed"
record.generated_at = datetime.now()
if record.video_url:
await record_generation_record_generated_resource(
db,
record,
resource_url=record.video_url,
storage_path=storage_path,
file_size_bytes=file_size_bytes,
remote_url=file_url,
generated_at=record.generated_at,
)
await sync_generation_record_media_token_snapshot(
db,
record,
provider_response=resp_data,
stage=PricingSnapshotStage.RESOURCE_DOWNLOAD_COMPLETED.value,
)
self._active.pop(record_id, None)
await db.commit()
logger.info(f"Video task completed: {record_id}")
elif status == "failed":
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=poll_result.get("error", "视频生成失败"),
)
self._active.pop(record_id, None)
await db.commit()
logger.info(f"Video task failed: {record_id}")
else:
count = self._active.get(record_id, 0) + 1
self._active[record_id] = count
if count >= MAX_POLLS:
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message="视频生成超时",
)
self._active.pop(record_id, None)
await db.commit()
logger.info(f"Video task timed out: {record_id}")
else:
await db.commit()
await asyncio.sleep(POLL_INTERVAL)
await self.queue.put(record_id)
async def _process_image(self, db, record):
"""Process image generation task - calls API directly."""
record_id = record.id
from app.services.image_gen import submit_image_task, _log_image_response
provider_call_started = False
provider_call_completed = False
try:
engine = await _get_runtime_engine(db, record)
provider_call_started = True
poll_result = await asyncio.to_thread(
submit_image_task,
db,
engine,
record,
include_media_references=False,
)
provider_call_completed = True
if poll_result["error"] == "":
remote_url = poll_result.get("image_url")
try:
provider_response = json.loads(poll_result.get("response_data") or "{}")
except (json.JSONDecodeError, TypeError):
provider_response = {}
record.provider_response_json = json.dumps(provider_response, ensure_ascii=False, default=str)
record.image_tokens_used = poll_result.get("image_tokens", 0)
# 火山图片接口为同步生成:最终响应返回后立即完成供应商成本核算。
await sync_generation_record_media_token_snapshot(
db,
record,
provider_response=provider_response,
stage=PricingSnapshotStage.PROVIDER_SYNC_COMPLETED.value,
)
storage_path = None
file_size_bytes = 0
if settings.STORAGE_TYPE == "local" and remote_url:
try:
date_dir = datetime.now().strftime("%Y/%m/%d")
dest_dir = os.path.join(settings.STORAGE_IMAGE_LOCAL_PATH, date_dir)
os.makedirs(dest_dir, exist_ok=True)
dest = os.path.join(dest_dir, f"{record_id}.png")
await download_image(remote_url, dest)
record.image_url = f"/generate/images/{date_dir}/{record_id}.png"
storage_path = dest
file_size_bytes = safe_file_size(dest)
except Exception as e:
logger.warning(f"Download failed, using remote URL: {e}")
record.image_url = remote_url
else:
record.image_url = remote_url
record.status = "completed"
record.generated_at = datetime.now()
if record.image_url:
await record_generation_record_generated_resource(
db,
record,
resource_url=record.image_url,
storage_path=storage_path,
file_size_bytes=file_size_bytes,
remote_url=remote_url,
generated_at=record.generated_at,
)
await sync_generation_record_media_token_snapshot(
db,
record,
provider_response=provider_response,
stage=PricingSnapshotStage.RESOURCE_DOWNLOAD_COMPLETED.value,
)
await db.commit()
logger.info(f"Image task completed: {record_id}")
else:
error_message = poll_result.get("error", "图片生成失败")
await mark_media_provider_cost_status(
db,
owner=record,
status=ProviderCostStatus.NOT_INCURRED.value,
reason=error_message,
usage_stage="provider_sync_failed",
)
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=error_message,
)
await db.commit()
logger.info(f"Image task failed: {record_id}")
_log_image_response(record_id, poll_result)
except Exception as e:
if provider_call_started:
uncertain = provider_call_completed or is_sync_image_provider_result_uncertain(e)
await mark_media_provider_cost_status(
db,
owner=record,
status=(
ProviderCostStatus.PROVIDER_RESULT_UNCERTAIN.value
if uncertain
else ProviderCostStatus.NOT_INCURRED.value
),
reason=str(e),
usage_stage="provider_sync_exception",
)
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=str(e),
)
_log_image_response(record_id, {}, str(e))
await db.commit()
logger.error(f"Image task failed: {record_id}, error: {e}")
def stop(self):
"""Signal the queue to stop."""
self.running = False
task_queue = TaskQueue()