Files
ai-safety-platform/app/api/v1/endpoints/providers.py
T
baozaotumao2025 14722be770 first commit
2026-07-18 21:00:26 +08:00

233 lines
8.6 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.
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(),
)