mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-14 23:21:04 +00:00
feat(core): add MCP tool integration and Ray-based parallel operations
This commit is contained in:
parent
1e7b8fbdad
commit
472e069bc5
5 changed files with 215 additions and 5 deletions
|
|
@ -21,7 +21,7 @@ class ServiceContext(BaseContext):
|
|||
self.language: str = ""
|
||||
self.thread_pool: ThreadPoolExecutor | None = None
|
||||
self.vector_store_dict: dict[str, dict] = {}
|
||||
self.external_mcp_tool_call_dict: dict = {}
|
||||
self.mcp_server_tool_call_mapping: dict = {}
|
||||
# Initialize a registry for every category defined in RegistryEnum
|
||||
self.registry_dict: dict[RegistryEnum, Registry] = {v: Registry() for v in RegistryEnum.__members__.values()}
|
||||
self.flow_dict: dict = {}
|
||||
|
|
|
|||
124
reme_ai/core/op/base_ray_op.py
Normal file
124
reme_ai/core/op/base_ray_op.py
Normal file
|
|
@ -0,0 +1,124 @@
|
|||
"""Base class for Ray-based parallel operations."""
|
||||
|
||||
from abc import ABCMeta
|
||||
from typing import Callable
|
||||
|
||||
import pandas as pd
|
||||
from loguru import logger
|
||||
from tqdm import tqdm
|
||||
|
||||
from .base_op import BaseOp
|
||||
from ..context import BaseContext, C
|
||||
|
||||
_RAY_IMPORT_ERROR = None
|
||||
|
||||
try:
|
||||
import ray
|
||||
except ImportError as e:
|
||||
_RAY_IMPORT_ERROR = e
|
||||
ray = None
|
||||
|
||||
|
||||
class BaseRayOp(BaseOp, metaclass=ABCMeta):
|
||||
"""Base class for Ray-based parallel operations."""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
if _RAY_IMPORT_ERROR:
|
||||
raise ImportError("Ray requires extra dependencies. Install with `pip install ray`")
|
||||
|
||||
super().__init__(**kwargs)
|
||||
self._ray_task_list: list = []
|
||||
|
||||
def submit_and_join_parallel_op(self, op: BaseOp, **kwargs) -> list:
|
||||
"""Submit a BaseOp to be executed in parallel via Ray."""
|
||||
return self.submit_and_join_ray_task(fn=op.call, task_desc=op.name, context=self.context, **kwargs)
|
||||
|
||||
def submit_and_join_ray_task(self, fn: Callable, parallel_key: str = "", task_desc: str = "", **kwargs) -> list:
|
||||
"""Divide data into chunks and execute them across Ray workers."""
|
||||
max_workers = C.service_config.ray_max_workers
|
||||
self._ray_task_list.clear()
|
||||
|
||||
# Automatically detect the key containing the list to parallelize
|
||||
if not parallel_key:
|
||||
for key, value in kwargs.items():
|
||||
if isinstance(value, list):
|
||||
parallel_key = key
|
||||
break
|
||||
|
||||
if not parallel_key:
|
||||
raise ValueError("No list found in kwargs to parallelize over.")
|
||||
|
||||
parallel_list = kwargs.pop(parallel_key)
|
||||
logger.info(f"Parallelizing '{parallel_key}' across {max_workers} workers")
|
||||
|
||||
# Put large shared objects into the Ray Object Store once
|
||||
optimized_kwargs = {
|
||||
k: (ray.put(v) if isinstance(v, (pd.DataFrame, pd.Series, dict, list, BaseContext)) else v)
|
||||
for k, v in kwargs.items()
|
||||
}
|
||||
|
||||
# Submit sliced chunks to reduce inter-node data transfer
|
||||
remote_task_loop = ray.remote(self._ray_task_loop)
|
||||
for i in range(max_workers):
|
||||
chunk = parallel_list[i::max_workers]
|
||||
if not chunk:
|
||||
continue
|
||||
|
||||
task = remote_task_loop.remote(
|
||||
fn,
|
||||
parallel_key,
|
||||
chunk,
|
||||
i,
|
||||
**optimized_kwargs,
|
||||
)
|
||||
self._ray_task_list.append(task)
|
||||
logger.info(f"Submitted task {i + 1}/{max_workers} for {task_desc}")
|
||||
|
||||
return self.join_ray_task(task_desc=task_desc)
|
||||
|
||||
@staticmethod
|
||||
def _ray_task_loop(internal_fn: Callable, parallel_key: str, chunk: list, actor_index: int, **kwargs) -> list:
|
||||
"""Execute the function over a specific chunk of data on a worker."""
|
||||
results = []
|
||||
for value in chunk:
|
||||
current_kwargs = {**kwargs, "actor_index": actor_index, parallel_key: value}
|
||||
t_result = internal_fn(**current_kwargs)
|
||||
|
||||
if t_result is not None:
|
||||
if isinstance(t_result, list):
|
||||
results.extend(t_result)
|
||||
else:
|
||||
results.append(t_result)
|
||||
return results
|
||||
|
||||
def submit_ray_task(self, fn, *args, **kwargs):
|
||||
"""Submit a single Ray task to the task list for later execution."""
|
||||
if not ray.is_initialized():
|
||||
ray.init(num_cpus=C.service_config.ray_max_workers, ignore_reinit_error=True)
|
||||
|
||||
remote_fn = ray.remote(fn)
|
||||
task = remote_fn.remote(*args, **kwargs)
|
||||
self._ray_task_list.append(task)
|
||||
return self
|
||||
|
||||
def join_ray_task(self, task_desc: str | None = None) -> list:
|
||||
"""Collect results from Ray workers using a progress bar."""
|
||||
results = []
|
||||
unfinished = list(self._ray_task_list)
|
||||
|
||||
with tqdm(total=len(unfinished), desc=task_desc or f"{self.name}_ray") as pbar:
|
||||
while unfinished:
|
||||
ready, unfinished = ray.wait(unfinished, num_returns=1)
|
||||
for obj_ref in ready:
|
||||
try:
|
||||
t_result = ray.get(obj_ref)
|
||||
if isinstance(t_result, list):
|
||||
results.extend(t_result)
|
||||
elif t_result is not None:
|
||||
results.append(t_result)
|
||||
except Exception as e:
|
||||
logger.error(f"Worker task failed: {e}")
|
||||
pbar.update(1)
|
||||
|
||||
self._ray_task_list.clear()
|
||||
return results
|
||||
|
|
@ -98,10 +98,7 @@ class ServiceConfig(BaseModel):
|
|||
ray_max_workers: int = Field(default=-1)
|
||||
disabled_flows: List[str] = Field(default_factory=list)
|
||||
enabled_flows: List[str] = Field(default_factory=list)
|
||||
external_mcp: Dict[str, dict] = Field(
|
||||
default_factory=dict,
|
||||
description="External MCP Server configuration",
|
||||
)
|
||||
mcp_servers: Dict[str, dict] = Field(default_factory=dict, description="External MCP Server configuration")
|
||||
|
||||
mcp: MCPConfig = Field(default_factory=MCPConfig)
|
||||
http: HttpConfig = Field(default_factory=HttpConfig)
|
||||
|
|
|
|||
7
reme_ai/core/tool/__init__.py
Normal file
7
reme_ai/core/tool/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""tool"""
|
||||
|
||||
from .mcp_tool import MCPTool
|
||||
|
||||
__all__ = [
|
||||
"MCPTool",
|
||||
]
|
||||
82
reme_ai/core/tool/mcp_tool.py
Normal file
82
reme_ai/core/tool/mcp_tool.py
Normal file
|
|
@ -0,0 +1,82 @@
|
|||
"""MCP (Model Context Protocol) tool integration for remote tool execution."""
|
||||
|
||||
from typing import List
|
||||
|
||||
from ..context import C
|
||||
from ..op import BaseOp
|
||||
from ..schema import ToolCall
|
||||
from ..utils import MCPClient
|
||||
|
||||
|
||||
@C.register_op()
|
||||
class MCPTool(BaseOp):
|
||||
"""Operator for calling remote MCP (Model Context Protocol) tools.
|
||||
|
||||
This class enables integration with external MCP servers to execute tools
|
||||
and retrieve their results. It supports parameter customization and retry logic.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
mcp_server: str = "",
|
||||
tool_name: str = "",
|
||||
enable_tool_response: bool = True,
|
||||
parameter_required: List[str] | None = None,
|
||||
parameter_optional: List[str] | None = None,
|
||||
parameter_deleted: List[str] | None = None,
|
||||
max_retries: int = 3,
|
||||
timeout: float | None = None,
|
||||
raise_exception: bool = False,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
super().__init__(
|
||||
enable_tool_response=enable_tool_response,
|
||||
max_retries=max_retries,
|
||||
raise_exception=raise_exception,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
self.mcp_server: str = mcp_server
|
||||
self.tool_name: str = tool_name
|
||||
self.parameter_required: List[str] | None = parameter_required
|
||||
self.parameter_optional: List[str] | None = parameter_optional
|
||||
self.parameter_deleted: List[str] | None = parameter_deleted
|
||||
self.timeout: float | None = timeout
|
||||
# Example MCP marketplace: https://bailian.console.aliyun.com/?tab=mcp#/mcp-market
|
||||
|
||||
self._client = MCPClient(C.service_config.mcp_servers)
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
tool_call_dict = C.mcp_server_tool_call_mapping[self.mcp_server]
|
||||
tool_call: ToolCall = tool_call_dict[self.tool_name].model_copy(deep=True)
|
||||
|
||||
# Initialize required list if not exists
|
||||
if tool_call.parameters.required is None:
|
||||
tool_call.parameters.required = []
|
||||
|
||||
if self.parameter_required:
|
||||
for name in self.parameter_required:
|
||||
if name not in tool_call.parameters.required:
|
||||
tool_call.parameters.required.append(name)
|
||||
|
||||
if self.parameter_optional:
|
||||
for name in self.parameter_optional:
|
||||
if name in tool_call.parameters.required:
|
||||
tool_call.parameters.required.remove(name)
|
||||
|
||||
if self.parameter_deleted:
|
||||
for name in self.parameter_deleted:
|
||||
tool_call.parameters.properties.pop(name, None)
|
||||
if tool_call.parameters.required and name in tool_call.parameters.required:
|
||||
tool_call.parameters.required.remove(name)
|
||||
|
||||
return tool_call
|
||||
|
||||
async def execute(self):
|
||||
self.output = await self._client.call_tool(
|
||||
server_name=self.mcp_server,
|
||||
tool_name=self.tool_name,
|
||||
arguments=self.input_dict,
|
||||
parse_text_result=True,
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue