243 lines
9.4 KiB
Python
243 lines
9.4 KiB
Python
import asyncio
|
|
import json
|
|
import logging
|
|
import os
|
|
from datetime import datetime
|
|
|
|
from sqlalchemy import select
|
|
|
|
from app.models.base import async_session
|
|
from app.models.generation_record import GenerationRecord
|
|
from app.services.video_gen import get_active_engine, poll_task_status, download_video, _log_video_response
|
|
from app.services.image_gen import get_active_image_engine, download_image
|
|
from app.services.resource_accounting_service import (
|
|
record_generation_record_generated_resource,
|
|
safe_file_size,
|
|
)
|
|
from app.config import settings
|
|
|
|
logger = logging.getLogger("videogen")
|
|
|
|
POLL_INTERVAL = 30 # seconds between polls
|
|
MAX_POLLS = 60 # max 30 minutes total
|
|
|
|
|
|
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),
|
|
)
|
|
)
|
|
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:
|
|
record.status = "failed"
|
|
record.error_message = "缺少外部任务ID"
|
|
await db.commit()
|
|
return
|
|
|
|
try:
|
|
engine = await get_active_engine(db)
|
|
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:
|
|
record.status = "failed"
|
|
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", "")
|
|
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"
|
|
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.video_tokens_used = poll_result.get("video_tokens", 0)
|
|
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,
|
|
)
|
|
self._active.pop(record_id, None)
|
|
await db.commit()
|
|
logger.info(f"Video task completed: {record_id}")
|
|
|
|
elif status == "failed":
|
|
record.status = "failed"
|
|
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:
|
|
record.status = "failed"
|
|
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
|
|
|
|
try:
|
|
engine = await get_active_image_engine(db)
|
|
poll_result = await asyncio.to_thread(submit_image_task, db, engine, record)
|
|
|
|
if poll_result["error"] == "":
|
|
remote_url = poll_result.get("image_url")
|
|
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.image_tokens_used = poll_result.get("image_tokens", 0)
|
|
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 db.commit()
|
|
logger.info(f"Image task completed: {record_id}")
|
|
else:
|
|
record.status = "failed"
|
|
record.error_message = poll_result.get("error", "图片生成失败")
|
|
await db.commit()
|
|
logger.info(f"Image task failed: {record_id}")
|
|
_log_image_response(record_id, poll_result)
|
|
|
|
except Exception as e:
|
|
record.status = "failed"
|
|
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()
|