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, 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_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, include_media_references=True, ) 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)