AISE1_Project_Irina_Livio/backend/agent/mcp_server_adapter.py
Livio Meuli cc14f25e69 refactor: add REVIEW annotations for dead/debug code across frontend and backend
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-05-24 10:38:59 +02:00

220 lines
9.4 KiB
Python

"""Adapter layer between the CodingAgent and one or more MCP tool servers.
MCPToolAdapter reads a JSON config file that lists MCP server processes, spawns
each process via stdio, queries its available tools, and stores them in a flat
registry. At call time it re-spawns the appropriate server process, executes
the requested tool, and returns the raw MCP result object.
Design note: connections are opened per-call (not kept alive) because Streamlit
reruns make it impractical to maintain long-lived async context managers across
the synchronous/asynchronous boundary.
"""
import asyncio
import json
import sys
from typing import List, Dict, Any
from pathlib import Path
from mcp import ClientSession, StdioServerParameters
from mcp.client.stdio import stdio_client
class MCPToolAdapter:
"""Discovers and dispatches MCP tools from one or more stdio-based MCP servers.
Workflow:
1. Call ``initialize_all_servers()`` once at startup to populate the
tool registry from every server listed in the config file.
2. Call ``get_all_tools()`` to retrieve the registry for building the
system-prompt tool description.
3. Call ``call_tool(name, arguments)`` whenever the agent wants to
execute a tool. The adapter resolves the owning server, opens a
fresh connection, and returns the MCP result object.
Attributes:
config_path: Path (relative to this file) of the JSON server config.
servers: Dict mapping server name → raw config params dict.
tool_registry: Flat list of registered tool dicts, each containing
"server", "tool_name", and "tool_description".
"""
def __init__(self, config_path: str = "mcp_server_config.json"):
self.config_path = config_path
self.servers: Dict[str, Dict] = {}
self.tool_registry: List[Dict[str, Any]] = []
def _load_config(self) -> Dict[str, Any]:
"""Load the MCP server configuration from the JSON file next to this module.
Returns:
Parsed config dict, or an empty dict if the file is missing or invalid.
"""
path = Path(__file__).parent / self.config_path
if not path.exists():
print(f"Config file not found: {path}")
return {}
try:
with open(path, 'r') as f:
return json.load(f)
except json.JSONDecodeError as e:
print(f"Error decoding JSON config: {e}")
return {}
async def initialize_all_servers(self):
"""Connect to every configured MCP server and register their tools.
Opens a short-lived stdio connection to each server, calls list_tools(),
and stores each discovered tool in ``self.tool_registry``. Servers that
fail to connect are skipped with a warning so a single broken server does
not prevent the others from loading.
"""
print("Initializing MCP sessions...")
config = self._load_config()
print(f"Loaded config for servers: {list(config.keys())}")
for server_name, params in config.items():
print(f"Testing connection to {server_name}...")
self.servers[server_name] = params
server_script = str(Path(__file__).parent / params["args"][0])
# Always use the current Python interpreter so the server runs in the
# same virtual environment as the adapter, regardless of the literal
# command string in the config ("py", "python", "python3").
if params.get("command") in ["py", "python", "python3"]:
server_command = sys.executable
else:
server_command = params["command"]
server_params = StdioServerParameters(
command=server_command,
args=[server_script],
)
try:
async with stdio_client(server_params) as (read_stream, write_stream):
print(f"Connected to {server_name}. Initializing session...")
async with ClientSession(read_stream, write_stream) as session:
await session.initialize()
print(f"Session initialized for {server_name}. Requesting tools...")
result = await session.list_tools()
# REVIEW: debug print — remove before shipping.
print(f"Tools received from {server_name}: {result}")
tools = result.tools
# REVIEW: duplicate print — identical message already printed inside the `async with` block above.
print(f"Tools received from {server_name}: {result}")
for tool in tools:
# Build a human-readable parameter description for the system prompt.
t_params = tool.inputSchema.get("properties", {})
if t_params:
param_lines = []
for pname, pinfo in t_params.items():
ptype = pinfo.get("type", "any")
pdesc = pinfo.get("description", "")
param_lines.append(f" - {pname} ({ptype}): {pdesc}")
param_str = "\n".join(param_lines)
else:
param_str = " (none)"
t_definition = f"- {tool.name}: {tool.description}\nParameters:\n{param_str}"
self.tool_registry.append({
"server": server_name,
"tool_name": tool.name,
"tool_description": t_definition
})
print(f"Registered tool '{tool.name}' from {server_name}.")
print(f"Session for {server_name} ready. {len(tools)} tools found.")
except Exception as e:
print(f"Failed to initialize {server_name}: {e}")
def get_all_tools(self) -> List[Dict[str, Any]]:
"""Return the full list of registered tools across all servers.
Returns:
List of dicts, each with keys "server", "tool_name", "tool_description".
"""
return self.tool_registry
async def call_tool(self, tool_name: str, arguments: Dict[str, Any]):
"""Look up a tool in the registry, connect to its server, and execute it.
Opens a fresh stdio connection for every call. This is intentionally
stateless so that server crashes or restarts are fully transparent.
Args:
tool_name: Name of the tool to call (must be in the registry).
arguments: Key-value arguments passed verbatim to the MCP server.
Returns:
The raw MCP ``CallToolResult`` object on success, or an error string
if the tool is not found or the server raises an exception.
"""
# Look up which server owns this tool.
tool_entry = next((t for t in self.tool_registry if t["tool_name"] == tool_name), None)
if not tool_entry:
print(f"Tool '{tool_name}' not found in MCP adapter registry.")
return f"Error: Tool '{tool_name}' not found in registry."
server_name = tool_entry["server"]
s_params = self.servers.get(server_name)
if s_params:
server_script = str(Path(__file__).parent / s_params["args"][0])
# Normalise the interpreter command the same way as in initialize_all_servers().
if s_params.get("command") in ["py", "python", "python3"]:
server_command = sys.executable
else:
server_command = s_params["command"]
server_params = StdioServerParameters(
command=server_command,
args=[server_script],
)
try:
async with stdio_client(server_params) as (read_stream, write_stream):
async with ClientSession(read_stream, write_stream) as session:
await session.initialize()
result = await session.call_tool(tool_name, arguments)
return result
except Exception as e:
return f"Error calling tool '{tool_name}' on server '{server_name}': {str(e)}"
return f"Error: Session for server '{server_name}' not active."
async def shutdown_all_sessions(self):
"""Close all open server connections gracefully.
Note: This method references ``self.exit_stack`` which is not currently
populated (connections are opened per-call). It is kept as a placeholder
for a future persistent-connection implementation.
"""
# REVIEW: self.exit_stack is never assigned in __init__ — calling this method will
# always raise AttributeError. Either remove this method or initialise exit_stack
# in __init__ as an empty dict.
for server_name, (transport_gen, session) in self.exit_stack.items():
try:
await session.__aexit__(None, None, None)
await transport_gen.__aexit__(None, None, None)
print(f"Session for {server_name} shut down.")
except Exception as e:
print(f"Error during shutdown of {server_name}: {e}")
def main():
adapter = MCPToolAdapter()
asyncio.run(adapter.initialize_all_servers())
print("All servers initialized. Registered tools:")
for tool in adapter.get_all_tools():
print(f"- {tool['tool_name']} (from {tool['server']})")
if __name__ == "__main__":
main()