feat(core): add MCP tool integration and Ray-based parallel operations

This commit is contained in:
jinli.yl 2025-12-31 23:43:56 +08:00
parent 1e7b8fbdad
commit 472e069bc5
5 changed files with 215 additions and 5 deletions

View file

@ -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 = {}

View 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

View file

@ -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)

View file

@ -0,0 +1,7 @@
"""tool"""
from .mcp_tool import MCPTool
__all__ = [
"MCPTool",
]

View 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,
)