mirror of
https://github.com/fastapiadmin/FastapiAdmin.git
synced 2026-09-21 20:55:14 +00:00
feat(mcp): 重构 MCP 服务器及智能助手模块
- 新增代码生成模块路由支持 - 重构 MCP 服务器API,支持增删改查接口和分页查询 - 实现 MCP 智能对话的流式响应和 WebSocket 通信 - 优化 MCP 数据模型,支持环境变量等配置项 - 改进 MCP CRUD 层,集成权限控制和统一操作方法 - 优化异常处理,提升接口错误信息准确性 - 修正依赖导入和数据校验,增强代码健壮性
This commit is contained in:
@@ -1,45 +1,35 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, Path, Query
|
||||
from starlette.responses import StreamingResponse
|
||||
|
||||
from backend.common.pagination import DependsPagination, PageData, paging_data
|
||||
from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base
|
||||
from backend.common.security.jwt import DependsJwtAuth
|
||||
from backend.common.security.permission import RequestPermission
|
||||
from backend.common.security.rbac import DependsRBAC
|
||||
from backend.database.db import CurrentSession
|
||||
from .schema import CreateMcpParam, GetMcpDetail, McpChatParam, UpdateMcpParam
|
||||
from .service import mcp_service
|
||||
|
||||
from fastapi import APIRouter, Depends, WebSocket
|
||||
from fastapi import APIRouter, Depends, Path, Query, Body, WebSocket
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
|
||||
from app.common.response import StreamResponse, SuccessResponse
|
||||
from app.common.request import PaginationService
|
||||
from app.core.base_params import PaginationQueryParam
|
||||
from app.core.dependencies import AuthPermission
|
||||
from app.core.router_class import OperationLogRoute
|
||||
from app.core.logger import logger
|
||||
from app.api.v1.module_system.auth.schema import AuthSchema
|
||||
from .service import MCPService
|
||||
from .schema import ChatQuerySchema
|
||||
from .param import McpQueryParam
|
||||
from .service import McpService
|
||||
from .schema import McpCreateSchema, McpUpdateSchema, ChatQuerySchema
|
||||
|
||||
|
||||
MCPRouter = APIRouter(route_class=OperationLogRoute, prefix="", tags=["MCP智能助手"])
|
||||
MCPRouter = APIRouter(route_class=OperationLogRoute, prefix="/mcp", tags=["MCP智能助手"])
|
||||
|
||||
|
||||
@MCPRouter.post("/mcp/chat", summary="智能对话", description="与MCP智能助手进行对话")
|
||||
@MCPRouter.post("/chat", summary="智能对话", description="与MCP智能助手进行对话")
|
||||
async def chat_controller(
|
||||
query: ChatQuerySchema,
|
||||
auth: AuthSchema = Depends(AuthPermission(permissions=["ai:mcp:chat"]))
|
||||
) -> StreamingResponse:
|
||||
"""智能对话接口"""
|
||||
logger.info(f"用户 {auth.user.name} 发起智能对话: {query.message[:50]}...")
|
||||
user_name = auth.user.name if auth.user else "未知用户"
|
||||
logger.info(f"用户 {user_name} 发起智能对话: {query.message[:50]}...")
|
||||
|
||||
async def generate_response():
|
||||
try:
|
||||
async for chunk in MCPService.chat_query(query.message):
|
||||
async for chunk in McpService.chat_query(query=query):
|
||||
# 确保返回的是字节串
|
||||
if chunk:
|
||||
yield chunk.encode('utf-8') if isinstance(chunk, str) else chunk
|
||||
@@ -47,11 +37,65 @@ async def chat_controller(
|
||||
logger.error(f"流式响应出错: {str(e)}")
|
||||
yield f"抱歉,处理您的请求时出现了错误: {str(e)}".encode('utf-8')
|
||||
|
||||
return StreamingResponse(generate_response(), media_type="text/plain; charset=utf-8")
|
||||
return StreamResponse(generate_response(), media_type="text/plain; charset=utf-8")
|
||||
|
||||
|
||||
@MCPRouter.websocket("/ws/mcp/chat", name="WebSocket聊天")
|
||||
@MCPRouter.get("/detail/{id}", summary="获取 MCP 服务器详情", description="获取 MCP 服务器详情")
|
||||
async def get_mcp_detail_controller(
|
||||
id: int = Path(..., description="MCP ID"),
|
||||
auth: AuthSchema = Depends(AuthPermission(permissions=["ai:mcp:query"]))
|
||||
) -> JSONResponse:
|
||||
result_dict = await McpService.get_mcp_detail_service(auth=auth, id=id)
|
||||
logger.info(f"获取 MCP 服务器详情成功 {id}")
|
||||
return SuccessResponse(data=result_dict, msg="获取 MCP 服务器详情成功")
|
||||
|
||||
|
||||
@MCPRouter.get("/list", summary="查询 MCP 服务器列表", description="查询 MCP 服务器列表")
|
||||
async def get_mcp_list_controller(
|
||||
page: PaginationQueryParam = Depends(),
|
||||
search: McpQueryParam = Depends(),
|
||||
auth: AuthSchema = Depends(AuthPermission(permissions=["ai:mcp:query"]))
|
||||
) -> JSONResponse:
|
||||
result_dict_list = await McpService.get_mcp_list_service(auth=auth, search=search, order_by=page.order_by)
|
||||
result_dict = await PaginationService.paginate(data_list=result_dict_list, page_no=page.page_no, page_size=page.page_size)
|
||||
logger.info(f"查询 MCP 服务器列表成功")
|
||||
return SuccessResponse(data=result_dict, msg="查询 MCP 服务器列表成功")
|
||||
|
||||
|
||||
@MCPRouter.post("/create", summary="创建 MCP 服务器", description="创建 MCP 服务器")
|
||||
async def create_mcp_controller(
|
||||
data: McpCreateSchema,
|
||||
auth: AuthSchema = Depends(AuthPermission(permissions=["ai:mcp:create"]))
|
||||
) -> JSONResponse:
|
||||
result_dict = await McpService.create_mcp_service(auth=auth, data=data)
|
||||
logger.info(f"创建 MCP 服务器成功: {result_dict}")
|
||||
return SuccessResponse(data=result_dict, msg="创建 MCP 服务器成功")
|
||||
|
||||
|
||||
@MCPRouter.put("/update/{id}", summary="修改 MCP 服务器", description="修改 MCP 服务器")
|
||||
async def update_mcp_controller(
|
||||
data: McpUpdateSchema,
|
||||
id: int = Path(..., description="MCP ID"),
|
||||
auth: AuthSchema = Depends(AuthPermission(permissions=["ai:mcp:update"]))
|
||||
) -> JSONResponse:
|
||||
result_dict = await McpService.update_mcp_service(auth=auth, id=id, data=data)
|
||||
logger.info(f"修改 MCP 服务器成功: {result_dict}")
|
||||
return SuccessResponse(data=result_dict, msg="修改 MCP 服务器成功")
|
||||
|
||||
|
||||
@MCPRouter.delete("/delete", summary="删除 MCP 服务器", description="删除 MCP 服务器")
|
||||
async def delete_mcp_controller(
|
||||
ids: list[int] = Body(..., description="ID列表"),
|
||||
auth: AuthSchema = Depends(AuthPermission(permissions=["ai:mcp:delete"]))
|
||||
) -> JSONResponse:
|
||||
await McpService.delete_mcp_service(auth=auth, ids=ids)
|
||||
logger.info(f"删除 MCP 服务器成功: {ids}")
|
||||
return SuccessResponse(msg="删除 MCP 服务器成功")
|
||||
|
||||
|
||||
@MCPRouter.websocket("/ws/chat", name="WebSocket聊天")
|
||||
async def websocket_chat_controller(
|
||||
query: ChatQuerySchema,
|
||||
websocket: WebSocket,
|
||||
):
|
||||
"""WebSocket聊天接口
|
||||
@@ -64,7 +108,7 @@ async def websocket_chat_controller(
|
||||
data = await websocket.receive_text()
|
||||
# 流式发送响应
|
||||
try:
|
||||
async for chunk in MCPService.chat_query(data):
|
||||
async for chunk in McpService.chat_query(query=query):
|
||||
if chunk:
|
||||
await websocket.send_text(chunk)
|
||||
except Exception as e:
|
||||
@@ -74,48 +118,3 @@ async def websocket_chat_controller(
|
||||
logger.error(f"WebSocket聊天出错: {str(e)}")
|
||||
finally:
|
||||
await websocket.close()
|
||||
|
||||
|
||||
@MCPRouter.get('/{pk}', summary='获取 MCP 服务器详情', dependencies=[DependsJwtAuth])
|
||||
async def get_mcp(pk: Annotated[int, Path(description='MCP ID')]) -> ResponseSchemaModel[GetMcpDetail]:
|
||||
mcp = await mcp_service.get(pk=pk)
|
||||
return response_base.success(data=mcp)
|
||||
|
||||
|
||||
@MCPRouter.get('', summary='分页获取所有 MCP 服务器', dependencies=[DependsJwtAuth, DependsPagination])
|
||||
async def get_pagination_mcps(
|
||||
db: CurrentSession,
|
||||
name: Annotated[str | None, Query(description='MCP 名称')] = None,
|
||||
type: Annotated[int | None, Query(description='MCP 类型')] = None,
|
||||
) -> ResponseSchemaModel[PageData[GetMcpDetail]]:
|
||||
mcp_select = await mcp_service.get_select(name=name, type=type)
|
||||
page_data = await paging_data(db, mcp_select)
|
||||
return response_base.success(data=page_data)
|
||||
|
||||
|
||||
@MCPRouter.post('', summary='创建 MCP 服务器', dependencies=[Depends(RequestPermission('sys:mcp:add')), DependsRBAC])
|
||||
async def create_mcp(obj: CreateMcpParam) -> ResponseModel:
|
||||
await mcp_service.create(obj=obj)
|
||||
return response_base.success()
|
||||
|
||||
|
||||
@MCPRouter.put('/{pk}', summary='更新 MCP 服务器', dependencies=[ Depends(RequestPermission('sys:mcp:edit')), DependsRBAC])
|
||||
async def update_mcp(pk: Annotated[int, Path(description='MCP ID')], obj: UpdateMcpParam) -> ResponseModel:
|
||||
count = await mcp_service.update(pk=pk, obj=obj)
|
||||
if count > 0:
|
||||
return response_base.success()
|
||||
return response_base.fail()
|
||||
|
||||
|
||||
@MCPRouter.delete('/{pk}',summary='删除 MCP 服务器',dependencies=[Depends(RequestPermission('sys:mcp:del')), DependsRBAC])
|
||||
async def delete_mcp(pk: Annotated[int, Path(description='MCP ID')]) -> ResponseModel:
|
||||
count = await mcp_service.delete(pk=pk)
|
||||
if count > 0:
|
||||
return response_base.success()
|
||||
return response_base.fail()
|
||||
|
||||
|
||||
@MCPRouter.post('/chat',summary='MCP ChatGPT', dependencies=[Depends(RequestPermission('sys:mcp:chat')), DependsRBAC])
|
||||
async def mcp_chat(obj: McpChatParam) -> StreamingResponse:
|
||||
data = await mcp_service.chat(obj=obj)
|
||||
return StreamingResponse(data, media_type='text/event-stream')
|
||||
@@ -1,81 +1,44 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from sqlalchemy import Select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy_crud_plus import CRUDPlus
|
||||
|
||||
from .model import Mcp
|
||||
from .schema import CreateMcpParam, UpdateMcpParam
|
||||
from typing import Dict, List, Optional, Sequence
|
||||
|
||||
from app.core.base_crud import CRUDBase
|
||||
from app.api.v1.module_system.auth.schema import AuthSchema
|
||||
from .model import McpModel
|
||||
from .schema import McpCreateSchema, McpUpdateSchema
|
||||
|
||||
|
||||
class CRUDMcp(CRUDPlus[Mcp]):
|
||||
"""MCP 服务器数据库操作类"""
|
||||
class McpCRUD(CRUDBase[McpModel, McpCreateSchema, McpUpdateSchema]):
|
||||
"""MCP 服务器数据层"""
|
||||
|
||||
async def get(self, db: AsyncSession, pk: int) -> Mcp | None:
|
||||
"""
|
||||
获取 MCP 服务器
|
||||
def __init__(self, auth: AuthSchema) -> None:
|
||||
"""初始化CRUD"""
|
||||
self.auth = auth
|
||||
super().__init__(model=McpModel(), auth=auth)
|
||||
|
||||
:param db: 数据库会话
|
||||
:param pk: MCP ID
|
||||
:return:
|
||||
"""
|
||||
return await self.select_model(db, pk)
|
||||
|
||||
async def get_by_name(self, db: AsyncSession, name: str) -> Mcp | None:
|
||||
"""
|
||||
通过名称获取 MCP 服务器
|
||||
|
||||
:param db: 数据库会话
|
||||
:param name: MCP 名称
|
||||
:return:
|
||||
"""
|
||||
return await self.select_model_by_column(db, name=name)
|
||||
|
||||
async def get_list(self, name: str | None, type: int | None) -> Select[Mcp]:
|
||||
"""
|
||||
获取 MCP 服务器列表
|
||||
|
||||
:param name: MCP 名称
|
||||
:param type: MCP 类型
|
||||
:return:
|
||||
"""
|
||||
filters = {}
|
||||
if name is not None:
|
||||
filters.update(name__like=f'%{name}%')
|
||||
if type is not None:
|
||||
filters.update(type=type)
|
||||
return await self.select_order('created_time', 'desc', **filters)
|
||||
|
||||
async def create(self, db: AsyncSession, obj: CreateMcpParam) -> None:
|
||||
"""
|
||||
创建 MCP 服务器
|
||||
|
||||
:param db: 数据库会话
|
||||
:param obj: 创建 MCP 服务器参数
|
||||
:return:
|
||||
"""
|
||||
await self.create_model(db, obj)
|
||||
|
||||
async def update(self, db: AsyncSession, pk: int, obj: UpdateMcpParam) -> int:
|
||||
"""
|
||||
更新 MCP 服务器
|
||||
|
||||
:param db: 数据库会话
|
||||
:param pk: MCP ID
|
||||
:param obj: 更新 MCP 服务器参数
|
||||
:return:
|
||||
"""
|
||||
return await self.update_model(db, pk, obj)
|
||||
|
||||
async def delete(self, db: AsyncSession, pk: int) -> int:
|
||||
"""
|
||||
删除 MCP 服务器
|
||||
|
||||
:param db: 数据库会话
|
||||
:param pk: MCP ID
|
||||
:return:
|
||||
"""
|
||||
return await self.delete_model(db, pk)
|
||||
async def get_by_id_crud(self, id: int) -> Optional[McpModel]:
|
||||
"""详情"""
|
||||
return await self.get(id=id)
|
||||
|
||||
async def get_by_name_crud(self, name: str) -> Optional[McpModel]:
|
||||
"""通过名称获取MCP服务器"""
|
||||
return await self.get(name=name)
|
||||
|
||||
async def get_list_crud(self, search: Optional[Dict] = None, order_by: Optional[List[Dict[str, str]]] = None) -> Sequence[McpModel]:
|
||||
"""列表查询"""
|
||||
return await self.list(search=search or {}, order_by=order_by or [{'id': 'asc'}])
|
||||
|
||||
async def create_crud(self, data: McpCreateSchema) -> Optional[McpModel]:
|
||||
"""创建"""
|
||||
return await self.create(data=data)
|
||||
|
||||
async def update_crud(self, id: int, data: McpUpdateSchema) -> Optional[McpModel]:
|
||||
"""更新"""
|
||||
return await self.update(id=id, data=data)
|
||||
|
||||
async def delete_crud(self, ids: List[int]) -> None:
|
||||
"""批量删除"""
|
||||
return await self.delete(ids=ids)
|
||||
|
||||
|
||||
mcp_dao: CRUDMcp = CRUDMcp(Mcp)
|
||||
mcp_crud: McpCRUD = McpCRUD(auth=AuthSchema())
|
||||
@@ -1,19 +1,23 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from sqlalchemy import JSON, String
|
||||
from typing import Optional, Dict, Any
|
||||
from sqlalchemy import JSON, String, Integer
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.core.base_model import CreatorMixin
|
||||
|
||||
class Mcp(CreatorMixin):
|
||||
"""MCP 服务器表"""
|
||||
|
||||
class McpModel(CreatorMixin):
|
||||
"""
|
||||
MCP 服务器表
|
||||
"""
|
||||
|
||||
__tablename__ = 'ai_mcp'
|
||||
__table_args__ = ({'comment': 'MCP 服务器表'})
|
||||
|
||||
name: Mapped[str] = mapped_column(String(50), unique=True, comment='MCP 名称')
|
||||
type: Mapped[int] = mapped_column(default=0, comment='MCP 类型(0stdio 1sse)')
|
||||
url: Mapped[str | None] = mapped_column(String(255), default=None, comment='远程 SSE 地址')
|
||||
command: Mapped[str | None] = mapped_column(String(255), default=None, comment='MCP 命令')
|
||||
args: Mapped[str | None] = mapped_column(String(255), default=None, comment='MCP 命令参数')
|
||||
env: Mapped[str | None] = mapped_column(JSON(), default=None, comment='MCP 环境变量')
|
||||
type: Mapped[int] = mapped_column(Integer, default=0, comment='MCP 类型(0:stdio 1:sse)')
|
||||
url: Mapped[Optional[str]] = mapped_column(String(255), default=None, comment='远程 SSE 地址')
|
||||
command: Mapped[Optional[str]] = mapped_column(String(255), default=None, comment='MCP 命令')
|
||||
args: Mapped[Optional[str]] = mapped_column(String(255), default=None, comment='MCP 命令参数')
|
||||
env: Mapped[Optional[Dict[str, Any]]] = mapped_column(JSON(), default=None, comment='MCP 环境变量')
|
||||
@@ -0,0 +1,35 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from typing import Optional, List
|
||||
from fastapi import Query
|
||||
from pydantic import Field
|
||||
|
||||
from app.core.base_schema import BaseSchema
|
||||
from app.common.enums import McpLLMProvider
|
||||
|
||||
|
||||
class McpQueryParam:
|
||||
"""MCP 服务器查询参数"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
name: Optional[str] = Query(None, description="MCP 名称"),
|
||||
type: Optional[int] = Query(None, description="MCP 类型"),
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
# 模糊查询字段
|
||||
self.name = ("like", name) if name else None
|
||||
|
||||
# 精确查询字段
|
||||
self.type = type
|
||||
|
||||
|
||||
class McpChatParam(BaseSchema):
|
||||
"""MCP 聊天参数"""
|
||||
pk: List[int] = Field(..., description='MCP ID 列表')
|
||||
provider: McpLLMProvider = Field(McpLLMProvider.openai, description='LLM 供应商')
|
||||
model: str = Field(..., description='LLM 名称')
|
||||
key: str = Field(..., description='LLM API Key')
|
||||
base_url: Optional[str] = Field(None, description='自定义 LLM API 地址,必须兼容 openai 供应商')
|
||||
prompt: str = Field(..., description='用户提示词')
|
||||
@@ -1,8 +1,7 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
from pydantic import ConfigDict, Field, HttpUrl, BaseModel, Field
|
||||
from typing import Optional, Dict, Any
|
||||
from pydantic import ConfigDict, Field, HttpUrl, BaseModel
|
||||
|
||||
from app.core.base_schema import BaseSchema
|
||||
from app.common.enums import McpLLMProvider, McpType
|
||||
@@ -10,48 +9,25 @@ from app.common.enums import McpLLMProvider, McpType
|
||||
|
||||
class ChatQuerySchema(BaseModel):
|
||||
"""聊天查询模型"""
|
||||
message: str = Field(..., min_length=1, max_length=4000, description="聊天消息", example="你好,你能帮我什么?")
|
||||
|
||||
class Config:
|
||||
json_schema_extra = {
|
||||
"example": {
|
||||
"message": "你好,你能帮我什么?"
|
||||
}
|
||||
}
|
||||
message: str = Field(..., min_length=1, max_length=4000, description="聊天消息")
|
||||
|
||||
|
||||
class McpSchemaBase(BaseSchema):
|
||||
name: str = Field(description='MCP 名称')
|
||||
type: McpType = Field(McpType.stdio, description='MCP 类型')
|
||||
description: str | None = Field(None, description='MCP 描述')
|
||||
url: HttpUrl | None = Field(None, description='远程 SSE 地址')
|
||||
command: str | None = Field(None, description='MCP 命令')
|
||||
args: str | None = Field(None, description='MCP 命令参数,多个参数用英文逗号隔开')
|
||||
env: dict[str, Any] | None = Field(None, description='MCP 环境变量')
|
||||
|
||||
|
||||
class CreateMcpParam(McpSchemaBase):
|
||||
class McpCreateSchema(BaseModel):
|
||||
"""创建 MCP 服务器参数"""
|
||||
name: str = Field(..., max_length=50, description='MCP 名称')
|
||||
type: McpType = Field(McpType.stdio, description='MCP 类型')
|
||||
description: Optional[str] = Field(None, max_length=255, description='MCP 描述')
|
||||
url: Optional[HttpUrl] = Field(None, description='远程 SSE 地址')
|
||||
command: Optional[str] = Field(None, max_length=255, description='MCP 命令')
|
||||
args: Optional[str] = Field(None, max_length=255, description='MCP 命令参数,多个参数用英文逗号隔开')
|
||||
env: Optional[Dict[str, Any]] = Field(None, description='MCP 环境变量')
|
||||
|
||||
|
||||
class UpdateMcpParam(McpSchemaBase):
|
||||
class McpUpdateSchema(McpCreateSchema):
|
||||
"""更新 MCP 服务器参数"""
|
||||
...
|
||||
|
||||
|
||||
class GetMcpDetail(McpSchemaBase):
|
||||
class McpOutSchema(McpCreateSchema, BaseSchema):
|
||||
"""MCP 服务器详情"""
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: int = Field(description='MCP ID')
|
||||
created_time: datetime = Field(description='创建时间')
|
||||
updated_time: datetime | None = Field(None, description='更新时间')
|
||||
|
||||
|
||||
class McpChatParam(BaseSchema):
|
||||
pk: list[int] = Field(description='MCP ID 列表')
|
||||
provider: McpLLMProvider = Field(McpLLMProvider.openai, description='LLM 供应商')
|
||||
model: str = Field(description='LLM 名称')
|
||||
key: str = Field(description='LLM API Key')
|
||||
base_url: str | None = Field(None, description='自定义 LLM API 地址,必须兼容 openai 供应商')
|
||||
prompt: str = Field(description='用户提示词')
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
@@ -1,169 +1,77 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
from typing import AsyncGenerator
|
||||
|
||||
from pydantic_ai import Agent
|
||||
from pydantic_ai.mcp import MCPServerHTTP, MCPServerStdio
|
||||
from pydantic_ai.messages import TextPart
|
||||
from pydantic_ai.models.anthropic import AnthropicModel
|
||||
from pydantic_ai.models.gemini import GeminiModel
|
||||
from pydantic_ai.models.openai import OpenAIModel
|
||||
from pydantic_ai.providers.anthropic import AnthropicProvider
|
||||
from pydantic_ai.providers.deepseek import DeepSeekProvider
|
||||
from pydantic_ai.providers.google_gla import GoogleGLAProvider
|
||||
from pydantic_ai.providers.openai import OpenAIProvider
|
||||
from sqlalchemy import Select
|
||||
from typing import AsyncGenerator, List, Dict, Optional, Any
|
||||
|
||||
from backend.common.exception import errors
|
||||
from backend.common.log import log
|
||||
from backend.common.response.response_schema import response_base
|
||||
from backend.database.db import async_db_session
|
||||
from backend.plugin.mcp.crud.crud_mcp import mcp_dao
|
||||
from backend.plugin.mcp.enums import McpLLMProvider, McpType
|
||||
from backend.plugin.mcp.schema.mcp import CreateMcpParam, McpChatParam, UpdateMcpParam
|
||||
from app.core.exceptions import CustomException
|
||||
from app.core.logger import logger
|
||||
from app.api.v1.module_system.auth.schema import AuthSchema
|
||||
from app.utils.ai_util import AIClient
|
||||
from .schema import McpCreateSchema, McpUpdateSchema, McpOutSchema, ChatQuerySchema
|
||||
from .param import McpQueryParam
|
||||
from .crud import McpCRUD
|
||||
from .model import McpModel
|
||||
|
||||
|
||||
class MCPService:
|
||||
"""MCP服务层 - 适配FastAPI-MCP"""
|
||||
class McpService:
|
||||
"""MCP服务层"""
|
||||
|
||||
@staticmethod
|
||||
async def get(*, pk: int):
|
||||
"""
|
||||
获取 MCP 服务器
|
||||
|
||||
:param pk: MCP ID
|
||||
:return:
|
||||
"""
|
||||
async with async_db_session() as db:
|
||||
mcp = await mcp_dao.get(db, pk)
|
||||
if not mcp:
|
||||
raise errors.NotFoundError(msg='MCP 服务器不存在')
|
||||
return mcp
|
||||
|
||||
@staticmethod
|
||||
async def get_select(*, name: str | None, type: int | None) -> Select:
|
||||
"""
|
||||
获取 MCP 服务器查询对象
|
||||
|
||||
:param name: MCP 名称
|
||||
:param type: MCP 类型
|
||||
:return:
|
||||
"""
|
||||
return await mcp_dao.get_list(name=name, type=type)
|
||||
|
||||
@staticmethod
|
||||
async def create(*, obj: CreateMcpParam) -> None:
|
||||
"""
|
||||
创建 MCP 服务器
|
||||
|
||||
:param obj: 创建 MCP 服务器参数
|
||||
:return:
|
||||
"""
|
||||
async with async_db_session.begin() as db:
|
||||
mcp = await mcp_dao.get_by_name(db, name=obj.name)
|
||||
if mcp:
|
||||
raise errors.ForbiddenError(msg='MCP 服务器已存在')
|
||||
await mcp_dao.create(db, obj)
|
||||
|
||||
@staticmethod
|
||||
async def update(*, pk: int, obj: UpdateMcpParam) -> int:
|
||||
"""
|
||||
更新 MCP 服务器
|
||||
|
||||
:param pk: MCP ID
|
||||
:param obj: 更新 MCP 服务器参数
|
||||
:return:
|
||||
"""
|
||||
async with async_db_session.begin() as db:
|
||||
mcp = await mcp_dao.get(db, pk)
|
||||
if mcp.name != obj.name:
|
||||
if mcp_dao.get_by_name(db, name=obj.name):
|
||||
raise errors.ForbiddenError(msg='MCP 服务器已存在')
|
||||
count = await mcp_dao.update(db, pk, obj)
|
||||
return count
|
||||
|
||||
@staticmethod
|
||||
async def delete(*, pk: int) -> int:
|
||||
"""
|
||||
删除 MCP 服务器
|
||||
|
||||
:param pk: MCP ID
|
||||
:return:
|
||||
"""
|
||||
async with async_db_session.begin() as db:
|
||||
count = await mcp_dao.delete(db, pk)
|
||||
return count
|
||||
|
||||
@staticmethod
|
||||
async def chat(*, obj: McpChatParam) -> AsyncGenerator:
|
||||
async with async_db_session() as db:
|
||||
mcp_servers = []
|
||||
for pk in obj.pk:
|
||||
mcp = await mcp_dao.get(db, pk)
|
||||
if not mcp:
|
||||
raise errors.NotFoundError(msg='MCP 服务器不存在')
|
||||
if mcp.type == McpType.sse:
|
||||
mcp_servers.append(MCPServerHTTP(url=mcp.url))
|
||||
else:
|
||||
mcp_servers.append(
|
||||
MCPServerStdio(
|
||||
command=mcp.command,
|
||||
args=mcp.args.split(',') if mcp.args is not None else None,
|
||||
env=mcp.env if mcp.env is not None else None,
|
||||
)
|
||||
)
|
||||
|
||||
if obj.provider == McpLLMProvider.deepseek:
|
||||
model = OpenAIModel(
|
||||
obj.model,
|
||||
provider=DeepSeekProvider(api_key=obj.key),
|
||||
)
|
||||
elif obj.provider == McpLLMProvider.anthropic:
|
||||
model = AnthropicModel(
|
||||
obj.model,
|
||||
provider=AnthropicProvider(api_key=obj.key),
|
||||
)
|
||||
elif obj.provider == McpLLMProvider.gemini:
|
||||
model = GeminiModel(
|
||||
obj.model,
|
||||
provider=GoogleGLAProvider(api_key=obj.key),
|
||||
)
|
||||
else:
|
||||
model = OpenAIModel(
|
||||
obj.model,
|
||||
provider=OpenAIProvider(
|
||||
base_url=obj.base_url,
|
||||
api_key=obj.key,
|
||||
),
|
||||
)
|
||||
|
||||
agent = Agent(model, mcp_servers=mcp_servers)
|
||||
|
||||
async def stream_messages():
|
||||
try:
|
||||
async with agent.run_mcp_servers():
|
||||
async with agent.run_stream(obj.prompt) as result:
|
||||
async for text in result.stream():
|
||||
yield TextPart(text).content.encode('utf-8') + b'\n'
|
||||
except Exception as e:
|
||||
log.error(e)
|
||||
yield response_base.fail(data=str(e)).model_dump_json()
|
||||
|
||||
return stream_messages()
|
||||
@classmethod
|
||||
async def get_mcp_detail_service(cls, auth: AuthSchema, id: int) -> Dict[str, Any]:
|
||||
"""详情"""
|
||||
obj = await McpCRUD(auth).get_by_id_crud(id=id)
|
||||
if not obj:
|
||||
raise CustomException(msg='MCP 服务器不存在')
|
||||
return McpOutSchema.model_validate(obj).model_dump()
|
||||
|
||||
@classmethod
|
||||
async def chat_query(cls, message: str):
|
||||
async def get_mcp_list_service(cls, auth: AuthSchema, search: Optional[McpQueryParam] = None, order_by: Optional[List[Dict[str, str]]] = None) -> List[Dict[str, Any]]:
|
||||
"""列表查询"""
|
||||
if order_by:
|
||||
order_by = eval(str(order_by))
|
||||
obj_list = await McpCRUD(auth).get_list_crud(search=search.__dict__ if search else {}, order_by=order_by)
|
||||
return [McpOutSchema.model_validate(obj).model_dump() for obj in obj_list]
|
||||
|
||||
@classmethod
|
||||
async def create_mcp_service(cls, auth: AuthSchema, data: McpCreateSchema) -> Dict[str, Any]:
|
||||
"""创建"""
|
||||
obj = await McpCRUD(auth).get_by_name_crud(name=data.name)
|
||||
if obj:
|
||||
raise CustomException(msg='创建失败,MCP 服务器已存在')
|
||||
obj = await McpCRUD(auth).create_crud(data=data)
|
||||
return McpOutSchema.model_validate(obj).model_dump()
|
||||
|
||||
@classmethod
|
||||
async def update_mcp_service(cls, auth: AuthSchema, id: int, data: McpUpdateSchema) -> Dict[str, Any]:
|
||||
"""更新"""
|
||||
obj = await McpCRUD(auth).get_by_id_crud(id=id)
|
||||
if not obj:
|
||||
raise CustomException(msg='更新失败,该数据不存在')
|
||||
exist_obj = await McpCRUD(auth).get_by_name_crud(name=data.name)
|
||||
if exist_obj and exist_obj.id != id:
|
||||
raise CustomException(msg='更新失败,MCP 服务器名称重复')
|
||||
obj = await McpCRUD(auth).update_crud(id=id, data=data)
|
||||
return McpOutSchema.model_validate(obj).model_dump()
|
||||
|
||||
@classmethod
|
||||
async def delete_mcp_service(cls, auth: AuthSchema, ids: List[int]) -> None:
|
||||
"""删除"""
|
||||
if len(ids) < 1:
|
||||
raise CustomException(msg='删除失败,删除对象不能为空')
|
||||
for id in ids:
|
||||
obj = await McpCRUD(auth).get_by_id_crud(id=id)
|
||||
if not obj:
|
||||
raise CustomException(msg='删除失败,该数据不存在')
|
||||
await McpCRUD(auth).delete_crud(ids=ids)
|
||||
|
||||
@classmethod
|
||||
async def chat_query(cls, query: ChatQuerySchema):
|
||||
"""处理聊天查询"""
|
||||
# 创建MCP客户端实例
|
||||
mcp_client = AIClient()
|
||||
try:
|
||||
# 处理消息
|
||||
async for response in mcp_client.process(message):
|
||||
async for response in mcp_client.process(query.message):
|
||||
yield response
|
||||
finally:
|
||||
# 确保关闭客户端连接
|
||||
await mcp_client.close()
|
||||
|
||||
|
||||
|
||||
mcp_service: McpService = McpService()
|
||||
Reference in New Issue
Block a user