mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-09 03:20:54 +00:00
feat(core): add operator framework and utility modules
This commit is contained in:
parent
91a07e4186
commit
90a53737c8
17 changed files with 1220 additions and 47 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
11
reme_ai/core/op/__init__.py
Normal file
11
reme_ai/core/op/__init__.py
Normal 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
336
reme_ai/core/op/base_op.py
Normal 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
|
||||
33
reme_ai/core/op/parallel_op.py
Normal file
33
reme_ai/core/op/parallel_op.py
Normal 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
|
||||
31
reme_ai/core/op/sequential_op.py
Normal file
31
reme_ai/core/op/sequential_op.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
184
reme_ai/core/utils/cache_handler.py
Normal file
184
reme_ai/core/utils/cache_handler.py
Normal 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),
|
||||
}
|
||||
91
reme_ai/core/utils/http_client.py
Normal file
91
reme_ai/core/utils/http_client.py
Normal 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
|
||||
107
reme_ai/core/utils/mcp_client.py
Normal file
107
reme_ai/core/utils/mcp_client.py
Normal 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)
|
||||
66
reme_ai/core/utils/pydantic_utils.py
Normal file
66
reme_ai/core/utils/pydantic_utils.py
Normal 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)
|
||||
37
tests/mcp_servers_demo.json
Normal file
37
tests/mcp_servers_demo.json
Normal 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
|
||||
}
|
||||
}
|
||||
}
|
||||
94
tests/test_cache_handler.py
Normal file
94
tests/test_cache_handler.py
Normal 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
31
tests/test_mcp_client.py
Normal 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
125
tests/test_mcp_server.py
Normal 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")
|
||||
Loading…
Add table
Reference in a new issue