Files
video-gen/video-gen-api/app/services/generation_provider_service.py
T

162 lines
6.4 KiB
Python

from __future__ import annotations
import asyncio
import json
import time
from types import SimpleNamespace
from typing import Any
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask
from app.models.image_engine import ImageEngine
from app.models.video_engine import VideoEngine
from app.services.generation_log_service import log_provider_call
from app.services.image_gen import poll_image_task_status, submit_image_task
from app.services.provider_limit import provider_limit
from app.services.video_gen import poll_task_status, submit_video_task
def _loads(data: str | None) -> dict:
if not data:
return {}
try:
obj = json.loads(data)
return obj if isinstance(obj, dict) else {}
except Exception:
return {}
async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any:
"""Use frozen snapshot for historical params, current DB row only for secret api_key."""
snapshot = _loads(task.engine_snapshot_json)
if not task.engine_id:
raise ValueError("缺少 engine_id")
if task.gen_type == "image":
result = await db.execute(select(ImageEngine).where(ImageEngine.id == task.engine_id))
else:
result = await db.execute(select(VideoEngine).where(VideoEngine.id == task.engine_id))
engine = result.scalar_one_or_none()
if not engine:
raise ValueError("引擎不存在或已删除")
return SimpleNamespace(
id=task.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"),
)
async def create_provider_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
if task.gen_type == "video":
return await _create_video_task(db, task)
if task.gen_type == "image":
return await _create_image_sync_task(db, task)
raise ValueError(f"不支持的生成类型: {task.gen_type}")
async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
"""Create video provider task through the original Ark SDK async task API."""
engine = await get_runtime_engine(db, task)
started = time.perf_counter()
async with provider_limit("ark_video_create", settings.ARK_VIDEO_CREATE_MAX_CONCURRENCY):
try:
provider_task_id = await submit_video_task(db, engine, task)
response = {"task_id": provider_task_id}
await log_provider_call(
task,
provider=engine.provider,
api_type="video_create",
model=engine.model_name,
engine_id=task.engine_id,
status="success",
latency_ms=int((time.perf_counter() - started) * 1000),
provider_task_id=provider_task_id,
response_data=response,
)
return {"task_id": provider_task_id, "response_data": response}
except Exception as exc:
await log_provider_call(
task,
provider=engine.provider,
api_type="video_create",
model=engine.model_name,
engine_id=task.engine_id,
status="failed",
latency_ms=int((time.perf_counter() - started) * 1000),
error_message=str(exc),
)
raise
async def _create_image_sync_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
"""Run the original synchronous image generation SDK under Celery control.
The legacy image SDK returns a final remote image URL immediately. We do
NOT use image_generation.tasks.create here, so image generation stays aligned
with the old working flow while no longer blocking the FastAPI request.
"""
engine = await get_runtime_engine(db, task)
started = time.perf_counter()
async with provider_limit("ark_image_sync_create", settings.ARK_IMAGE_CREATE_MAX_CONCURRENCY):
try:
result = await asyncio.to_thread(submit_image_task, db, engine, task)
if result.get("error"):
raise RuntimeError(result.get("error"))
response_data = _try_json(result.get("response_data")) or result
await log_provider_call(
task,
provider=engine.provider,
api_type="image_sync_create",
model=engine.model_name,
engine_id=task.engine_id,
status="success",
latency_ms=int((time.perf_counter() - started) * 1000),
provider_task_id=None,
response_data=response_data,
)
return {
"task_id": None,
"remote_result_url": result.get("image_url"),
"image_tokens": result.get("image_tokens", 0) or 0,
"response_data": response_data,
}
except Exception as exc:
await log_provider_call(
task,
provider=engine.provider,
api_type="image_sync_create",
model=engine.model_name,
engine_id=task.engine_id,
status="failed",
latency_ms=int((time.perf_counter() - started) * 1000),
error_message=str(exc),
)
raise
def _try_json(text: Any) -> Any:
if not isinstance(text, str):
return text
try:
return json.loads(text)
except Exception:
return None
async def poll_provider_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
engine = await get_runtime_engine(db, task)
task_id = task.seedance_task_id or task.provider_task_id
if task.gen_type == "video":
async with provider_limit("ark_video_poll", settings.ARK_VIDEO_POLL_MAX_CONCURRENCY):
return await poll_task_status(engine, task_id)
async with provider_limit("ark_image_poll", settings.ARK_IMAGE_POLL_MAX_CONCURRENCY):
return await poll_image_task_status(engine, task_id)