feat(core): add operator framework and utility modules

This commit is contained in:
jinli.yl 2025-12-31 17:16:53 +08:00
parent 91a07e4186
commit 90a53737c8
17 changed files with 1220 additions and 47 deletions

View file

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

View file

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

View file

@ -0,0 +1,11 @@
"""op"""
from .base_op import BaseOp
from .parallel_op import ParallelOp
from .sequential_op import SequentialOp
__all__ = [
"BaseOp",
"ParallelOp",
"SequentialOp",
]

336
reme_ai/core/op/base_op.py Normal file
View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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",
]

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

31
tests/test_mcp_client.py Normal file
View file

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

125
tests/test_mcp_server.py Normal file
View file

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