first commit

This commit is contained in:
baozaotumao2025
2026-07-18 21:00:26 +08:00
commit 14722be770
105 changed files with 25004 additions and 0 deletions
+232
View File
@@ -0,0 +1,232 @@
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(),
)