105 lines
4.1 KiB
Python
105 lines
4.1 KiB
Python
import asyncio
|
|
import json
|
|
from typing import List, Dict, Any
|
|
from pathlib import Path
|
|
|
|
from mcp import ClientSession, StdioServerParameters
|
|
from mcp.client.stdio import stdio_client
|
|
|
|
class MCPToolAdapter:
|
|
def __init__(self, config_path: str = "mcp_server_config.json"):
|
|
self.config_path = config_path
|
|
self.sessions: Dict[str, ClientSession] = {}
|
|
self.exit_stack: Dict[str, Any] = {}
|
|
self.tool_registry: List[Dict[str, Any]] = []
|
|
|
|
def _load_config(self) -> Dict[str, Any]:
|
|
"""Lädt die Server-Konfiguration aus der JSON-Datei."""
|
|
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_sessions(self):
|
|
"""Initialisiert alle konfigurierten MCP-Sessions und registriert die Tools."""
|
|
config = self._load_config()
|
|
|
|
for server_name, params in config.items():
|
|
print(f"Initializing session for {server_name}...")
|
|
|
|
server_params = StdioServerParameters(
|
|
command=params["command"],
|
|
args=params.get("args", []),
|
|
env=params.get("env", None),
|
|
)
|
|
|
|
try:
|
|
# Verbindung aufbauen
|
|
transport_gen = stdio_client(server_params)
|
|
read, write = await transport_gen.__aenter__()
|
|
session = ClientSession(read, write)
|
|
await session.__aenter__()
|
|
await session.initialize()
|
|
|
|
self.sessions[server_name] = session
|
|
# Speichern für den Shutdown
|
|
self.exit_stack[server_name] = (transport_gen, session)
|
|
|
|
# Tools abrufen und registrieren
|
|
result = await session.list_tools()
|
|
# result ist oft ein Objekt, wir greifen auf das .tools Attribut zu
|
|
tools = getattr(result, 'tools', [])
|
|
|
|
for tool in tools:
|
|
# 'tool' ist hier meist ein Tool-Objekt vom MCP SDK
|
|
self.tool_registry.append({
|
|
"server": server_name,
|
|
"name": tool.name,
|
|
"definition": tool
|
|
})
|
|
|
|
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]]:
|
|
"""Gibt alle gesammelten Tools zurück."""
|
|
return self.tool_registry
|
|
|
|
async def call_tool(self, tool_name: str, arguments: Dict[str, Any]):
|
|
"""Findet den richtigen Server für ein Tool und führt es aus."""
|
|
# Suche in der Registry nach dem passenden Server
|
|
tool_entry = next((t for t in self.tool_registry if t["name"] == tool_name), None)
|
|
|
|
if not tool_entry:
|
|
return f"Error: Tool '{tool_name}' not found in registry."
|
|
|
|
server_name = tool_entry["server"]
|
|
session = self.sessions.get(server_name)
|
|
|
|
if session:
|
|
try:
|
|
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):
|
|
"""Schließt alle offenen Verbindungen sauber."""
|
|
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}") |