生成记录接口兼容图片类型|CHAT对话生成视频/图片任务API+celery异步模块完成
This commit is contained in:
@@ -0,0 +1,161 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user