Files
video-gen/video-gen-api/app/services/generation/provider_service.py
T

227 lines
8.7 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 ImageProviderError, 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
from app.types.generation.provider import ImageProviderBatchResult
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 {}
def _try_json(value: Any) -> Any:
if not isinstance(value, str):
return value
try:
return json.loads(value)
except Exception:
return None
async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any:
"""使用任务快照冻结历史参数,只从当前引擎记录读取密钥。"""
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"),
multi_generation_enabled=bool(
snapshot.get("multi_generation_enabled")
if snapshot.get("multi_generation_enabled") is not None
else getattr(engine, "multi_generation_enabled", False)
),
max_generation_count=int(
snapshot.get("max_generation_count")
or getattr(engine, "max_generation_count", 1)
or 1
),
multi_image_max_images=int(
snapshot.get("multi_image_max_images")
or getattr(engine, "multi_image_max_images", 15)
or 15
),
max_reference_image_count=int(
snapshot.get("max_reference_image_count")
if snapshot.get("max_reference_image_count") is not None
else getattr(engine, "max_reference_image_count", 14)
),
output_format=(
snapshot.get("output_format")
if snapshot.get("output_format") is not None
else getattr(engine, "output_format", "")
) or "",
)
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_result(db, task)
raise ValueError(f"不支持的生成类型: {task.gen_type}")
async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
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(None, engine, task, include_media_references=True)
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_batch_result(
db: AsyncSession,
task: ChatGenerationTask,
*,
generation_count: int,
) -> ImageProviderBatchResult:
engine = await get_runtime_engine(db, task)
return await create_image_sync_batch_result_with_engine(
task,
engine,
generation_count=generation_count,
)
async def create_image_sync_batch_result_with_engine(
task: ChatGenerationTask,
engine: Any,
*,
generation_count: int,
) -> ImageProviderBatchResult:
"""执行一次同步图片请求。
generation_count > 1 时是一次组图 API 调用;失败后绝不退化为多次单图调用。
"""
count = max(1, int(generation_count or 1))
started = time.perf_counter()
api_type = "image_sync_batch_create" if count > 1 else "image_sync_create"
async with provider_limit("ark_image_sync_create", settings.ARK_IMAGE_CREATE_MAX_CONCURRENCY):
try:
result = await asyncio.to_thread(
submit_image_task,
None,
engine,
task,
include_media_references=True,
generation_count=count,
)
response_data = result.get("response_data") or result
await log_provider_call(
task,
provider=engine.provider,
api_type=api_type,
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,
total_tokens=int(result.get("image_tokens", 0) or 0),
)
return result
except Exception as exc:
error_message = exc.safe_message if isinstance(exc, ImageProviderError) else str(exc)
await log_provider_call(
task,
provider=engine.provider,
api_type=api_type,
model=engine.model_name,
engine_id=task.engine_id,
status="failed",
latency_ms=int((time.perf_counter() - started) * 1000),
error_message=error_message,
response_data=exc.as_dict() if isinstance(exc, ImageProviderError) else None,
)
raise
async def create_image_sync_result(db: AsyncSession, task: ChatGenerationTask) -> dict:
result = await create_image_sync_batch_result(db, task, generation_count=1)
items = result.get("items") or []
if len(items) != 1:
raise RuntimeError(f"图片供应商单图返回数量异常,期望 1,实际 {len(items)}")
item = items[0]
if item.get("error_message"):
raise RuntimeError(item.get("error_message") or "图片生成失败")
image_url = item.get("remote_result_url")
if not image_url:
raise RuntimeError("图片供应商未返回有效图片地址")
return {
"task_id": None,
"remote_result_url": image_url,
"image_tokens": int(result.get("image_tokens", 0) or 0),
"response_data": result.get("response_data") or {},
}
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)