additional mcp-adapter without RAG logic

This commit is contained in:
Irina Rueegg 2026-05-07 11:36:56 +02:00
parent 0a329139ca
commit 582e0bd711
3 changed files with 201 additions and 82 deletions

View File

@ -383,7 +383,7 @@ class CodingAgent:
self.api_key = os.getenv("API_KEY") self.api_key = os.getenv("API_KEY")
self.model = os.getenv("MODEL") self.model = os.getenv("MODEL")
async def _call_api(self, messages: list, relevant_tools: list = None) -> str: async def _call_api(self, messages: list) -> str:
"""Make a raw API call and return the response content string.""" """Make a raw API call and return the response content string."""
headers = {"Content-Type": "application/json"} headers = {"Content-Type": "application/json"}
@ -393,7 +393,6 @@ class CodingAgent:
payload = { payload = {
"model": self.model, "model": self.model,
"messages": messages, "messages": messages,
"tools": relevant_tools,
"temperature": 0.2, "temperature": 0.2,
"max_tokens": 4096, "max_tokens": 4096,
"stream": False, "stream": False,
@ -447,10 +446,25 @@ class CodingAgent:
current_context = self.messages[-1]["content"] if self.messages else "" current_context = self.messages[-1]["content"] if self.messages else ""
self.iteration += 1 self.iteration += 1
self.messages = trim_messages(self.messages) self.messages = trim_messages(self.messages)
relevant_tools = await get_tools_for_prompt(current_context)
tool_prompt = await get_tools_for_prompt(current_context)
enhanced_messages = self.messages.copy()
enhanced_messages.append({
"role": "system",
"content": f"""
Available tools for this step:
{tool_prompt}
You MUST choose one of these tools.
"""
})
try: try:
raw = await self._call_api(self.messages, relevant_tools) raw = await self._call_api(enhanced_messages)
raw = _strip_code_fences(raw) raw = _strip_code_fences(raw)
action = json.loads(raw) action = json.loads(raw)
except json.JSONDecodeError: except json.JSONDecodeError:

View File

@ -1,114 +1,105 @@
import asyncio import asyncio
import json import json
# import os
import numpy as np
from typing import List, Dict, Any from typing import List, Dict, Any
from pathlib import Path from pathlib import Path
from sentence_transformers import SentenceTransformer # embedder
from mcp import ClientSession, StdioServerParameters from mcp import ClientSession, StdioServerParameters
from mcp.client.stdio import stdio_client from mcp.client.stdio import stdio_client
class MCPToolRAGAdapter: class MCPToolAdapter:
def __init__ (self, config_path: str = "mcp_server_config.json"): def __init__(self, config_path: str = "mcp_server_config.json"):
self.config_path = config_path self.config_path = config_path
self.tools = [] self.sessions: Dict[str, ClientSession] = {}
self.toolnames = [] self.exit_stack: Dict[str, Any] = {}
self.embedder = SentenceTransformer('all-MiniLM-L6-v2') # for embedding tool descriptions self.tool_registry: List[Dict[str, Any]] = []
self.sessions = {}
self.exit_stack = {}
self.tool_registry = {}
self.tool_embeddings = None
def _load_config(self) -> Dict[str, Any]: def _load_config(self) -> Dict[str, Any]:
config_path = Path(__file__).parent / self.config_path """Lädt die Server-Konfiguration aus der JSON-Datei."""
if not config_path.exists(): path = Path(__file__).parent / self.config_path
if not path.exists():
print(f"Config file not found: {path}")
return {} return {}
try: try:
with open(self.config_path, 'r') as f: with open(path, 'r') as f:
return json.load(f) return json.load(f)
except json.JSONDecodeError as e: except json.JSONDecodeError as e:
print(f"Error decoding JSON config: {e}") print(f"Error decoding JSON config: {e}")
return {} return {}
async def initialize_all_sessions(self): async def initialize_all_sessions(self):
"""Initialize all MCP sessions defined in the config file and index their tools.""" """Initialisiert alle konfigurierten MCP-Sessions und registriert die Tools."""
config = self._load_config() config = self._load_config()
for server_name, params in config.items(): for server_name, params in config.items():
print(f"initializing session for {server_name} with params: {params}") print(f"Initializing session for {server_name}...")
server_params = StdioServerParameters( server_params = StdioServerParameters(
commanf=params["command"], command=params["command"],
args=params.get("args", []), args=params.get("args", []),
# env=params.get("env", {}), env=params.get("env", None),
) )
# Verbindung aufbauen (Kontext-Manager manuell handhaben für Langzeit-Sessions) try:
transport_gen = stdio_client(server_params) # Verbindung aufbauen
read, write = await transport_gen.__aenter__() transport_gen = stdio_client(server_params)
session = ClientSession(read, write) read, write = await transport_gen.__aenter__()
await session.__aenter__() session = ClientSession(read, write)
await session.initialize() await session.__aenter__()
await session.initialize()
self.sessions[server_name] = session self.sessions[server_name] = session
self.exit_stack[server_name] = (transport_gen, session) # Zum späteren sauberen Schließen speichern # Speichern für den Shutdown
print(f"Session for {server_name} initialized successfully.") self.exit_stack[server_name] = (transport_gen, session)
# call tools and index thme # Tools abrufen und registrieren
result = await session.list_tools() result = await session.list_tools()
tools = result.get("tools", []) # result ist oft ein Objekt, wir greifen auf das .tools Attribut zu
tools = getattr(result, 'tools', [])
for tool in tools: for tool in tools:
self.tool_registry.append({ # 'tool' ist hier meist ein Tool-Objekt vom MCP SDK
"server": server_name, self.tool_registry.append({
"tool_name": tool["name"], "server": server_name,
"definition": tool, "name": tool.name,
"search_text": f"{tool['name']}: {tool.get('description', '')}", "definition": tool
}) })
self.tool_names.append(tool["name"])
print(f"Session for {server_name} ready. {len(tools)} tools found.")
# 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]]: except Exception as e:
"""Given a user query, return the most relevant tools based on semantic similarity.""" print(f"Failed to initialize {server_name}: {e}")
if not self.tool_embeddings or not self.tool_registry:
print("No tools indexed yet.") def get_all_tools(self) -> List[Dict[str, Any]]:
return [] """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)
query_embedding = self.embedder.encode([query]) if not tool_entry:
similarities = np.dot(self.tool_embeddings, query_embedding.T).flatten() return f"Error: Tool '{tool_name}' not found in registry."
top_indices = np.argsort(similarities)[-top_k:][::-1]
relevant_tools = [self.tool_registry[i] for i in top_indices] server_name = tool_entry["server"]
session = self.sessions.get(server_name)
return relevant_tools if session:
try:
async def call_tool(self, tool_name: str, arguments: Dict): result = await session.call_tool(tool_name, arguments)
""" Finds the right server for the tool and calls it with the provided arguments. """ return result
for item in self.tool_registry: except Exception as e:
if item["definition"].name == tool_name: return f"Error calling tool '{tool_name}' on server '{server_name}': {str(e)}"
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." return f"Error: Session for server '{server_name}' not active."
async def shutdown_all_sessions(self): async def shutdown_all_sessions(self):
"""Gracefully shutdown all MCP sessions.""" """Schließt alle offenen Verbindungen sauber."""
for server_name, (transport_gen, session) in self.exit_stack.items(): for server_name, (transport_gen, session) in self.exit_stack.items():
try: try:
await session.__aexit__(None, None, None) await session.__aexit__(None, None, None)
await transport_gen.__aexit__(None, None, None) await transport_gen.__aexit__(None, None, None)
print(f"Session for {server_name} shut down successfully.") print(f"Session for {server_name} shut down.")
except Exception as e: except Exception as e:
print(f"Error shutting down session for {server_name}: {e}") print(f"Error during shutdown of {server_name}: {e}")

View File

@ -0,0 +1,114 @@
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}")