85 lines
2.2 KiB
Python
85 lines
2.2 KiB
Python
from __future__ import annotations
|
|
|
|
import os
|
|
from functools import lru_cache
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import yaml
|
|
from fastapi import HTTPException
|
|
from pydantic import BaseModel, Field
|
|
|
|
_ROOT_DIR = Path(__file__).resolve().parents[1]
|
|
_MODELS_FILE = _ROOT_DIR / "config" / "models.yaml"
|
|
|
|
|
|
class ModelSpec(BaseModel):
|
|
id: str
|
|
label: str
|
|
provider: str = "deepseek"
|
|
api_model: str
|
|
base_url: str
|
|
api_key_env: str
|
|
max_tokens: int = 4000
|
|
# 直接并入请求体的模型专属参数,例如 thinking / reasoning_effort。
|
|
params: dict[str, Any] = Field(default_factory=dict)
|
|
|
|
|
|
class ModelsFile(BaseModel):
|
|
default: str
|
|
models: list[ModelSpec] = Field(default_factory=list)
|
|
|
|
|
|
class ModelOption(BaseModel):
|
|
id: str
|
|
label: str
|
|
|
|
|
|
class ModelsListResponse(BaseModel):
|
|
default: str
|
|
models: list[ModelOption]
|
|
|
|
|
|
@lru_cache
|
|
def load_models_file() -> ModelsFile:
|
|
if not _MODELS_FILE.is_file():
|
|
raise RuntimeError(f"缺少模型配置文件:{_MODELS_FILE}")
|
|
raw = yaml.safe_load(_MODELS_FILE.read_text(encoding="utf-8")) or {}
|
|
data = ModelsFile.model_validate(raw)
|
|
if not data.models:
|
|
raise RuntimeError("models.yaml 中 models 不能为空")
|
|
ids = {m.id for m in data.models}
|
|
if data.default not in ids:
|
|
raise RuntimeError(f"models.yaml 的 default「{data.default}」不在 models 列表中")
|
|
return data
|
|
|
|
|
|
def list_model_options() -> ModelsListResponse:
|
|
data = load_models_file()
|
|
return ModelsListResponse(
|
|
default=data.default,
|
|
models=[ModelOption(id=m.id, label=m.label) for m in data.models],
|
|
)
|
|
|
|
|
|
def get_model_spec(model_id: str | None = None) -> ModelSpec:
|
|
data = load_models_file()
|
|
chosen = (model_id or "").strip() or data.default
|
|
for item in data.models:
|
|
if item.id == chosen:
|
|
return item
|
|
raise HTTPException(
|
|
status_code=400,
|
|
detail=f"不支持的模型「{chosen}」,请从 /api/ai/models 列表中选择",
|
|
)
|
|
|
|
|
|
def resolve_api_key(spec: ModelSpec) -> str:
|
|
key = (os.getenv(spec.api_key_env) or "").strip()
|
|
if not key:
|
|
raise HTTPException(
|
|
status_code=500,
|
|
detail=f"未配置密钥环境变量:{spec.api_key_env}",
|
|
)
|
|
return key
|