additional mcp-adapter without RAG logic
This commit is contained in:
parent
0a329139ca
commit
582e0bd711
@ -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:
|
||||||
|
|||||||
@ -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}")
|
||||||
114
backend/agent/mcp_server_adapter_RAG.py
Normal file
114
backend/agent/mcp_server_adapter_RAG.py
Normal 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}")
|
||||||
Loading…
x
Reference in New Issue
Block a user