diff --git a/reme/core/op/mcp_tool.py b/reme/core/op/mcp_tool.py index c59e105c..68e78946 100644 --- a/reme/core/op/mcp_tool.py +++ b/reme/core/op/mcp_tool.py @@ -1,7 +1,5 @@ """MCP (Model Context Protocol) tool integration for remote tool execution.""" -from typing import List - from mcp.types import CallToolResult, TextContent from .base_tool import BaseTool @@ -16,9 +14,9 @@ class MCPTool(BaseTool): self, mcp_server: str = "", tool_name: str = "", - parameter_required: List[str] | None = None, - parameter_optional: List[str] | None = None, - parameter_deleted: List[str] | None = None, + 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, @@ -28,9 +26,9 @@ class MCPTool(BaseTool): 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.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 diff --git a/reme/core/schema/tool_call.py b/reme/core/schema/tool_call.py index e7ba9b78..22577817 100644 --- a/reme/core/schema/tool_call.py +++ b/reme/core/schema/tool_call.py @@ -1,7 +1,7 @@ """MCP Tool Schema definitions for recursive JSON Schema representation.""" import json -from typing import Any, Dict, List, Optional, Union +from typing import Any, Union, Optional from mcp.types import Tool from pydantic import BaseModel, ConfigDict, Field, model_validator, field_validator @@ -16,10 +16,10 @@ class ToolAttr(BaseModel): type: str = Field(default=str(JsonSchemaEnum.STRING), description="The data type of the attribute") description: Optional[str] = Field(default=None, description="Description of the attribute") - required: Optional[List[str]] = Field(default=None, description="Required property names for object types") - properties: Optional[Dict[str, "ToolAttr"]] = Field(default=None, description="Child properties for objects") - items: Optional[Union[Dict[str, Any], "ToolAttr"]] = Field(default=None, description="Schema for array items") - enum: Optional[List[str]] = Field(default=None, description="Allowed values for the attribute") + required: Optional[list[str]] = Field(default=None, description="Required property names for object types") + properties: Optional[dict[str, "ToolAttr"]] = Field(default=None, description="Child properties for objects") + items: Optional[Union[dict[str, Any], "ToolAttr"]] = Field(default=None, description="Schema for array items") + enum: Optional[list[str]] = Field(default=None, description="Allowed values for the attribute") @field_validator("type") @classmethod @@ -129,9 +129,13 @@ class ToolCall(BaseModel): return data - def simple_input_dump(self) -> dict: - """Returns a standardized tool definition dictionary.""" - return { + def simple_input_dump(self, as_dict: bool = True) -> dict | str: + """Returns a standardized tool definition dictionary or JSON string. + + Args: + as_dict: If True, returns dict; if False, returns JSON string. + """ + result = { "type": self.type, self.type: { "name": self.name, @@ -139,10 +143,15 @@ class ToolCall(BaseModel): "parameters": self.parameters.simple_input_dump(), }, } + return result if as_dict else json.dumps(result) - def simple_output_dump(self) -> dict: - """Convert ToolCall to output format dictionary for API responses.""" - return { + def simple_output_dump(self, as_dict: bool = True) -> dict | str: + """Convert ToolCall to output format dictionary or JSON string for API responses. + + Args: + as_dict: If True, returns dict; if False, returns JSON string. + """ + result = { "index": self.index, "id": self.id, self.type: { @@ -151,6 +160,7 @@ class ToolCall(BaseModel): }, "type": self.type, } + return result if as_dict else json.dumps(result) @property def argument_dict(self) -> dict: diff --git a/reme/core/schema/vector_node.py b/reme/core/schema/vector_node.py index 937ef4be..8ee55649 100644 --- a/reme/core/schema/vector_node.py +++ b/reme/core/schema/vector_node.py @@ -1,6 +1,5 @@ """Defines the data structure for individual vector embedding nodes within a retrieval system.""" -from typing import List, Dict from uuid import uuid4 from pydantic import BaseModel, Field @@ -11,5 +10,5 @@ class VectorNode(BaseModel): vector_id: str = Field(default_factory=lambda: uuid4().hex) content: str = Field(default="") - vector: List[float] | None = Field(default=None) - metadata: Dict[str, str | bool | int | float] = Field(default_factory=dict) + vector: list[float] | None = Field(default=None) + metadata: dict[str, str | bool | int | float] = Field(default_factory=dict) diff --git a/reme/core/utils/http_client.py b/reme/core/utils/http_client.py index 8c8e92c3..99c94306 100644 --- a/reme/core/utils/http_client.py +++ b/reme/core/utils/http_client.py @@ -2,7 +2,6 @@ import json from collections.abc import AsyncIterator -from typing import Optional import httpx from loguru import logger @@ -45,7 +44,7 @@ class HttpClient: response.raise_for_status() return response.json() - async def execute_flow(self, flow_name: str, **kwargs) -> Optional[Response]: + async def execute_flow(self, flow_name: str, **kwargs) -> Response | None: """Execute a flow with automated retry logic.""" endpoint = f"{self.base_url}/{flow_name}"