diff --git a/backend/agent/coding_agent.py b/backend/agent/coding_agent.py index b751b88..1617044 100644 --- a/backend/agent/coding_agent.py +++ b/backend/agent/coding_agent.py @@ -383,7 +383,7 @@ class CodingAgent: self.api_key = os.getenv("API_KEY") 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.""" headers = {"Content-Type": "application/json"} @@ -393,7 +393,6 @@ class CodingAgent: payload = { "model": self.model, "messages": messages, - "tools": relevant_tools, "temperature": 0.2, "max_tokens": 4096, "stream": False, @@ -447,10 +446,25 @@ class CodingAgent: current_context = self.messages[-1]["content"] if self.messages else "" self.iteration += 1 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: - raw = await self._call_api(self.messages, relevant_tools) + raw = await self._call_api(enhanced_messages) raw = _strip_code_fences(raw) action = json.loads(raw) except json.JSONDecodeError: diff --git a/backend/agent/mcp_server_adapter.py b/backend/agent/mcp_server_adapter.py index 6583c9b..03124e3 100644 --- a/backend/agent/mcp_server_adapter.py +++ b/backend/agent/mcp_server_adapter.py @@ -1,114 +1,105 @@ -import asyncio +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"): +class MCPToolAdapter: + 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 - + 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]: - config_path = Path(__file__).parent / self.config_path - if not config_path.exists(): + """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(self.config_path, 'r') as f: - return json.load(f) + 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): - """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() + 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( - commanf=params["command"], + command=params["command"], args=params.get("args", []), - # env=params.get("env", {}), + env=params.get("env", None), ) - # 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() + 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 - self.exit_stack[server_name] = (transport_gen, session) # Zum späteren sauberen Schließen speichern - print(f"Session for {server_name} initialized successfully.") + self.sessions[server_name] = session + # Speichern für den Shutdown + self.exit_stack[server_name] = (transport_gen, session) - # call tools and index thme - result = await session.list_tools() - tools = result.get("tools", []) + # 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: - 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.") + 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.") - 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 [] + 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) - query_embedding = self.embedder.encode([query]) - similarities = np.dot(self.tool_embeddings, query_embedding.T).flatten() - top_indices = np.argsort(similarities)[-top_k:][::-1] + if not tool_entry: + return f"Error: Tool '{tool_name}' not found in registry." - 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 - - 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}" + 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"Tool '{tool_name}' not found in registry." - + return f"Error: Session for server '{server_name}' not active." + 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(): try: await session.__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: - print(f"Error shutting down session for {server_name}: {e}") + print(f"Error during shutdown of {server_name}: {e}") \ No newline at end of file diff --git a/backend/agent/mcp_server_adapter_RAG.py b/backend/agent/mcp_server_adapter_RAG.py new file mode 100644 index 0000000..6583c9b --- /dev/null +++ b/backend/agent/mcp_server_adapter_RAG.py @@ -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}")