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(), )