mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-09 03:20:54 +00:00
feat(core): add service metadata and prompt formatting capabilities to base operator
This commit is contained in:
parent
90a53737c8
commit
1e7b8fbdad
3 changed files with 134 additions and 3 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue