feat: add MCP tool integration

This commit is contained in:
Xin Wang
2026-07-18 00:00:06 +08:00
parent bdf3d3dd9c
commit e39bb48ba8
21 changed files with 1641 additions and 62 deletions

View File

@@ -27,6 +27,7 @@ from routes import (
conversations,
health,
knowledge_bases,
mcp_servers,
model_registry,
node_types,
tools,
@@ -58,6 +59,7 @@ app.include_router(auth.router)
app.include_router(conversations.router)
app.include_router(assistants.router)
app.include_router(knowledge_bases.router)
app.include_router(mcp_servers.router)
app.include_router(model_registry.router)
app.include_router(node_types.router)
app.include_router(tools.router)

View File

@@ -207,15 +207,55 @@ class AssistantModelBinding(Base):
)
class McpServer(Base):
"""Workspace-level remote MCP connection used to discover executable tools."""
__tablename__ = "mcp_servers"
id: Mapped[str] = mapped_column(String(40), primary_key=True)
name: Mapped[str] = mapped_column(String(128))
description: Mapped[str] = mapped_column(String(2048), default="")
transport: Mapped[str] = mapped_column(String(32), default="streamable_http")
url: Mapped[str] = mapped_column(String(2048))
config: Mapped[dict] = mapped_column(JSONB, default=dict)
secrets: Mapped[dict] = mapped_column(JSONB, default=dict)
status: Mapped[str] = mapped_column(String(16), index=True, default="active")
last_synced_at: Mapped[datetime | None] = mapped_column(
DateTime(timezone=True), nullable=True
)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), server_default=func.now()
)
updated_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), server_default=func.now(), onupdate=func.now()
)
class Tool(Base):
"""Reusable LLM tool definition; supported types are executed at runtime."""
__tablename__ = "tools"
__table_args__ = (
UniqueConstraint(
"mcp_server_id",
"remote_tool_name",
name="uq_tools_mcp_server_remote_name",
),
)
id: Mapped[str] = mapped_column(String(40), primary_key=True)
name: Mapped[str] = mapped_column(String(128))
function_name: Mapped[str] = mapped_column(String(64), unique=True, index=True)
type: Mapped[str] = mapped_column(String(24), index=True)
mcp_server_id: Mapped[str | None] = mapped_column(
String(40),
ForeignKey("mcp_servers.id", ondelete="CASCADE"),
nullable=True,
index=True,
)
remote_tool_name: Mapped[str | None] = mapped_column(
String(255), nullable=True
)
description: Mapped[str] = mapped_column(String(2048), default="")
definition: Mapped[dict] = mapped_column(JSONB, default=dict)
secrets: Mapped[dict] = mapped_column(JSONB, default=dict)

View File

@@ -0,0 +1,107 @@
"""add MCP server connections and discovered tools
Revision ID: 20260717_0008
Revises: 20260712_0007
"""
from collections.abc import Sequence
from alembic import op
import sqlalchemy as sa
from sqlalchemy.dialects import postgresql
revision: str = "20260717_0008"
down_revision: str | Sequence[str] | None = "20260712_0007"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
op.create_table(
"mcp_servers",
sa.Column("id", sa.String(length=40), nullable=False),
sa.Column("name", sa.String(length=128), nullable=False),
sa.Column(
"description",
sa.String(length=2048),
server_default="",
nullable=False,
),
sa.Column(
"transport",
sa.String(length=32),
server_default="streamable_http",
nullable=False,
),
sa.Column("url", sa.String(length=2048), nullable=False),
sa.Column(
"config",
postgresql.JSONB(astext_type=sa.Text()),
server_default=sa.text("'{}'::jsonb"),
nullable=False,
),
sa.Column(
"secrets",
postgresql.JSONB(astext_type=sa.Text()),
server_default=sa.text("'{}'::jsonb"),
nullable=False,
),
sa.Column(
"status",
sa.String(length=16),
server_default="active",
nullable=False,
),
sa.Column("last_synced_at", sa.DateTime(timezone=True), nullable=True),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
server_default=sa.text("now()"),
nullable=False,
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
server_default=sa.text("now()"),
nullable=False,
),
sa.PrimaryKeyConstraint("id"),
)
op.create_index("ix_mcp_servers_status", "mcp_servers", ["status"])
op.add_column(
"tools",
sa.Column("mcp_server_id", sa.String(length=40), nullable=True),
)
op.add_column(
"tools",
sa.Column("remote_tool_name", sa.String(length=255), nullable=True),
)
op.create_index("ix_tools_mcp_server_id", "tools", ["mcp_server_id"])
op.create_foreign_key(
"fk_tools_mcp_server_id",
"tools",
"mcp_servers",
["mcp_server_id"],
["id"],
ondelete="CASCADE",
)
op.create_unique_constraint(
"uq_tools_mcp_server_remote_name",
"tools",
["mcp_server_id", "remote_tool_name"],
)
def downgrade() -> None:
op.drop_constraint(
"uq_tools_mcp_server_remote_name",
"tools",
type_="unique",
)
op.drop_constraint("fk_tools_mcp_server_id", "tools", type_="foreignkey")
op.drop_index("ix_tools_mcp_server_id", table_name="tools")
op.drop_column("tools", "remote_tool_name")
op.drop_column("tools", "mcp_server_id")
op.drop_index("ix_mcp_servers_status", table_name="mcp_servers")
op.drop_table("mcp_servers")

