from __future__ import annotations import json from fastapi import HTTPException from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.enums.common import MAX_GENERATION_COUNT, MIN_GENERATION_COUNT from app.enums.generation_provider import IMAGE_MULTI_OUTPUT_MAX, IMAGE_MULTI_REFERENCE_MAX from app.models.image_engine import ImageEngine from app.models.video_engine import VideoEngine IMAGE_DEFAULT_SIZE = "2K" IMAGE_DEFAULT_PROPORTION = "1:1" IMAGE_DEFAULT_PX = "2048x2048" VIDEO_DEFAULT_DURATION = 4 VIDEO_DEFAULT_RATIO = "16:9" VIDEO_DEFAULT_RESOLUTION = "480p" def normalize_px(value: str | None) -> str | None: if not value: return value return value.replace("×", "x").replace("X", "x").replace("×x", "x").replace("x×", "x") def parse_json_list(value: str | None, fallback: list): try: parsed = json.loads(value or "") return parsed if isinstance(parsed, list) else fallback except Exception: return fallback def image_supported_sizes(engine: ImageEngine) -> dict: try: data = json.loads(engine.supported_sizes or "{}") return data if isinstance(data, dict) else {} except Exception: return {} def normalize_generation_count(value: int | None) -> int: try: count = int(value or MIN_GENERATION_COUNT) except (TypeError, ValueError): count = MIN_GENERATION_COUNT return min(MAX_GENERATION_COUNT, max(MIN_GENERATION_COUNT, count)) async def get_image_engine(db: AsyncSession, engine_id: str | None) -> ImageEngine: query = select(ImageEngine).where(ImageEngine.is_active == True, ImageEngine.deleted_at.is_(None)) if engine_id: query = query.where(ImageEngine.id == engine_id) else: query = query.order_by(ImageEngine.priority.desc()).limit(1) result = await db.execute(query) engine = result.scalar_one_or_none() if not engine: raise HTTPException(status_code=400, detail="没有可用的图片引擎") return engine async def get_video_engine(db: AsyncSession, engine_id: str | None) -> VideoEngine: query = select(VideoEngine).where(VideoEngine.is_active == True, VideoEngine.deleted_at.is_(None)) if engine_id: query = query.where(VideoEngine.id == engine_id) else: query = query.order_by(VideoEngine.priority.desc()) result = await db.execute(query.limit(1)) engine = result.scalar_one_or_none() if not engine: raise HTTPException(status_code=400, detail="没有可用的视频引擎") return engine def build_image_snapshot(engine: ImageEngine, size: str, proportion: str, px: str) -> dict: return { "engine_type": "image", "id": engine.id, "name": engine.name, "provider": engine.provider, "api_base": engine.api_base, "api_key_masked": "****" if engine.api_key else "", "model_name": engine.model_name, "generate_url": engine.generate_url, "supported_models": parse_json_list(engine.supported_models, []), "default_size": engine.default_size, "multi_generation_enabled": bool(getattr(engine, "multi_generation_enabled", False)), "max_generation_count": normalize_generation_count(getattr(engine, "max_generation_count", 1)), "multi_image_max_images": int(getattr(engine, "multi_image_max_images", IMAGE_MULTI_OUTPUT_MAX) or IMAGE_MULTI_OUTPUT_MAX), "max_reference_image_count": int(getattr(engine, "max_reference_image_count", IMAGE_MULTI_REFERENCE_MAX) or 0), "output_format": (getattr(engine, "output_format", "") or "").lower().strip(), "selected_size": size, "selected_proportion": proportion, "selected_px": px, } def build_video_snapshot(engine: VideoEngine, ratio: str, resolution: str, duration: int) -> dict: return { "engine_type": "video", "id": engine.id, "name": engine.name, "provider": engine.provider, "api_base": engine.api_base, "api_key_masked": "****" if engine.api_key else "", "model_name": engine.model_name, "generate_url": engine.generate_url, "query_url": engine.query_url, "supported_ratios": parse_json_list(engine.supported_ratios, []), "supported_resolutions": parse_json_list(engine.supported_resolutions, []), "supported_durations": parse_json_list(engine.supported_durations, []), "max_duration": engine.max_duration, "max_image_count": int(getattr(engine, "max_image_count", 0) or 0), "max_video_count": int(getattr(engine, "max_video_count", 0) or 0), "max_audio_count": int(getattr(engine, "max_audio_count", 0) or 0), "supports_universal_reference": bool(getattr(engine, "supports_universal_reference", False)), "supports_first_last_frame": bool(getattr(engine, "supports_first_last_frame", False)), "multi_generation_enabled": bool(getattr(engine, "multi_generation_enabled", False)), "max_generation_count": normalize_generation_count(getattr(engine, "max_generation_count", 1)), "selected_ratio": ratio, "selected_resolution": resolution, "selected_duration": duration, }