From 90a53737c85784766b61d8d501db88662a2b8a0a Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 31 Dec 2025 17:16:53 +0800 Subject: [PATCH] feat(core): add operator framework and utility modules --- reme_ai/core/context/runtime_context.py | 86 +++-- reme_ai/core/enumeration/json_schema_enum.py | 17 +- reme_ai/core/op/__init__.py | 11 + reme_ai/core/op/base_op.py | 336 +++++++++++++++++++ reme_ai/core/op/parallel_op.py | 33 ++ reme_ai/core/op/sequential_op.py | 31 ++ reme_ai/core/schema/request.py | 6 - reme_ai/core/schema/tool_call.py | 4 +- reme_ai/core/utils/__init__.py | 8 + reme_ai/core/utils/cache_handler.py | 184 ++++++++++ reme_ai/core/utils/http_client.py | 91 +++++ reme_ai/core/utils/mcp_client.py | 107 ++++++ reme_ai/core/utils/pydantic_utils.py | 66 ++++ tests/mcp_servers_demo.json | 37 ++ tests/test_cache_handler.py | 94 ++++++ tests/test_mcp_client.py | 31 ++ tests/test_mcp_server.py | 125 +++++++ 17 files changed, 1220 insertions(+), 47 deletions(-) create mode 100644 reme_ai/core/op/__init__.py create mode 100644 reme_ai/core/op/base_op.py create mode 100644 reme_ai/core/op/parallel_op.py create mode 100644 reme_ai/core/op/sequential_op.py create mode 100644 reme_ai/core/utils/cache_handler.py create mode 100644 reme_ai/core/utils/http_client.py create mode 100644 reme_ai/core/utils/mcp_client.py create mode 100644 reme_ai/core/utils/pydantic_utils.py create mode 100644 tests/mcp_servers_demo.json create mode 100644 tests/test_cache_handler.py create mode 100644 tests/test_mcp_client.py create mode 100644 tests/test_mcp_server.py diff --git a/reme_ai/core/context/runtime_context.py b/reme_ai/core/context/runtime_context.py index ded35ded..b2362161 100644 --- a/reme_ai/core/context/runtime_context.py +++ b/reme_ai/core/context/runtime_context.py @@ -1,15 +1,14 @@ -"""Module providing a runtime context for managing response states and asynchronous data streaming.""" +"""Runtime context for managing response states and asynchronous data streaming.""" import asyncio from .base_context import BaseContext from ..enumeration import ChunkEnum -from ..schema import Response -from ..schema import StreamChunk +from ..schema import Response, StreamChunk class RuntimeContext(BaseContext): - """A context class for handling execution state, including response metadata and stream queues.""" + """Context for execution state, response metadata, and stream queues.""" def __init__( self, @@ -17,40 +16,67 @@ class RuntimeContext(BaseContext): stream_queue: asyncio.Queue | None = None, **kwargs, ): - """Initialize the runtime context with optional response objects and message queues.""" + """Initialize the context with optional response and queue.""" super().__init__(**kwargs) + self.response = response or Response() + self.stream_queue = stream_queue - self.response: Response | None = response if response is not None else Response() - self.stream_queue: asyncio.Queue | None = stream_queue + @classmethod + def from_context(cls, context: "RuntimeContext | None" = None, **kwargs) -> "RuntimeContext": + """Create a new context from an existing instance or keywords.""" + if context is None: + return cls(**kwargs) - async def add_stream_string_and_type(self, chunk: str, chunk_type: ChunkEnum): - """Create and enqueue a stream chunk from a raw string and specific type.""" - if self.stream_queue is None: - return self + new_instance = cls(response=context.response, stream_queue=context.stream_queue) + new_instance.update(context) + if kwargs: + new_instance.update(kwargs) + return new_instance - # Package raw data into a StreamChunk schema - stream_chunk = StreamChunk(chunk_type=chunk_type, chunk=chunk) - await self.stream_queue.put(stream_chunk) + async def _enqueue(self, chunk: StreamChunk) -> None: + """Internal helper to put a chunk into the queue if it exists.""" + if self.stream_queue: + await self.stream_queue.put(chunk) + + async def add_stream_string(self, chunk: str, chunk_type: ChunkEnum) -> "RuntimeContext": + """Enqueue a stream chunk from a raw string and type.""" + await self._enqueue(StreamChunk(chunk_type=chunk_type, chunk=chunk)) return self - async def add_stream_chunk(self, stream_chunk: StreamChunk): - """Directly enqueue an existing stream chunk into the stream queue.""" - if self.stream_queue is None: - return self - await self.stream_queue.put(stream_chunk) + async def add_stream_chunk(self, stream_chunk: StreamChunk) -> "RuntimeContext": + """Enqueue an existing stream chunk.""" + await self._enqueue(stream_chunk) return self - async def add_stream_done(self): - """Enqueue a termination chunk to signal the end of the data stream.""" - if self.stream_queue is None: - return self - - # Create a special chunk representing the completion state - done_chunk = StreamChunk(chunk_type=ChunkEnum.DONE, chunk="", done=True) - await self.stream_queue.put(done_chunk) + async def add_stream_done(self) -> "RuntimeContext": + """Enqueue a termination chunk to signal the end of the stream.""" + await self._enqueue(StreamChunk(chunk_type=ChunkEnum.DONE, chunk="", done=True)) return self - def add_response_error(self, e: Exception): - """Update the internal response object to reflect a failure state using exception details.""" + def add_response_error(self, e: Exception) -> "RuntimeContext": + """Record an exception into the response object.""" self.response.success = False - self.response.answer = str(e.args) + self.response.answer = str(e) + return self + + def apply_mapping(self, mapping: dict[str, str]) -> "RuntimeContext": + """Copy internal values based on a source-to-target key map.""" + if not mapping: + return self + + for source, target in mapping.items(): + if source in self: + self[target] = self[source] + return self + + def validate_required_keys(self, required_keys: dict[str, bool], context_name: str = "context") -> "RuntimeContext": + """Ensure all required keys are present in the context. + + Args: + required_keys: Dictionary mapping key names to boolean indicating if required + context_name: Name of the context for error messages (e.g., operator name) + """ + for key, is_required in required_keys.items(): + if is_required and key not in self: + raise ValueError(f"{context_name}: missing required input '{key}'") + return self diff --git a/reme_ai/core/enumeration/json_schema_enum.py b/reme_ai/core/enumeration/json_schema_enum.py index 17b59380..507645f4 100644 --- a/reme_ai/core/enumeration/json_schema_enum.py +++ b/reme_ai/core/enumeration/json_schema_enum.py @@ -3,17 +3,16 @@ from enum import Enum -class JsonSchemaEnum(str, Enum): +class JsonSchemaEnum(Enum): """Enumeration of valid JSON Schema data types.""" - STRING = "string" - NUMBER = "number" - INTEGER = "integer" - OBJECT = "object" - ARRAY = "array" - BOOLEAN = "boolean" - NULL = "null" + STRING = str + NUMBER = float + INTEGER = int + OBJECT = dict + ARRAY = list + BOOLEAN = bool def __str__(self) -> str: """Returns the string representation of the enum value.""" - return self.value + return self.name.lower() diff --git a/reme_ai/core/op/__init__.py b/reme_ai/core/op/__init__.py new file mode 100644 index 00000000..a8ba9a12 --- /dev/null +++ b/reme_ai/core/op/__init__.py @@ -0,0 +1,11 @@ +"""op""" + +from .base_op import BaseOp +from .parallel_op import ParallelOp +from .sequential_op import SequentialOp + +__all__ = [ + "BaseOp", + "ParallelOp", + "SequentialOp", +] diff --git a/reme_ai/core/op/base_op.py b/reme_ai/core/op/base_op.py new file mode 100644 index 00000000..7a131a28 --- /dev/null +++ b/reme_ai/core/op/base_op.py @@ -0,0 +1,336 @@ +"""Base operator class for LLM workflow execution and composition.""" + +import asyncio +import copy +import inspect +from pathlib import Path +from typing import Callable, Any, Union + +from loguru import logger +from tqdm import tqdm + +from ..context import RuntimeContext, PromptHandler, C, BaseContext +from ..embedding import BaseEmbeddingModel +from ..llm import BaseLLM +from ..schema import ToolCall, ToolAttr +from ..token_counter import BaseTokenCounter +from ..utils import camel_to_snake, CacheHandler, timer +from ..vector_store import BaseVectorStore + + +class BaseOp: + """Base operator class for LLM workflow execution and composition.""" + + def __new__(cls, *args, **kwargs): + """Capture initialization arguments for object cloning.""" + instance = super().__new__(cls) + instance._init_args = copy.copy(args) + instance._init_kwargs = copy.copy(kwargs) + return instance + + def __init__( + self, + name: str = "", + async_mode: bool = True, + language: str = "", + prompt_name: str = "", + llm: str | BaseLLM = "default", + embedding_model: str | BaseEmbeddingModel = "default", + vector_store: str | BaseVectorStore = "default", + token_counter: str | BaseTokenCounter = "default", + enable_cache: bool = False, + cache_path: str = "cache/op", + sub_ops: Union[list["BaseOp"], dict[str, "BaseOp"], "BaseOp", None] = None, + input_mapping: dict[str, str] | None = None, + output_mapping: dict[str, str] | None = None, + enable_tool_response: bool = False, + enable_sync_thread_pool: bool = True, + max_retries: int = 1, + raise_exception: bool = False, + **kwargs, + ): + """Initialize operator configurations and internal state.""" + self.name = name or camel_to_snake(self.__class__.__name__) + self.async_mode = async_mode + self.language = language or C.language + self.prompt = self._get_prompt_handler(prompt_name) + + self._llm = llm + self._embedding_model = embedding_model + self._vector_store = vector_store + self._token_counter = token_counter + + self.enable_cache = enable_cache + self.cache_path = cache_path + self.sub_ops = BaseContext[str, BaseOp]() + self.add_sub_ops(sub_ops) + + self.input_mapping = input_mapping + self.output_mapping = output_mapping + self.enable_tool_response = enable_tool_response + self.enable_sync_thread_pool = enable_sync_thread_pool + self.max_retries = max(1, max_retries) + self.raise_exception = raise_exception + self.op_params = kwargs + + self._pending_tasks: list = [] + self.context: RuntimeContext | None = None + self._cache: CacheHandler | None = None + self._tool_call: ToolCall | None = None + + def _get_prompt_handler(self, prompt_name: str) -> PromptHandler: + """Load prompt configuration from the associated YAML file.""" + path = Path(inspect.getfile(self.__class__)) + path = path.with_stem(prompt_name) if prompt_name else path + return PromptHandler(language=self.language).load_prompt_by_file(path.with_suffix(".yaml")) + + def _build_tool_call(self) -> ToolCall | None: + """Build and return the tool call schema; override in subclasses.""" + + def _validate_inputs(self): + """Ensure all required tool inputs are present in context.""" + if self.tool_call: + parameters = self.tool_call.parameters + if parameters.type == "object" and parameters.properties: + required_list = parameters.required or [] + required_keys = {k: (k in required_list) for k in parameters.properties.keys()} + self.context.validate_required_keys(required_keys, self.name) + + def _handle_failure(self, e: Exception, attempt: int): + """Log failures and handle final retry logic.""" + logger.exception(f"{self.name} failed (attempt {attempt + 1}): {e}") + if attempt == self.max_retries - 1: + if self.raise_exception: + raise e + + if self.tool_call: + self.output = f"{self.name} failed: {e}" + + @property + def tool_call(self) -> ToolCall: + """Lazily construct and return the tool call metadata.""" + if self._tool_call is None: + self._tool_call = self._build_tool_call() + assert self._tool_call, "tool_call is not defined!" + self._tool_call.name = self._tool_call.name or self.name + if not self._tool_call.output.properties: + self._tool_call.output = ToolAttr( + type="object", + properties={ + f"{self.name}_result": ToolAttr(type="string", description=f"Execution result of {self.name}"), + }, + ) + return self._tool_call + + @property + def input_dict(self) -> dict: + """Extract required and optional inputs from context based on schema.""" + parameters = self.tool_call.parameters + if parameters.type != "object" or not parameters.properties: + return {} + required_keys = set(parameters.required or []) + return {k: self.context[k] for k in parameters.properties.keys() if (k in required_keys or k in self.context)} + + @property + def output(self) -> Any: + """Get the single output value from context.""" + output_properties = self.tool_call.output.properties + if not output_properties: + return None + keys = list(output_properties.keys()) + return self.context[keys[0]] + + @output.setter + def output(self, value: Any): + """Set the single output value into context.""" + output_properties = self.tool_call.output.properties + if not output_properties: + return + keys = list(output_properties.keys()) + self.context[keys[0]] = value + + @property + def cache(self) -> CacheHandler: + """Access the operator-specific cache handler.""" + assert self.enable_cache, "Cache is disabled!" + if not self._cache: + self._cache = CacheHandler(f"{self.cache_path}/{self.name}") + return self._cache + + @property + def llm(self) -> BaseLLM: + """Lazily initialize and return the LLM instance.""" + if isinstance(self._llm, str): + cfg = C.service_config.llm[self._llm] + self._llm = C.get_llm_class(cfg.backend)(model_name=cfg.model_name, **cfg.model_extra) + return self._llm + + @property + def embedding_model(self) -> BaseEmbeddingModel: + """Lazily initialize and return the embedding model instance.""" + if isinstance(self._embedding_model, str): + cfg = C.service_config.embedding_model[self._embedding_model] + self._embedding_model = C.get_embedding_model_class(cfg.backend)( + model_name=cfg.model_name, + **cfg.model_extra, + ) + return self._embedding_model + + @property + def vector_store(self) -> BaseVectorStore: + """Lazily initialize and return the vector store instance.""" + if isinstance(self._vector_store, str): + self._vector_store = C.get_vector_store(self._vector_store) + return self._vector_store + + @property + def token_counter(self) -> BaseTokenCounter: + """Lazily initialize and return the token counter instance.""" + if isinstance(self._token_counter, str): + cfg = C.service_config.token_counter[self._token_counter] + self._token_counter = C.get_token_counter_class(cfg.backend)( + model_name=cfg.model_name, + **cfg.model_extra, + ) + return self._token_counter + + async def before_execute(self): + """Prepare context and validate before async execution.""" + self.context.apply_mapping(self.input_mapping) + self._validate_inputs() + + async def execute(self): + """Define core async logic in subclasses.""" + + async def after_execute(self): + """Finalize context and mappings after async execution.""" + self.context.apply_mapping(self.output_mapping) + if self.tool_call and self.enable_tool_response: + self.context.response.answer = self.output + + if not isinstance(self._llm, str) and hasattr(self._llm, "close"): + await self._llm.close() + if not isinstance(self._embedding_model, str) and hasattr(self._embedding_model, "close"): + await self._embedding_model.close() + + def before_execute_sync(self): + """Prepare context and validate before sync execution.""" + self.context.apply_mapping(self.input_mapping) + self._validate_inputs() + + def execute_sync(self): + """Define core sync logic in subclasses.""" + + def after_execute_sync(self): + """Finalize context and mappings after sync execution.""" + self.context.apply_mapping(self.output_mapping) + if self.tool_call and self.enable_tool_response: + self.context.response.answer = self.output + + if not isinstance(self._llm, str) and hasattr(self._llm, "close_sync"): + self._llm.close_sync() + if not isinstance(self._embedding_model, str) and hasattr(self._embedding_model, "close_sync"): + self._embedding_model.close_sync() + + @timer + def call_sync(self, context: RuntimeContext = None, **kwargs): + """Execute the operator synchronously with retry logic.""" + self.context = RuntimeContext.from_context(context, **kwargs) + for i in range(self.max_retries): + try: + self.before_execute_sync() + self.execute_sync() + self.after_execute_sync() + break + except Exception as e: + self._handle_failure(e, i) + return self.output if self.tool_call else None + + async def call(self, context: RuntimeContext = None, **kwargs): + """Execute the operator asynchronously with retry logic.""" + self.context = RuntimeContext.from_context(context, **kwargs) + for i in range(self.max_retries): + try: + await self.before_execute() + await self.execute() + await self.after_execute() + break + except Exception as e: + self._handle_failure(e, i) + return self.output if self.tool_call else None + + def submit_sync_task(self, fn: Callable, *args, **kwargs) -> "BaseOp": + """Submit a task to the thread pool or local queue.""" + task = C.thread_pool.submit(fn, *args, **kwargs) if self.enable_sync_thread_pool else (fn, args, kwargs) + self._pending_tasks.append(task) + return self + + def submit_async_task(self, coro_fn: Callable, *args, **kwargs) -> "BaseOp": + """Submit an async task to the pending tasks queue.""" + task = coro_fn(*args, **kwargs) + self._pending_tasks.append(task) + return self + + def join_sync_tasks(self, task_desc: str = None) -> list: + """Wait for all pending sync tasks and return flattened results.""" + results = [] + for task in tqdm(self._pending_tasks, desc=task_desc or self.name): + res = task.result() if self.enable_sync_thread_pool else task[0](*task[1], **task[2]) + if res: + results.extend(res if isinstance(res, list) else [res]) + self._pending_tasks.clear() + return results + + async def join_async_tasks(self, return_exceptions: bool = True) -> list: + """Wait for all pending async tasks and aggregate results.""" + try: + raw_results = await asyncio.gather(*self._pending_tasks, return_exceptions=return_exceptions) + results = [] + for res in raw_results: + if isinstance(res, Exception): + logger.error(f"Async task failed: {res}") + continue + if res: + results.extend(res if isinstance(res, list) else [res]) + return results + finally: + self._pending_tasks.clear() + + def add_sub_ops(self, sub_ops: Union[list["BaseOp"], dict[str, "BaseOp"], "BaseOp", None]): + """Add child operators to this operator's sub_ops context.""" + if not sub_ops: + return + + if isinstance(sub_ops, dict): + ops_dict = sub_ops + else: + ops_dict = {op.name: op for op in (sub_ops if isinstance(sub_ops, list) else [sub_ops])} + + for name, op in ops_dict.items(): + assert self.async_mode == op.async_mode, "Async mode mismatch!" + self.sub_ops[name] = op + + def add_sub_op(self, sub_op: "BaseOp"): + """Add a single child operator to this operator's sub_ops context.""" + self.add_sub_ops(sub_op) + + def __lshift__(self, ops): + """Operator overload for adding sub-operators.""" + self.add_sub_ops(ops) + return self + + def __rshift__(self, op: "BaseOp"): + """Operator overload for sequential execution composition.""" + from .sequential_op import SequentialOp + + seq = SequentialOp(sub_ops=[self], async_mode=self.async_mode) + seq.add_sub_ops(op.sub_ops if isinstance(op, SequentialOp) else op) + return seq + + def __or__(self, op: "BaseOp"): + """Operator overload for parallel execution composition.""" + from .parallel_op import ParallelOp + + 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 diff --git a/reme_ai/core/op/parallel_op.py b/reme_ai/core/op/parallel_op.py new file mode 100644 index 00000000..746485ca --- /dev/null +++ b/reme_ai/core/op/parallel_op.py @@ -0,0 +1,33 @@ +"""Module providing the ParallelOp class for concurrent operation execution.""" + +from .base_op import BaseOp + + +class ParallelOp(BaseOp): + """Operation class that executes multiple sub-operations in parallel.""" + + async def execute(self): + """Executes all sub-operations concurrently using asynchronous tasks.""" + for op in self.sub_ops.values(): + assert op.async_mode + self.submit_async_task(op.call, context=self.context) + await self.join_async_tasks() + + def execute_sync(self): + """Executes all sub-operations concurrently using synchronous task management.""" + for op in self.sub_ops.values(): + assert not op.async_mode + self.submit_sync_task(op.call_sync, context=self.context) + self.join_sync_tasks() + + def __lshift__(self, op: dict[str, BaseOp] | list[BaseOp] | BaseOp): + """Raises RuntimeError as the shift operator is not supported for parallel operations.""" + raise RuntimeError(f"`<<` is not supported in `{self.name}`") + + def __or__(self, op: BaseOp): + """Adds sub-operations to the current parallel group using the bitwise OR operator.""" + if isinstance(op, ParallelOp) and op.sub_ops: + self.add_sub_ops(op.sub_ops) + else: + self.add_sub_op(op) + return self diff --git a/reme_ai/core/op/sequential_op.py b/reme_ai/core/op/sequential_op.py new file mode 100644 index 00000000..79243cff --- /dev/null +++ b/reme_ai/core/op/sequential_op.py @@ -0,0 +1,31 @@ +"""Module providing the SequentialOp class for serial operation execution.""" + +from .base_op import BaseOp + + +class SequentialOp(BaseOp): + """Operation class that executes sub-operations one after another in order.""" + + async def execute(self): + """Executes sub-operations sequentially using asynchronous awaits.""" + for op in self.sub_ops.values(): + assert op.async_mode + await op.call(context=self.context) + + def execute_sync(self): + """Executes sub-operations sequentially in a synchronous blocking manner.""" + for op in self.sub_ops.values(): + assert not op.async_mode + op.call_sync(context=self.context) + + def __lshift__(self, op: dict[str, BaseOp] | list[BaseOp] | BaseOp): + """Raises RuntimeError as the left shift operator is not supported.""" + raise RuntimeError(f"`<<` is not supported in `{self.name}`") + + def __rshift__(self, op: BaseOp): + """Appends operations to the sequence using the bitwise right shift operator.""" + if isinstance(op, SequentialOp) and op.sub_ops: + self.add_sub_ops(op.sub_ops) + else: + self.add_sub_op(op) + return self diff --git a/reme_ai/core/schema/request.py b/reme_ai/core/schema/request.py index aa054219..ece942b9 100644 --- a/reme_ai/core/schema/request.py +++ b/reme_ai/core/schema/request.py @@ -1,17 +1,11 @@ """Defines the data structure for processing incoming user requests and message history.""" -from typing import List - from pydantic import Field, BaseModel, ConfigDict -from .message import Message - class Request(BaseModel): """Represents a structured request payload containing a query, message list, and metadata.""" model_config = ConfigDict(extra="allow") - query: str = Field(default="") - messages: List[Message] = Field(default_factory=list) metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/tool_call.py b/reme_ai/core/schema/tool_call.py index a94dfb1d..0a96f8c4 100644 --- a/reme_ai/core/schema/tool_call.py +++ b/reme_ai/core/schema/tool_call.py @@ -16,7 +16,7 @@ class ToolAttr(BaseModel): model_config = ConfigDict(extra="allow") - type: str = Field(default=JsonSchemaEnum.STRING.value, description="The data type of the attribute") + 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") @@ -27,7 +27,7 @@ class ToolAttr(BaseModel): @classmethod def validate_type_is_valid_enum(cls, v: str) -> str: """Validates that the provided type string exists within JsonSchemaEnum values.""" - valid_types = [e.value for e in JsonSchemaEnum] + valid_types = [str(e) for e in JsonSchemaEnum] if v not in valid_types: raise ValueError(f"Invalid type: '{v}'. Must be one of {valid_types}") diff --git a/reme_ai/core/utils/__init__.py b/reme_ai/core/utils/__init__.py index ed1916b2..6b217d21 100644 --- a/reme_ai/core/utils/__init__.py +++ b/reme_ai/core/utils/__init__.py @@ -1,14 +1,22 @@ """utils""" +from .cache_handler import CacheHandler from .case_converter import snake_to_camel, camel_to_snake from .env_utils import load_env +from .http_client import HttpClient +from .mcp_client import MCPClient +from .pydantic_utils import create_pydantic_model from .singleton import singleton from .timer import timer __all__ = [ + "CacheHandler", "snake_to_camel", "camel_to_snake", "load_env", + "HttpClient", + "MCPClient", + "create_pydantic_model", "singleton", "timer", ] diff --git a/reme_ai/core/utils/cache_handler.py b/reme_ai/core/utils/cache_handler.py new file mode 100644 index 00000000..70c8585a --- /dev/null +++ b/reme_ai/core/utils/cache_handler.py @@ -0,0 +1,184 @@ +"""Local file-based cache utility for DataFrames, lists, dicts, and strings.""" + +import json +from datetime import datetime, timedelta +from pathlib import Path +from typing import Any + +import pandas as pd +from loguru import logger + + +class CacheHandler: + """Handles persistent data caching with expiration and type support.""" + + _EXTENSIONS = { + pd.DataFrame: ".csv", + dict: ".json", + list: ".json", + str: ".txt", + } + + _TYPE_NAMES = { + "DataFrame": pd.DataFrame, + "dict": dict, + "list": list, + "str": str, + } + + def __init__(self, cache_dir: str | Path = "cache"): + """Initialize cache directory and load existing metadata.""" + self.cache_dir = Path(cache_dir) + self.cache_dir.mkdir(parents=True, exist_ok=True) + self.metadata_file = self.cache_dir / "metadata.json" + self.metadata: dict[str, Any] = self._load_metadata() + + def set_cache_dir(self, cache_dir: str | Path) -> None: + """Change the cache directory and reload metadata.""" + self.cache_dir = Path(cache_dir) + self.cache_dir.mkdir(parents=True, exist_ok=True) + self.metadata_file = self.cache_dir / "metadata.json" + self.metadata = self._load_metadata() + logger.info(f"Cache directory moved to: {self.cache_dir}") + + def _load_metadata(self) -> dict[str, Any]: + """Load metadata from the JSON file.""" + if self.metadata_file.exists(): + try: + with open(self.metadata_file, "r", encoding="utf-8") as f: + return json.load(f) + except (json.JSONDecodeError, OSError) as e: + logger.warning(f"Metadata load failed: {e}") + return {} + + def _save_metadata(self) -> None: + """Persist metadata to the disk.""" + try: + with open(self.metadata_file, "w", encoding="utf-8") as f: + json.dump(self.metadata, f, ensure_ascii=False, indent=2) + except OSError as e: + logger.error(f"Metadata save failed: {e}") + + def _get_path(self, key: str, data_type: type | None = None) -> Path: + """Resolve the file path based on data type or metadata.""" + ext = ".dat" + if data_type in self._EXTENSIONS: + ext = self._EXTENSIONS[data_type] + elif key in self.metadata: + stored_type = self.metadata[key].get("data_type") + ext = self._EXTENSIONS.get(self._TYPE_NAMES.get(stored_type, None), ".dat") + return self.cache_dir / f"{key}{ext}" + + @staticmethod + def _execute_save(data: Any, path: Path, dtype: type, **kwargs) -> dict: + """Execute type-specific save operations.""" + if dtype is pd.DataFrame: + data.to_csv(path, index=kwargs.get("index", False), encoding="utf-8") + return {"row_count": len(data), "file_size": path.stat().st_size} + + if dtype in (dict, list): + with open(path, "w", encoding="utf-8") as f: + json.dump(data, f, ensure_ascii=False, indent=2) + return {"item_count": len(data), "file_size": path.stat().st_size} + + if dtype is str: + path.write_text(data, encoding=kwargs.get("encoding", "utf-8")) + return {"char_count": len(data), "file_size": path.stat().st_size} + + raise ValueError(f"Unsupported type: {dtype}") + + @staticmethod + def _execute_load(path: Path, type_name: str, **kwargs) -> Any: + """Execute type-specific load operations.""" + if type_name == "DataFrame": + return pd.read_csv(path, encoding=kwargs.get("encoding", "utf-8")) + if type_name in ("dict", "list"): + with open(path, "r", encoding="utf-8") as f: + return json.load(f) + if type_name == "str": + return path.read_text(encoding=kwargs.get("encoding", "utf-8")) + raise ValueError(f"Unknown data type in metadata: {type_name}") + + def save(self, key: str, data: Any, expire_hours: float | None = None, **kwargs) -> bool: + """Save data to cache with optional expiration.""" + try: + dtype = type(data) + path = self._get_path(key, dtype) + stats = self._execute_save(data, path, dtype, **kwargs) + + now = datetime.now() + self.metadata[key] = { + "created_at": now.isoformat(), + "expire_at": (now + timedelta(hours=expire_hours)).isoformat() if expire_hours else None, + "data_type": dtype.__name__, + **stats, + } + self._save_metadata() + return True + except Exception as e: + logger.error(f"Save failed for {key}: {e}") + return False + + def load(self, key: str, auto_clean: bool = True, **kwargs) -> Any | None: + """Load data from cache if not expired.""" + if self._is_expired(key): + if auto_clean: + self.delete(key) + return None + + path = self._get_path(key) + if not path.exists() or key not in self.metadata: + return None + + try: + return self._execute_load(path, self.metadata[key]["data_type"], **kwargs) + except Exception as e: + logger.error(f"Load failed for {key}: {e}") + return None + + def _is_expired(self, key: str) -> bool: + """Check if the cached entry has expired.""" + entry = self.metadata.get(key) + if not entry or not entry.get("expire_at"): + return False + return datetime.now() > datetime.fromisoformat(entry["expire_at"]) + + def delete(self, key: str) -> bool: + """Remove a specific cache entry and its file.""" + try: + path = self._get_path(key) + if path.exists(): + path.unlink() + if key in self.metadata: + del self.metadata[key] + self._save_metadata() + return True + except OSError as e: + logger.error(f"Delete failed for {key}: {e}") + return False + + def exists(self, key: str) -> bool: + """Check if a valid cache entry exists.""" + return key in self.metadata and not self._is_expired(key) + + def clear_all(self) -> bool: + """Purge all cache files and reset metadata.""" + try: + for file in self.cache_dir.iterdir(): + if file.is_file(): + file.unlink() + self.metadata = {} + self._save_metadata() + return True + except OSError as e: + logger.error(f"Clear all failed: {e}") + return False + + def get_stats(self) -> dict[str, Any]: + """Return cache usage statistics.""" + total_size = sum(f.stat().st_size for f in self.cache_dir.glob("*") if f.is_file()) + return { + "count": len(self.metadata), + "size_mb": round(total_size / (1024 * 1024), 2), + "dir": str(self.cache_dir), + } diff --git a/reme_ai/core/utils/http_client.py b/reme_ai/core/utils/http_client.py new file mode 100644 index 00000000..8c8e92c3 --- /dev/null +++ b/reme_ai/core/utils/http_client.py @@ -0,0 +1,91 @@ +"""Asynchronous HTTP client for executing flows with built-in retry logic.""" + +import json +from collections.abc import AsyncIterator +from typing import Optional + +import httpx +from loguru import logger + +from ..schema import Response + + +class HttpClient: + """Async client for flow endpoints with automated retries and error handling.""" + + def __init__( + self, + base_url: str = "http://localhost:8001", + timeout: float = 3600.0, + max_retries: int = 3, + raise_exception: bool = True, + ): + """Initialize the client with base configuration.""" + self.base_url = base_url.rstrip("/") + self.timeout = timeout + self.max_retries = max_retries + self.raise_exception = raise_exception + self.client = httpx.AsyncClient(timeout=timeout) + + async def __aenter__(self): + """Enter async context manager.""" + return self + + async def __aexit__(self, exc_type, exc_val, exc_tb): + """Exit async context manager and close connection.""" + await self.close() + + async def close(self): + """Close the underlying HTTP client.""" + await self.client.aclose() + + async def health_check(self) -> dict[str, str]: + """Check the health status of the flow service.""" + response = await self.client.get(f"{self.base_url}/health") + response.raise_for_status() + return response.json() + + async def execute_flow(self, flow_name: str, **kwargs) -> Optional[Response]: + """Execute a flow with automated retry logic.""" + endpoint = f"{self.base_url}/{flow_name}" + + for attempt in range(self.max_retries): + try: + response = await self.client.post(endpoint, json=kwargs) + response.raise_for_status() + return Response(**response.json()) + + except (httpx.HTTPError, Exception) as e: + logger.error(f"Flow {flow_name} failed (attempt {attempt + 1}/{self.max_retries}): {e}") + if attempt == self.max_retries - 1 and self.raise_exception: + raise e + return None + + async def list_endpoints(self) -> dict: + """Retrieve available endpoints from OpenAPI specification.""" + response = await self.client.get(f"{self.base_url}/openapi.json") + response.raise_for_status() + return response.json() + + async def execute_stream_flow(self, flow_name: str, **kwargs) -> AsyncIterator[dict[str, str]]: + """Execute a flow and yield parsed SSE stream chunks.""" + endpoint = f"{self.base_url}/{flow_name}" + + async with self.client.stream("POST", endpoint, json=kwargs) as response: + response.raise_for_status() + async for line in response.aiter_lines(): + if not line or not line.startswith("data:"): + continue + + content = line.removeprefix("data:").strip() + if content == "[DONE]": + break + + try: + data = json.loads(content) + yield { + "type": data.get("chunk_type", "answer"), + "content": data.get("chunk", ""), + } + except json.JSONDecodeError: + continue diff --git a/reme_ai/core/utils/mcp_client.py b/reme_ai/core/utils/mcp_client.py new file mode 100644 index 00000000..3a1c70e7 --- /dev/null +++ b/reme_ai/core/utils/mcp_client.py @@ -0,0 +1,107 @@ +"""Module for managing Model Context Protocol (MCP) server connections.""" + +import os +import re +from contextlib import asynccontextmanager +from typing import Any + +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 ..schema import ToolCall + + +class MCPClient: + """A client manager for handling multiple MCP transport protocols.""" + + def __init__(self, config: dict): + """Initialize the client with server configuration.""" + self.config = config + + @staticmethod + def _infer_transport_type(cfg: dict[str, Any]) -> str: + """Infer the transport type based on configuration fields.""" + if "command" in cfg: + return "stdio" + + if "url" in cfg: + url = cfg["url"].lower() + if url.endswith("/sse") or "sse" in url: + return "sse" + return "streamable-http" + + raise ValueError(f"Could not infer transport type for: {cfg}") + + def _replace_env_vars(self, data: str | dict | list) -> Any: + """Replace environment variable placeholders in configuration.""" + if isinstance(data, str): + return re.sub(r"\$\{(\w+)\}", lambda m: os.getenv(m.group(1), m.group(0)), data) + if isinstance(data, dict): + return {k: self._replace_env_vars(v) for k, v in data.items()} + if isinstance(data, list): + return [self._replace_env_vars(i) for i in data] + return data + + @asynccontextmanager + async def _get_transport(self, cfg: dict[str, Any]): + """Context manager to yield the appropriate MCP transport.""" + # Pop 'type' if present, otherwise infer it + t_type = cfg.pop("type", None) or self._infer_transport_type(cfg) + + try: + if t_type == "stdio": + params = StdioServerParameters( + command=cfg["command"], + args=cfg.get("args", []), + env=cfg.get("env", None), + ) + async with stdio_client(params) as transport: + yield transport + elif t_type == "sse": + async with sse_client(**cfg) as transport: + yield transport + elif t_type == "streamable-http": + async with streamablehttp_client(**cfg) as transport: + yield transport + else: + raise NotImplementedError(f"Unsupported transport: {t_type}") + finally: + pass # Ensure proper cleanup + + @asynccontextmanager + async def connect_to_server(self, server_name: str): + """Establish a session with the specified MCP server.""" + server_config = self.config.get("mcpServers", {}).get(server_name) + if not server_config: + raise ValueError(f"Config for '{server_name}' not found.") + + # Process environment variables and transport selection + cfg = self._replace_env_vars(server_config) + + async with self._get_transport(cfg) as (read, write): + async with ClientSession(read, write) as session: + await session.initialize() + yield session + + async def list_tools(self, server_name: str) -> list[Tool]: + """Retrieve available tools from a specific server.""" + async with self.connect_to_server(server_name) as session: + result = await session.list_tools() + return result.tools + + async def list_tool_calls(self, server_name: str, return_dict: bool = True) -> list[dict | ToolCall]: + """Retrieve available tools from a specific server.""" + tools = await self.list_tools(server_name) + tool_calls: list[ToolCall] = [ToolCall.from_mcp_tool(tool) for tool in tools] + if return_dict: + return [tool_call.simple_input_dump() for tool_call in tool_calls] + + return tool_calls + + async def call_tool(self, server_name: str, tool_name: str, arguments: dict[str, Any]) -> CallToolResult: + """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) diff --git a/reme_ai/core/utils/pydantic_utils.py b/reme_ai/core/utils/pydantic_utils.py new file mode 100644 index 00000000..06b1f4fe --- /dev/null +++ b/reme_ai/core/utils/pydantic_utils.py @@ -0,0 +1,66 @@ +""" +Utility module for dynamic Pydantic model generation based on schema definitions. +""" + +from typing import Any, Literal + +from pydantic import create_model, Field + +from . import snake_to_camel +from ..enumeration import JsonSchemaEnum +from ..schema import ToolAttr, Request + +TYPE_MAPPING = {str(t): t.value for t in JsonSchemaEnum} + + +def create_pydantic_model(name: str, parameters: ToolAttr | None = None) -> type[Request]: + """ + Recursively generates a Pydantic model from a ToolAttr schema definition. + """ + fields = {} + + if not parameters or not parameters.properties: + return create_model(f"{snake_to_camel(name)}Model", __base__=Request) + + for field_name, attr in parameters.properties.items(): + # 1. Determine the base field type + if attr.type == "object" and attr.properties: + # Handle nested objects recursively + field_type = create_pydantic_model(field_name, attr) + + elif attr.type == "array" and attr.items: + # Handle array/list types + if isinstance(attr.items, ToolAttr): + if attr.items.type == "object": + inner_type = create_pydantic_model(f"{field_name}_item", attr.items) + else: + inner_type = TYPE_MAPPING.get(attr.items.type, Any) + field_type = list[inner_type] + else: + # Fallback for simple dictionary item definitions + field_type = list[Any] + + else: + # Handle primitive types + field_type = TYPE_MAPPING.get(attr.type, Any) + + # 2. Handle enumeration constraints + if attr.enum: + # Dynamically create a Literal type from the enum list + field_type = Literal[tuple(attr.enum)] # type: ignore + + # 3. Determine requirement status and default values + is_required = False + if parameters.required and field_name in parameters.required: + is_required = True + + # 4. Construct Field metadata + field_info = Field(default=... if is_required else None, description=attr.description) + + if not is_required: + field_type = field_type | None + + fields[field_name] = (field_type, field_info) + + # Dynamically construct the final Pydantic model class + return create_model(f"{snake_to_camel(name)}Model", **fields, __base__=Request) diff --git a/tests/mcp_servers_demo.json b/tests/mcp_servers_demo.json new file mode 100644 index 00000000..e8f9a858 --- /dev/null +++ b/tests/mcp_servers_demo.json @@ -0,0 +1,37 @@ +{ + "mcpServers": { + "sqlite-explorer": { + "type": "stdio", + "command": "uv", + "args": [ + "run", + "--with", + "mcp-server-sqlite", + "mcp-server-sqlite", + "--db-path", + "/path/to/your/database.db" + ], + "env": { + "CUSTOM_VAR": "optional_value" + } + }, + "remote-fetcher": { + "type": "sse", + "url": "https://mcp-server.example.com/sse", + "headers": { + "Authorization": "Bearer {BAILIAN_MCP_API_KEY}", + "Content-Type": "application/json" + }, + "timeout": 5, + "sse_read_timeout": 300 + }, + "my-modern-remote": { + "type": "streamable-http", + "url": "https://api.example.com/mcp", + "headers": { + "Authorization": "Bearer {BAILIAN_MCP_API_KEY}" + }, + "timeout": 5 + } + } +} \ No newline at end of file diff --git a/tests/test_cache_handler.py b/tests/test_cache_handler.py new file mode 100644 index 00000000..ddcac86f --- /dev/null +++ b/tests/test_cache_handler.py @@ -0,0 +1,94 @@ +""" +Self-contained script for CacheHandler's comprehensive test suite. +""" + +import shutil +from datetime import datetime, timedelta +from pathlib import Path + +import pandas as pd +from loguru import logger + +from reme_ai.core.utils.cache_handler import CacheHandler + + +def run_tests(): + """Execute comprehensive tests for CacheHandler.""" + test_dir = Path("test_cache_system") + if test_dir.exists(): + shutil.rmtree(test_dir) + + handler = CacheHandler(cache_dir=test_dir) + logger.info("Starting CacheHandler tests...") + + # 1. Test Data Types + logger.info("Testing data types support...") + + # DataFrame + df = pd.DataFrame({"a": [1, 2], "b": [3, 4]}) + assert handler.save("df_test", df) + assert isinstance(handler.load("df_test"), pd.DataFrame) + assert handler.load("df_test").shape == (2, 2) + + # Dict & List + d = {"key": "value", "nested": [1, 2]} + l_value = [1, "string", {"a": 1}] + assert handler.save("dict_test", d) + assert handler.save("list_test", l_value) + assert handler.load("dict_test")["key"] == "value" + assert handler.load("list_test")[1] == "string" + + # String + s = "Hello World" + assert handler.save("str_test", s) + assert handler.load("str_test") == "Hello World" + + # 2. Test Expiration + logger.info("Testing expiration logic...") + # Save with 1 second expiry (approx 0.00027 hours) + handler.save("exp_test", {"data": 1}, expire_hours=0.00001) + assert handler.exists("exp_test") is True + + # Manually modify metadata to force expiration for instant test + handler.metadata["exp_test"]["expire_at"] = (datetime.now() - timedelta(seconds=1)).isoformat() + assert handler.exists("exp_test") is False + assert handler.load("exp_test") is None + assert "exp_test" not in handler.metadata # Auto-cleaned + + # 3. Test Existence and Deletion + logger.info("Testing delete and exists...") + handler.save("del_test", "delete me") + assert handler.exists("del_test") is True + handler.delete("del_test") + assert handler.exists("del_test") is False + assert not (test_dir / "del_test.txt").exists() + + # 4. Test Persistence (Reload handler) + logger.info("Testing persistence...") + handler.save("persist_test", [1, 2, 3]) + new_handler = CacheHandler(cache_dir=test_dir) + assert new_handler.exists("persist_test") is True + assert new_handler.load("persist_test") == [1, 2, 3] + + # 5. Test Statistics and Clear + logger.info("Testing stats and clear...") + stats = handler.get_stats() + assert stats["count"] > 0 + handler.clear_all() + assert handler.get_stats()["count"] == 0 + assert len(list(test_dir.glob("*"))) == 1 # Only metadata.json remains + + # 6. Test Error Handling + logger.info("Testing error handling...") + assert handler.load("non_existent_key") is None + # Test unsupported type + assert handler.save("invalid", {1, 2}) is False + + logger.success("All tests passed successfully!") + + # Cleanup after tests + shutil.rmtree(test_dir) + + +if __name__ == "__main__": + run_tests() diff --git a/tests/test_mcp_client.py b/tests/test_mcp_client.py new file mode 100644 index 00000000..7db2ad61 --- /dev/null +++ b/tests/test_mcp_client.py @@ -0,0 +1,31 @@ +"""Test module for demonstrating MCPClient functionality.""" + +import asyncio +import json + +from reme_ai.core.utils import MCPClient + + +async def main(): + """Execute demonstration of the MCPClient.""" + test_mcp = "test_mcp" + config_data = { + "mcpServers": { + test_mcp: { + "url": "http://127.0.0.1:8010/sse", + }, + }, + } + + client = MCPClient(config_data) + + try: + t_list = await client.list_tool_calls(test_mcp) + for t in t_list: + print(json.dumps(t, ensure_ascii=False, indent=2)) + except Exception as e: + print(f"Error occurred: {e}") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/test_mcp_server.py b/tests/test_mcp_server.py new file mode 100644 index 00000000..67f9c542 --- /dev/null +++ b/tests/test_mcp_server.py @@ -0,0 +1,125 @@ +"""Dynamic MCP server implementation with JSON-schema based tool registration.""" + +from typing import Any + +from fastmcp import FastMCP +from fastmcp.tools import FunctionTool + +from reme_ai.core.schema import ToolCall +from reme_ai.core.utils import create_pydantic_model + +mcp = FastMCP("DynamicSchemaServer", port=8010) + +# Configuration including enum examples +MODES_CONFIG = { + "register_user": ToolCall( + **{ + "name": "register_user", + "description": "Register a new user with metadata, tags, and roles.", + "parameters": { + "type": "object", + "properties": { + "username": {"type": "string", "description": "Unique username"}, + "role": { + "type": "string", + "enum": ["admin", "editor", "viewer"], + "description": "User access level", + }, + "metadata": { + "type": "object", + "description": "User metadata", + "properties": { + "age": {"type": "integer"}, + "location": {"type": "string"}, + }, + "required": ["age"], + }, + "tags": { + "type": "array", + "description": "User tags", + "items": { + "type": "object", + "properties": { + "tag_id": {"type": "string"}, + "level": {"type": "number"}, + }, + "required": ["tag_id"], + }, + }, + }, + "required": ["username", "metadata", "role"], + }, + }, + ), + "create_order": ToolCall( + **{ + "name": "create_order", + "description": "创建订单", + "parameters": { + "type": "object", + "properties": { + "order_id": {"type": "string", "description": "订单ID"}, + "amount": {"type": "number", "description": "订单金额"}, + "customer": { + "type": "object", + "description": "客户信息", + "properties": { + "name": {"type": "string", "description": "客户姓名"}, + "email": {"type": "string", "description": "客户邮箱"}, + "phone": {"type": "string", "description": "联系电话"}, + }, + "required": ["name", "email"], + }, + }, + "required": ["order_id", "customer"], + }, + }, + ), +} + + +async def core_handler(mode: str, **kwargs: Any) -> dict[str, Any]: + """Process dynamic tool requests and return execution results.""" + print(f"Executing Mode: {mode}, Parameters: {kwargs}") + return { + "status": "success", + "mode": mode, + "received_data": kwargs, + } + + +def register_dynamic_tools() -> None: + """Iterate over tool configurations and register them to the MCP instance.""" + for mode_name, tool_call in MODES_CONFIG.items(): + # Create Pydantic model from tool parameters + request_model = create_pydantic_model(tool_call.name, tool_call.parameters) + + # Create execution function with closure to capture current mode and model + def create_tool_func(current_mode: str, model: type): + async def execute_tool(**kwargs: Any) -> dict[str, Any]: + # Validate and normalize input using Pydantic model + validated_data = model(**kwargs).model_dump(exclude_none=True) + return await core_handler(current_mode, **validated_data) + + return execute_tool + + tool_fn = create_tool_func(mode_name, request_model) + + # Extract parameters schema + tool_call_schema = tool_call.simple_input_dump() + parameters = tool_call_schema[tool_call_schema["type"]]["parameters"] + + # Create FunctionTool and register + tool = FunctionTool( + name=tool_call.name, + description=tool_call.description, + fn=tool_fn, + parameters=parameters, + ) + + mcp.add_tool(tool) + + +if __name__ == "__main__": + register_dynamic_tools() + mcp.run(transport="sse")