View File

@@ -15,6 +15,18 @@ from pydantic import BaseModel, Field
RuntimeMode = Literal["pipeline", "realtime"]
class RuntimeMcpServer(BaseModel):
"""Resolved MCP connection containing secrets for one runtime session."""
id: str
name: str = ""
transport: str = "streamable_http"
url: str
timeout_seconds: int = 30
headers: dict[str, str] = Field(default_factory=dict)
secret_headers: dict[str, str] = Field(default_factory=dict)
class RuntimeTool(BaseModel):
"""Tool data resolved from an assistant binding for one runtime session."""
@@ -25,6 +37,7 @@ class RuntimeTool(BaseModel):
description: str = ""
definition: dict = Field(default_factory=dict)
secrets: dict = Field(default_factory=dict)
mcp_server: RuntimeMcpServer | None = None
class RuntimeModelResource(BaseModel):

View File

@@ -2,7 +2,7 @@
# webrtc -> SmallWebRTCTransport / SmallWebRTCConnection + aiortc
# silero -> 本地 VAD(判断用户说话起止),语音必备
# openai -> OpenAI 兼容的 LLM/STT/TTS 客户端(DeepSeek、SenseVoice、CosyVoice 都走它)
pipecat-ai[webrtc,websocket,silero,openai]==1.5.0
pipecat-ai[webrtc,websocket,silero,openai,mcp]==1.5.0
Pillow>=11.1.0,<13
# FastGPT 类型助手:本地 SDK(包 /api/v1/chat/completions 流式 + chatId 会话)

View File

