120 lines
4.6 KiB
Python
120 lines
4.6 KiB
Python
import json
|
|
import logging
|
|
|
|
from fastapi import APIRouter, Depends
|
|
from fastapi.responses import JSONResponse
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.dependencies import get_db
|
|
from app.models.image_engine import ImageEngine
|
|
from app.models.video_engine import VideoEngine
|
|
from app.schemas.api_v3.model import ApiModelInfo, ApiModelsResponse
|
|
from app.services.api_v3 import auth_service
|
|
from app.services.api_v3.pricing_service import get_priced_models
|
|
|
|
logger = logging.getLogger("videogen")
|
|
|
|
router = APIRouter(prefix="/models", tags=["api-v3-models"])
|
|
|
|
|
|
@router.get(
|
|
"",
|
|
summary="获取可用模型列表",
|
|
description="获取当前 API Key 可调用的所有视频和图片模型(仅返回已配置价格的模型)",
|
|
)
|
|
async def list_models(
|
|
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
|
|
db: AsyncSession = Depends(get_db),
|
|
) -> JSONResponse:
|
|
"""获取当前 API Key 可用的模型列表。"""
|
|
models: list[ApiModelInfo] = []
|
|
|
|
# 获取所有已配置价格的引擎 ID 集合
|
|
priced_engine_ids = await get_priced_models(db)
|
|
|
|
# 获取 API Key 的白名单引擎 ID 集合
|
|
allowed_engine_ids = {m.get("engine_id", "") for m in key_context.callable_models} if key_context.callable_models else set()
|
|
|
|
# 确定要返回的引擎 ID 列表
|
|
target_engine_ids = priced_engine_ids if not allowed_engine_ids else (allowed_engine_ids & priced_engine_ids)
|
|
|
|
# 构建引擎信息映射
|
|
engine_info_map = {m.get("engine_id", ""): m for m in key_context.callable_models}
|
|
|
|
for engine_id in target_engine_ids:
|
|
engine_type = engine_info_map.get(engine_id, {}).get("engine_type", "")
|
|
model_name = engine_info_map.get(engine_id, {}).get("model_name", "")
|
|
|
|
# 如果没有从白名单获取到类型,尝试从数据库加载
|
|
if not engine_type:
|
|
video_result = await db.execute(
|
|
select(VideoEngine).where(VideoEngine.id == engine_id, VideoEngine.deleted_at.is_(None)).limit(1)
|
|
)
|
|
if video_result.scalar_one_or_none():
|
|
engine_type = "video"
|
|
else:
|
|
image_result = await db.execute(
|
|
select(ImageEngine).where(ImageEngine.id == engine_id, ImageEngine.deleted_at.is_(None)).limit(1)
|
|
)
|
|
if image_result.scalar_one_or_none():
|
|
engine_type = "image"
|
|
|
|
# 加载引擎详情
|
|
supported_ratios = None
|
|
supported_resolutions = None
|
|
supported_durations = None
|
|
supported_sizes = None
|
|
|
|
try:
|
|
if engine_type == "video":
|
|
result = await db.execute(
|
|
select(VideoEngine).where(VideoEngine.id == engine_id, VideoEngine.deleted_at.is_(None)).limit(1)
|
|
)
|
|
engine = result.scalar_one_or_none()
|
|
if engine:
|
|
if not model_name:
|
|
model_name = engine.model_name
|
|
supported_ratios = _parse_json_list(engine.supported_ratios)
|
|
supported_resolutions = _parse_json_list(engine.supported_resolutions)
|
|
supported_durations = _parse_json_list(engine.supported_durations)
|
|
elif engine_type == "image":
|
|
result = await db.execute(
|
|
select(ImageEngine).where(ImageEngine.id == engine_id, ImageEngine.deleted_at.is_(None)).limit(1)
|
|
)
|
|
engine = result.scalar_one_or_none()
|
|
if engine:
|
|
if not model_name:
|
|
model_name = engine.model_name
|
|
supported_sizes = _parse_json_list(engine.supported_sizes)
|
|
except Exception:
|
|
pass
|
|
|
|
info = ApiModelInfo(
|
|
model=model_name,
|
|
engine_type=engine_type,
|
|
engine_id=engine_id,
|
|
supported_ratios=supported_ratios,
|
|
supported_resolutions=supported_resolutions,
|
|
supported_durations=supported_durations,
|
|
supported_sizes=supported_sizes,
|
|
)
|
|
|
|
models.append(info)
|
|
|
|
return JSONResponse(
|
|
content={"code": 0, "data": {"models": [m.model_dump() for m in models]}, "message": "ok"},
|
|
status_code=200,
|
|
)
|
|
|
|
|
|
def _parse_json_list(value: str | None) -> list[str | int] | None:
|
|
"""解析 JSON 列表字段。"""
|
|
if not value:
|
|
return None
|
|
try:
|
|
parsed = json.loads(value)
|
|
return parsed if isinstance(parsed, list) else None
|
|
except (json.JSONDecodeError, TypeError):
|
|
return None
|