diff --git a/reme_ai/core/op/base_op.py b/reme_ai/core/op/base_op.py index 7a131a28..e5f8450e 100644 --- a/reme_ai/core/op/base_op.py +++ b/reme_ai/core/op/base_op.py @@ -194,6 +194,11 @@ class BaseOp: ) return self._token_counter + @property + def service_metadata(self) -> dict: + """Get service configuration metadata.""" + return C.service_config.model_extra + async def before_execute(self): """Prepare context and validate before async execution.""" self.context.apply_mapping(self.input_mapping) @@ -334,3 +339,19 @@ class BaseOp: par = ParallelOp(sub_ops=[self], async_mode=self.async_mode) par.add_sub_ops(op.sub_ops if isinstance(op, ParallelOp) else op) return par + + def prompt_format(self, prompt_name: str, **kwargs) -> str: + """Format a prompt template with provided keyword arguments.""" + return self.prompt.prompt_format(prompt_name=prompt_name, **kwargs) + + def get_prompt(self, prompt_name: str) -> str: + """Get a prompt template by name.""" + return self.prompt.get_prompt(prompt_name=prompt_name) + + def copy(self, **kwargs): + """Create a copy of this operator with optional parameter overrides.""" + copy_op = self.__class__(*self._init_args, **self._init_kwargs, **kwargs) + if self.sub_ops: + copy_op.sub_ops.clear() + copy_op.add_sub_ops(self.sub_ops) + return copy_op diff --git a/reme_ai/core/utils/mcp_client.py b/reme_ai/core/utils/mcp_client.py index 3a1c70e7..4c13be9f 100644 --- a/reme_ai/core/utils/mcp_client.py +++ b/reme_ai/core/utils/mcp_client.py @@ -9,7 +9,7 @@ from mcp import ClientSession, StdioServerParameters, Tool from mcp.client.sse import sse_client from mcp.client.stdio import stdio_client from mcp.client.streamable_http import streamablehttp_client -from mcp.types import CallToolResult +from mcp.types import CallToolResult, TextContent from ..schema import ToolCall @@ -101,7 +101,21 @@ class MCPClient: return tool_calls - async def call_tool(self, server_name: str, tool_name: str, arguments: dict[str, Any]) -> CallToolResult: + async def call_tool( + self, + server_name: str, + tool_name: str, + arguments: dict[str, Any], + parse_text_result: bool = False, + ) -> CallToolResult | str: """Execute a tool on a specific server.""" async with self.connect_to_server(server_name) as session: - return await session.call_tool(tool_name, arguments) + tool_results: CallToolResult = await session.call_tool(tool_name, arguments) + if not parse_text_result: + return tool_results + + text_result = [] + for block in tool_results.content: + if isinstance(block, TextContent): + text_result.append(block.text) + return "\n".join(text_result) diff --git a/tests/test_mcp_client.py b/tests/test_mcp_client.py index 7db2ad61..d2fef40e 100644 --- a/tests/test_mcp_client.py +++ b/tests/test_mcp_client.py @@ -1,5 +1,7 @@ """Test module for demonstrating MCPClient functionality.""" +# pylint: disable=too-many-return-statements,too-many-statements + import asyncio import json @@ -20,11 +22,105 @@ async def main(): client = MCPClient(config_data) try: + # List all available tools + print("=" * 50) + print("Listing available tools:") + print("=" * 50) t_list = await client.list_tool_calls(test_mcp) for t in t_list: print(json.dumps(t, ensure_ascii=False, indent=2)) + + # Helper function to build default values + def build_default_value(param_info: dict) -> any: + """Build a default value for a parameter based on its schema.""" + param_type = param_info.get("type", "string") + + # Handle enum types - use the first enum value + if "enum" in param_info and param_info["enum"]: + return param_info["enum"][0] + + # Handle different types + if param_type == "string": + return "example_string" + elif param_type == "number": + return 0.0 + elif param_type == "integer": + return 0 + elif param_type == "boolean": + return False + elif param_type == "array": + return [] + elif param_type == "object": + # Recursively build nested objects + obj = {} + nested_properties = param_info.get("properties", {}) + nested_required = param_info.get("required", []) + + for nested_param_name in nested_required: + if nested_param_name in nested_properties: + nested_param_info = nested_properties[nested_param_name] + obj[nested_param_name] = build_default_value(nested_param_info) + + return obj + else: + return None + + # Call tools if available + if t_list: + # Execute the first two tools (or fewer if not enough tools available) + tools_to_execute = min(2, len(t_list)) + + for idx in range(tools_to_execute): + print("\n" + "=" * 50) + print(f"Calling tool #{idx + 1}:") + print("=" * 50) + + # Get the tool's information + current_tool = t_list[idx] + tool_type = current_tool.get("type", "function") + tool_body = current_tool.get(tool_type, {}) + + tool_name = tool_body.get("name") + tool_description = tool_body.get("description", "") + + # Prepare arguments based on the tool's input schema + tool_arguments = {} + parameters = tool_body.get("parameters", {}) + properties = parameters.get("properties", {}) + required = parameters.get("required", []) + + # Build minimal arguments for required parameters + for param_name in required: + if param_name in properties: + param_info = properties[param_name] + tool_arguments[param_name] = build_default_value(param_info) + + # Validate tool name exists + if not tool_name: + print("Error: Tool name not found in the tool definition") + print(f"Tool structure: {json.dumps(current_tool, ensure_ascii=False, indent=2)}") + else: + print(f"Tool name: {tool_name}") + print(f"Tool description: {tool_description}") + print(f"Arguments: {json.dumps(tool_arguments, ensure_ascii=False, indent=2)}") + + # Call the tool + result = await client.call_tool(test_mcp, tool_name, tool_arguments, parse_text_result=True) + + print("\n" + "-" * 50) + print("Tool call result:") + print("-" * 50) + print(f"Content: {result}") + if hasattr(result, "isError"): + print(f"Is Error: {result.isError}") + else: + print("\nNo tools available to call.") + except Exception as e: print(f"Error occurred: {e}") + import traceback + + traceback.print_exc() if __name__ == "__main__":