From 472e069bc51428b452eae6aabd18820396b319ce Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 31 Dec 2025 23:43:56 +0800 Subject: [PATCH] feat(core): add MCP tool integration and Ray-based parallel operations --- reme_ai/core/context/service_context.py | 2 +- reme_ai/core/op/base_ray_op.py | 124 ++++++++++++++++++++++++ reme_ai/core/schema/service_config.py | 5 +- reme_ai/core/tool/__init__.py | 7 ++ reme_ai/core/tool/mcp_tool.py | 82 ++++++++++++++++ 5 files changed, 215 insertions(+), 5 deletions(-) create mode 100644 reme_ai/core/op/base_ray_op.py create mode 100644 reme_ai/core/tool/__init__.py create mode 100644 reme_ai/core/tool/mcp_tool.py diff --git a/reme_ai/core/context/service_context.py b/reme_ai/core/context/service_context.py index 65702be1..97426594 100644 --- a/reme_ai/core/context/service_context.py +++ b/reme_ai/core/context/service_context.py @@ -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 = {} diff --git a/reme_ai/core/op/base_ray_op.py b/reme_ai/core/op/base_ray_op.py new file mode 100644 index 00000000..88cc2985 --- /dev/null +++ b/reme_ai/core/op/base_ray_op.py @@ -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 diff --git a/reme_ai/core/schema/service_config.py b/reme_ai/core/schema/service_config.py index 67cac68f..023eb3a7 100644 --- a/reme_ai/core/schema/service_config.py +++ b/reme_ai/core/schema/service_config.py @@ -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) diff --git a/reme_ai/core/tool/__init__.py b/reme_ai/core/tool/__init__.py new file mode 100644 index 00000000..02dbd468 --- /dev/null +++ b/reme_ai/core/tool/__init__.py @@ -0,0 +1,7 @@ +"""tool""" + +from .mcp_tool import MCPTool + +__all__ = [ + "MCPTool", +] diff --git a/reme_ai/core/tool/mcp_tool.py b/reme_ai/core/tool/mcp_tool.py new file mode 100644 index 00000000..f9afa029 --- /dev/null +++ b/reme_ai/core/tool/mcp_tool.py @@ -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, + )