@@ -0,0 +1,304 @@
"""Workspace MCP Server CRUD and explicit remote tool synchronization."""
from __future__ import annotations
import re
import uuid
from datetime import datetime, timezone
from db.models import AssistantToolBinding, McpServer, Tool
from db.session import get_session
from fastapi import APIRouter, Depends, HTTPException
from models import RuntimeMcpServer
from schemas import McpServerOut, McpServerUpsert, McpSyncResult, ToolOut
from services.auth import require_admin
from services.masking import mask_secrets, merge_secrets
from services.tool_resources import tool_to_out
from services.tools import McpClientError, McpToolClient
from sqlalchemy import func, select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
router = APIRouter(
prefix="/api/mcp-servers",
tags=["mcp-servers"],
dependencies=[Depends(require_admin)],
)
def _runtime_server(server: McpServer) -> RuntimeMcpServer:
config = server.config or {}
secrets = server.secrets or {}
return RuntimeMcpServer(
id=server.id,
name=server.name,
transport=server.transport,
url=server.url,
timeout_seconds=int(config.get("timeout_seconds") or 30),
headers={
str(key): str(value)
for key, value in (config.get("headers") or {}).items()
},
secret_headers={
str(key): str(value)
for key, value in (secrets.get("headers") or {}).items()
},
)
async def _to_out(session: AsyncSession, server: McpServer) -> McpServerOut:
tool_count = int(
(
await session.execute(
select(func.count(Tool.id)).where(Tool.mcp_server_id == server.id)
)
).scalar_one()
)
config = server.config or {}
secrets = server.secrets or {}
return McpServerOut(
id=server.id,
name=server.name,
description=server.description,
transport=server.transport, # type: ignore[arg-type]
url=server.url,
timeout_seconds=int(config.get("timeout_seconds") or 30),
headers=dict(config.get("headers") or {}),
secret_headers=mask_secrets(secrets.get("headers") or {}),
status=server.status, # type: ignore[arg-type]
tool_count=tool_count,
last_synced_at=(
server.last_synced_at.isoformat() if server.last_synced_at else None
),
updated_at=server.updated_at.isoformat() if server.updated_at else None,
)
def _apply_body(
server: McpServer,
body: McpServerUpsert,
) -> None:
server.name = body.name.strip()
server.description = body.description.strip()
server.transport = body.transport
server.url = body.url.strip()
server.config = {
"timeout_seconds": body.timeout_seconds,
"headers": body.headers,
}
server.secrets = merge_secrets(
{"headers": body.secret_headers},
server.secrets or {},
)
server.status = body.status
def _function_slug(value: str) -> str:
slug = re.sub(r"[^a-z0-9]+", "_", value.lower()).strip("_")
if not slug or not slug[0].isalpha():
slug = f"tool_{slug}"
return slug[:64]
async def _available_function_name(
session: AsyncSession,
server: McpServer,
remote_name: str,
) -> str:
base = _function_slug(f"mcp_{server.name}_{remote_name}")
candidate = base
suffix = 2
while (
await session.execute(
select(Tool.id).where(Tool.function_name == candidate).limit(1)
)
).scalar_one_or_none():
tail = f"_{suffix}"
candidate = f"{base[: 64 - len(tail)]}{tail}"
suffix += 1
return candidate
@router.get("", response_model=list[McpServerOut])
async def list_mcp_servers(session: AsyncSession = Depends(get_session)):
rows = (
await session.execute(
select(McpServer).order_by(McpServer.updated_at.desc())
)
).scalars().all()
return [await _to_out(session, server) for server in rows]
@router.post("", response_model=McpServerOut)
async def create_mcp_server(
body: McpServerUpsert,
session: AsyncSession = Depends(get_session),
):
server = McpServer(
id=f"mcp_{uuid.uuid4().hex[:12]}",
name=body.name.strip(),
url=body.url.strip(),
)
_apply_body(server, body)
session.add(server)
await session.commit()
await session.refresh(server)
return await _to_out(session, server)
@router.get("/{server_id}", response_model=McpServerOut)
async def get_mcp_server(
server_id: str,
session: AsyncSession = Depends(get_session),
):
server = await session.get(McpServer, server_id)
if not server:
raise HTTPException(404, "MCP Server 不存在")
return await _to_out(session, server)
@router.put("/{server_id}", response_model=McpServerOut)
async def update_mcp_server(
server_id: str,
body: McpServerUpsert,
session: AsyncSession = Depends(get_session),
):
server = await session.get(McpServer, server_id)
if not server:
raise HTTPException(404, "MCP Server 不存在")
_apply_body(server, body)
await session.commit()
await session.refresh(server)
return await _to_out(session, server)
@router.post("/{server_id}/sync", response_model=McpSyncResult)
async def sync_mcp_tools(
server_id: str,
session: AsyncSession = Depends(get_session),
):
server = await session.get(McpServer, server_id)
if not server:
raise HTTPException(404, "MCP Server 不存在")
if server.status != "active":
raise HTTPException(400, "请先启用 MCP Server")
try:
discovered = await McpToolClient().list_tools(_runtime_server(server))
except McpClientError as exc:
raise HTTPException(400, str(exc)) from exc
existing_rows = (
await session.execute(
select(Tool).where(Tool.mcp_server_id == server.id)
)
).scalars().all()
existing = {
str(tool.remote_tool_name): tool
for tool in existing_rows
if tool.remote_tool_name
}
created = 0
updated = 0
synchronized: list[Tool] = []
for remote in discovered:
remote_name = str(remote["name"])
tool = existing.get(remote_name)
if tool is None:
tool = Tool(
id=f"tool_{uuid.uuid4().hex[:12]}",
name=f"{server.name} · {remote_name}",
function_name=await _available_function_name(
session,
server,
remote_name,
),
type="mcp",
mcp_server_id=server.id,
remote_tool_name=remote_name,
description=str(remote.get("description") or ""),
definition={},
secrets={},
status="active",
)
session.add(tool)
created += 1
else:
updated += 1
previous_config = (tool.definition or {}).get("config") or {}
tool.description = str(remote.get("description") or tool.description)
tool.definition = {
"schema_version": 1,
"type": "mcp",
"config": {
"remote_tool_name": remote_name,
"input_schema": remote.get("input_schema") or {},
"schema_hash": remote.get("schema_hash") or "",
"dynamic_variable_assignments": previous_config.get(
"dynamic_variable_assignments"
)
or {},
},
}
synchronized.append(tool)
server.last_synced_at = datetime.now(timezone.utc)
try:
await session.commit()
except IntegrityError as exc:
await session.rollback()
raise HTTPException(409, "同步工具时发生名称冲突,请重试") from exc
for tool in synchronized:
await session.refresh(tool)
await session.refresh(server)
server_out = await _to_out(session, server)
tool_outputs: list[ToolOut] = [tool_to_out(tool) for tool in synchronized]
return McpSyncResult(
server=server_out,
created=created,
updated=updated,
tools=tool_outputs,
)
@router.get("/{server_id}/tools", response_model=list[ToolOut])
async def list_mcp_tools(
server_id: str,
session: AsyncSession = Depends(get_session),
):
if not await session.get(McpServer, server_id):
raise HTTPException(404, "MCP Server 不存在")
rows = (
await session.execute(
select(Tool)
.where(Tool.mcp_server_id == server_id)
.order_by(Tool.name)
)
).scalars().all()
return [tool_to_out(tool) for tool in rows]
@router.delete("/{server_id}")
async def delete_mcp_server(
server_id: str,
session: AsyncSession = Depends(get_session),
):
server = await session.get(McpServer, server_id)
if not server:
raise HTTPException(404, "MCP Server 不存在")
in_use = (
await session.execute(
select(AssistantToolBinding.tool_id)
.join(Tool, Tool.id == AssistantToolBinding.tool_id)
.where(Tool.mcp_server_id == server_id)
.limit(1)
)
).scalar_one_or_none()
if in_use:
raise HTTPException(409, "MCP 工具正被助手引用,请先解绑")
await session.delete(server)
await session.commit()
return {"ok": True}

