mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-04 23:17:20 +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.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]:
|
||||
|
||||
+27
-302
@@ -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),
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user