233 lines
8.6 KiB
Python
233 lines
8.6 KiB
Python
import logging
|
||
from typing import Literal
|
||
|
||
from fastapi import APIRouter, Depends, Path, Request, Response, status
|
||
from sqlalchemy.exc import IntegrityError
|
||
from sqlalchemy.orm import Session
|
||
|
||
from app.api.deps import get_database, get_provider_factory
|
||
from app.core.config import Settings, get_settings
|
||
from app.core.exceptions import ConflictError, ForbiddenError, NotFoundError
|
||
from app.db.repositories.repositories import ProviderConfigRepository, UserRepository
|
||
from app.providers.model_provider import ProviderFactory
|
||
from app.schemas.api import (
|
||
DiscoverModelsResponse,
|
||
ErrorResponse,
|
||
ProviderCheckResponse,
|
||
ProviderConfigCreate,
|
||
ProviderConfigResponse,
|
||
ProviderConfigUpdate,
|
||
)
|
||
|
||
router = APIRouter()
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
def _response(row) -> ProviderConfigResponse:
|
||
return ProviderConfigResponse(
|
||
provider_id=row.profile_type,
|
||
source="database",
|
||
base_url=row.endpoint,
|
||
chat_path=row.chat_path,
|
||
models_path=row.models_path,
|
||
model_name=row.model_name,
|
||
auth_type=row.auth_type,
|
||
auth_header=row.auth_header,
|
||
auth_prefix=row.auth_prefix,
|
||
verify_ssl=row.verify_ssl,
|
||
api_key_configured=bool(row.api_key),
|
||
created_at=row.created_at,
|
||
updated_at=row.updated_at,
|
||
)
|
||
|
||
|
||
def _env_response(settings: Settings, profile_type: Literal["target", "judge"]):
|
||
prefix = f"{profile_type}_"
|
||
return ProviderConfigResponse(
|
||
provider_id=profile_type,
|
||
source="env",
|
||
base_url=getattr(settings, prefix + "base_url"),
|
||
chat_path=getattr(settings, prefix + "chat_path"),
|
||
models_path=getattr(settings, prefix + "models_path"),
|
||
model_name=getattr(settings, prefix + "model"),
|
||
auth_type=getattr(settings, prefix + "auth_type"),
|
||
auth_header=getattr(settings, prefix + "auth_header"),
|
||
auth_prefix=getattr(settings, prefix + "auth_prefix").strip(),
|
||
verify_ssl=settings.request_verify_ssl,
|
||
api_key_configured=bool(getattr(settings, prefix + "api_key")),
|
||
created_at=None,
|
||
updated_at=None,
|
||
)
|
||
|
||
|
||
def _admin(request: Request, db: Session) -> None:
|
||
user = UserRepository(db).get_active(request.state.user_id)
|
||
if not user or not user.is_admin:
|
||
raise ForbiddenError("需要管理员权限")
|
||
|
||
|
||
@router.get(
|
||
"/providers",
|
||
response_model=list[ProviderConfigResponse],
|
||
summary="列出当前生效的模型提供商配置",
|
||
)
|
||
def list_provider_configs(
|
||
db: Session = Depends(get_database), settings: Settings = Depends(get_settings)
|
||
):
|
||
persisted = {row.profile_type: row for row in ProviderConfigRepository(db).list()}
|
||
return [
|
||
_response(persisted[profile_type])
|
||
if profile_type in persisted
|
||
else _env_response(settings, profile_type)
|
||
for profile_type in ("target", "judge")
|
||
]
|
||
|
||
|
||
@router.get(
|
||
"/providers/{provider_id}",
|
||
response_model=ProviderConfigResponse,
|
||
summary="读取模型提供商配置",
|
||
)
|
||
def get_provider_config(
|
||
provider_id: Literal["target", "judge"] = Path(description="模型提供商配置 ID。"),
|
||
db: Session = Depends(get_database),
|
||
settings: Settings = Depends(get_settings),
|
||
):
|
||
row = ProviderConfigRepository(db).get_by_type(provider_id)
|
||
return _response(row) if row else _env_response(settings, provider_id)
|
||
|
||
|
||
@router.post(
|
||
"/providers", status_code=status.HTTP_201_CREATED,
|
||
response_model=ProviderConfigResponse, summary="创建模型提供商配置"
|
||
)
|
||
def create_provider_config(
|
||
payload: ProviderConfigCreate, request: Request, db: Session = Depends(get_database)
|
||
):
|
||
_admin(request, db)
|
||
values = payload.model_dump()
|
||
values["profile_type"] = values.pop("provider_id")
|
||
values["endpoint"] = values.pop("base_url")
|
||
try:
|
||
row = ProviderConfigRepository(db).create(**values)
|
||
except IntegrityError as exc:
|
||
db.rollback()
|
||
raise ConflictError(f"{payload.provider_id} 模型提供商配置已存在") from exc
|
||
logger.info("Provider config created db_id=%s provider_id=%s by_user_id=%s", row.id, row.profile_type, request.state.user_id)
|
||
return _response(row)
|
||
|
||
|
||
@router.patch(
|
||
"/providers/{provider_id}",
|
||
response_model=ProviderConfigResponse,
|
||
summary="更新模型提供商配置",
|
||
)
|
||
def update_provider_config(
|
||
provider_id: Literal["target", "judge"], payload: ProviderConfigUpdate, request: Request,
|
||
db: Session = Depends(get_database),
|
||
):
|
||
_admin(request, db)
|
||
repo = ProviderConfigRepository(db)
|
||
row = repo.get_by_type(provider_id)
|
||
if not row:
|
||
raise NotFoundError("数据库中不存在该模型提供商配置;请先创建覆盖配置")
|
||
values = payload.model_dump(exclude_unset=True)
|
||
if "base_url" in values:
|
||
values["endpoint"] = values.pop("base_url")
|
||
auth_type = values.get("auth_type", row.auth_type)
|
||
api_key = values.get("api_key", row.api_key)
|
||
if auth_type != "none" and not api_key:
|
||
raise ConflictError("启用认证时必须提供 API Key")
|
||
row = repo.update(row, **values)
|
||
logger.info("Provider config updated db_id=%s provider_id=%s by_user_id=%s", row.id, row.profile_type, request.state.user_id)
|
||
return _response(row)
|
||
|
||
|
||
@router.delete(
|
||
"/providers/{provider_id}",
|
||
status_code=status.HTTP_204_NO_CONTENT,
|
||
summary="删除模型提供商配置",
|
||
)
|
||
def delete_provider_config(
|
||
provider_id: Literal["target", "judge"], request: Request, db: Session = Depends(get_database)
|
||
) -> Response:
|
||
_admin(request, db)
|
||
repo = ProviderConfigRepository(db)
|
||
row = repo.get_by_type(provider_id)
|
||
if not row:
|
||
raise NotFoundError("数据库中不存在该模型提供商配置;环境变量配置不能删除")
|
||
profile_type = row.profile_type
|
||
repo.delete(row)
|
||
logger.info("Provider config deleted db_id=%s provider_id=%s by_user_id=%s", row.id, profile_type, request.state.user_id)
|
||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||
|
||
PROVIDER_CHECK_RESPONSES = {
|
||
500: {
|
||
"model": ErrorResponse,
|
||
"description": "服务端配置缺失,例如启用了认证但未设置对应的 API Key。",
|
||
},
|
||
502: {
|
||
"model": ErrorResponse,
|
||
"description": "上游模型服务不可达、认证失败或未返回可解析的模型响应。",
|
||
},
|
||
}
|
||
|
||
|
||
@router.post(
|
||
"/providers/{provider_id}/check",
|
||
response_model=ProviderCheckResponse,
|
||
summary="检查模型提供商连接",
|
||
description="""检查目标模型或裁判模型的认证、连通性和最小推理。
|
||
|
||
路径中的 `provider_id` 指定配置:`target` 使用数据库覆盖或 `TARGET_*` 回退,`judge` 使用
|
||
数据库覆盖或 `JUDGE_*` 回退。两者均支持以下配置项:
|
||
|
||
* `*_BASE_URL`:服务根地址(例如 `https://host/v1`)
|
||
* `*_CHAT_PATH`:聊天接口路径,默认 `/chat/completions`
|
||
* `*_MODEL`:本次最小推理使用的模型名
|
||
* `*_AUTH_TYPE`:认证类型;`none` 时不发送认证头
|
||
* `*_API_KEY`、`*_AUTH_HEADER`、`*_AUTH_PREFIX`:认证头由 prefix 与 key 组合;密钥不会出现在响应中
|
||
|
||
还可通过 `REQUEST_TIMEOUT_SECONDS` 和 `REQUEST_VERIFY_SSL` 控制超时和 TLS 校验。
|
||
修改 `.env` 后需重启服务,配置才会重新加载。成功表示服务接受认证,并能完成一次
|
||
`只回复 OK` 的最小聊天请求。""",
|
||
responses=PROVIDER_CHECK_RESPONSES,
|
||
)
|
||
async def check_provider(
|
||
provider_id: Literal["target", "judge"] = Path(description="模型提供商配置 ID。"),
|
||
factory: ProviderFactory = Depends(get_provider_factory),
|
||
):
|
||
provider = factory.create(provider_id)
|
||
reply = await provider.chat([{"role": "user", "content": "只回复 OK"}])
|
||
return ProviderCheckResponse(
|
||
provider_id=provider_id,
|
||
endpoint=provider.endpoint,
|
||
authentication={"configured": True},
|
||
model=provider.model,
|
||
reachable=True,
|
||
response_preview=reply.content[:200],
|
||
)
|
||
|
||
|
||
@router.get(
|
||
"/providers/{provider_id}/models",
|
||
response_model=DiscoverModelsResponse,
|
||
summary="查询模型提供商的可用模型",
|
||
description=(
|
||
"查询 `target` 或 `judge` 配置对应的上游 `/models` 接口。"
|
||
"结果为上游声明的模型 ID,不代表本平台已验证其可用性。"
|
||
),
|
||
responses=PROVIDER_CHECK_RESPONSES,
|
||
)
|
||
async def discover_models(
|
||
provider_id: Literal["target", "judge"] = Path(
|
||
description="模型提供商配置 ID:`target` 为被测模型,`judge` 为裁判模型。"
|
||
),
|
||
factory: ProviderFactory = Depends(get_provider_factory),
|
||
):
|
||
provider = factory.create(provider_id)
|
||
return DiscoverModelsResponse(
|
||
provider_id=provider_id,
|
||
models=await provider.list_models(),
|
||
)
|