diff --git a/backend/app/api/v1/module_ai/mcp/controller.py b/backend/app/api/v1/module_ai/mcp/controller.py index bc8316fa..da78733e 100644 --- a/backend/app/api/v1/module_ai/mcp/controller.py +++ b/backend/app/api/v1/module_ai/mcp/controller.py @@ -1,5 +1,19 @@ # -*- 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.responses import JSONResponse, StreamingResponse @@ -59,4 +73,49 @@ async def websocket_chat_controller( except Exception as e: logger.error(f"WebSocket聊天出错: {str(e)}") finally: - await websocket.close() \ No newline at end of file + 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') \ No newline at end of file diff --git a/backend/app/api/v1/module_ai/mcp/crud.py b/backend/app/api/v1/module_ai/mcp/crud.py new file mode 100644 index 00000000..5eef8ece --- /dev/null +++ b/backend/app/api/v1/module_ai/mcp/crud.py @@ -0,0 +1,81 @@ +#!/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 + + +class CRUDMcp(CRUDPlus[Mcp]): + """MCP 服务器数据库操作类""" + + async def get(self, db: AsyncSession, pk: int) -> Mcp | None: + """ + 获取 MCP 服务器 + + :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) + + +mcp_dao: CRUDMcp = CRUDMcp(Mcp) \ No newline at end of file diff --git a/backend/app/api/v1/module_ai/mcp/model.py b/backend/app/api/v1/module_ai/mcp/model.py new file mode 100644 index 00000000..94c113ef --- /dev/null +++ b/backend/app/api/v1/module_ai/mcp/model.py @@ -0,0 +1,21 @@ +# -*- coding: utf-8 -*- + +from sqlalchemy import JSON, String +from sqlalchemy.dialects.mysql import LONGTEXT +from sqlalchemy.dialects.postgresql import TEXT +from sqlalchemy.orm import Mapped, mapped_column + +from app.core.base_model import CreatorMixin + +class Mcp(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 环境变量') \ No newline at end of file diff --git a/backend/app/api/v1/module_ai/mcp/param.py b/backend/app/api/v1/module_ai/mcp/param.py new file mode 100644 index 00000000..e69de29b diff --git a/backend/app/api/v1/module_ai/mcp/schema.py b/backend/app/api/v1/module_ai/mcp/schema.py index e4e6ed99..c38cf714 100644 --- a/backend/app/api/v1/module_ai/mcp/schema.py +++ b/backend/app/api/v1/module_ai/mcp/schema.py @@ -1,7 +1,11 @@ # -*- coding: utf-8 -*- -from pydantic import BaseModel, Field -from typing import Optional +from datetime import datetime +from typing import Any, Optional +from pydantic import ConfigDict, Field, HttpUrl, BaseModel, Field + +from app.core.base_schema import BaseSchema +from app.common.enums import McpLLMProvider, McpType class ChatQuerySchema(BaseModel): @@ -13,4 +17,41 @@ class ChatQuerySchema(BaseModel): "example": { "message": "你好,你能帮我什么?" } - } \ No newline at end of file + } + + +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): + """创建 MCP 服务器参数""" + + +class UpdateMcpParam(McpSchemaBase): + """更新 MCP 服务器参数""" + + +class GetMcpDetail(McpSchemaBase): + """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='用户提示词') \ No newline at end of file diff --git a/backend/app/api/v1/module_ai/mcp/service.py b/backend/app/api/v1/module_ai/mcp/service.py index 3a9393a1..358ba4f0 100644 --- a/backend/app/api/v1/module_ai/mcp/service.py +++ b/backend/app/api/v1/module_ai/mcp/service.py @@ -1,11 +1,156 @@ # -*- 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 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.utils.ai_util import AIClient class MCPService: """MCP服务层 - 适配FastAPI-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 chat_query(cls, message: str): """处理聊天查询""" @@ -17,4 +162,8 @@ class MCPService: yield response finally: # 确保关闭客户端连接 - await mcp_client.close() \ No newline at end of file + await mcp_client.close() + + + +mcp_service: McpService = McpService() \ No newline at end of file diff --git a/backend/app/common/enums.py b/backend/app/common/enums.py index 0478d3d2..4ab7c075 100644 --- a/backend/app/common/enums.py +++ b/backend/app/common/enums.py @@ -55,4 +55,21 @@ class RedisInitKeyConfig(Enum): @property def remark(self) -> str: """获取Redis键名说明""" - return self.value.get('remark', '') \ No newline at end of file + return self.value.get('remark', '') + + +class McpType(Enum): + """Mcp 服务器类型""" + + stdio = 0 + sse = 1 + + +class McpLLMProvider(Enum): + """MCP 大语言模型供应商""" + + openai = 'openai' + deepseek = 'deepseek' + anthropic = 'anthropic' + gemini = 'gemini' + qwen = 'qwen' \ No newline at end of file