import asyncio import json # import os import numpy as np from typing import List, Dict, Any from pathlib import Path from sentence_transformers import SentenceTransformer # embedder from mcp import ClientSession, StdioServerParameters from mcp.client.stdio import stdio_client class MCPToolRAGAdapter: def __init__ (self, config_path: str = "mcp_server_config.json"): self.config_path = config_path self.tools = [] self.toolnames = [] self.embedder = SentenceTransformer('all-MiniLM-L6-v2') # for embedding tool descriptions self.sessions = {} self.exit_stack = {} self.tool_registry = {} self.tool_embeddings = None def _load_config(self) -> Dict[str, Any]: config_path = Path(__file__).parent / self.config_path if not config_path.exists(): return {} try: with open(self.config_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): """Initialize all MCP sessions defined in the config file and index their tools.""" config = self._load_config() for server_name, params in config.items(): print(f"initializing session for {server_name} with params: {params}") server_params = StdioServerParameters( commanf=params["command"], args=params.get("args", []), # env=params.get("env", {}), ) # Verbindung aufbauen (Kontext-Manager manuell handhaben für Langzeit-Sessions) 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 self.exit_stack[server_name] = (transport_gen, session) # Zum späteren sauberen Schließen speichern print(f"Session for {server_name} initialized successfully.") # call tools and index thme result = await session.list_tools() tools = result.get("tools", []) for tool in tools: self.tool_registry.append({ "server": server_name, "tool_name": tool["name"], "definition": tool, "search_text": f"{tool['name']}: {tool.get('description', '')}", }) self.tool_names.append(tool["name"]) # embeddings for all tools in this session if self.tool_registry: texts = [t["search_text"] for t in self.tool_registry] self.tool_embeddings = self.embedder.encode(texts) print(f"Indexing completed. {len(texts)} tools ready.") def get_relevant_tools(self, query: str, top_k: int = 5) -> List[Dict[str, Any]]: """Given a user query, return the most relevant tools based on semantic similarity.""" if not self.tool_embeddings or not self.tool_registry: print("No tools indexed yet.") return [] query_embedding = self.embedder.encode([query]) similarities = np.dot(self.tool_embeddings, query_embedding.T).flatten() top_indices = np.argsort(similarities)[-top_k:][::-1] relevant_tools = [self.tool_registry[i] for i in top_indices] return relevant_tools async def call_tool(self, tool_name: str, arguments: Dict): """ Finds the right server for the tool and calls it with the provided arguments. """ for item in self.tool_registry: if item["definition"].name == tool_name: server_name = item["server"] session = self.sessions.get(server_name) if session: try: result = await session.call_tool(tool_name, arguments) return result except Exception as e: print(f"Error calling tool {tool_name} on server {server_name}: {e}") return f"Error calling tool: {e}" return f"Tool '{tool_name}' not found in registry." async def shutdown_all_sessions(self): """Gracefully shutdown all MCP sessions.""" 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 successfully.") except Exception as e: print(f"Error shutting down session for {server_name}: {e}")