Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -4,11 +4,20 @@
tool listing and tool execution requests and responses.
"""

from enum import Enum
from typing import Any

from pydantic import BaseModel


class MCPTransport(str, Enum):
"""Supported transports for connecting to an MCP server."""

AUTO = "auto"
STREAMABLE_HTTP = "streamable_http"
SSE = "sse"


# MCPToolList Request and Response
class MCPToolListRequest(BaseModel):
"""Request model for listing available MCP tools from servers.
Expand All @@ -19,6 +28,7 @@ class MCPToolListRequest(BaseModel):

mcp_server_ids: list[str] | None = None
mcp_server_urls: list[str] | None = None
transport: MCPTransport = MCPTransport.AUTO


class MCPInfo(BaseModel):
Expand Down Expand Up @@ -80,6 +90,7 @@ class MCPCallToolRequest(BaseModel):
mcp_server_url: str | None = None
tool_name: str
tool_args: dict[str, Any] | None = None
transport: MCPTransport = MCPTransport.AUTO


class MCPTextResponse(BaseModel):
Expand Down
4 changes: 3 additions & 1 deletion core/plugin/link/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,8 @@ requires-python = ">=3.11"
dependencies = [
"fastapi>=0.111.0",
"jsonschema>=4.22.0",
"anyio>=4.5,<5",
"httpx>=0.27.1,<1",
"sqlalchemy>=2.0.30",
"sqlmodel>=0.0.18",
"loguru>=0.7.2",
Expand All @@ -23,7 +25,7 @@ dependencies = [
"opentelemetry-exporter-otlp>=1.22.0",
"redis-py-cluster==2.1.3",
"orjson>=3.10.15",
"mcp==1.6.0",
"mcp>=1.28.1,<2",
"aiohttp>=3.12.15",
"rich>=14.1.0",
"toml>=0.10.2",
Expand Down
207 changes: 93 additions & 114 deletions core/plugin/link/service/community/tools/mcp/mcp_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,8 +13,6 @@
from common.otlp.trace.span import Span
from fastapi import Body
from loguru import logger
from mcp import ClientSession
from mcp.client.sse import sse_client
from opentelemetry.trace import Status as OTelStatus
from opentelemetry.trace import StatusCode
from plugin.link.api.schemas.community.tools.mcp.mcp_tools_schema import (
Expand All @@ -28,18 +26,25 @@
MCPToolListData,
MCPToolListRequest,
MCPToolListResponse,
MCPTransport,
)
from plugin.link.consts import const
from plugin.link.domain.models.manager import get_db_engine
from plugin.link.infra.kafka_telemetry import send_telemetry_sync
from plugin.link.infra.tool_crud.process import ToolCrudOperation
from plugin.link.service.community.tools.mcp.mcp_transport import (
MCPTransportError,
initialized_mcp_session,
)
from plugin.link.utils.errors.code import ErrCode
from plugin.link.utils.security.access_interceptor import is_in_blacklist, is_local_url
from plugin.link.utils.sid.sid_generator2 import new_sid


async def _process_mcp_server_by_id(
mcp_server_id: str, span_context: Any
mcp_server_id: str,
span_context: Any,
transport: MCPTransport = MCPTransport.AUTO,
) -> MCPItemInfo:
"""Process a single MCP server by ID and return its tools."""
err, url = get_mcp_server_url(mcp_server_id=mcp_server_id, span=span_context)
Expand All @@ -60,10 +65,14 @@ async def _process_mcp_server_by_id(
tools=[],
)

return await _connect_and_get_tools(url, server_id=mcp_server_id)
return await _connect_and_get_tools(
url, server_id=mcp_server_id, transport=transport
)


async def _process_mcp_server_by_url(url: str) -> MCPItemInfo:
async def _process_mcp_server_by_url(
url: str, transport: MCPTransport = MCPTransport.AUTO
) -> MCPItemInfo:
"""Process a single MCP server by URL and return its tools."""
if is_local_url(url):
err = ErrCode.MCP_SERVER_LOCAL_URL_ERR
Expand All @@ -83,71 +92,67 @@ async def _process_mcp_server_by_url(url: str) -> MCPItemInfo:
tools=[],
)

return await _connect_and_get_tools(url, server_url=url)
return await _connect_and_get_tools(url, server_url=url, transport=transport)


def _transport_error_code(error: MCPTransportError) -> ErrCode:
"""Map a transport setup phase to the existing public MCP error codes."""
if error.phase == "initialization":
return ErrCode.MCP_SERVER_INITIAL_ERR
if error.phase == "session creation":
return ErrCode.MCP_SERVER_SESSION_ERR
return ErrCode.MCP_SERVER_CONNECT_ERR


async def _connect_and_get_tools(
url: str, server_id: Optional[str] = None, server_url: Optional[str] = None
url: str,
server_id: Optional[str] = None,
server_url: Optional[str] = None,
transport: MCPTransport = MCPTransport.AUTO,
) -> MCPItemInfo:
"""Connect to MCP server and retrieve tools."""
try:
async with sse_client(url=url) as (read, write):
async with initialized_mcp_session(url, transport) as (session, _):
try:
async with ClientSession(read, write, logging_callback=None) as session:
try:
await session.initialize()
except Exception:
err = ErrCode.MCP_SERVER_INITIAL_ERR
return MCPItemInfo(
server_id=server_id,
server_url=server_url,
server_status=err.code,
server_message=err.msg,
tools=[],
)

try:
tools_result = await session.list_tools()
tools_dict = tools_result.model_dump()["tools"]
tools = []
for tool in tools_dict:
tool_info = MCPInfo(
name=tool.get("name", "No name available"),
description=tool.get(
"description", "No description available"
),
inputSchema=tool.get("inputSchema"),
)
tools.append(tool_info)

success = ErrCode.SUCCESSES
return MCPItemInfo(
server_id=server_id,
server_url=server_url,
server_status=success.code,
server_message=success.msg,
tools=tools,
)
except Exception:
err = ErrCode.MCP_SERVER_TOOL_LIST_ERR
return MCPItemInfo(
server_id=server_id,
server_url=server_url,
server_status=err.code,
server_message=err.msg,
tools=[],
)
tools_result = await session.list_tools()
tools_dict = tools_result.model_dump()["tools"]
tools = []
for tool in tools_dict:
tool_info = MCPInfo(
name=tool.get("name", "No name available"),
description=tool.get("description", "No description available"),
inputSchema=tool.get("inputSchema"),
)
tools.append(tool_info)

success = ErrCode.SUCCESSES
return MCPItemInfo(
server_id=server_id,
server_url=server_url,
server_status=success.code,
server_message=success.msg,
tools=tools,
)
except Exception:
err = ErrCode.MCP_SERVER_SESSION_ERR
err = ErrCode.MCP_SERVER_TOOL_LIST_ERR
return MCPItemInfo(
server_id=server_id,
server_url=server_url,
server_status=err.code,
server_message=err.msg,
tools=[],
)
except MCPTransportError as error:
err = _transport_error_code(error)
return MCPItemInfo(
server_id=server_id,
server_url=server_url,
server_status=err.code,
server_message=err.msg,
tools=[],
)
except Exception:
err = ErrCode.MCP_SERVER_CONNECT_ERR
err = ErrCode.MCP_SERVER_SESSION_ERR
return MCPItemInfo(
server_id=server_id,
server_url=server_url,
Expand Down Expand Up @@ -197,15 +202,17 @@ async def tool_list(list_info: MCPToolListRequest = Body()) -> MCPToolListRespon
# Process IDs
if mcp_server_ids:
for mcp_server_id in mcp_server_ids:
item = await _process_mcp_server_by_id(mcp_server_id, span_context)
item = await _process_mcp_server_by_id(
mcp_server_id, span_context, list_info.transport
)
items.append(item)

# Process URLs
if mcp_server_urls:
for url in mcp_server_urls:
if not url.strip():
continue
item = await _process_mcp_server_by_url(url)
item = await _process_mcp_server_by_url(url, list_info.transport)
items.append(item)

success = ErrCode.SUCCESSES
Expand Down Expand Up @@ -254,26 +261,6 @@ def _log_error_to_kafka(
send_telemetry_sync(node_trace)


async def _initialize_session(
session: Any,
session_id: str,
span_context: Any,
node_trace: NodeTraceLog,
mcp_server_id: str,
m: Meter,
) -> Optional[MCPCallToolResponse]:
"""Initialize MCP session with error handling."""
try:
await session.initialize()
except Exception:
err = ErrCode.MCP_SERVER_INITIAL_ERR
span_context.add_error_event(err.msg)
span_context.set_status(OTelStatus(StatusCode.ERROR))
_log_error_to_kafka(err, node_trace, mcp_server_id, m)
return _create_error_response(err, session_id)
return None


async def _execute_tool_call(
session: Any,
tool_name: str,
Expand Down Expand Up @@ -320,50 +307,41 @@ async def _call_mcp_tool(
node_trace: NodeTraceLog,
mcp_server_id: str,
m: Meter,
transport: MCPTransport = MCPTransport.AUTO,
) -> MCPCallToolResponse:
"""Execute the actual MCP tool call with proper error handling."""
try:
async with sse_client(url=url) as (read, write):
try:
async with ClientSession(read, write, logging_callback=None) as session:
# Initialize session
init_result = await _initialize_session(
session, session_id, span_context, node_trace, mcp_server_id, m
)
if init_result:
return init_result

# Execute tool call
call_result = await _execute_tool_call(
session,
tool_name,
tool_args,
session_id,
span_context,
node_trace,
mcp_server_id,
m,
)
async with initialized_mcp_session(url, transport) as (session, _):
call_result = await _execute_tool_call(
session,
tool_name,
tool_args,
session_id,
span_context,
node_trace,
mcp_server_id,
m,
)

if isinstance(call_result[0], MCPCallToolResponse):
return call_result[0]
if isinstance(call_result[0], MCPCallToolResponse):
return call_result[0]

is_error, content = call_result
success = ErrCode.SUCCESSES
return MCPCallToolResponse(
code=success.code,
message=success.msg,
sid=session_id,
data=MCPCallToolData(isError=is_error, content=content),
)
except Exception:
err = ErrCode.MCP_SERVER_SESSION_ERR
span_context.add_error_event(err.msg)
span_context.set_status(OTelStatus(StatusCode.ERROR))
_log_error_to_kafka(err, node_trace, mcp_server_id, m)
return _create_error_response(err, session_id)
is_error, content = call_result
success = ErrCode.SUCCESSES
return MCPCallToolResponse(
code=success.code,
message=success.msg,
sid=session_id,
data=MCPCallToolData(isError=is_error, content=content),
)
except MCPTransportError as error:
err = _transport_error_code(error)
span_context.add_error_event(err.msg)
span_context.set_status(OTelStatus(StatusCode.ERROR))
_log_error_to_kafka(err, node_trace, mcp_server_id, m)
return _create_error_response(err, session_id)
except Exception:
err = ErrCode.MCP_SERVER_CONNECT_ERR
err = ErrCode.MCP_SERVER_SESSION_ERR
span_context.add_error_event(err.msg)
span_context.set_status(OTelStatus(StatusCode.ERROR))
_log_error_to_kafka(err, node_trace, mcp_server_id, m)
Expand Down Expand Up @@ -456,6 +434,7 @@ async def call_tool(call_info: MCPCallToolRequest = Body()) -> MCPCallToolRespon
node_trace,
mcp_server_id,
m,
call_info.transport,
)
span_context.add_info_events({"call_tool_result": result.model_dump_json()})
# Log success if the call succeeded
Expand Down
Loading
Loading