mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-05 23:47:41 +08:00
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:
+150
-1
@@ -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
@@ -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),
|
|
||||||
)
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
|
||||||
@@ -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")
|
||||||
@@ -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)
|
|
||||||
Reference in New Issue
Block a user