Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@ dependencies = [
]

[project.optional-dependencies]
mcp = ["mcp>=1.28.1"]
mcp = ["mcp>=2.0"]
dev = ["pre-commit", "pytest", "pytest-cov", "pytest-doctestplus", "pytest-asyncio", "ruff", "docutils"]
doc = [
"jupyter-book>=2.1.6,<3",
Expand Down
78 changes: 47 additions & 31 deletions src/gdm/mcp/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,9 +13,9 @@

import typer
from gdm.distribution import DistributionSystem
from mcp.server import Server
from mcp.server import Server, ServerRequestContext
from mcp.server.stdio import stdio_server
from mcp.types import TextContent, Tool
from mcp.types import CallToolResult, ListToolsResult, TextContent, Tool

from gdm.mcp import __version__
from gdm.mcp.exceptions import GDMMCPException
Expand Down Expand Up @@ -51,9 +51,6 @@
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("gdm_mcp")

# Create MCP server instance
app = Server("grid-data-models-mcp")

# Runtime toggle for serving tool calls. Control tools remain available.
_TOOL_CALLS_ENABLED = True
_CONTROL_TOOLS = {"set_tool_calls_enabled", "get_tool_calls_enabled"}
Expand Down Expand Up @@ -185,10 +182,9 @@
}


@app.list_tools()
async def list_tools() -> list[Tool]:
async def list_tools(ctx: ServerRequestContext, params=None) -> ListToolsResult:

Check warning on line 185 in src/gdm/mcp/server.py

View check run for this annotation

SonarQubeCloud / SonarCloud Code Analysis

Use asynchronous features in this function or remove the `async` keyword.

See more on https://sonarcloud.io/project/issues?id=NREL-Distribution-Suites_grid-data-models&issues=AZ-0-m5QwJvOFR3RlF6O&open=AZ-0-m5QwJvOFR3RlF6O&pullRequest=203
"""List all available MCP tools."""
return [
tools = [
# Validation tools
Tool(
name="diagnose_system",
Expand Down Expand Up @@ -686,6 +682,7 @@
},
),
]
return ListToolsResult(tools=tools)


# Tool dispatch map
Expand Down Expand Up @@ -717,46 +714,58 @@
}


@app.call_tool()
async def call_tool(name: str, arguments: Any) -> list[TextContent]:
async def call_tool(ctx: ServerRequestContext, params: Any) -> CallToolResult:
"""Handle tool calls from MCP clients."""
name = params.name
arguments = params.arguments or {}
try:
logger.info(f"Tool called: {name} with arguments: {arguments}")

if not _TOOL_CALLS_ENABLED and name not in _CONTROL_TOOLS:
return [
TextContent(
type="text",
text=json.dumps(
{
"error": (
"Tool calls are currently disabled. "
"Use set_tool_calls_enabled to re-enable."
)
},
indent=2,
),
)
]
return CallToolResult(
content=[
TextContent(
type="text",
text=json.dumps(
{
"error": (
"Tool calls are currently disabled. "
"Use set_tool_calls_enabled to re-enable."
)
},
indent=2,
),
)
]
)

handler = _TOOL_HANDLERS.get(name)
if handler is not None:
result = await handler(arguments)
else:
result = {"error": f"Unknown tool: {name}"}

return [TextContent(type="text", text=json.dumps(result, indent=2, default=str))]
return CallToolResult(
content=[TextContent(type="text", text=json.dumps(result, indent=2, default=str))]
)

except GDMMCPException as e:
logger.error(f"GDM MCP error in {name}: {str(e)}")
return [TextContent(type="text", text=json.dumps({"error": str(e)}, indent=2))]
return CallToolResult(
content=[TextContent(type="text", text=json.dumps({"error": str(e)}, indent=2))],
is_error=True,
)
except Exception as e:
logger.error(f"Unexpected error in {name}: {str(e)}", exc_info=True)
return [
TextContent(
type="text", text=json.dumps({"error": f"Unexpected error: {str(e)}"}, indent=2)
)
]
return CallToolResult(
content=[
TextContent(
type="text",
text=json.dumps({"error": f"Unexpected error: {str(e)}"}, indent=2),
)
],
is_error=True,
)


# Tool implementations
Expand Down Expand Up @@ -1079,6 +1088,13 @@
# Run the server
import asyncio

