127 lines
5.1 KiB
Python
127 lines
5.1 KiB
Python
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,
|
||
}
|