Files
video-gen/video-gen-api/app/services/generation/ai/engine_service.py
T

127 lines
5.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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,
}