app = Server(
"grid-data-models-mcp",
version=__version__,
on_list_tools=list_tools,
on_call_tool=call_tool,
)

async def run():
async with stdio_server() as (read_stream, write_stream):
await app.run(read_stream, write_stream, app.create_initialization_options())
Expand Down
50 changes: 32 additions & 18 deletions tests/test_mcp_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,30 @@
import json
import os
import sqlite3
from unittest.mock import MagicMock

import gdm.mcp.server as mcp_server
import pytest
from gdm.mcp.server import _load_system_with_fallback_name
from gdm.distribution.components import DistributionBus, DistributionVoltageSource


def _make_call_tool_params(name: str, arguments: dict | None = None):
"""Create a mock CallToolRequestParams for testing."""
params = MagicMock()
params.name = name
params.arguments = arguments or {}
return params


def _call_tool(name: str, arguments: dict | None = None):
"""Helper to call tool_handler with mock context."""
ctx = MagicMock()
params = _make_call_tool_params(name, arguments)
result = asyncio.run(mcp_server.call_tool(ctx, params))
return result


def test_load_system_with_fallback_name_for_null_name(simple_system, tmp_path):
"""Falls back to file stem when serialized system name is null."""
system_path = tmp_path / "null_name_system.json"
Expand Down Expand Up @@ -95,8 +112,8 @@ def test_get_tool_calls_enabled_reports_current_state():
"""Control status tool should report current runtime toggle state."""
mcp_server._TOOL_CALLS_ENABLED = True

response = asyncio.run(mcp_server.call_tool("get_tool_calls_enabled", {}))
payload = json.loads(response[0].text)
response = _call_tool("get_tool_calls_enabled", {})
payload = json.loads(response.content[0].text)

assert payload["tool_calls_enabled"] is True

Expand All @@ -105,42 +122,39 @@ def test_set_tool_calls_enabled_disables_non_control_calls():
"""Disabling should block normal tools while allowing control tools."""
mcp_server._TOOL_CALLS_ENABLED = True

disable_response = asyncio.run(
mcp_server.call_tool("set_tool_calls_enabled", {"enabled": False})
)
disable_payload = json.loads(disable_response[0].text)
disable_response = _call_tool("set_tool_calls_enabled", {"enabled": False})
disable_payload = json.loads(disable_response.content[0].text)
assert disable_payload["tool_calls_enabled"] is False

blocked_response = asyncio.run(mcp_server.call_tool("unknown_normal_tool", {}))
blocked_payload = json.loads(blocked_response[0].text)
blocked_response = _call_tool("unknown_normal_tool", {})
blocked_payload = json.loads(blocked_response.content[0].text)
assert "disabled" in blocked_payload["error"].lower()

# Control tools remain callable so clients can re-enable.
status_response = asyncio.run(mcp_server.call_tool("get_tool_calls_enabled", {}))
status_payload = json.loads(status_response[0].text)
status_response = _call_tool("get_tool_calls_enabled", {})
status_payload = json.loads(status_response.content[0].text)
assert status_payload["tool_calls_enabled"] is False


def test_set_tool_calls_enabled_can_reenable():
"""Re-enabling should restore normal call flow."""
mcp_server._TOOL_CALLS_ENABLED = False

enable_response = asyncio.run(
mcp_server.call_tool("set_tool_calls_enabled", {"enabled": True})
)
enable_payload = json.loads(enable_response[0].text)
enable_response = _call_tool("set_tool_calls_enabled", {"enabled": True})
enable_payload = json.loads(enable_response.content[0].text)

assert enable_payload["tool_calls_enabled"] is True

unknown_response = asyncio.run(mcp_server.call_tool("unknown_normal_tool", {}))
unknown_payload = json.loads(unknown_response[0].text)
unknown_response = _call_tool("unknown_normal_tool", {})
unknown_payload = json.loads(unknown_response.content[0].text)
assert "unknown tool" in unknown_payload["error"].lower()


def test_list_tools_includes_reduce_system():
"""Tool list should expose model-reduction capability."""
tools = asyncio.run(mcp_server.list_tools())
tool_names = {tool.name for tool in tools}
ctx = MagicMock()
result = asyncio.run(mcp_server.list_tools(ctx, None))
tool_names = {tool.name for tool in result.tools}

assert "reduce_system" in tool_names
assert "save_system" in tool_names
Expand Down
Loading