162 lines
6.4 KiB
Python
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).limit(1))
|
|
else:
|
|
result = await db.execute(select(VideoEngine).where(VideoEngine.id == task.engine_id).limit(1))
|
|
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)
|