feat MCP adapter: functional adapter and tool calls. Tools need to be fine grained
This commit is contained in:
parent
88c647a428
commit
e5831f69a8
@ -21,11 +21,13 @@ import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
import asyncio
|
||||
import pprint
|
||||
|
||||
import requests
|
||||
import httpx
|
||||
from dotenv import load_dotenv
|
||||
#from backend.agent.mcp_server_adapter import MCPToolRAGAdapter
|
||||
#from mcp_server_adapter import MCPToolAdapter # Import from current directory for easier testing without package structure
|
||||
from backend.agent.mcp_server_adapter import MCPToolAdapter
|
||||
|
||||
# ── mcp server initialization ────────────────────────────────────────────────────────────────
|
||||
@ -206,20 +208,23 @@ def build_all_tool_description() -> str:
|
||||
|
||||
descriptions = []
|
||||
for tool in all_tools:
|
||||
params = tool.inputSchema.get("properties", {})
|
||||
if params:
|
||||
param_lines = []
|
||||
for pname, pinfo in params.items():
|
||||
ptype = pinfo.get("type", "any")
|
||||
pdesc = pinfo.get("description", "")
|
||||
param_lines.append(f" - {pname} ({ptype}): {pdesc}")
|
||||
param_str = "\n".join(param_lines)
|
||||
else:
|
||||
param_str = " (none)"
|
||||
descriptions.append(
|
||||
f"- {tool.name}: {tool.description}\n"
|
||||
f"Parameters:\n{param_str}"
|
||||
)
|
||||
pprint.pprint(f"{tool}")
|
||||
descriptions.append(f"- {tool['tool_name']}: {tool['tool_description']}")
|
||||
#
|
||||
#params = tool.tool_description.inputSchema.get("properties", {})
|
||||
#if params:
|
||||
# param_lines = []
|
||||
# for pname, pinfo in params.items():
|
||||
# ptype = pinfo.get("type", "any")
|
||||
# pdesc = pinfo.get("description", "")
|
||||
# param_lines.append(f" - {pname} ({ptype}): {pdesc}")
|
||||
# param_str = "\n".join(param_lines)
|
||||
#else:
|
||||
# param_str = " (none)"
|
||||
#descriptions.append(
|
||||
# f"- {tool.name}: {tool.description}\n"
|
||||
# f"Parameters:\n{param_str}"
|
||||
#)
|
||||
|
||||
return "\n".join(descriptions)
|
||||
|
||||
@ -232,6 +237,7 @@ async def dispatch_tool(tool_name: str, arguments: dict) -> str:
|
||||
return f"DONE: {summary}"
|
||||
|
||||
try:
|
||||
print(f"Trying to call tool '{tool_name}' in dispatch_tool through MCPToolAdapter...")
|
||||
result = await adapter.call_tool(tool_name, arguments)
|
||||
|
||||
if result.isError:
|
||||
@ -455,28 +461,11 @@ class CodingAgent:
|
||||
return {"thought": "Max iterations reached.", "tool": "done",
|
||||
"arguments": {"summary": "Stopped: max iterations reached."}}
|
||||
|
||||
#current_context = self.messages[-1]["content"] if self.messages else ""
|
||||
self.iteration += 1
|
||||
self.messages = trim_messages(self.messages)
|
||||
|
||||
#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)
|
||||
raw = self._call_api(self.messages)
|
||||
raw = _strip_code_fences(raw)
|
||||
action = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
@ -591,4 +580,38 @@ class CodingAgent:
|
||||
})
|
||||
self.pending_action = None
|
||||
|
||||
def main():
|
||||
"""Example of how to use the CodingAgent in a simple loop."""
|
||||
agent = CodingAgent()
|
||||
task = "Write a Python function that returns the nth Fibonacci number."
|
||||
agent.start_task(task)
|
||||
|
||||
if agent.pending_action:
|
||||
print(f"Initial proposed action: {agent.pending_action['action']}")
|
||||
|
||||
|
||||
while not agent.is_done:
|
||||
action = asyncio.run(agent.propose_next_action())
|
||||
print(f"Proposed action: {action}")
|
||||
|
||||
if action["tool"] == "done":
|
||||
print("Task completed.")
|
||||
break
|
||||
else:
|
||||
user_feedback = input("Approve this action? (y/n) ")
|
||||
if user_feedback.lower() == "y":
|
||||
result = asyncio.run(agent.approve())
|
||||
print(f"Tool result: {result}")
|
||||
elif user_feedback.lower() == "n":
|
||||
feedback = input("Enter feedback for the agent: ")
|
||||
agent.reject(feedback)
|
||||
|
||||
|
||||
if result["is_done"]:
|
||||
print("Task completed.")
|
||||
break
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
@ -1,5 +1,6 @@
|
||||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
from typing import List, Dict, Any
|
||||
from pathlib import Path
|
||||
|
||||
@ -37,17 +38,21 @@ class MCPToolAdapter:
|
||||
print(f"Testing connection to {server_name}...")
|
||||
|
||||
self.servers[server_name] = params
|
||||
server_script = str(Path(__file__).parent / params["args"][0])
|
||||
if params.get("command") in ["py", "python", "python3"]:
|
||||
server_command = sys.executable
|
||||
else:
|
||||
server_command = params["command"]
|
||||
|
||||
server_params = StdioServerParameters(
|
||||
command=params["command"],
|
||||
args=params.get("args", [])
|
||||
command=server_command,
|
||||
args=[server_script],
|
||||
)
|
||||
|
||||
try:
|
||||
# Verbindung aufbauen
|
||||
async with stdio_client(server_params) as (read_stream, write_stream):
|
||||
print(f"Connected to {server_name}. Initializing session...")
|
||||
#print(f"read_stream: {read_stream}\nwrite_stream: {write_stream}")
|
||||
async with ClientSession(read_stream, write_stream) as session:
|
||||
await session.initialize()
|
||||
print(f"Session initialized for {server_name}. Requesting tools...")
|
||||
@ -61,7 +66,7 @@ class MCPToolAdapter:
|
||||
t_params = tool.inputSchema.get("properties", {})
|
||||
if t_params:
|
||||
param_lines = []
|
||||
for pname, pinfo in params.items():
|
||||
for pname, pinfo in t_params.items():
|
||||
ptype = pinfo.get("type", "any")
|
||||
pdesc = pinfo.get("description", "")
|
||||
param_lines.append(f" - {pname} ({ptype}): {pdesc}")
|
||||
@ -93,18 +98,25 @@ class MCPToolAdapter:
|
||||
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)
|
||||
tool_entry = next((t for t in self.tool_registry if t["tool_name"] == tool_name), None)
|
||||
|
||||
if not tool_entry:
|
||||
print(f"Tool '{tool_name}' not found in MCP adapter registry.")
|
||||
return f"Error: Tool '{tool_name}' not found in registry."
|
||||
|
||||
server_name = tool_entry["server"]
|
||||
server = self.servers.get(server_name)
|
||||
s_params = self.servers.get(server_name)
|
||||
|
||||
if server:
|
||||
if s_params:
|
||||
server_script = str(Path(__file__).parent / s_params["args"][0])
|
||||
if s_params.get("command") in ["py", "python", "python3"]:
|
||||
server_command = sys.executable
|
||||
else:
|
||||
server_command = s_params["command"]
|
||||
|
||||
server_params = StdioServerParameters(
|
||||
command=server["command"],
|
||||
args=server.get("args", [])
|
||||
command=server_command,
|
||||
args=[server_script],
|
||||
)
|
||||
|
||||
try:
|
||||
@ -133,7 +145,7 @@ def main():
|
||||
asyncio.run(adapter.initialize_all_servers())
|
||||
print("All servers initialized. Registered tools:")
|
||||
for tool in adapter.get_all_tools():
|
||||
print(f"- {tool['name']} (from {tool['server']})")
|
||||
print(f"- {tool['tool_name']} (from {tool['server']})")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@ -1,10 +1,10 @@
|
||||
{"FileSearchServer": {
|
||||
"command": "python3",
|
||||
"command": "py",
|
||||
"args": ["servers/mcp_server_file_search.py"]
|
||||
},
|
||||
|
||||
"WebSearchServer": {
|
||||
"command": "python3",
|
||||
"command": "py",
|
||||
"args": ["servers/mcp_server_web_search.py"],
|
||||
"env": {
|
||||
"DDGS_API_KEY": "your_ddgs_api_key_here"
|
||||
@ -12,7 +12,7 @@
|
||||
},
|
||||
|
||||
"CodeExecutionServer": {
|
||||
"command": "python3",
|
||||
"command": "py",
|
||||
"args": ["servers/mcp_server_code_execution.py"]
|
||||
}
|
||||
|
||||
|
||||
@ -136,4 +136,5 @@ def search_files(query: str) -> str:
|
||||
# ── Run the server ───────────────────────────────────────────────────────────
|
||||
|
||||
if __name__ == "__main__":
|
||||
mcp.run(transport="stdio")
|
||||
mcp.run(transport="stdio")
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user