248 lines
9.1 KiB
Python
248 lines
9.1 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.services.generation.pipeline.owner_service import (
|
|
GenerationOwner,
|
|
owner_include_media_references,
|
|
owner_provider_task_id,
|
|
)
|
|
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
|
|
from app.types.generation.provider import ImageProviderBatchResult
|
|
from app.utils.id_gen import generate_id
|
|
|
|
|
|
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
|
|
|
|
|
|
def _snapshot_owner(task: GenerationOwner) -> SimpleNamespace:
|
|
"""Copy loaded scalar fields before commit closes the current transaction."""
|
|
values = {
|
|
key: value
|
|
for key, value in vars(task).items()
|
|
if key != "_sa_instance_state"
|
|
}
|
|
values.setdefault("generation_mode", getattr(task, "generation_mode", None) or "generation_record")
|
|
return SimpleNamespace(**values)
|
|
|
|
|
|
async def get_runtime_engine(db: AsyncSession, task: GenerationOwner) -> 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: GenerationOwner) -> 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: GenerationOwner) -> dict:
|
|
engine = await get_runtime_engine(db, task)
|
|
task_snapshot = _snapshot_owner(task)
|
|
include_references = owner_include_media_references(task_snapshot)
|
|
# Close the engine lookup transaction before the long provider HTTP call.
|
|
await db.commit()
|
|
async with provider_limit("ark_video_create", settings.ARK_VIDEO_CREATE_MAX_CONCURRENCY):
|
|
provider_task_id = await submit_video_task(
|
|
None,
|
|
engine,
|
|
task_snapshot,
|
|
include_media_references=include_references,
|
|
)
|
|
return {"task_id": provider_task_id, "response_data": {"task_id": provider_task_id}}
|
|
|
|
|
|
async def create_image_sync_batch_result(
|
|
db: AsyncSession,
|
|
task: GenerationOwner,
|
|
*,
|
|
generation_count: int,
|
|
) -> ImageProviderBatchResult:
|
|
engine = await get_runtime_engine(db, task)
|
|
task_snapshot = _snapshot_owner(task)
|
|
# Do not keep a database transaction open while the synchronous provider call runs.
|
|
await db.commit()
|
|
return await create_image_sync_batch_result_with_engine(
|
|
task_snapshot,
|
|
engine,
|
|
generation_count=generation_count,
|
|
)
|
|
|
|
|
|
async def create_image_sync_batch_result_with_engine(
|
|
task: GenerationOwner,
|
|
engine: Any,
|
|
*,
|
|
generation_count: int,
|
|
) -> ImageProviderBatchResult:
|
|
"""执行一次同步图片请求。
|
|
|
|
generation_count > 1 时是一次组图 API 调用;失败后绝不退化为多次单图调用。
|
|
"""
|
|
count = max(1, int(generation_count or 1))
|
|
async with provider_limit("ark_image_sync_create", settings.ARK_IMAGE_CREATE_MAX_CONCURRENCY):
|
|
return await asyncio.to_thread(
|
|
submit_image_task,
|
|
None,
|
|
engine,
|
|
task,
|
|
include_media_references=owner_include_media_references(task),
|
|
generation_count=count,
|
|
)
|
|
|
|
|
|
async def create_image_sync_result(db: AsyncSession, task: GenerationOwner) -> 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: GenerationOwner) -> dict:
|
|
engine = await get_runtime_engine(db, task)
|
|
task_snapshot = _snapshot_owner(task)
|
|
task_id = owner_provider_task_id(task_snapshot)
|
|
# Polling may block on the remote provider; release the lookup transaction first.
|
|
await db.commit()
|
|
if not task_id:
|
|
raise ValueError("缺少供应商任务ID")
|
|
|
|
api_type = f"{task_snapshot.gen_type}_poll"
|
|
call_id = generate_id()
|
|
await log_provider_call(
|
|
task_snapshot,
|
|
provider=engine.provider,
|
|
api_type=api_type,
|
|
model=engine.model_name,
|
|
engine_id=task_snapshot.engine_id,
|
|
status="request",
|
|
provider_task_id=task_id,
|
|
request_data={"provider_task_id": task_id},
|
|
call_id=call_id,
|
|
)
|
|
started = time.perf_counter()
|
|
try:
|
|
if task_snapshot.gen_type == "video":
|
|
async with provider_limit("ark_video_poll", settings.ARK_VIDEO_POLL_MAX_CONCURRENCY):
|
|
result = await poll_task_status(engine, task_id)
|
|
else:
|
|
async with provider_limit("ark_image_poll", settings.ARK_IMAGE_POLL_MAX_CONCURRENCY):
|
|
result = await poll_image_task_status(engine, task_id)
|
|
await log_provider_call(
|
|
task_snapshot,
|
|
provider=engine.provider,
|
|
api_type=api_type,
|
|
model=engine.model_name,
|
|
engine_id=task_snapshot.engine_id,
|
|
status="success",
|
|
latency_ms=int((time.perf_counter() - started) * 1000),
|
|
provider_task_id=task_id,
|
|
response_data=_try_json(result.get("response_data")) or result,
|
|
total_tokens=int(result.get("video_tokens", 0) or result.get("image_tokens", 0) or 0),
|
|
call_id=call_id,
|
|
)
|
|
return result
|
|
except Exception as exc:
|
|
await log_provider_call(
|
|
task_snapshot,
|
|
provider=engine.provider,
|
|
api_type=api_type,
|
|
model=engine.model_name,
|
|
engine_id=task_snapshot.engine_id,
|
|
status="failed",
|
|
latency_ms=int((time.perf_counter() - started) * 1000),
|
|
provider_task_id=task_id,
|
|
error_message=str(exc),
|
|
call_id=call_id,
|
|
)
|
|
raise
|