View File

@@ -7,7 +7,8 @@ from db.session import get_session
from fastapi import APIRouter, Depends, HTTPException
from schemas import ToolOut, ToolUpsert
from services.auth import require_admin
from services.masking import mask_secrets, merge_secrets
from services.masking import merge_secrets
from services.tool_resources import tool_to_out
from sqlalchemy import select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
@@ -20,21 +21,10 @@ router = APIRouter(
)
def _to_out(tool: Tool) -> ToolOut:
return ToolOut(
id=tool.id,
name=tool.name,
function_name=tool.function_name,
type=tool.type,
description=tool.description,
definition=tool.definition,
secrets=mask_secrets(tool.secrets or {}),
status=tool.status,
updated_at=tool.updated_at.isoformat() if tool.updated_at else None,
)
def _payload(body: ToolUpsert, stored_secrets: dict | None = None) -> dict:
def _payload(
body: ToolUpsert,
stored_secrets: dict | None = None,
) -> dict:
definition = body.definition.model_dump()
secrets = (
merge_secrets(body.secrets, stored_secrets or {})
@@ -45,6 +35,12 @@ def _payload(body: ToolUpsert, stored_secrets: dict | None = None) -> dict:
"name": body.name.strip(),
"function_name": body.function_name,
"type": definition["type"],
"mcp_server_id": body.mcp_server_id if definition["type"] == "mcp" else None,
"remote_tool_name": (
definition["config"]["remote_tool_name"]
if definition["type"] == "mcp"
else None
),
"description": body.description.strip(),
"definition": definition,
"secrets": secrets,
@@ -66,7 +62,7 @@ async def _commit(session: AsyncSession, tool: Tool) -> ToolOut:
await session.rollback()
raise HTTPException(409, "工具函数名已存在") from exc
await session.refresh(tool)
return _to_out(tool)
return tool_to_out(tool)
@router.get("", response_model=list[ToolOut])
@@ -74,11 +70,13 @@ async def list_tools(session: AsyncSession = Depends(get_session)):
rows = (
await session.execute(select(Tool).order_by(Tool.updated_at.desc()))
).scalars().all()
return [_to_out(tool) for tool in rows]
return [tool_to_out(tool) for tool in rows]
@router.post("", response_model=ToolOut)
async def create_tool(body: ToolUpsert, session: AsyncSession = Depends(get_session)):
if body.definition.type == "mcp":
raise HTTPException(400, "MCP 工具请通过 MCP Server 同步创建")
tool = Tool(id=f"tool_{uuid.uuid4().hex[:12]}", **_payload(body))
session.add(tool)
return await _commit(session, tool)
@@ -89,7 +87,7 @@ async def get_tool(tool_id: str, session: AsyncSession = Depends(get_session)):
tool = await session.get(Tool, tool_id)
if not tool:
raise HTTPException(404, "工具不存在")
return _to_out(tool)
return tool_to_out(tool)
@router.put("/{tool_id}", response_model=ToolOut)
@@ -101,6 +99,14 @@ async def update_tool(
tool = await session.get(Tool, tool_id)
if not tool:
raise HTTPException(404, "工具不存在")
if tool.type == "mcp":
if body.definition.type != "mcp":
raise HTTPException(400, "同步生成的 MCP 工具不能修改类型")
if (
body.mcp_server_id != tool.mcp_server_id
or body.definition.config.remote_tool_name != tool.remote_tool_name
):
raise HTTPException(400, "MCP Server 和远端工具名称不能修改")
for key, value in _payload(body, tool.secrets or {}).items():
setattr(tool, key, value)
return await _commit(session, tool)
@@ -111,6 +117,8 @@ async def duplicate_tool(tool_id: str, session: AsyncSession = Depends(get_sessi
source = await session.get(Tool, tool_id)
if not source:
raise HTTPException(404, "工具不存在")
if source.type == "mcp":
raise HTTPException(400, "MCP 工具由 Server 同步维护,不能单独复制")
function_name = _duplicate_function_name(source.function_name)
existing = (
await session.execute(select(Tool).where(Tool.function_name == function_name))
@@ -124,6 +132,8 @@ async def duplicate_tool(tool_id: str, session: AsyncSession = Depends(get_sessi
name=f"{source.name} 副本",
function_name=function_name,
type=source.type,
mcp_server_id=source.mcp_server_id,
remote_tool_name=source.remote_tool_name,
description=source.description,
definition=dict(source.definition or {}),
secrets=dict(source.secrets or {}),

View File

@@ -19,8 +19,10 @@ ModelType = Literal["LLM", "ASR", "TTS", "Realtime", "Embedding", "Agent"]
AssistantType = Literal["prompt", "workflow", "dify", "fastgpt", "opencode"]
TurnEndStrategy = Literal["silence", "smart_turn"]
KnowledgeRetrievalMode = Literal["automatic", "on_demand"]
ToolType = Literal["end_call", "http"]
ToolType = Literal["end_call", "http", "mcp"]
ToolStatus = Literal["active", "archived", "draft"]
McpServerStatus = Literal["active", "archived", "draft"]
McpTransport = Literal["streamable_http"]
ToolParameterType = Literal["string", "number", "integer", "boolean", "object", "array"]
ToolParameterLocation = Literal["path", "query", "body", "header"]
DynamicVariableType = Literal["string", "number", "boolean"]
@@ -219,8 +221,22 @@ class HttpToolDefinition(CamelModel):
config: HttpToolConfig
class McpToolConfig(CamelModel):
remote_tool_name: str = Field(min_length=1, max_length=255)
input_schema: dict[str, Any] = Field(default_factory=dict)
schema_hash: str = ""
dynamic_variable_assignments: dict[str, str] = Field(default_factory=dict)
class McpToolDefinition(CamelModel):
schema_version: int = 1
type: Literal["mcp"] = "mcp"
config: McpToolConfig
ToolDefinition = Annotated[
Union[EndCallToolDefinition, HttpToolDefinition], Field(discriminator="type")
Union[EndCallToolDefinition, HttpToolDefinition, McpToolDefinition],
Field(discriminator="type"),
]
@@ -231,14 +247,49 @@ class ToolUpsert(CamelModel):
definition: ToolDefinition
secrets: dict[str, Any] = Field(default_factory=dict)
status: ToolStatus = "active"
mcp_server_id: str | None = None
class ToolOut(ToolUpsert):
id: str
type: ToolType
remote_tool_name: str | None = None
updated_at: str | None = None
# ---------- MCP Server ----------
class McpServerUpsert(CamelModel):
name: str = Field(min_length=1, max_length=128)
description: str = Field(default="", max_length=2048)
transport: McpTransport = "streamable_http"
url: str = Field(min_length=1, max_length=2048)
timeout_seconds: int = Field(default=30, ge=1, le=120)
headers: dict[str, str] = Field(default_factory=dict)
secret_headers: dict[str, str] = Field(default_factory=dict)
status: McpServerStatus = "active"
@field_validator("url")
@classmethod
def validate_url(cls, value: str) -> str:
if not value.startswith(("http://", "https://")):
raise ValueError("MCP Server URL 必须使用 http:// 或 https://")
return value
class McpServerOut(McpServerUpsert):
id: str
tool_count: int = 0
last_synced_at: str | None = None
updated_at: str | None = None
class McpSyncResult(CamelModel):
server: McpServerOut
created: int
updated: int
tools: list[ToolOut]
# ---------- 知识库 ----------
class KnowledgeBaseUpsert(CamelModel):
name: str

View File

@@ -55,8 +55,8 @@ class PromptBrain(BaseBrain):
for tool in cfg.tools:
if tool.type == "end_call":
schema, handler = self._make_end_call_tool(tool, runtime)
elif tool.type == "http":
schema, handler = self._make_http_tool(tool, runtime)
elif tool.type in {"http", "mcp"}:
schema, handler = self._make_remote_tool(tool, runtime)
else:
continue
schemas.append(schema)
@@ -96,7 +96,7 @@ class PromptBrain(BaseBrain):
if self._dynamic_enabled and self._runtime is not None:
self._runtime.set_system_prompt(self._store.render(self._cfg.prompt))
def _make_http_tool(self, tool, runtime: BrainRuntime):
def _make_remote_tool(self, tool, runtime: BrainRuntime):
properties, required = self._tools.schema_parts(tool)
self._tools.register_secrets(tool)
@@ -108,7 +108,7 @@ class PromptBrain(BaseBrain):
await params.result_callback(result)
except (ToolExecutionError, ValueError) as exc:
await params.result_callback(
{"status": "error", "message": f"HTTP 工具调用失败: {exc}"}
{"status": "error", "message": f"工具调用失败: {exc}"}
)
schema = FunctionSchema(

View File

@@ -328,7 +328,7 @@ class WorkflowBrain(BaseBrain):
functions: list[FlowsFunctionSchema] = []
for tool_id in stage.tool_ids:
tool = self._tool_by_id.get(str(tool_id))
if tool and tool.type == "http":
if tool and tool.type in {"http", "mcp"}:
functions.append(self._flow_tool(tool, node_id))
knowledge_function = self._knowledge_function(node_id)
if knowledge_function:

View File

@@ -9,12 +9,14 @@ from db.models import (
AssistantModelBinding,
AssistantToolBinding,
KnowledgeBase,
McpServer,
ModelResource,
Tool,
)
from models import (
AssistantConfig,
RuntimeKnowledgeBase,
RuntimeMcpServer,
RuntimeModelResource,
RuntimeTool,
)
@@ -102,18 +104,55 @@ async def _tools_for(session: AsyncSession, assistant: Assistant) -> list[Runtim
.order_by(AssistantToolBinding.created_at, Tool.id)
)
).scalars().all()
return [
RuntimeTool(
id=tool.id,
name=tool.name,
function_name=tool.function_name,
type=tool.type,
description=tool.description,
definition=tool.definition or {},
secrets=tool.secrets or {},
)
server_ids = {
str(tool.mcp_server_id)
for tool in tools
]
if tool.type == "mcp" and tool.mcp_server_id
}
server_rows = (
(
await session.execute(
select(McpServer).where(
McpServer.id.in_(server_ids),
McpServer.status == "active",
)
)
).scalars().all()
if server_ids
else []
)
servers = {server.id: server for server in server_rows}
resolved: list[RuntimeTool] = []
for tool in tools:
mcp_server = None
if tool.type == "mcp":
server = servers.get(str(tool.mcp_server_id or ""))
if server is None:
continue
config = server.config or {}
secrets = server.secrets or {}
mcp_server = RuntimeMcpServer(
id=server.id,
name=server.name,
transport=server.transport,
url=server.url,
timeout_seconds=int(config.get("timeout_seconds") or 30),
headers=dict(config.get("headers") or {}),
secret_headers=dict(secrets.get("headers") or {}),
)
resolved.append(
RuntimeTool(
id=tool.id,
name=tool.name,
function_name=tool.function_name,
type=tool.type,
description=tool.description,
definition=tool.definition or {},
secrets=tool.secrets or {},
mcp_server=mcp_server,
)
)
return resolved
async def resolve_runtime_config(

View File

@@ -14,6 +14,7 @@ from services.runtime_variables import (
DynamicVariableStore,
value_at_path,
)
from services.tools import McpClientError, McpToolClient
class ToolExecutionError(RuntimeError):
@@ -21,8 +22,16 @@ class ToolExecutionError(RuntimeError):
class ToolExecutor:
def __init__(self, store: DynamicVariableStore):
"""Execute supported tools and apply their results to session variables."""
def __init__(
self,
store: DynamicVariableStore,
*,
mcp_client: McpToolClient | None = None,
):
self.store = store
self._mcp_client = mcp_client or McpToolClient()
def register_secrets(self, tool: RuntimeTool) -> None:
dynamic = (tool.secrets or {}).get("dynamic_variables") or {}
@@ -34,6 +43,12 @@ class ToolExecutor:
@staticmethod
def schema_parts(tool: RuntimeTool) -> tuple[dict[str, Any], list[str]]:
config = (tool.definition or {}).get("config") or {}
if tool.type == "mcp":
schema = config.get("input_schema") or {}
return (
dict(schema.get("properties") or {}),
[str(name) for name in schema.get("required") or []],
)
parameters = list(config.get("parameters") or [])
properties = {
str(parameter.get("name")): {
@@ -58,11 +73,16 @@ class ToolExecutor:
result_assignments: dict[str, str] | None = None,
) -> dict[str, Any]:
self.register_secrets(tool)
if tool.type != "http":
raise ToolExecutionError(f"Action 暂不支持工具类型: {tool.type}")
return await self._execute_http(
normalized_arguments = dict(arguments or {})
if tool.type == "http":
result = await self._execute_http(tool, normalized_arguments)
elif tool.type == "mcp":
result = await self._execute_mcp(tool, normalized_arguments)
else:
raise ToolExecutionError(f"不支持工具类型: {tool.type}")
return self._apply_result_assignments(
tool,
dict(arguments or {}),
result,
result_assignments=result_assignments,
)
@@ -70,8 +90,6 @@ class ToolExecutor:
self,
tool: RuntimeTool,
arguments: dict[str, Any],
*,
result_assignments: dict[str, str] | None,
) -> dict[str, Any]:
config = (tool.definition or {}).get("config") or {}
parameters = list(config.get("parameters") or [])
@@ -130,11 +148,59 @@ class ToolExecutor:
except ValueError:
payload = {"text": response.text[:8000]}
return {
"status": "ok",
"status_code": response.status_code,
"data": payload,
}
async def _execute_mcp(
self,
tool: RuntimeTool,
arguments: dict[str, Any],
) -> dict[str, Any]:
server = tool.mcp_server
if server is None:
raise ToolExecutionError("MCP 工具缺少可用的 Server 连接")
config = (tool.definition or {}).get("config") or {}
remote_tool_name = str(config.get("remote_tool_name") or "")
if not remote_tool_name:
raise ToolExecutionError("MCP 工具缺少远端工具名称")
resolved_server = server.model_copy(
update={
"url": self.store.render(server.url),
"headers": self.store.render_data(server.headers),
"secret_headers": self.store.render_data(
server.secret_headers,
allow_secrets=True,
),
}
)
try:
payload = await self._mcp_client.call_tool(
resolved_server,
remote_tool_name,
arguments,
)
except McpClientError as exc:
raise ToolExecutionError(str(exc)) from exc
return {"status": "ok", "data": payload}
def _apply_result_assignments(
self,
tool: RuntimeTool,
result: dict[str, Any],
*,
result_assignments: dict[str, str] | None,
) -> dict[str, Any]:
config = (tool.definition or {}).get("config") or {}
assignments = (
result_assignments
if result_assignments is not None
else config.get("dynamic_variable_assignments") or {}
)
payload = result.get("data")
updated: list[str] = []
for variable_name, path in assignments.items():
try:
@@ -146,9 +212,4 @@ class ToolExecutor:
continue
self.store.assign(str(variable_name), value)
updated.append(str(variable_name))
return {
"status": "ok",
"status_code": response.status_code,
"data": payload,
"updated_variables": updated,
}
return {**result, "updated_variables": updated}

View File

@@ -0,0 +1,21 @@
"""Shared serialization helpers for reusable tool resources."""
from db.models import Tool
from schemas import ToolOut
from services.masking import mask_secrets
def tool_to_out(tool: Tool) -> ToolOut:
return ToolOut(
id=tool.id,
name=tool.name,
function_name=tool.function_name,
type=tool.type, # type: ignore[arg-type]
description=tool.description,
definition=tool.definition,
secrets=mask_secrets(tool.secrets or {}),
status=tool.status, # type: ignore[arg-type]
mcp_server_id=tool.mcp_server_id,
remote_tool_name=tool.remote_tool_name,
updated_at=tool.updated_at.isoformat() if tool.updated_at else None,
)

View File

@@ -0,0 +1,5 @@
"""Small protocol adapters used by the shared ToolExecutor."""
from services.tools.mcp_client import McpClientError, McpToolClient
__all__ = ["McpClientError", "McpToolClient"]

View File

@@ -0,0 +1,152 @@
"""One-shot Streamable HTTP MCP discovery and execution.
Connections intentionally live for one operation in the MVP. This keeps the
async context ownership simple and correct; a session pool can be introduced
later if measurements show MCP handshakes are a meaningful source of latency.
"""
from __future__ import annotations
import hashlib
import json
from contextlib import asynccontextmanager
from datetime import timedelta
from typing import Any, AsyncIterator
import httpx
from mcp import ClientSession
from mcp.client.streamable_http import streamable_http_client
from models import RuntimeMcpServer
MAX_MCP_RESULT_BYTES = 1_000_000
class McpClientError(RuntimeError):
"""Safe error surfaced to tool callers and the administration API."""
def schema_hash(input_schema: dict[str, Any]) -> str:
serialized = json.dumps(
input_schema,
ensure_ascii=False,
sort_keys=True,
separators=(",", ":"),
)
return hashlib.sha256(serialized.encode("utf-8")).hexdigest()
class McpToolClient:
"""Minimal public MCP SDK wrapper shared by discovery and runtime calls."""
@asynccontextmanager
async def _session(
self,
server: RuntimeMcpServer,
) -> AsyncIterator[ClientSession]:
if server.transport != "streamable_http":
raise McpClientError(f"暂不支持 MCP transport: {server.transport}")
headers = {**server.headers, **server.secret_headers}
timeout = max(1, min(int(server.timeout_seconds), 120))
try:
async with httpx.AsyncClient(
headers=headers,
timeout=httpx.Timeout(timeout),
follow_redirects=False,
) as http_client:
async with streamable_http_client(
server.url,
http_client=http_client,
) as streams:
read_stream, write_stream, _ = streams
async with ClientSession(
read_stream,
write_stream,
read_timeout_seconds=timedelta(seconds=timeout),
) as session:
await session.initialize()
yield session
except McpClientError:
raise
except httpx.TimeoutException as exc:
raise McpClientError("MCP Server 连接或调用超时") from exc
except Exception as exc: # MCP SDK exposes several transport exceptions.
raise McpClientError(f"MCP Server 调用失败: {exc}") from exc
async def list_tools(self, server: RuntimeMcpServer) -> list[dict[str, Any]]:
"""Discover all remote tools, following MCP pagination cursors."""
discovered: list[dict[str, Any]] = []
async with self._session(server) as session:
cursor: str | None = None
while True:
result = await session.list_tools(cursor=cursor)
for tool in result.tools:
input_schema = dict(tool.inputSchema or {})
discovered.append(
{
"name": tool.name,
"description": tool.description or "",
"input_schema": input_schema,
"schema_hash": schema_hash(input_schema),
}
)
cursor = result.nextCursor
if not cursor:
break
return discovered
async def call_tool(
self,
server: RuntimeMcpServer,
remote_tool_name: str,
arguments: dict[str, Any],
) -> dict[str, Any]:
"""Execute one remote tool and normalize the result for local callers."""
async with self._session(server) as session:
result = await session.call_tool(
remote_tool_name,
arguments=arguments,
read_timeout_seconds=timedelta(seconds=server.timeout_seconds),
)
data = self._normalize_result(result)
encoded = json.dumps(data, ensure_ascii=False, default=str).encode("utf-8")
if len(encoded) > MAX_MCP_RESULT_BYTES:
raise McpClientError("MCP 工具响应超过 1 MB 限制")
if bool(result.isError):
message = data.get("text") or "MCP 工具返回执行错误"
raise McpClientError(str(message)[:2048])
return data
@staticmethod
def _normalize_result(result: Any) -> dict[str, Any]:
"""Keep useful text/JSON while deliberately omitting binary payloads."""
content: list[dict[str, Any]] = []
text_parts: list[str] = []
for item in result.content or []:
item_type = str(getattr(item, "type", "unknown"))
if item_type == "text":
text = str(getattr(item, "text", ""))
content.append({"type": "text", "text": text})
if text:
text_parts.append(text)
continue
if item_type == "resource":
resource = getattr(item, "resource", None)
text = getattr(resource, "text", None)
if isinstance(text, str):
content.append({"type": "resource", "text": text})
text_parts.append(text)
else:
content.append({"type": "resource", "omitted": True})
continue
# Audio and image base64 data are not sent into the LLM context in MVP.
content.append({"type": item_type, "omitted": True})
return {
"text": "\n".join(text_parts),
"structuredContent": result.structuredContent,
"content": content,
}

View File

@@ -0,0 +1,105 @@
import unittest
from models import RuntimeMcpServer, RuntimeTool
from services.runtime_variables import DynamicVariableStore
from services.tool_executor import ToolExecutionError, ToolExecutor
from services.tools import McpClientError
def mcp_tool() -> RuntimeTool:
return RuntimeTool(
id="tool_mcp_order",
name="查询订单",
function_name="query_order",
type="mcp",
description="查询订单状态",
definition={
"schema_version": 1,
"type": "mcp",
"config": {
"remote_tool_name": "get_order",
"input_schema": {
"type": "object",
"properties": {
"order_id": {
"type": "string",
"description": "订单编号",
}
},
"required": ["order_id"],
},
"dynamic_variable_assignments": {
"order_status": "structuredContent.status"
},
},
},
mcp_server=RuntimeMcpServer(
id="mcp_orders",
name="订单服务",
url="https://mcp.example.com/{{tenant}}",
headers={"X-Tenant": "{{tenant}}"},
secret_headers={"Authorization": "Bearer secret"},
),
)
class FakeMcpClient:
def __init__(self):
self.calls = []
self.error: Exception | None = None
async def call_tool(self, server, remote_tool_name, arguments):
self.calls.append((server, remote_tool_name, arguments))
if self.error:
raise self.error
return {
"text": "订单已发货",
"structuredContent": {"status": "shipped"},
"content": [{"type": "text", "text": "订单已发货"}],
}
class ToolExecutorMcpTests(unittest.IsolatedAsyncioTestCase):
def test_mcp_schema_is_exposed_to_llm(self):
properties, required = ToolExecutor.schema_parts(mcp_tool())
self.assertEqual(properties["order_id"]["type"], "string")
self.assertEqual(required, ["order_id"])
async def test_mcp_execution_renders_connection_and_assigns_result(self):
store = DynamicVariableStore(
{"tenant": "school-a"},
variable_types={"order_status": "string"},
)
client = FakeMcpClient()
executor = ToolExecutor(store, mcp_client=client)
result = await executor.execute(
mcp_tool(),
{"order_id": "A-100"},
)
self.assertEqual(result["status"], "ok")
self.assertEqual(result["updated_variables"], ["order_status"])
self.assertEqual(store.values["order_status"], "shipped")
server, remote_name, arguments = client.calls[0]
self.assertEqual(server.url, "https://mcp.example.com/school-a")
self.assertEqual(server.headers["X-Tenant"], "school-a")
self.assertEqual(server.secret_headers["Authorization"], "Bearer secret")
self.assertEqual(remote_name, "get_order")
self.assertEqual(arguments, {"order_id": "A-100"})
async def test_mcp_client_errors_use_shared_tool_error(self):
client = FakeMcpClient()
client.error = McpClientError("远端不可用")
executor = ToolExecutor(
DynamicVariableStore({"tenant": "school-a"}),
mcp_client=client,
)
with self.assertRaisesRegex(ToolExecutionError, "远端不可用"):
await executor.execute(mcp_tool(), {"order_id": "A-100"})
if __name__ == "__main__":
unittest.main()