refactor(api): LLM 提供商管理端点收敛为通用 manage 接口,端点层零特色

- schemas 新增 LlmProviderAction 公共动作词汇表
- LLMProviderManager 新增 provider_manage 统一入口,默认值填充、
  API Key 豁免判断、密钥脱敏与错误归因改写全部下沉封闭,
  不再硬编码 chatgpt/github-copilot 等产品名
- endpoint 收敛为 POST /llm/manage(ManageRequest 透传);
  OAuth 回跳地址由端点按具名回调路由统一构造后注入动作参数,
  /provider-auth/callback/{provider_id} 因浏览器回跳协议约束保留具名路由
- 新增 14 项 provider_manage 契约守护测试,
  替换原针对端点函数的两个旧测试文件
This commit is contained in:
jxxghp
2026-08-16 07:50:16 +08:00
parent a6dbd799d5
commit 441ec9475e
7 changed files with 461 additions and 675 deletions
+150 -1
View File
@@ -23,7 +23,7 @@ import jwt
from app.runtime.config import settings from app.runtime.config import settings
from app.db.oper.systemconfig import SystemConfigOper from app.db.oper.systemconfig import SystemConfigOper
from app.runtime.log import logger from app.runtime.log import logger
from app.schemas.types import SystemConfigKey from app.schemas.types import LlmProviderAction, SystemConfigKey
from app.foundation.singleton import Singleton from app.foundation.singleton import Singleton
@@ -2959,6 +2959,155 @@ class LLMProviderManager(metaclass=Singleton):
self._mark_session_error(session, str(err)) self._mark_session_error(session, str(err))
return self.get_session_status(session_id) return self.get_session_status(session_id)
async def provider_manage(self, provider: str, action: str, **params: Any) -> Dict[str, Any]:
"""
LLM 提供商统一管理入口。
按公共动作词汇表分发,统一返回 {"success", "message", "data"}
临时配置默认值填充、密钥脱敏与错误归因改写均封闭在此,
上层透传时无需感知任何提供商特色。
"""
normalized = action.value if isinstance(action, LlmProviderAction) else str(action)
try:
if normalized == LlmProviderAction.LIST_PROVIDERS.value:
return {"success": True, "message": "", "data": await self.list_providers_async()}
if normalized == LlmProviderAction.LIST_MODELS.value:
return await self._manage_list_models(provider, **params)
if normalized == LlmProviderAction.START_AUTH.value:
data = await self.start_auth(
provider, str(params.get("method") or ""), params.get("callback_url")
)
return {"success": True, "message": "", "data": data}
if normalized == LlmProviderAction.AUTH_STATUS.value:
data = self.get_session_status(str(params.get("session_id") or ""))
return {"success": True, "message": "", "data": data}
if normalized == LlmProviderAction.POLL_AUTH.value:
data = await self.poll_auth_session(str(params.get("session_id") or ""))
return {"success": True, "message": "", "data": data}
if normalized == LlmProviderAction.DISCONNECT.value:
await self.clear_auth(provider)
return {"success": True, "message": "", "data": None}
if normalized == LlmProviderAction.TEST.value:
return await self._manage_test(provider, **params)
return {"success": False, "message": f"不支持的管理动作:{normalized}", "data": None}
except Exception as err:
return {"success": False, "message": self._sanitize_error(str(err)), "data": None}
async def _manage_list_models(self, provider: str, **params: Any) -> Dict[str, Any]:
"""管理动作:查询模型目录,附带授权状态摘要。"""
from app.agent.llm.helper import LLMHelper
api_key = params.get("api_key")
try:
models = await LLMHelper().get_models(
provider=provider,
api_key=api_key,
base_url=params.get("base_url"),
base_url_preset=params.get("base_url_preset"),
user_agent=params.get("user_agent"),
use_proxy=params.get("use_proxy"),
force_refresh=bool(params.get("force_refresh", False)),
)
except Exception as err:
return {"success": False, "message": self._sanitize_error(str(err), api_key), "data": None}
return {
"success": True,
"message": "",
"data": {
"provider": provider,
"models": models,
"auth_status": self.get_auth_status(provider),
},
}
def _requires_api_key(self, provider_id: str) -> bool:
"""判断测试调用是否必须 API Key:支持 OAuth 授权或已有保存凭据的提供商可豁免。"""
try:
spec = self.get_provider(provider_id)
except Exception:
return True
if self.get_saved_auth(provider_id):
return False
return not spec.oauth_methods
async def _manage_test(self, provider: str, **params: Any) -> Dict[str, Any]:
"""管理动作:使用传入配置或当前已保存配置执行一次最小 LLM 调用。"""
from app.agent.llm.helper import LLMHelper, LLMTestTimeout
provider_name = provider or settings.LLM_PROVIDER
model = params.get("model") if params.get("model") is not None else settings.LLM_MODEL
enabled = params.get("enabled")
enabled = bool(enabled) if enabled is not None else bool(settings.AI_AGENT_ENABLE)
api_key = params.get("api_key") if params.get("api_key") is not None else settings.LLM_API_KEY
data = {"provider": provider_name, "model": model}
if not provider_name:
return {"success": False, "message": "请配置LLM提供商和模型", "data": None}
if not model or not model.strip():
return {"success": False, "message": "请先配置 LLM 模型", "data": None}
if not enabled:
return {"success": False, "message": "请先启用智能助手", "data": data}
if self._requires_api_key(provider_name) and (not api_key or not api_key.strip()):
return {"success": False, "message": "请先配置 LLM API Key", "data": data}
test_kwargs: Dict[str, Any] = {
"provider": provider_name,
"model": model,
"thinking_level": params.get("thinking_level"),
"api_key": api_key,
"base_url": params.get("base_url"),
"base_url_preset": params.get("base_url_preset"),
"user_agent": params.get("user_agent"),
"use_proxy": params.get("use_proxy"),
"api_protocol": params.get("api_protocol"),
"web_search_mode": params.get("web_search_mode"),
}
if params.get("temperature") is not None:
test_kwargs["temperature"] = params.get("temperature")
try:
result = await LLMHelper.test_current_settings(**test_kwargs)
except (LLMTestTimeout, TimeoutError) as err:
logger.warning(err)
return {"success": False, "message": "LLM 调用超时", "data": None}
except Exception as err:
return {"success": False, "message": self._sanitize_error(str(err), api_key), "data": None}
if not result.get("reply_preview"):
return {"success": False, "message": "模型响应为空", "data": result}
return {"success": True, "message": "", "data": result}
@staticmethod
def _sanitize_error(message: str, api_key: Optional[str] = None) -> str:
"""清理错误信息中的敏感字段,并把 SDK 内部解析错误改写为可定位的基础地址提示。"""
if not message:
return "LLM 没有返回任何内容"
sanitized = message
if api_key:
sanitized = sanitized.replace(api_key, "***")
sanitized = re.sub(
r"(?i)(api[_-]?key\s*[:=]\s*)([^\s,;]+)",
r"\1***",
sanitized,
)
sanitized = re.sub(
r"(?i)authorization\s*:\s*bearer\s+[^\s,;]+",
"Authorization: ***",
sanitized,
)
normalized_message = sanitized.lower().replace("_", "").replace(" ", "")
if "str" in normalized_message and (
"modeldump" in normalized_message
or "setprivateattributes" in normalized_message
):
return (
"服务返回内容不是兼容的模型响应,请检查基础地址是否填写为 "
"API Base URL,如果服务要求 /v1 等版本路径,请包含在基础地址中,"
"不要填写网页地址或完整的 chat/completions 路径"
)
return sanitized
async def _exchange_chatgpt_code_for_tokens( async def _exchange_chatgpt_code_for_tokens(
self, code: str, redirect_uri: str, code_verifier: str self, code: str, redirect_uri: str, code_verifier: str
) -> dict[str, Any]: ) -> dict[str, Any]:
+27 -302
View File
@@ -1,236 +1,45 @@
import re from typing import Any, Dict, Optional
from typing import Annotated, Optional
from fastapi import Body, Depends, Request, Response from fastapi import Depends, Request, Response
from fastapi.responses import HTMLResponse from fastapi.responses import HTMLResponse
from pydantic import BaseModel
from app import schemas from app import schemas
from app.api.response import ResponseAPIRouter from app.api.response import ResponseAPIRouter
from app.agent.llm import ( from app.agent.llm import LLMProviderManager, render_auth_result_html
LLMHelper,
LLMProviderManager,
LLMTestTimeout,
render_auth_result_html,
)
from app.runtime.config import settings
from app.db.models import User from app.db.models import User
from app.api.deps import get_current_active_superuser_async, get_current_active_user_async from app.api.deps import get_current_active_superuser_async
from app.runtime.log import logger
router = ResponseAPIRouter() router = ResponseAPIRouter()
class LlmTestRequest(BaseModel):
"""
LLM 测试调用请求参数。
"""
enabled: Optional[bool] = None
provider: Optional[str] = None
model: Optional[str] = None
thinking_level: Optional[str] = None
api_key: Optional[str] = None
base_url: Optional[str] = None
base_url_preset: Optional[str] = None
user_agent: Optional[str] = None
temperature: Optional[float] = None
use_proxy: Optional[bool] = None
api_protocol: Optional[str] = None
web_search_mode: Optional[str] = None
class LlmProviderAuthStartRequest(BaseModel):
"""
LLM 提供商授权启动请求参数。
"""
provider: str
method: str
def _sanitize_llm_error(message: str, api_key: Optional[str] = None) -> str:
"""
清理错误信息中的敏感字段,避免回显密钥。
"""
if not message:
return "LLM 没有返回任何内容"
sanitized = message
if api_key:
sanitized = sanitized.replace(api_key, "***")
sanitized = re.sub(
r"(?i)(api[_-]?key\s*[:=]\s*)([^\s,;]+)",
r"\1***",
sanitized,
)
sanitized = re.sub(
r"(?i)authorization\s*:\s*bearer\s+[^\s,;]+",
"Authorization: ***",
sanitized,
)
normalized_message = sanitized.lower().replace("_", "").replace(" ", "")
if "str" in normalized_message and (
"modeldump" in normalized_message
or "setprivateattributes" in normalized_message
):
return (
"服务返回内容不是兼容的模型响应,请检查基础地址是否填写为 "
"API Base URL,如果服务要求 /v1 等版本路径,请包含在基础地址中,"
"不要填写网页地址或完整的 chat/completions 路径"
)
return sanitized
@router.get(
"/models",
summary="获取LLM模型列表",
response_model=schemas.Response[schemas.LLMModelCatalogData],
)
async def get_llm_models(
provider: str,
api_key: Optional[str] = None,
base_url: Optional[str] = None,
base_url_preset: Optional[str] = None,
user_agent: Optional[str] = None,
use_proxy: Optional[bool] = None,
force_refresh: Optional[bool] = False,
_: User = Depends(get_current_active_user_async),
):
"""
获取指定 provider 的模型目录。
"""
try:
provider_manager = LLMProviderManager()
models = await LLMHelper().get_models(
provider=provider,
api_key=api_key,
base_url=base_url,
base_url_preset=base_url_preset,
user_agent=user_agent,
use_proxy=use_proxy,
force_refresh=bool(force_refresh),
)
return schemas.Response(
success=True,
data={
"provider": provider,
"models": models,
"auth_status": provider_manager.get_auth_status(provider),
},
)
except Exception as err:
return schemas.Response(
success=False,
message=_sanitize_llm_error(str(err), api_key),
)
@router.get(
"/providers",
summary="获取LLM提供商目录",
response_model=schemas.Response[list[schemas.LLMProviderInfo]],
)
async def get_llm_providers(
_: User = Depends(get_current_active_user_async),
):
"""
返回前端可直接渲染的 provider 目录。
"""
try:
providers = await LLMProviderManager().list_providers_async()
return schemas.Response(success=True, data=providers)
except Exception as err:
return schemas.Response(success=False, message=str(err))
@router.post( @router.post(
"/provider-auth/start", "/manage",
summary="启动LLM提供商授权", summary="LLM提供商统一管理",
response_model=schemas.Response[schemas.LLMProviderAuthSession], response_model=schemas.Response[Dict[str, Any]],
) )
async def start_llm_provider_auth( async def manage_provider(
payload: LlmProviderAuthStartRequest, request: Request,
request: Request, payload: schemas.ManageRequest,
_: User = Depends(get_current_active_superuser_async), _: User = Depends(get_current_active_superuser_async),
): ):
""" """
启动 provider 授权会话。 LLM 提供商统一管理入口:前端上送 target/action/params 原样透传,
端点不定义任何提供商特定的名称、参数或响应字段;
OAuth 回跳地址由具名回调路由统一构造后注入动作参数
""" """
try: params = dict(payload.params)
callback_url = None params.setdefault(
if payload.provider == "chatgpt" and payload.method == "browser_oauth": "callback_url",
callback_url = str( str(request.url_for("llm_provider_auth_callback", provider_id=payload.target)),
request.url_for( )
"llm_provider_auth_callback", provider_id=payload.provider result = await LLMProviderManager().provider_manage(
) payload.target, payload.action, **params
) )
result = await LLMProviderManager().start_auth( return schemas.Response(
payload.provider, success=bool(result.get("success")),
payload.method, message=result.get("message"),
callback_url, data=result.get("data"),
) )
return schemas.Response(success=True, data=result)
except Exception as err:
return schemas.Response(success=False, message=str(err))
@router.get(
"/provider-auth/{session_id}",
summary="获取LLM提供商授权会话状态",
response_model=schemas.Response[schemas.LLMProviderAuthSession],
)
async def get_llm_provider_auth_session(
session_id: str,
_: User = Depends(get_current_active_superuser_async),
):
"""
查询授权会话状态。
"""
try:
result = LLMProviderManager().get_session_status(session_id)
return schemas.Response(success=True, data=result)
except Exception as err:
return schemas.Response(success=False, message=str(err))
@router.post(
"/provider-auth/{session_id}/poll",
summary="轮询LLM提供商授权会话",
response_model=schemas.Response[schemas.LLMProviderAuthSession],
)
async def poll_llm_provider_auth_session(
session_id: str,
_: User = Depends(get_current_active_superuser_async),
):
"""
轮询 device code / OAuth 会话状态。
"""
try:
result = await LLMProviderManager().poll_auth_session(session_id)
return schemas.Response(success=True, data=result)
except Exception as err:
return schemas.Response(success=False, message=str(err))
@router.delete(
"/provider-auth/{provider_id}",
summary="断开LLM提供商授权",
response_model=schemas.Response[None],
)
async def delete_llm_provider_auth(
provider_id: str,
_: User = Depends(get_current_active_superuser_async),
):
"""
删除已保存的 provider 授权信息。
"""
try:
await LLMProviderManager().clear_auth(provider_id)
return schemas.Response(success=True)
except Exception as err:
return schemas.Response(success=False, message=str(err))
@router.get( @router.get(
@@ -264,87 +73,3 @@ async def llm_provider_auth_callback(
error_description, error_description,
) )
return HTMLResponse(content=render_auth_result_html(success, message)) return HTMLResponse(content=render_auth_result_html(success, message))
@router.post(
"/test",
summary="测试LLM调用",
response_model=schemas.Response[schemas.LLMTestResult],
)
async def llm_test(
payload: Annotated[Optional[LlmTestRequest], Body()] = None,
_: User = Depends(get_current_active_superuser_async),
):
"""
使用传入配置或当前已保存配置执行一次最小 LLM 调用。
"""
payload = payload or LlmTestRequest(
enabled=settings.AI_AGENT_ENABLE,
provider=settings.LLM_PROVIDER,
model=settings.LLM_MODEL,
thinking_level=settings.LLM_THINKING_LEVEL,
api_key=settings.LLM_API_KEY,
base_url=settings.LLM_BASE_URL,
base_url_preset=settings.LLM_BASE_URL_PRESET,
user_agent=settings.LLM_USER_AGENT,
use_proxy=settings.LLM_USE_PROXY,
api_protocol=settings.LLM_API_PROTOCOL,
web_search_mode=settings.LLM_WEB_SEARCH_MODE,
)
if not payload.provider:
return schemas.Response(success=False, message="请配置LLM提供商和模型")
if not payload.model or not payload.model.strip():
return schemas.Response(success=False, message="请先配置 LLM 模型")
data = {
"provider": payload.provider,
"model": payload.model,
}
if not payload.enabled:
return schemas.Response(success=False, message="请先启用智能助手", data=data)
if payload.provider not in {"chatgpt", "github-copilot"} and (
not payload.api_key or not payload.api_key.strip()
):
return schemas.Response(
success=False,
message="请先配置 LLM API Key",
data=data,
)
try:
test_kwargs = {
"provider": payload.provider,
"model": payload.model,
"thinking_level": payload.thinking_level,
"api_key": payload.api_key,
"base_url": payload.base_url,
"base_url_preset": payload.base_url_preset,
"user_agent": payload.user_agent,
"use_proxy": payload.use_proxy,
"api_protocol": payload.api_protocol,
"web_search_mode": payload.web_search_mode,
}
if payload.temperature is not None:
test_kwargs["temperature"] = payload.temperature
result = await LLMHelper.test_current_settings(**test_kwargs)
if not result.get("reply_preview"):
return schemas.Response(
success=False,
message="模型响应为空",
data=result,
)
return schemas.Response(success=True, data=result)
except (LLMTestTimeout, TimeoutError) as err:
logger.warning(err)
return schemas.Response(
success=False,
message="LLM 调用超时",
)
except Exception as err:
return schemas.Response(
success=False,
message=_sanitize_llm_error(str(err), payload.api_key),
)
+23
View File
@@ -531,6 +531,29 @@ class StorageAction(str, Enum):
SUPPORT_TRANSTYPE = "support_transtype" SUPPORT_TRANSTYPE = "support_transtype"
# LLM 提供商通用管理动作
class LlmProviderAction(str, Enum):
"""
LLM 提供商通用管理动作
作为提供商管理契约的公共词汇表,具体动作的支持范围与参数语义由提供商实现自行解释
"""
# 查询提供商目录
LIST_PROVIDERS = "list_providers"
# 查询模型目录
LIST_MODELS = "list_models"
# 启动授权会话
START_AUTH = "start_auth"
# 查询授权会话状态
AUTH_STATUS = "auth_status"
# 轮询授权会话
POLL_AUTH = "poll_auth"
# 断开授权
DISCONNECT = "disconnect"
# 测试调用
TEST = "test"
# 下载器类型 # 下载器类型
class DownloaderType(Enum): class DownloaderType(Enum):
# Qbittorrent # Qbittorrent
+9
View File
@@ -230,6 +230,15 @@ common `schemas.ManageRequest` body (`target` + `action` + `params`) and must
never define target-specific names, parameters or response fields — the never define target-specific names, parameters or response fields — the
frontend supplies them and the endpoint passes them through untouched. frontend supplies them and the endpoint passes them through untouched.
LLM providers follow the same contract: `LLMProviderManager.provider_manage`
dispatches actions from the shared `schemas.types.LlmProviderAction`
vocabulary, seals default-value filling, key sanitization and error rewriting
inside, and the endpoint layer exposes a single `POST /api/v1/llm/manage` with
the same `ManageRequest` body. The only exception is the named OAuth callback
route (`GET /api/v1/llm/provider-auth/callback/{provider_id}`), which stays
named because external browsers redirect to that URL; the endpoint builds the
callback URL from that route name and injects it as an action parameter.
### DB / Oper layer ### DB / Oper layer
SQLAlchemy models stay under `app/db/models/`; the data access classes live in SQLAlchemy models stay under `app/db/models/`; the data access classes live in
-77
View File
@@ -1,77 +0,0 @@
import asyncio
from unittest.mock import AsyncMock, patch
from app.api.endpoints import llm as llm_endpoint
def test_llm_test_maps_internal_model_dump_error_to_base_url_hint():
"""LLM 测试遇到 SDK 内部响应解析错误时应提示检查基础地址。"""
with patch.object(llm_endpoint.settings, "AI_AGENT_ENABLE", True), patch.object(
llm_endpoint.settings, "LLM_PROVIDER", "openai"
), patch.object(llm_endpoint.settings, "LLM_MODEL", "gpt-4o-mini"), patch.object(
llm_endpoint.settings, "LLM_API_KEY", "sk-test"
), patch.object(
llm_endpoint.settings, "LLM_BASE_URL", "https://example.com/not-api"
), patch.object(
llm_endpoint.LLMHelper,
"test_current_settings",
AsyncMock(side_effect=RuntimeError("'str' object has no attribute 'model_dump'")),
create=True,
):
resp = asyncio.run(llm_endpoint.llm_test(_="token"))
assert not resp.success
assert "基础地址" in resp.message
assert "API Base URL" in resp.message
assert "model_dump" not in resp.message
def test_llm_test_maps_internal_private_attribute_error_to_base_url_hint():
"""LLM 测试遇到 SDK 内部属性错误时应提示检查基础地址。"""
with patch.object(llm_endpoint.settings, "AI_AGENT_ENABLE", True), patch.object(
llm_endpoint.settings, "LLM_PROVIDER", "openai"
), patch.object(llm_endpoint.settings, "LLM_MODEL", "gpt-4o-mini"), patch.object(
llm_endpoint.settings, "LLM_API_KEY", "sk-test"
), patch.object(
llm_endpoint.LLMHelper,
"test_current_settings",
AsyncMock(
side_effect=AttributeError(
"'str' object has no attribute '_set_private_attributes'"
)
),
create=True,
):
resp = asyncio.run(llm_endpoint.llm_test(_="token"))
assert not resp.success
assert "基础地址" in resp.message
assert "API Base URL" in resp.message
assert "_set_private_attributes" not in resp.message
def test_llm_models_maps_internal_private_attribute_error_to_base_url_hint():
"""LLM 模型列表遇到 SDK 内部属性错误时应提示检查基础地址。"""
with patch.object(
llm_endpoint.LLMHelper,
"get_models",
AsyncMock(
side_effect=AttributeError(
"'str' object has no attribute '_set_private_attributes'"
)
),
create=True,
):
resp = asyncio.run(
llm_endpoint.get_llm_models(
provider="openai",
api_key="sk-test",
base_url="https://example.com",
_="token",
)
)
assert not resp.success
assert "基础地址" in resp.message
assert "API Base URL" in resp.message
assert "_set_private_attributes" not in resp.message
+252
View File
@@ -0,0 +1,252 @@
"""
LLM 提供商通用管理契约(provider_manage)守护测试
验证通用模式的三条核心性质:
1. 动作词汇表由 schemas 契约层统一定义,未支持动作返回统一错误结构
2. 动作标识兼容枚举与原始字符串,统一返回 {"success", "message", "data"}
3. 端点层零提供商特色:ManageRequest 原样透传,默认值填充/校验/脱敏封闭在 Manager 内
"""
import asyncio
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
import pytest
from app import schemas
from app.agent.llm.helper import LLMHelper
from app.agent.llm.provider import LLMProviderManager
from app.runtime.config import settings
from app.schemas.types import LlmProviderAction
@pytest.fixture
def manager():
return LLMProviderManager()
def test_provider_manage_rejects_unknown_action(manager):
"""动作词汇表之外的请求返回统一错误结构。"""
result = asyncio.run(manager.provider_manage("openai", "not_an_action"))
assert result["success"] is False
assert "不支持" in result["message"]
assert "data" in result
def test_provider_manage_accepts_enum_and_string_action(manager, monkeypatch):
"""动作标识兼容枚举对象与原始字符串,两种形式等价。"""
clear_mock = AsyncMock()
monkeypatch.setattr(manager, "clear_auth", clear_mock)
for action in (LlmProviderAction.DISCONNECT, "disconnect"):
result = asyncio.run(manager.provider_manage("openai", action))
assert result["success"] is True
assert result["data"] is None
assert clear_mock.await_count == 2
def test_provider_manage_test_requires_ai_agent_enabled(manager):
"""智能助手未启用时测试动作返回提示。"""
with patch.object(settings, "AI_AGENT_ENABLE", False):
result = asyncio.run(manager.provider_manage("deepseek", "test"))
assert result["success"] is False
assert result["message"] == "请先启用智能助手"
def test_provider_manage_test_requires_model(manager):
"""未配置模型时测试动作返回提示。"""
with patch.object(settings, "AI_AGENT_ENABLE", True), patch.object(
settings, "LLM_API_KEY", "sk-test"
), patch.object(settings, "LLM_MODEL", ""):
result = asyncio.run(manager.provider_manage("deepseek", "test"))
assert result["success"] is False
assert result["message"] == "请先配置 LLM 模型"
def test_provider_manage_test_requires_api_key(manager, monkeypatch):
"""无 OAuth 授权方式且无已保存凭据的提供商必须配置 API Key。"""
monkeypatch.setattr(manager, "get_saved_auth", lambda provider_id: None)
with patch.object(settings, "AI_AGENT_ENABLE", True), patch.object(
settings, "LLM_API_KEY", None
), patch.object(settings, "LLM_MODEL", "deepseek-chat"):
result = asyncio.run(manager.provider_manage("deepseek", "test"))
assert result["success"] is False
assert result["message"] == "请先配置 LLM API Key"
assert result["data"]["model"] == "deepseek-chat"
def test_provider_manage_test_exempts_oauth_providers_from_api_key(manager, monkeypatch):
"""支持 OAuth 授权的提供商无需 API Key,端点与 Manager 均不硬编码提供商名。"""
monkeypatch.setattr(manager, "get_saved_auth", lambda provider_id: None)
monkeypatch.setattr(
manager, "get_provider", lambda provider_id: SimpleNamespace(oauth_methods=("browser_oauth",))
)
test_mock = AsyncMock(return_value={"provider": "chatgpt", "model": "gpt-4o", "reply_preview": "OK"})
with patch.object(settings, "AI_AGENT_ENABLE", True), patch.object(
settings, "LLM_API_KEY", None
), patch.object(LLMHelper, "test_current_settings", test_mock):
result = asyncio.run(
manager.provider_manage("chatgpt", "test", model="gpt-4o")
)
assert result["success"] is True
test_mock.assert_awaited_once()
def test_provider_manage_test_returns_reply_preview(manager, monkeypatch):
"""测试成功时返回模型响应预览,显式参数优先于已保存配置。"""
monkeypatch.setattr(manager, "get_saved_auth", lambda provider_id: None)
test_mock = AsyncMock(
return_value={"provider": "openai", "model": "gpt-4.1-mini", "duration_ms": 123, "reply_preview": "OK"}
)
with patch.object(settings, "AI_AGENT_ENABLE", False), patch.object(
LLMHelper, "test_current_settings", test_mock
):
result = asyncio.run(
manager.provider_manage(
"openai",
"test",
enabled=True,
model="gpt-4.1-mini",
thinking_level="high",
api_key="sk-live",
base_url="https://example.com/v1",
use_proxy=False,
)
)
test_mock.assert_awaited_once_with(
provider="openai",
model="gpt-4.1-mini",
thinking_level="high",
api_key="sk-live",
base_url="https://example.com/v1",
base_url_preset=None,
user_agent=None,
use_proxy=False,
api_protocol=None,
web_search_mode=None,
)
assert result["success"] is True
assert result["data"]["reply_preview"] == "OK"
def test_provider_manage_test_rejects_empty_reply(manager, monkeypatch):
"""模型响应为空时返回失败但保留结果详情。"""
monkeypatch.setattr(manager, "get_saved_auth", lambda provider_id: None)
with patch.object(settings, "AI_AGENT_ENABLE", True), patch.object(
LLMHelper,
"test_current_settings",
AsyncMock(return_value={"provider": "deepseek", "model": "deepseek-chat", "duration_ms": 12}),
):
result = asyncio.run(
manager.provider_manage("deepseek", "test", model="deepseek-chat", api_key="sk-test")
)
assert result["success"] is False
assert result["message"] == "模型响应为空"
assert result["data"]["duration_ms"] == 12
def test_provider_manage_test_maps_timeout_error(manager, monkeypatch):
"""调用超时返回统一的超时提示。"""
monkeypatch.setattr(manager, "get_saved_auth", lambda provider_id: None)
with patch.object(settings, "AI_AGENT_ENABLE", True), patch.object(
LLMHelper,
"test_current_settings",
AsyncMock(side_effect=TimeoutError("request timed out")),
):
result = asyncio.run(
manager.provider_manage("deepseek", "test", model="deepseek-chat", api_key="sk-test")
)
assert result["success"] is False
assert result["message"] == "LLM 调用超时"
def test_provider_manage_test_sanitizes_error_message(manager, monkeypatch):
"""错误信息中的密钥与授权头必须脱敏。"""
monkeypatch.setattr(manager, "get_saved_auth", lambda provider_id: None)
raw_error = (
"request failed api_key=sk-secret "
"Authorization: Bearer sk-secret "
"base error sk-secret"
)
with patch.object(settings, "AI_AGENT_ENABLE", True), patch.object(
LLMHelper, "test_current_settings", AsyncMock(side_effect=RuntimeError(raw_error))
):
result = asyncio.run(
manager.provider_manage("deepseek", "test", model="deepseek-chat", api_key="sk-secret")
)
assert result["success"] is False
assert "sk-secret" not in result["message"]
assert "Authorization: Bearer" not in result["message"]
assert "***" in result["message"]
def test_provider_manage_test_maps_internal_error_to_base_url_hint(manager, monkeypatch):
"""SDK 内部响应解析错误应改写为可定位的基础地址提示。"""
monkeypatch.setattr(manager, "get_saved_auth", lambda provider_id: None)
with patch.object(settings, "AI_AGENT_ENABLE", True), patch.object(
LLMHelper,
"test_current_settings",
AsyncMock(side_effect=RuntimeError("'str' object has no attribute 'model_dump'")),
):
result = asyncio.run(
manager.provider_manage("openai", "test", model="gpt-4o-mini", api_key="sk-test")
)
assert result["success"] is False
assert "基础地址" in result["message"]
assert "API Base URL" in result["message"]
assert "model_dump" not in result["message"]
def test_provider_manage_list_models_sanitizes_base_url_hint(manager):
"""模型列表查询遇到 SDK 内部错误时同样给出基础地址提示。"""
with patch.object(
LLMHelper,
"get_models",
AsyncMock(side_effect=AttributeError("'str' object has no attribute '_set_private_attributes'")),
):
result = asyncio.run(
manager.provider_manage("openai", "list_models", api_key="sk-test", base_url="https://example.com")
)
assert result["success"] is False
assert "基础地址" in result["message"]
assert "API Base URL" in result["message"]
assert "_set_private_attributes" not in result["message"]
def test_provider_manage_list_models_returns_catalog_with_auth_status(manager):
"""模型目录查询成功时附带授权状态摘要。"""
models = [{"id": "gpt-4o"}]
with patch.object(LLMHelper, "get_models", AsyncMock(return_value=models)), patch.object(
manager, "get_auth_status", lambda provider_id: {"connected": False}
):
result = asyncio.run(manager.provider_manage("openai", "list_models"))
assert result["success"] is True
assert result["data"]["provider"] == "openai"
assert result["data"]["models"] == models
assert result["data"]["auth_status"] == {"connected": False}
def test_llm_manage_endpoint_passes_through_manage_request(monkeypatch):
"""端点仅透传 ManageRequest,并按具名回调路由注入 OAuth 回跳地址。"""
from app.api.endpoints import llm as llm_endpoint
captured = {}
async def fake_manage(self, provider, action, **params):
captured["provider"] = provider
captured["action"] = action
captured["params"] = params
return {"success": True, "message": "", "data": {"ok": True}}
monkeypatch.setattr(LLMProviderManager, "provider_manage", fake_manage)
request = SimpleNamespace(
url_for=lambda name, **kwargs: f"https://host/api/v1/llm/provider-auth/callback/{kwargs['provider_id']}"
)
payload = schemas.ManageRequest(target="chatgpt", action="start_auth", params={"method": "browser_oauth"})
resp = asyncio.run(llm_endpoint.manage_provider(request, payload, _="token"))
assert resp.success is True
assert resp.data == {"ok": True}
assert captured["provider"] == "chatgpt"
assert captured["action"] == "start_auth"
assert captured["params"]["method"] == "browser_oauth"
assert captured["params"]["callback_url"].endswith("/callback/chatgpt")
-295
View File
@@ -1,295 +0,0 @@
import asyncio
import unittest
from types import ModuleType
from unittest.mock import AsyncMock, patch
from app.testing import stub_modules
def _stub(name: str, **attrs) -> tuple:
"""构造带指定属性的占位模块,返回 ``(模块名, 模块)`` 供 :func:`stub_modules` 使用。"""
module = ModuleType(name)
for key, value in attrs.items():
setattr(module, key, value)
return name, module
class _Dummy:
def __init__(self, *args, **kwargs):
pass
def __getattr__(self, _name):
return lambda *args, **kwargs: None
class _DummyError(Exception):
def __init__(self, message="", duration_ms=None):
super().__init__(message)
self.duration_ms = duration_ms
# 在 import 期用占位模块替换重依赖/外部模块,import 完由 stub_modules 精确还原,避免污染其它用例
_STUB_MODULES = dict([
_stub("pillow_avif"),
_stub("aiofiles"),
_stub("psutil"),
_stub("app.application.site.sites", SitesHelper=_Dummy),
_stub("app.chain.mediaserver", MediaServerChain=_Dummy),
_stub("app.chain.search", SearchChain=_Dummy),
_stub("app.chain.system", SystemChain=_Dummy),
_stub("app.agent.llm", LLMHelper=_Dummy, LLMProviderManager=_Dummy,
LLMTestError=_DummyError, LLMTestTimeout=_DummyError,
render_auth_result_html=lambda success, message: message),
_stub("app.runtime.events", eventmanager=_Dummy(), Event=_Dummy, EventManager=_Dummy),
_stub("app.domain.metainfo", MetaInfo=_Dummy),
_stub("app.runtime.extensions.module_manager", ModuleManager=_Dummy),
_stub("app.application.security.access", verify_apitoken=_Dummy, verify_resource_token=_Dummy, verify_token=_Dummy),
_stub("app.db.models", User=_Dummy),
_stub("app.db.oper.systemconfig", SystemConfigOper=_Dummy),
_stub("app.api.deps", get_current_active_superuser=_Dummy,
get_current_active_superuser_async=_Dummy, get_current_active_user_async=_Dummy),
_stub("app.application.mediaserver", MediaServerHelper=_Dummy),
_stub("app.application.messaging.message", MessageHelper=_Dummy),
_stub("app.runtime.progress", ProgressHelper=_Dummy),
_stub("app.application.filter", RuleHelper=_Dummy),
_stub("app.adapters.external.server", MoviePilotServerHelper=_Dummy),
_stub("app.runtime.state", SystemHelper=_Dummy),
_stub("app.application.image", ImageHelper=_Dummy),
_stub("app.scheduler", Scheduler=_Dummy),
_stub("app.runtime.log", logger=_Dummy(), log_settings=_Dummy(),
LogConfigModel=type("LogConfigModel", (), {})),
_stub("app.foundation.crypto", HashUtils=_Dummy),
_stub("app.adapters.network.http", RequestUtils=_Dummy, AsyncRequestUtils=_Dummy),
_stub("version", APP_VERSION="test"),
])
with stub_modules(_STUB_MODULES):
from app.api.endpoints import llm as system_endpoint
class LlmTestEndpointTest(unittest.TestCase):
def test_llm_test_requires_ai_agent_enabled(self):
with patch.object(system_endpoint.settings, "AI_AGENT_ENABLE", False):
resp = asyncio.run(system_endpoint.llm_test(_="token"))
self.assertFalse(resp.success)
self.assertEqual(resp.message, "请先启用智能助手")
def test_llm_test_requires_api_key(self):
with patch.object(system_endpoint.settings, "AI_AGENT_ENABLE", True), patch.object(
system_endpoint.settings, "LLM_API_KEY", None
), patch.object(system_endpoint.settings, "LLM_MODEL", "deepseek-chat"):
resp = asyncio.run(system_endpoint.llm_test(_="token"))
self.assertFalse(resp.success)
self.assertEqual(resp.message, "请先配置 LLM API Key")
self.assertEqual(resp.data["model"], "deepseek-chat")
def test_llm_test_requires_model(self):
with patch.object(system_endpoint.settings, "AI_AGENT_ENABLE", True), patch.object(
system_endpoint.settings, "LLM_API_KEY", "sk-test"
), patch.object(system_endpoint.settings, "LLM_MODEL", ""):
resp = asyncio.run(system_endpoint.llm_test(_="token"))
self.assertFalse(resp.success)
self.assertEqual(resp.message, "请先配置 LLM 模型")
def test_llm_test_returns_successful_reply_preview(self):
llm_test_mock = AsyncMock(
return_value={
"provider": "deepseek",
"model": "deepseek-chat",
"duration_ms": 321,
"reply_preview": "OK",
}
)
with patch.object(system_endpoint.settings, "AI_AGENT_ENABLE", True), patch.object(
system_endpoint.settings, "LLM_PROVIDER", "deepseek"
), patch.object(system_endpoint.settings, "LLM_MODEL", "deepseek-chat"), patch.object(
system_endpoint.settings, "LLM_THINKING_LEVEL", "max"
), patch.object(
system_endpoint.settings, "LLM_API_KEY", "sk-test"
), patch.object(
system_endpoint.settings, "LLM_BASE_URL", "https://api.deepseek.com"
), patch.object(
system_endpoint.settings, "LLM_BASE_URL_PRESET", "deepseek-default"
), patch.object(
system_endpoint.settings, "LLM_USER_AGENT", "MoviePilot-Test/1.0"
), patch.object(
system_endpoint.settings, "LLM_USE_PROXY", True
), patch.object(
system_endpoint.settings, "LLM_API_PROTOCOL", "responses"
), patch.object(
system_endpoint.LLMHelper,
"test_current_settings",
llm_test_mock,
create=True,
):
resp = asyncio.run(system_endpoint.llm_test(_="token"))
llm_test_mock.assert_awaited_once_with(
provider="deepseek",
model="deepseek-chat",
thinking_level="max",
api_key="sk-test",
base_url="https://api.deepseek.com",
base_url_preset="deepseek-default",
user_agent="MoviePilot-Test/1.0",
use_proxy=True,
api_protocol="responses",
web_search_mode="local",
)
self.assertTrue(resp.success)
self.assertEqual(resp.data["provider"], "deepseek")
self.assertEqual(resp.data["model"], "deepseek-chat")
self.assertEqual(resp.data["duration_ms"], 321)
self.assertEqual(resp.data["reply_preview"], "OK")
def test_llm_test_prefers_request_payload_over_saved_settings(self):
llm_test_mock = AsyncMock(
return_value={
"provider": "openai",
"model": "gpt-4.1-mini",
"duration_ms": 123,
"reply_preview": "OK",
}
)
payload = system_endpoint.LlmTestRequest(
enabled=True,
provider="openai",
model="gpt-4.1-mini",
thinking_level="high",
api_key="sk-live",
base_url="https://example.com/v1",
base_url_preset="openai-default",
user_agent="MoviePilot-Custom/1.0",
use_proxy=False,
)
with patch.object(system_endpoint.settings, "AI_AGENT_ENABLE", False), patch.object(
system_endpoint.settings, "LLM_PROVIDER", "deepseek"
), patch.object(system_endpoint.settings, "LLM_MODEL", "deepseek-chat"), patch.object(
system_endpoint.settings, "LLM_API_KEY", "sk-saved"
), patch.object(
system_endpoint.settings, "LLM_BASE_URL", "https://api.deepseek.com"
), patch.object(
system_endpoint.LLMHelper,
"test_current_settings",
llm_test_mock,
create=True,
):
resp = asyncio.run(system_endpoint.llm_test(payload=payload, _="token"))
llm_test_mock.assert_awaited_once_with(
provider="openai",
model="gpt-4.1-mini",
thinking_level="high",
api_key="sk-live",
base_url="https://example.com/v1",
base_url_preset="openai-default",
user_agent="MoviePilot-Custom/1.0",
use_proxy=False,
api_protocol=None,
web_search_mode=None,
)
self.assertTrue(resp.success)
self.assertEqual(resp.data["provider"], "openai")
self.assertEqual(resp.data["model"], "gpt-4.1-mini")
def test_llm_test_supports_legacy_thinking_payload(self):
llm_test_mock = AsyncMock(
return_value={
"provider": "deepseek",
"model": "deepseek-v4-pro",
"duration_ms": 123,
"reply_preview": "OK",
}
)
payload = system_endpoint.LlmTestRequest(
enabled=True,
provider="deepseek",
model="deepseek-v4-pro",
api_key="sk-live",
base_url="https://api.deepseek.com",
base_url_preset="deepseek-default",
user_agent=None,
use_proxy=None,
)
with patch.object(system_endpoint.settings, "AI_AGENT_ENABLE", False), patch.object(
system_endpoint.LLMHelper,
"test_current_settings",
llm_test_mock,
create=True,
):
resp = asyncio.run(system_endpoint.llm_test(payload=payload, _="token"))
llm_test_mock.assert_awaited_once_with(
provider="deepseek",
model="deepseek-v4-pro",
thinking_level=None,
api_key="sk-live",
base_url="https://api.deepseek.com",
base_url_preset="deepseek-default",
user_agent=None,
use_proxy=None,
api_protocol=None,
web_search_mode=None,
)
self.assertTrue(resp.success)
def test_llm_test_rejects_empty_reply(self):
with patch.object(system_endpoint.settings, "AI_AGENT_ENABLE", True), patch.object(
system_endpoint.settings, "LLM_PROVIDER", "deepseek"
), patch.object(system_endpoint.settings, "LLM_MODEL", "deepseek-chat"), patch.object(
system_endpoint.settings, "LLM_API_KEY", "sk-test"
), patch.object(
system_endpoint.LLMHelper,
"test_current_settings",
AsyncMock(return_value={"provider": "deepseek", "model": "deepseek-chat", "duration_ms": 12}),
create=True,
):
resp = asyncio.run(system_endpoint.llm_test(_="token"))
self.assertFalse(resp.success)
self.assertEqual(resp.message, "模型响应为空")
self.assertEqual(resp.data["duration_ms"], 12)
def test_llm_test_maps_timeout_error(self):
with patch.object(system_endpoint.settings, "AI_AGENT_ENABLE", True), patch.object(
system_endpoint.settings, "LLM_PROVIDER", "deepseek"
), patch.object(system_endpoint.settings, "LLM_MODEL", "deepseek-chat"), patch.object(
system_endpoint.settings, "LLM_API_KEY", "sk-test"
), patch.object(
system_endpoint.LLMHelper,
"test_current_settings",
AsyncMock(side_effect=TimeoutError("request timed out")),
create=True,
):
resp = asyncio.run(system_endpoint.llm_test(_="token"))
self.assertFalse(resp.success)
self.assertEqual(resp.message, "LLM 调用超时")
def test_llm_test_sanitizes_error_message(self):
raw_error = (
"request failed api_key=sk-secret "
"Authorization: Bearer sk-secret "
"base error sk-secret"
)
with patch.object(system_endpoint.settings, "AI_AGENT_ENABLE", True), patch.object(
system_endpoint.settings, "LLM_API_KEY", "sk-secret"
), patch.object(system_endpoint.settings, "LLM_PROVIDER", "deepseek"), patch.object(
system_endpoint.settings, "LLM_MODEL", "deepseek-chat"
), patch.object(
system_endpoint.LLMHelper,
"test_current_settings",
AsyncMock(side_effect=RuntimeError(raw_error)),
create=True,
):
resp = asyncio.run(system_endpoint.llm_test(_="token"))
self.assertFalse(resp.success)
self.assertNotIn("sk-secret", resp.message)
self.assertNotIn("Authorization: Bearer", resp.message)
self.assertIn("***", resp.message)