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" } # 只有原对象确实有 generation_mode 时才保留,避免 SimpleNamespace 快照被误判为 ChatGenerationTask if getattr(task, "generation_mode", None): values.setdefault("generation_mode", task.generation_mode) 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