Merge pull request #1061 from cyzus/feat-include-mcp-in-manus-agent
Manus Agent to include MCP tools
This commit is contained in:
+100
-11
@@ -1,24 +1,24 @@
|
||||
from typing import Optional
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from pydantic import Field, model_validator
|
||||
|
||||
from app.agent.browser import BrowserContextHelper
|
||||
from app.agent.toolcall import ToolCallAgent
|
||||
from app.config import config
|
||||
from app.logger import logger
|
||||
from app.prompt.manus import NEXT_STEP_PROMPT, SYSTEM_PROMPT
|
||||
from app.tool import Terminate, ToolCollection
|
||||
from app.tool.browser_use_tool import BrowserUseTool
|
||||
from app.tool.mcp import MCPClients, MCPClientTool
|
||||
from app.tool.python_execute import PythonExecute
|
||||
from app.tool.str_replace_editor import StrReplaceEditor
|
||||
|
||||
|
||||
class Manus(ToolCallAgent):
|
||||
"""A versatile general-purpose agent."""
|
||||
"""A versatile general-purpose agent with support for both local and MCP tools."""
|
||||
|
||||
name: str = "Manus"
|
||||
description: str = (
|
||||
"A versatile agent that can solve various tasks using multiple tools"
|
||||
)
|
||||
description: str = "A versatile agent that can solve various tasks using multiple tools including MCP-based tools"
|
||||
|
||||
system_prompt: str = SYSTEM_PROMPT.format(directory=config.workspace_root)
|
||||
next_step_prompt: str = NEXT_STEP_PROMPT
|
||||
@@ -26,6 +26,9 @@ class Manus(ToolCallAgent):
|
||||
max_observe: int = 10000
|
||||
max_steps: int = 20
|
||||
|
||||
# MCP clients for remote tool access
|
||||
mcp_clients: MCPClients = Field(default_factory=MCPClients)
|
||||
|
||||
# Add general-purpose tools to the tool collection
|
||||
available_tools: ToolCollection = Field(
|
||||
default_factory=lambda: ToolCollection(
|
||||
@@ -34,16 +37,107 @@ class Manus(ToolCallAgent):
|
||||
)
|
||||
|
||||
special_tool_names: list[str] = Field(default_factory=lambda: [Terminate().name])
|
||||
|
||||
browser_context_helper: Optional[BrowserContextHelper] = None
|
||||
|
||||
# Track connected MCP servers
|
||||
connected_servers: Dict[str, str] = Field(
|
||||
default_factory=dict
|
||||
) # server_id -> url/command
|
||||
_initialized: bool = False
|
||||
|
||||
@model_validator(mode="after")
|
||||
def initialize_helper(self) -> "Manus":
|
||||
"""Initialize basic components synchronously."""
|
||||
self.browser_context_helper = BrowserContextHelper(self)
|
||||
return self
|
||||
|
||||
@classmethod
|
||||
async def create(cls, **kwargs) -> "Manus":
|
||||
"""Factory method to create and properly initialize a Manus instance."""
|
||||
instance = cls(**kwargs)
|
||||
await instance.initialize_mcp_servers()
|
||||
instance._initialized = True
|
||||
return instance
|
||||
|
||||
async def initialize_mcp_servers(self) -> None:
|
||||
"""Initialize connections to configured MCP servers."""
|
||||
for server_id, server_config in config.mcp_config.servers.items():
|
||||
try:
|
||||
if server_config.type == "sse":
|
||||
if server_config.url:
|
||||
await self.connect_mcp_server(server_config.url, server_id)
|
||||
logger.info(
|
||||
f"Connected to MCP server {server_id} at {server_config.url}"
|
||||
)
|
||||
elif server_config.type == "stdio":
|
||||
if server_config.command:
|
||||
await self.connect_mcp_server(
|
||||
server_config.command,
|
||||
server_id,
|
||||
use_stdio=True,
|
||||
stdio_args=server_config.args,
|
||||
)
|
||||
logger.info(
|
||||
f"Connected to MCP server {server_id} using command {server_config.command}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to connect to MCP server {server_id}: {e}")
|
||||
|
||||
async def connect_mcp_server(
|
||||
self,
|
||||
server_url: str,
|
||||
server_id: str = "",
|
||||
use_stdio: bool = False,
|
||||
stdio_args: List[str] = None,
|
||||
) -> None:
|
||||
"""Connect to an MCP server and add its tools."""
|
||||
if use_stdio:
|
||||
await self.mcp_clients.connect_stdio(
|
||||
server_url, stdio_args or [], server_id
|
||||
)
|
||||
self.connected_servers[server_id or server_url] = server_url
|
||||
else:
|
||||
await self.mcp_clients.connect_sse(server_url, server_id)
|
||||
self.connected_servers[server_id or server_url] = server_url
|
||||
|
||||
# Update available tools with only the new tools from this server
|
||||
new_tools = [
|
||||
tool for tool in self.mcp_clients.tools if tool.server_id == server_id
|
||||
]
|
||||
self.available_tools.add_tools(*new_tools)
|
||||
|
||||
async def disconnect_mcp_server(self, server_id: str = "") -> None:
|
||||
"""Disconnect from an MCP server and remove its tools."""
|
||||
await self.mcp_clients.disconnect(server_id)
|
||||
if server_id:
|
||||
self.connected_servers.pop(server_id, None)
|
||||
else:
|
||||
self.connected_servers.clear()
|
||||
|
||||
# Rebuild available tools without the disconnected server's tools
|
||||
base_tools = [
|
||||
tool
|
||||
for tool in self.available_tools.tools
|
||||
if not isinstance(tool, MCPClientTool)
|
||||
]
|
||||
self.available_tools = ToolCollection(*base_tools)
|
||||
self.available_tools.add_tools(*self.mcp_clients.tools)
|
||||
|
||||
async def cleanup(self):
|
||||
"""Clean up Manus agent resources."""
|
||||
if self.browser_context_helper:
|
||||
await self.browser_context_helper.cleanup_browser()
|
||||
# Disconnect from all MCP servers only if we were initialized
|
||||
if self._initialized:
|
||||
await self.disconnect_mcp_server()
|
||||
self._initialized = False
|
||||
|
||||
async def think(self) -> bool:
|
||||
"""Process current state and decide next actions with appropriate context."""
|
||||
if not self._initialized:
|
||||
await self.initialize_mcp_servers()
|
||||
self._initialized = True
|
||||
|
||||
original_prompt = self.next_step_prompt
|
||||
recent_messages = self.memory.messages[-3:] if self.memory.messages else []
|
||||
browser_in_use = any(
|
||||
@@ -64,8 +158,3 @@ class Manus(ToolCallAgent):
|
||||
self.next_step_prompt = original_prompt
|
||||
|
||||
return result
|
||||
|
||||
async def cleanup(self):
|
||||
"""Clean up Manus agent resources."""
|
||||
if self.browser_context_helper:
|
||||
await self.browser_context_helper.cleanup_browser()
|
||||
|
||||
+43
-1
@@ -1,3 +1,4 @@
|
||||
import json
|
||||
import threading
|
||||
import tomllib
|
||||
from pathlib import Path
|
||||
@@ -98,12 +99,51 @@ class SandboxSettings(BaseModel):
|
||||
)
|
||||
|
||||
|
||||
class MCPServerConfig(BaseModel):
|
||||
"""Configuration for a single MCP server"""
|
||||
|
||||
type: str = Field(..., description="Server connection type (sse or stdio)")
|
||||
url: Optional[str] = Field(None, description="Server URL for SSE connections")
|
||||
command: Optional[str] = Field(None, description="Command for stdio connections")
|
||||
args: List[str] = Field(
|
||||
default_factory=list, description="Arguments for stdio command"
|
||||
)
|
||||
|
||||
|
||||
class MCPSettings(BaseModel):
|
||||
"""Configuration for MCP (Model Context Protocol)"""
|
||||
|
||||
server_reference: str = Field(
|
||||
"app.mcp.server", description="Module reference for the MCP server"
|
||||
)
|
||||
servers: Dict[str, MCPServerConfig] = Field(
|
||||
default_factory=dict, description="MCP server configurations"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def load_server_config(cls) -> Dict[str, MCPServerConfig]:
|
||||
"""Load MCP server configuration from JSON file"""
|
||||
config_path = PROJECT_ROOT / "config" / "mcp.json"
|
||||
|
||||
try:
|
||||
config_file = config_path if config_path.exists() else None
|
||||
if not config_file:
|
||||
return {}
|
||||
|
||||
with config_file.open() as f:
|
||||
data = json.load(f)
|
||||
servers = {}
|
||||
|
||||
for server_id, server_config in data.get("mcpServers", {}).items():
|
||||
servers[server_id] = MCPServerConfig(
|
||||
type=server_config["type"],
|
||||
url=server_config.get("url"),
|
||||
command=server_config.get("command"),
|
||||
args=server_config.get("args", []),
|
||||
)
|
||||
return servers
|
||||
except Exception as e:
|
||||
raise ValueError(f"Failed to load MCP server config: {e}")
|
||||
|
||||
|
||||
class AppConfig(BaseModel):
|
||||
@@ -223,9 +263,11 @@ class Config:
|
||||
mcp_config = raw_config.get("mcp", {})
|
||||
mcp_settings = None
|
||||
if mcp_config:
|
||||
# Load server configurations from JSON
|
||||
mcp_config["servers"] = MCPSettings.load_server_config()
|
||||
mcp_settings = MCPSettings(**mcp_config)
|
||||
else:
|
||||
mcp_settings = MCPSettings()
|
||||
mcp_settings = MCPSettings(servers=MCPSettings.load_server_config())
|
||||
|
||||
config_dict = {
|
||||
"llm": {
|
||||
|
||||
+94
-42
@@ -1,5 +1,5 @@
|
||||
from contextlib import AsyncExitStack
|
||||
from typing import List, Optional
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from mcp import ClientSession, StdioServerParameters
|
||||
from mcp.client.sse import sse_client
|
||||
@@ -15,6 +15,8 @@ class MCPClientTool(BaseTool):
|
||||
"""Represents a tool proxy that can be called on the MCP server from the client side."""
|
||||
|
||||
session: Optional[ClientSession] = None
|
||||
server_id: str = "" # Add server identifier
|
||||
original_name: str = ""
|
||||
|
||||
async def execute(self, **kwargs) -> ToolResult:
|
||||
"""Execute the tool by making a remote call to the MCP server."""
|
||||
@@ -22,7 +24,8 @@ class MCPClientTool(BaseTool):
|
||||
return ToolResult(error="Not connected to MCP server")
|
||||
|
||||
try:
|
||||
result = await self.session.call_tool(self.name, kwargs)
|
||||
logger.info(f"Executing tool: {self.original_name}")
|
||||
result = await self.session.call_tool(self.original_name, kwargs)
|
||||
content_str = ", ".join(
|
||||
item.text for item in result.content if isinstance(item, TextContent)
|
||||
)
|
||||
@@ -33,83 +36,132 @@ class MCPClientTool(BaseTool):
|
||||
|
||||
class MCPClients(ToolCollection):
|
||||
"""
|
||||
A collection of tools that connects to an MCP server and manages available tools through the Model Context Protocol.
|
||||
A collection of tools that connects to multiple MCP servers and manages available tools through the Model Context Protocol.
|
||||
"""
|
||||
|
||||
session: Optional[ClientSession] = None
|
||||
exit_stack: AsyncExitStack = None
|
||||
sessions: Dict[str, ClientSession] = {}
|
||||
exit_stacks: Dict[str, AsyncExitStack] = {}
|
||||
description: str = "MCP client tools for server interaction"
|
||||
|
||||
def __init__(self):
|
||||
super().__init__() # Initialize with empty tools list
|
||||
self.name = "mcp" # Keep name for backward compatibility
|
||||
self.exit_stack = AsyncExitStack()
|
||||
|
||||
async def connect_sse(self, server_url: str) -> None:
|
||||
async def connect_sse(self, server_url: str, server_id: str = "") -> None:
|
||||
"""Connect to an MCP server using SSE transport."""
|
||||
if not server_url:
|
||||
raise ValueError("Server URL is required.")
|
||||
if self.session:
|
||||
await self.disconnect()
|
||||
|
||||
server_id = server_id or server_url
|
||||
|
||||
# Always ensure clean disconnection before new connection
|
||||
if server_id in self.sessions:
|
||||
await self.disconnect(server_id)
|
||||
|
||||
exit_stack = AsyncExitStack()
|
||||
self.exit_stacks[server_id] = exit_stack
|
||||
|
||||
streams_context = sse_client(url=server_url)
|
||||
streams = await self.exit_stack.enter_async_context(streams_context)
|
||||
self.session = await self.exit_stack.enter_async_context(
|
||||
ClientSession(*streams)
|
||||
)
|
||||
streams = await exit_stack.enter_async_context(streams_context)
|
||||
session = await exit_stack.enter_async_context(ClientSession(*streams))
|
||||
self.sessions[server_id] = session
|
||||
|
||||
await self._initialize_and_list_tools()
|
||||
await self._initialize_and_list_tools(server_id)
|
||||
|
||||
async def connect_stdio(self, command: str, args: List[str]) -> None:
|
||||
async def connect_stdio(
|
||||
self, command: str, args: List[str], server_id: str = ""
|
||||
) -> None:
|
||||
"""Connect to an MCP server using stdio transport."""
|
||||
if not command:
|
||||
raise ValueError("Server command is required.")
|
||||
if self.session:
|
||||
await self.disconnect()
|
||||
|
||||
server_id = server_id or command
|
||||
|
||||
# Always ensure clean disconnection before new connection
|
||||
if server_id in self.sessions:
|
||||
await self.disconnect(server_id)
|
||||
|
||||
exit_stack = AsyncExitStack()
|
||||
self.exit_stacks[server_id] = exit_stack
|
||||
|
||||
server_params = StdioServerParameters(command=command, args=args)
|
||||
stdio_transport = await self.exit_stack.enter_async_context(
|
||||
stdio_transport = await exit_stack.enter_async_context(
|
||||
stdio_client(server_params)
|
||||
)
|
||||
read, write = stdio_transport
|
||||
self.session = await self.exit_stack.enter_async_context(
|
||||
ClientSession(read, write)
|
||||
)
|
||||
session = await exit_stack.enter_async_context(ClientSession(read, write))
|
||||
self.sessions[server_id] = session
|
||||
|
||||
await self._initialize_and_list_tools()
|
||||
await self._initialize_and_list_tools(server_id)
|
||||
|
||||
async def _initialize_and_list_tools(self) -> None:
|
||||
async def _initialize_and_list_tools(self, server_id: str) -> None:
|
||||
"""Initialize session and populate tool map."""
|
||||
if not self.session:
|
||||
raise RuntimeError("Session not initialized.")
|
||||
session = self.sessions.get(server_id)
|
||||
if not session:
|
||||
raise RuntimeError(f"Session not initialized for server {server_id}")
|
||||
|
||||
await self.session.initialize()
|
||||
response = await self.session.list_tools()
|
||||
|
||||
# Clear existing tools
|
||||
self.tools = tuple()
|
||||
self.tool_map = {}
|
||||
await session.initialize()
|
||||
response = await session.list_tools()
|
||||
|
||||
# Create proper tool objects for each server tool
|
||||
for tool in response.tools:
|
||||
original_name = tool.name
|
||||
# Always prefix with server_id to ensure uniqueness
|
||||
tool_name = f"mcp_{server_id}_{original_name}"
|
||||
|
||||
server_tool = MCPClientTool(
|
||||
name=tool.name,
|
||||
name=tool_name,
|
||||
description=tool.description,
|
||||
parameters=tool.inputSchema,
|
||||
session=self.session,
|
||||
session=session,
|
||||
server_id=server_id,
|
||||
original_name=original_name,
|
||||
)
|
||||
self.tool_map[tool.name] = server_tool
|
||||
self.tool_map[tool_name] = server_tool
|
||||
|
||||
# Update tools tuple
|
||||
self.tools = tuple(self.tool_map.values())
|
||||
logger.info(
|
||||
f"Connected to server with tools: {[tool.name for tool in response.tools]}"
|
||||
f"Connected to server {server_id} with tools: {[tool.name for tool in response.tools]}"
|
||||
)
|
||||
|
||||
async def disconnect(self) -> None:
|
||||
"""Disconnect from the MCP server and clean up resources."""
|
||||
if self.session and self.exit_stack:
|
||||
await self.exit_stack.aclose()
|
||||
self.session = None
|
||||
self.tools = tuple()
|
||||
async def disconnect(self, server_id: str = "") -> None:
|
||||
"""Disconnect from a specific MCP server or all servers if no server_id provided."""
|
||||
if server_id:
|
||||
if server_id in self.sessions:
|
||||
try:
|
||||
exit_stack = self.exit_stacks.get(server_id)
|
||||
|
||||
# Close the exit stack which will handle session cleanup
|
||||
if exit_stack:
|
||||
try:
|
||||
await exit_stack.aclose()
|
||||
except RuntimeError as e:
|
||||
if "cancel scope" in str(e).lower():
|
||||
logger.warning(
|
||||
f"Cancel scope error during disconnect from {server_id}, continuing with cleanup: {e}"
|
||||
)
|
||||
else:
|
||||
raise
|
||||
|
||||
# Clean up references
|
||||
self.sessions.pop(server_id, None)
|
||||
self.exit_stacks.pop(server_id, None)
|
||||
|
||||
# Remove tools associated with this server
|
||||
self.tool_map = {
|
||||
k: v
|
||||
for k, v in self.tool_map.items()
|
||||
if v.server_id != server_id
|
||||
}
|
||||
self.tools = tuple(self.tool_map.values())
|
||||
logger.info(f"Disconnected from MCP server {server_id}")
|
||||
except Exception as e:
|
||||
logger.error(f"Error disconnecting from server {server_id}: {e}")
|
||||
else:
|
||||
# Disconnect from all servers in a deterministic order
|
||||
for sid in sorted(list(self.sessions.keys())):
|
||||
await self.disconnect(sid)
|
||||
self.tool_map = {}
|
||||
logger.info("Disconnected from MCP server")
|
||||
self.tools = tuple()
|
||||
logger.info("Disconnected from all MCP servers")
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from app.exceptions import ToolError
|
||||
from app.logger import logger
|
||||
from app.tool.base import BaseTool, ToolFailure, ToolResult
|
||||
|
||||
|
||||
@@ -48,11 +49,23 @@ class ToolCollection:
|
||||
return self.tool_map.get(name)
|
||||
|
||||
def add_tool(self, tool: BaseTool):
|
||||
"""Add a single tool to the collection.
|
||||
|
||||
If a tool with the same name already exists, it will be skipped and a warning will be logged.
|
||||
"""
|
||||
if tool.name in self.tool_map:
|
||||
logger.warning(f"Tool {tool.name} already exists in collection, skipping")
|
||||
return self
|
||||
|
||||
self.tools += (tool,)
|
||||
self.tool_map[tool.name] = tool
|
||||
return self
|
||||
|
||||
def add_tools(self, *tools: BaseTool):
|
||||
"""Add multiple tools to the collection.
|
||||
|
||||
If any tool has a name conflict with an existing tool, it will be skipped and a warning will be logged.
|
||||
"""
|
||||
for tool in tools:
|
||||
self.add_tool(tool)
|
||||
return self
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
{
|
||||
"mcpServers": {
|
||||
"server1": {
|
||||
"type": "sse",
|
||||
"url": "http://localhost:8000/sse"
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user