From 441ec9475e02a9012620c19b787274708e799c1b Mon Sep 17 00:00:00 2001 From: jxxghp Date: Sun, 16 Aug 2026 07:50:16 +0800 Subject: [PATCH] =?UTF-8?q?refactor(api):=20LLM=20=E6=8F=90=E4=BE=9B?= =?UTF-8?q?=E5=95=86=E7=AE=A1=E7=90=86=E7=AB=AF=E7=82=B9=E6=94=B6=E6=95=9B?= =?UTF-8?q?=E4=B8=BA=E9=80=9A=E7=94=A8=20manage=20=E6=8E=A5=E5=8F=A3?= =?UTF-8?q?=EF=BC=8C=E7=AB=AF=E7=82=B9=E5=B1=82=E9=9B=B6=E7=89=B9=E8=89=B2?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - schemas 新增 LlmProviderAction 公共动作词汇表 - LLMProviderManager 新增 provider_manage 统一入口,默认值填充、 API Key 豁免判断、密钥脱敏与错误归因改写全部下沉封闭, 不再硬编码 chatgpt/github-copilot 等产品名 - endpoint 收敛为 POST /llm/manage(ManageRequest 透传); OAuth 回跳地址由端点按具名回调路由统一构造后注入动作参数, /provider-auth/callback/{provider_id} 因浏览器回跳协议约束保留具名路由 - 新增 14 项 provider_manage 契约守护测试, 替换原针对端点函数的两个旧测试文件 --- app/agent/llm/provider.py | 151 +++++++++- app/api/endpoints/llm.py | 329 ++-------------------- app/schemas/types.py | 23 ++ docs/rules/05-architecture.md | 9 + tests/test_llm_endpoint_error_messages.py | 77 ----- tests/test_llm_provider_manage.py | 252 +++++++++++++++++ tests/test_system_llm_test.py | 295 ------------------- 7 files changed, 461 insertions(+), 675 deletions(-) delete mode 100644 tests/test_llm_endpoint_error_messages.py create mode 100644 tests/test_llm_provider_manage.py delete mode 100644 tests/test_system_llm_test.py diff --git a/app/agent/llm/provider.py b/app/agent/llm/provider.py index d837e5082..d50da01a1 100644 --- a/app/agent/llm/provider.py +++ b/app/agent/llm/provider.py @@ -23,7 +23,7 @@ import jwt from app.runtime.config import settings from app.db.oper.systemconfig import SystemConfigOper 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 @@ -2959,6 +2959,155 @@ class LLMProviderManager(metaclass=Singleton): self._mark_session_error(session, str(err)) 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( self, code: str, redirect_uri: str, code_verifier: str ) -> dict[str, Any]: diff --git a/app/api/endpoints/llm.py b/app/api/endpoints/llm.py index 1ed38f720..93882e609 100644 --- a/app/api/endpoints/llm.py +++ b/app/api/endpoints/llm.py @@ -1,236 +1,45 @@ -import re -from typing import Annotated, Optional +from typing import Any, Dict, Optional -from fastapi import Body, Depends, Request, Response +from fastapi import Depends, Request, Response from fastapi.responses import HTMLResponse -from pydantic import BaseModel from app import schemas from app.api.response import ResponseAPIRouter -from app.agent.llm import ( - LLMHelper, - LLMProviderManager, - LLMTestTimeout, - render_auth_result_html, -) -from app.runtime.config import settings +from app.agent.llm import LLMProviderManager, render_auth_result_html from app.db.models import User -from app.api.deps import get_current_active_superuser_async, get_current_active_user_async -from app.runtime.log import logger +from app.api.deps import get_current_active_superuser_async 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( - "/provider-auth/start", - summary="启动LLM提供商授权", - response_model=schemas.Response[schemas.LLMProviderAuthSession], + "/manage", + summary="LLM提供商统一管理", + response_model=schemas.Response[Dict[str, Any]], ) -async def start_llm_provider_auth( - payload: LlmProviderAuthStartRequest, - request: Request, - _: User = Depends(get_current_active_superuser_async), +async def manage_provider( + request: Request, + payload: schemas.ManageRequest, + _: User = Depends(get_current_active_superuser_async), ): """ - 启动 provider 授权会话。 + LLM 提供商统一管理入口:前端上送 target/action/params 原样透传, + 端点不定义任何提供商特定的名称、参数或响应字段; + OAuth 回跳地址由具名回调路由统一构造后注入动作参数 """ - try: - callback_url = None - if payload.provider == "chatgpt" and payload.method == "browser_oauth": - callback_url = str( - request.url_for( - "llm_provider_auth_callback", provider_id=payload.provider - ) - ) - result = await LLMProviderManager().start_auth( - payload.provider, - payload.method, - callback_url, - ) - 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)) + params = dict(payload.params) + params.setdefault( + "callback_url", + str(request.url_for("llm_provider_auth_callback", provider_id=payload.target)), + ) + result = await LLMProviderManager().provider_manage( + payload.target, payload.action, **params + ) + return schemas.Response( + success=bool(result.get("success")), + message=result.get("message"), + data=result.get("data"), + ) @router.get( @@ -264,87 +73,3 @@ async def llm_provider_auth_callback( error_description, ) 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), - ) diff --git a/app/schemas/types.py b/app/schemas/types.py index 943504a4c..d46efbdf4 100644 --- a/app/schemas/types.py +++ b/app/schemas/types.py @@ -531,6 +531,29 @@ class StorageAction(str, Enum): 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): # Qbittorrent diff --git a/docs/rules/05-architecture.md b/docs/rules/05-architecture.md index 7ae7d643c..9c263ab3a 100644 --- a/docs/rules/05-architecture.md +++ b/docs/rules/05-architecture.md @@ -230,6 +230,15 @@ common `schemas.ManageRequest` body (`target` + `action` + `params`) and must never define target-specific names, parameters or response fields — the 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 SQLAlchemy models stay under `app/db/models/`; the data access classes live in diff --git a/tests/test_llm_endpoint_error_messages.py b/tests/test_llm_endpoint_error_messages.py deleted file mode 100644 index 708dc596d..000000000 --- a/tests/test_llm_endpoint_error_messages.py +++ /dev/null @@ -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 diff --git a/tests/test_llm_provider_manage.py b/tests/test_llm_provider_manage.py new file mode 100644 index 000000000..bc57142e0 --- /dev/null +++ b/tests/test_llm_provider_manage.py @@ -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") diff --git a/tests/test_system_llm_test.py b/tests/test_system_llm_test.py deleted file mode 100644 index 9d03775e2..000000000 --- a/tests/test_system_llm_test.py +++ /dev/null @@ -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)