diff --git a/.github/workflows/pre-commit.yml b/.github/workflows/pre-commit.yml
new file mode 100644
index 00000000..9685f577
--- /dev/null
+++ b/.github/workflows/pre-commit.yml
@@ -0,0 +1,38 @@
+name: Pre-commit
+
+on: [ push, pull_request ]
+
+jobs:
+ run:
+ runs-on: ${{ matrix.os }}
+ strategy:
+ fail-fast: True
+ matrix:
+ os: [ ubuntu-latest ]
+ env:
+ OS: ${{ matrix.os }}
+ PYTHON: '3.10'
+ steps:
+ - uses: actions/checkout@master
+ - name: Setup Python
+ uses: actions/setup-python@master
+ with:
+ python-version: '3.10'
+ - name: Update setuptools
+ run: |
+ pip install -U setuptools wheel
+ - name: Install
+ run: |
+ pip install -q -e .[dev]
+ - name: Install pre-commit
+ run: |
+ pre-commit install
+ - name: Pre-commit starts
+ run: |
+ pre-commit run --all-files > pre-commit.log 2>&1 || true
+ cat pre-commit.log
+ if grep -q Failed pre-commit.log; then
+ echo -e "\e[41m [**FAIL**] Please install pre-commit and format your code first. \e[0m"
+ exit 1
+ fi
+ echo -e "\e[46m ********************************Passed******************************** \e[0m"
\ No newline at end of file
diff --git a/README.md b/README.md
index 2becbcff..2423c643 100644
--- a/README.md
+++ b/README.md
@@ -5,6 +5,8 @@
+
+
diff --git a/README_ZH.md b/README_ZH.md
index c59b7186..5a39d3e6 100644
--- a/README_ZH.md
+++ b/README_ZH.md
@@ -5,6 +5,8 @@
+
+
diff --git a/docs/deprecated.txt b/docs/deprecated.txt
new file mode 100644
index 00000000..6d7cb14f
--- /dev/null
+++ b/docs/deprecated.txt
@@ -0,0 +1,11 @@
+from loguru import logger
+
+用英文注释,完善module/class/function docstring,要一句话简洁,不要变更代码
+用英文注释,完善module/class/function docstring,要一句话简洁,代码要简洁,符合pep和pylint规范
+C0114: Missing module docstring (missing-module-docstring)
+C0115: Missing class docstring (missing-class-docstring)
+C0116: Missing function or method docstring (missing-function-docstring)
+done: { for f in ./*.py; do [[ "$f" != "./__init__.py" ]] && grep -v '^[[:space:]]*#' "$f"; done; } | pbcopy
+
+然后是一个完整的tests,但是不要用其他的包,只是test开头的函数或者类,要求from loguru import logger
+写一个测试文件,不要使用pytest,普通的test,要求英文注释
\ No newline at end of file
diff --git a/pyproject.toml b/pyproject.toml
index 60de58d9..abb263da 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -43,7 +43,8 @@ dev = [
"myst-nb",
"sphinxcontrib-bibtex",
"furo",
- "sphinxcontrib-mermaid"
+ "sphinxcontrib-mermaid",
+ "pre-commit",
]
token = [
diff --git a/reme_ai/core/context/__init__.py b/reme_ai/core/context/__init__.py
new file mode 100644
index 00000000..7f26d600
--- /dev/null
+++ b/reme_ai/core/context/__init__.py
@@ -0,0 +1,16 @@
+"""context"""
+
+from .base_context import BaseContext
+from .prompt_handler import PromptHandler
+from .registry import Registry
+from .runtime_context import RuntimeContext
+from .service_context import ServiceContext, C
+
+__all__ = [
+ "BaseContext",
+ "PromptHandler",
+ "Registry",
+ "RuntimeContext",
+ "ServiceContext",
+ "C",
+]
diff --git a/reme_ai/core/context/base_context.py b/reme_ai/core/context/base_context.py
new file mode 100644
index 00000000..dabd8cdb
--- /dev/null
+++ b/reme_ai/core/context/base_context.py
@@ -0,0 +1,41 @@
+"""Module providing a dictionary subclass with attribute-style access and pickling support."""
+
+from typing import Generic, TypeVar
+
+_KT = TypeVar("_KT")
+_VT = TypeVar("_VT")
+
+
+class BaseContext(dict, Generic[_KT, _VT]):
+ """A dictionary subclass that enables accessing and modifying keys as attributes."""
+
+ def __getattr__(self, name: str) -> _VT:
+ """Retrieve a dictionary item as an attribute."""
+ try:
+ return self[name]
+ except KeyError as e:
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e
+
+ def __setattr__(self, name: str, value: _VT) -> None:
+ """Assign a value to a dictionary item using attribute syntax."""
+ self[name] = value
+
+ def __delattr__(self, name: str) -> None:
+ """Remove a dictionary item using attribute syntax."""
+ try:
+ # Delete item from dict via key
+ del self[name]
+ except KeyError as e:
+ raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e
+
+ def __getstate__(self) -> dict:
+ """Return the dictionary representation for pickling."""
+ return dict(self)
+
+ def __setstate__(self, state: dict) -> None:
+ """Restore the dictionary state from a pickled object."""
+ self.update(state)
+
+ def __reduce__(self):
+ """Define the reconstruction logic for pickling processes."""
+ return self.__class__, (), self.__getstate__()
diff --git a/reme_ai/core/context/prompt_handler.py b/reme_ai/core/context/prompt_handler.py
new file mode 100644
index 00000000..e48f98ac
--- /dev/null
+++ b/reme_ai/core/context/prompt_handler.py
@@ -0,0 +1,95 @@
+"""Module for managing and formatting prompt templates from files or dictionaries."""
+
+from pathlib import Path
+
+import yaml
+from loguru import logger
+
+from .base_context import BaseContext
+from .service_context import C
+
+
+class PromptHandler(BaseContext):
+ """A context-aware handler for loading, retrieving, and formatting prompt templates."""
+
+ def __init__(self, language: str = "", **kwargs):
+ """Initialize the handler with a specific language and optional context data."""
+ super().__init__(**kwargs)
+ self.language: str = language or C.language
+
+ def load_prompt_by_file(self, prompt_file_path: Path | str = None):
+ """Load prompt configurations from a YAML file into the context."""
+ if prompt_file_path is None:
+ return self
+
+ if isinstance(prompt_file_path, str):
+ prompt_file_path = Path(prompt_file_path)
+
+ if not prompt_file_path.exists():
+ return self
+
+ with prompt_file_path.open(encoding="utf-8") as f:
+ # Load YAML content using the full loader
+ prompt_dict = yaml.load(f, yaml.FullLoader)
+ self.load_prompt_dict(prompt_dict)
+ return self
+
+ def load_prompt_dict(self, prompt_dict: dict = None):
+ """Merge a dictionary of prompt strings into the current context."""
+ if not prompt_dict:
+ return self
+
+ for key, value in prompt_dict.items():
+ if isinstance(value, str):
+ if key in self:
+ logger.warning(f"Overwriting prompt key={key}, old_value={self[key]}, new_value={value}")
+ else:
+ logger.debug(f"Adding new prompt key={key}, value={value}")
+ self[key] = value
+ return self
+
+ def get_prompt(self, prompt_name: str):
+ """Retrieve a prompt by name, automatically appending the language suffix if needed."""
+ key: str = prompt_name
+ if self.language and not key.endswith(self.language.strip()):
+ key += "_" + self.language.strip()
+
+ assert key in self, f"prompt_name={key} not found."
+ return self[key]
+
+ def prompt_format(self, prompt_name: str, **kwargs) -> str:
+ """Format a prompt by filtering flagged lines and filling template variables."""
+ prompt = self.get_prompt(prompt_name)
+
+ # Separate boolean flags from string formatting arguments
+ flag_kwargs = {k: v for k, v in kwargs.items() if isinstance(v, bool)}
+ other_kwargs = {k: v for k, v in kwargs.items() if not isinstance(v, bool)}
+
+ if flag_kwargs:
+ split_prompt = []
+ for line in prompt.strip().split("\n"):
+ hit = False
+ hit_flag = True
+ for key, flag in flag_kwargs.items():
+ if not line.startswith(f"[{key}]"):
+ continue
+
+ hit = True
+ hit_flag = flag
+ # Remove the flag prefix from the line
+ line = line.strip(f"[{key}]")
+ break
+
+ # Include line if no flag is present or if the flag evaluates to True
+ if not hit:
+ split_prompt.append(line)
+ elif hit_flag:
+ split_prompt.append(line)
+
+ prompt = "\n".join(split_prompt)
+
+ if other_kwargs:
+ # Apply standard Python string formatting
+ prompt = prompt.format(**other_kwargs)
+
+ return prompt
diff --git a/reme_ai/core/context/registry.py b/reme_ai/core/context/registry.py
new file mode 100644
index 00000000..f403037d
--- /dev/null
+++ b/reme_ai/core/context/registry.py
@@ -0,0 +1,19 @@
+"""Module providing a registry class for managing class-to-name mappings via decorators."""
+
+from .base_context import BaseContext
+
+
+class Registry(BaseContext):
+ """A registry container that uses decorators to map and store class references."""
+
+ def register(self, name: str = "", add_cls: bool = True):
+ """Return a decorator that registers a class under a specific name in the registry."""
+
+ def decorator(cls):
+ if add_cls:
+ # Use provided name or default to the class name as the key
+ key = name or cls.__name__
+ self[key] = cls
+ return cls
+
+ return decorator
diff --git a/reme_ai/core/context/runtime_context.py b/reme_ai/core/context/runtime_context.py
new file mode 100644
index 00000000..ded35ded
--- /dev/null
+++ b/reme_ai/core/context/runtime_context.py
@@ -0,0 +1,56 @@
+"""Module providing a 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
+
+
+class RuntimeContext(BaseContext):
+ """A context class for handling execution state, including response metadata and stream queues."""
+
+ def __init__(
+ self,
+ response: Response | None = None,
+ stream_queue: asyncio.Queue | None = None,
+ **kwargs,
+ ):
+ """Initialize the runtime context with optional response objects and message queues."""
+ super().__init__(**kwargs)
+
+ self.response: Response | None = response if response is not None else Response()
+ self.stream_queue: asyncio.Queue | None = stream_queue
+
+ 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
+
+ # Package raw data into a StreamChunk schema
+ stream_chunk = StreamChunk(chunk_type=chunk_type, chunk=chunk)
+ await self.stream_queue.put(stream_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)
+ 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)
+ return self
+
+ def add_response_error(self, e: Exception):
+ """Update the internal response object to reflect a failure state using exception details."""
+ self.response.success = False
+ self.response.answer = str(e.args)
diff --git a/reme_ai/core/context/service_context.py b/reme_ai/core/context/service_context.py
new file mode 100644
index 00000000..af99a870
--- /dev/null
+++ b/reme_ai/core/context/service_context.py
@@ -0,0 +1,105 @@
+"""Module for managing global service configurations and component registries via a singleton context."""
+
+from concurrent.futures import ThreadPoolExecutor
+from typing import Dict
+
+from .base_context import BaseContext
+from .registry import Registry
+from ..enumeration import RegistryEnum
+from ..schema import ServiceConfig
+from ..utils import singleton
+
+
+@singleton
+class ServiceContext(BaseContext):
+ """A singleton container for global application state, thread pools, and component registries."""
+
+ def __init__(self, **kwargs):
+ """Initialize the global context with configuration objects and specialized registries."""
+ super().__init__(**kwargs)
+
+ self.service_config: ServiceConfig | None = None
+ self.language: str = ""
+ self.thread_pool: ThreadPoolExecutor | None = None
+ self.vector_store_dict: Dict[str, dict] = {}
+ self.external_mcp_tool_call_dict: dict = {}
+ # Initialize a registry for every category defined in RegistryEnum
+ self.registry_dict: Dict[RegistryEnum, Registry] = {v: Registry() for v in RegistryEnum.__members__.values()}
+ self.flow_dict: dict = {}
+
+ def register(self, name: str, register_type: RegistryEnum):
+ """Return a decorator to register a component within a specific registry category."""
+ return self.registry_dict[register_type].register(name=name)
+
+ def register_llm(self, name: str = ""):
+ """Register a Large Language Model class."""
+ return self.register(name=name, register_type=RegistryEnum.LLM)
+
+ def register_embedding_model(self, name: str = ""):
+ """Register an embedding model class."""
+ return self.register(name=name, register_type=RegistryEnum.EMBEDDING_MODEL)
+
+ def register_vector_store(self, name: str = ""):
+ """Register a vector store implementation class."""
+ return self.register(name=name, register_type=RegistryEnum.VECTOR_STORE)
+
+ def register_op(self, name: str = ""):
+ """Register an operation (Op) class."""
+ return self.register(name=name, register_type=RegistryEnum.OP)
+
+ def register_flow(self, name: str = ""):
+ """Register a workflow or logic flow class."""
+ return self.register(name=name, register_type=RegistryEnum.FLOW)
+
+ def register_service(self, name: str = ""):
+ """Register a backend service class."""
+ return self.register(name=name, register_type=RegistryEnum.SERVICE)
+
+ def register_token_counter(self, name: str = ""):
+ """Register a token counting utility class."""
+ return self.register(name=name, register_type=RegistryEnum.TOKEN_COUNTER)
+
+ def get_model_class(self, name: str, register_type: RegistryEnum):
+ """Retrieve a registered class by name from a specific registry category."""
+ assert name in self.registry_dict[register_type], f"{name} not in registry_dict[{register_type}]"
+ return self.registry_dict[register_type][name]
+
+ def get_embedding_model_class(self, name: str):
+ """Get the embedding model class registered under the given name."""
+ return self.get_model_class(name, RegistryEnum.EMBEDDING_MODEL)
+
+ def get_llm_class(self, name: str):
+ """Get the LLM class registered under the given name."""
+ return self.get_model_class(name, RegistryEnum.LLM)
+
+ def get_vector_store_class(self, name: str):
+ """Get the vector store class registered under the given name."""
+ return self.get_model_class(name, RegistryEnum.VECTOR_STORE)
+
+ def get_op_class(self, name: str):
+ """Get the operation class registered under the given name."""
+ return self.get_model_class(name, RegistryEnum.OP)
+
+ def get_flow_class(self, name: str):
+ """Get the flow class registered under the given name."""
+ return self.get_model_class(name, RegistryEnum.FLOW)
+
+ def get_service_class(self, name: str):
+ """Get the service class registered under the given name."""
+ return self.get_model_class(name, RegistryEnum.SERVICE)
+
+ def get_token_counter_class(self, name: str):
+ """Get the token counter class registered under the given name."""
+ return self.get_model_class(name, RegistryEnum.TOKEN_COUNTER)
+
+ def get_vector_store(self, name: str):
+ """Retrieve a specific vector store instance by name."""
+ return self.vector_store_dict[name]
+
+ def get_flow(self, name: str):
+ """Retrieve a specific flow instance by name."""
+ return self.flow_dict[name]
+
+
+# Export a global instance for easy access across the application
+C = ServiceContext()
diff --git a/reme_ai/core/enumeration/__init__.py b/reme_ai/core/enumeration/__init__.py
new file mode 100644
index 00000000..3323bdc9
--- /dev/null
+++ b/reme_ai/core/enumeration/__init__.py
@@ -0,0 +1,15 @@
+"""enumeration"""
+
+from .chunk_enum import ChunkEnum
+from .http_enum import HttpEnum
+from .json_schema_enum import JsonSchemaEnum
+from .registry_enum import RegistryEnum
+from .role import Role
+
+__all__ = [
+ "ChunkEnum",
+ "HttpEnum",
+ "JsonSchemaEnum",
+ "RegistryEnum",
+ "Role",
+]
diff --git a/reme_ai/core/enumeration/chunk_enum.py b/reme_ai/core/enumeration/chunk_enum.py
new file mode 100644
index 00000000..dbe37106
--- /dev/null
+++ b/reme_ai/core/enumeration/chunk_enum.py
@@ -0,0 +1,25 @@
+"""Defines the types of data chunks used in streaming responses."""
+
+from enum import Enum
+
+
+class ChunkEnum(str, Enum):
+ """Enumeration of possible chunk categories for stream processing."""
+
+ # Internal reasoning or chain-of-thought process
+ THINK = "think"
+
+ # The final generated response content
+ ANSWER = "answer"
+
+ # Metadata or calls related to external tools
+ TOOL = "tool"
+
+ # Resource consumption and token usage statistics
+ USAGE = "usage"
+
+ # Error messages or exception details
+ ERROR = "error"
+
+ # Final signal indicating the completion of the stream
+ DONE = "done"
diff --git a/reme_ai/core/enumeration/http_enum.py b/reme_ai/core/enumeration/http_enum.py
new file mode 100644
index 00000000..19622242
--- /dev/null
+++ b/reme_ai/core/enumeration/http_enum.py
@@ -0,0 +1,22 @@
+"""Provides a collection of standard HTTP request methods."""
+
+from enum import Enum
+
+
+class HttpEnum(str, Enum):
+ """Enumeration of supported HTTP methods for network requests."""
+
+ # Retrieves data from a specified resource
+ GET = "get"
+
+ # Submits data to be processed to a specified resource
+ POST = "post"
+
+ # Identical to GET but only retrieves the response headers
+ HEAD = "head"
+
+ # Uploads or replaces the representation of a target resource
+ PUT = "put"
+
+ # Deletes the specified resource from the server
+ DELETE = "delete"
diff --git a/reme_ai/core/enumeration/json_schema_enum.py b/reme_ai/core/enumeration/json_schema_enum.py
new file mode 100644
index 00000000..17b59380
--- /dev/null
+++ b/reme_ai/core/enumeration/json_schema_enum.py
@@ -0,0 +1,19 @@
+"""Defines the standard data types supported by JSON Schema."""
+
+from enum import Enum
+
+
+class JsonSchemaEnum(str, Enum):
+ """Enumeration of valid JSON Schema data types."""
+
+ STRING = "string"
+ NUMBER = "number"
+ INTEGER = "integer"
+ OBJECT = "object"
+ ARRAY = "array"
+ BOOLEAN = "boolean"
+ NULL = "null"
+
+ def __str__(self) -> str:
+ """Returns the string representation of the enum value."""
+ return self.value
diff --git a/reme_ai/core/enumeration/registry_enum.py b/reme_ai/core/enumeration/registry_enum.py
new file mode 100644
index 00000000..876c06b8
--- /dev/null
+++ b/reme_ai/core/enumeration/registry_enum.py
@@ -0,0 +1,28 @@
+"""Defines the registry categories for core components of the system."""
+
+from enum import Enum
+
+
+class RegistryEnum(str, Enum):
+ """Enumeration of component types registered within the application lifecycle."""
+
+ # Large Language Model interfaces
+ LLM = "llm"
+
+ # Models used for generating vector embeddings
+ EMBEDDING_MODEL = "embedding_model"
+
+ # Databases or storage systems for vector search
+ VECTOR_STORE = "vector_store"
+
+ # Atomic operations or functional units
+ OP = "op"
+
+ # Orchestrated sequences of operations or workflows
+ FLOW = "flow"
+
+ # External APIs or shared internal services
+ SERVICE = "service"
+
+ # Utilities for tracking and limiting token consumption
+ TOKEN_COUNTER = "token_counter"
diff --git a/reme_ai/core/enumeration/role.py b/reme_ai/core/enumeration/role.py
new file mode 100644
index 00000000..4acad7e5
--- /dev/null
+++ b/reme_ai/core/enumeration/role.py
@@ -0,0 +1,19 @@
+"""Defines the participant roles in a chat completion sequence."""
+
+from enum import Enum
+
+
+class Role(str, Enum):
+ """Enumeration of standard personas involved in a conversation flow."""
+
+ # High-level instructions to guide the model's behavior
+ SYSTEM = "system"
+
+ # Input or queries provided by the human user
+ USER = "user"
+
+ # Responses or messages generated by the AI model
+ ASSISTANT = "assistant"
+
+ # Output or results returned from external tool executions
+ TOOL = "tool"
diff --git a/reme_ai/core/llm/__init__.py b/reme_ai/core/llm/__init__.py
new file mode 100644
index 00000000..57578f49
--- /dev/null
+++ b/reme_ai/core/llm/__init__.py
@@ -0,0 +1,15 @@
+"""llm"""
+
+from .base_llm import BaseLLM
+from .lite_llm import LiteLLM
+from .lite_llm_sync import LiteLLMSync
+from .openai_llm import OpenAILLM
+from .openai_llm_sync import OpenAILLMSync
+
+__all__ = [
+ "BaseLLM",
+ "LiteLLM",
+ "LiteLLMSync",
+ "OpenAILLM",
+ "OpenAILLMSync",
+]
diff --git a/reme_ai/core/llm/base_llm.py b/reme_ai/core/llm/base_llm.py
new file mode 100644
index 00000000..6f1da121
--- /dev/null
+++ b/reme_ai/core/llm/base_llm.py
@@ -0,0 +1,364 @@
+"""Abstract base interface for ReMe LLM implementations."""
+
+import asyncio
+import json
+import time
+from abc import ABC
+from typing import List, Callable, Generator, AsyncGenerator, Any, Optional, Dict
+
+from loguru import logger
+
+from ..enumeration import ChunkEnum, Role
+from ..schema import Message
+from ..schema import StreamChunk
+from ..schema import ToolCall
+
+
+class BaseLLM(ABC):
+ """Abstract base class defining the standard interface for LLM interactions."""
+
+ def __init__(self, model_name: str, max_retries: int = 3, raise_exception: bool = False, **kwargs):
+ """Initialize the LLM client with model configurations and retry policies."""
+ self.model_name: str = model_name
+ self.max_retries: int = max_retries
+ self.raise_exception: bool = raise_exception
+ self.kwargs: dict = kwargs
+
+ @staticmethod
+ def _process_stream_chunk(
+ stream_chunk: StreamChunk,
+ state: dict,
+ enable_stream_print: bool = False,
+ ) -> None:
+ """Update the aggregation state by processing an individual stream chunk."""
+ if stream_chunk.chunk_type is ChunkEnum.USAGE:
+ if enable_stream_print:
+ print(f"\n{json.dumps(stream_chunk.chunk, ensure_ascii=False, indent=2)}", flush=True)
+
+ elif stream_chunk.chunk_type is ChunkEnum.THINK:
+ if enable_stream_print:
+ if not state["enter_think"]:
+ state["enter_think"] = True
+ print("\n", end="", flush=True)
+ print(stream_chunk.chunk, end="", flush=True)
+ state["reasoning_content"] += stream_chunk.chunk
+
+ elif stream_chunk.chunk_type is ChunkEnum.ANSWER:
+ if enable_stream_print:
+ if not state["enter_answer"]:
+ state["enter_answer"] = True
+ if state["enter_think"]:
+ print("\n", flush=True)
+ print(stream_chunk.chunk, end="", flush=True)
+ state["answer_content"] += stream_chunk.chunk
+
+ elif stream_chunk.chunk_type is ChunkEnum.TOOL:
+ if enable_stream_print:
+ print(f"\n{json.dumps(stream_chunk.chunk, ensure_ascii=False, indent=2)}", flush=True)
+ state["tool_calls"].append(stream_chunk.chunk)
+
+ elif stream_chunk.chunk_type is ChunkEnum.ERROR:
+ if enable_stream_print:
+ print(f"\n{stream_chunk.chunk}", flush=True)
+
+ @staticmethod
+ def _create_message_from_state(state: dict) -> Message:
+ """Construct a Message object from the accumulated aggregation state."""
+ return Message(
+ role=Role.ASSISTANT,
+ reasoning_content=state["reasoning_content"],
+ content=state["answer_content"],
+ tool_calls=state["tool_calls"],
+ )
+
+ @staticmethod
+ def _accumulate_tool_call_chunk(
+ tool_call,
+ ret_tools: List[ToolCall],
+ ) -> None:
+ """Assemble incremental tool call fragments into complete ToolCall objects."""
+ index = tool_call.index
+
+ # Ensure we have a ToolCall object at this index
+ while len(ret_tools) <= index:
+ ret_tools.append(ToolCall(index=index))
+
+ # Accumulate tool call parts (id, name, arguments)
+ if tool_call.id:
+ ret_tools[index].id += tool_call.id
+
+ if tool_call.function and tool_call.function.name:
+ ret_tools[index].name += tool_call.function.name
+
+ if tool_call.function and tool_call.function.arguments:
+ ret_tools[index].arguments += tool_call.function.arguments
+
+ @staticmethod
+ def _validate_and_serialize_tools(
+ ret_tools: List[ToolCall],
+ tools: Optional[List[ToolCall]],
+ ) -> List[Dict]:
+ """Validate tool call integrity and return serialized tool dictionaries."""
+ if not ret_tools:
+ return []
+
+ # Create lookup dict for tool validation
+ tool_dict: Dict[str, ToolCall] = {x.name: x for x in tools} if tools else {}
+ validated_tools = []
+
+ for tool in ret_tools:
+ # Skip tools that weren't in the provided tool list
+ if tool.name not in tool_dict:
+ continue
+
+ # Validate tool arguments are valid JSON
+ if not tool.check_argument():
+ raise ValueError(
+ f"Tool call {tool.name} has invalid JSON arguments: {tool.arguments}",
+ )
+
+ validated_tools.append(tool.simple_output_dump())
+
+ return validated_tools
+
+ def _build_stream_kwargs(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ log_params: bool = True,
+ **kwargs,
+ ) -> dict:
+ """Construct provider-specific parameters for streaming API requests."""
+ raise NotImplementedError
+
+ async def _stream_chat(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ stream_kwargs: Optional[dict] = None,
+ ) -> AsyncGenerator[StreamChunk, None]:
+ """Internal async generator for streaming raw response chunks."""
+ raise NotImplementedError
+
+ def _stream_chat_sync(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ stream_kwargs: Optional[dict] = None,
+ ) -> Generator[StreamChunk, None, None]:
+ """Internal synchronous generator for streaming raw response chunks."""
+ raise NotImplementedError
+
+ async def _stream_with_retry(
+ self,
+ operation_name: str,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]],
+ stream_kwargs: dict,
+ ) -> AsyncGenerator[StreamChunk, None]:
+ """Execute the async streaming operation with retry logic and error recovery."""
+ for i in range(self.max_retries):
+ try:
+ async for chunk in self._stream_chat(messages=messages, tools=tools, stream_kwargs=stream_kwargs):
+ yield chunk
+ return
+
+ except Exception as e:
+ logger.exception(f"{operation_name} with model={self.model_name} encounter error with e={e.args}")
+
+ if i == self.max_retries - 1:
+ if self.raise_exception:
+ raise e
+ yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e))
+ return
+
+ yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e))
+ await asyncio.sleep(i + 1)
+
+ def _stream_with_retry_sync(
+ self,
+ operation_name: str,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]],
+ stream_kwargs: dict,
+ ) -> Generator[StreamChunk, None, None]:
+ """Execute the synchronous streaming operation with retry logic and error recovery."""
+ for i in range(self.max_retries):
+ try:
+ yield from self._stream_chat_sync(messages=messages, tools=tools, stream_kwargs=stream_kwargs)
+ return
+
+ except Exception as e:
+ logger.exception(f"{operation_name} with model={self.model_name} encounter error with e={e.args}")
+
+ if i == self.max_retries - 1:
+ if self.raise_exception:
+ raise e
+ yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e))
+ return
+
+ yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e))
+ time.sleep(i + 1)
+
+ async def stream_chat(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ **kwargs,
+ ) -> AsyncGenerator[StreamChunk, None]:
+ """Public async interface for streaming chat completions with retries."""
+ stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs)
+ async for chunk in self._stream_with_retry("stream chat", messages, tools, stream_kwargs):
+ yield chunk
+
+ def stream_chat_sync(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ **kwargs,
+ ) -> Generator[StreamChunk, None, None]:
+ """Public synchronous interface for streaming chat completions with retries."""
+ stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs)
+ yield from self._stream_with_retry_sync("stream chat sync", messages, tools, stream_kwargs)
+
+ async def _chat(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ enable_stream_print: bool = False,
+ **kwargs,
+ ) -> Message:
+ """Internal async method to aggregate a full response by consuming the stream."""
+ state = {
+ "enter_think": False,
+ "enter_answer": False,
+ "reasoning_content": "",
+ "answer_content": "",
+ "tool_calls": [],
+ }
+
+ stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs)
+ async for stream_chunk in self._stream_chat(messages=messages, tools=tools, stream_kwargs=stream_kwargs):
+ self._process_stream_chunk(stream_chunk, state, enable_stream_print)
+
+ return self._create_message_from_state(state)
+
+ def _chat_sync(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ enable_stream_print: bool = False,
+ **kwargs,
+ ) -> Message:
+ """Internal synchronous method to aggregate a full response by consuming the stream."""
+ state = {
+ "enter_think": False,
+ "enter_answer": False,
+ "reasoning_content": "",
+ "answer_content": "",
+ "tool_calls": [],
+ }
+
+ stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs)
+ for stream_chunk in self._stream_chat_sync(messages=messages, tools=tools, stream_kwargs=stream_kwargs):
+ self._process_stream_chunk(stream_chunk, state, enable_stream_print)
+
+ return self._create_message_from_state(state)
+
+ async def _execute_with_retry(
+ self,
+ operation_name: str,
+ operation_fn: Callable[[], Any],
+ callback_fn: Optional[Callable[[Message], Any]] = None,
+ default_value: Any = None,
+ ) -> Message | Any:
+ """Execute a generic async operation with error handling and retry logic."""
+ for i in range(self.max_retries):
+ try:
+ result = await operation_fn()
+ return callback_fn(result) if callback_fn else result
+
+ except Exception as e:
+ logger.exception(f"{operation_name} with model={self.model_name} encounter error with e={e.args}")
+
+ if i == self.max_retries - 1:
+ if self.raise_exception:
+ raise e
+ return default_value
+
+ await asyncio.sleep(1 + i)
+ return default_value
+
+ def _execute_with_retry_sync(
+ self,
+ operation_name: str,
+ operation_fn: Callable[[], Message],
+ callback_fn: Optional[Callable[[Message], Any]] = None,
+ default_value: Any = None,
+ ) -> Message | Any:
+ """Execute a generic synchronous operation with error handling and retry logic."""
+ for i in range(self.max_retries):
+ try:
+ result = operation_fn()
+ return callback_fn(result) if callback_fn else result
+
+ except Exception as e:
+ logger.exception(f"{operation_name} with model={self.model_name} encounter error with e={e.args}")
+
+ if i == self.max_retries - 1:
+ if self.raise_exception:
+ raise e
+ return default_value
+
+ time.sleep(1 + i)
+ return default_value
+
+ async def chat(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ enable_stream_print: bool = False,
+ callback_fn: Optional[Callable[[Message], Any]] = None,
+ default_value: Any = None,
+ **kwargs,
+ ) -> Message | Any:
+ """Perform an async chat completion with integrated retries and error handling."""
+ return await self._execute_with_retry(
+ operation_name="chat",
+ operation_fn=lambda: self._chat(
+ messages=messages,
+ tools=tools,
+ enable_stream_print=enable_stream_print,
+ **kwargs,
+ ),
+ callback_fn=callback_fn,
+ default_value=default_value,
+ )
+
+ def chat_sync(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ enable_stream_print: bool = False,
+ callback_fn: Optional[Callable[[Message], Any]] = None,
+ default_value: Any = None,
+ **kwargs,
+ ) -> Message | Any:
+ """Perform a synchronous chat completion with integrated retries and error handling."""
+ return self._execute_with_retry_sync(
+ operation_name="chat sync",
+ operation_fn=lambda: self._chat_sync(
+ messages=messages,
+ tools=tools,
+ enable_stream_print=enable_stream_print,
+ **kwargs,
+ ),
+ callback_fn=callback_fn,
+ default_value=default_value,
+ )
+
+ async def close(self):
+ """Release any asynchronous resources or connections held by the client."""
+
+ def close_sync(self):
+ """Release any synchronous resources or connections held by the client."""
diff --git a/reme_ai/core/llm/lite_llm.py b/reme_ai/core/llm/lite_llm.py
new file mode 100644
index 00000000..32d3005c
--- /dev/null
+++ b/reme_ai/core/llm/lite_llm.py
@@ -0,0 +1,118 @@
+"""LiteLLM asynchronous implementation for ReMe."""
+
+import os
+from typing import List, AsyncGenerator, Optional
+
+import litellm
+from loguru import logger
+
+from .base_llm import BaseLLM
+from ..context import C
+from ..enumeration import ChunkEnum
+from ..schema import Message
+from ..schema import StreamChunk
+from ..schema import ToolCall
+
+
+@C.register_llm("litellm")
+class LiteLLM(BaseLLM):
+ """Async LLM implementation using LiteLLM to support multiple providers."""
+
+ def __init__(
+ self,
+ api_key: Optional[str] = None,
+ base_url: Optional[str] = None,
+ custom_llm_provider: str = "openai",
+ **kwargs,
+ ):
+ """Initialize the LiteLLM client with API configuration and provider settings."""
+ super().__init__(**kwargs)
+ self.api_key: Optional[str] = api_key or os.getenv("REME_LLM_API_KEY")
+ self.base_url: Optional[str] = base_url or os.getenv("REME_LLM_BASE_URL")
+ self.custom_llm_provider: str = custom_llm_provider
+
+ def _build_stream_kwargs(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ log_params: bool = True,
+ **kwargs,
+ ) -> dict:
+ """Construct and log the parameters dictionary for LiteLLM API calls."""
+ # Construct the API parameters by merging multiple sources
+ llm_kwargs = {
+ "model": self.model_name,
+ "messages": [x.simple_dump() for x in messages],
+ "tools": [x.simple_input_dump() for x in tools] if tools else None,
+ "stream": True,
+ "custom_llm_provider": self.custom_llm_provider,
+ **self.kwargs,
+ **kwargs,
+ }
+
+ # Add API key and base URL if provided
+ if self.api_key:
+ llm_kwargs["api_key"] = self.api_key
+ if self.base_url:
+ llm_kwargs["base_url"] = self.base_url
+
+ # Log parameters for debugging, with message/tool counts instead of full content
+ if log_params:
+ log_kwargs: dict = {}
+ for k, v in llm_kwargs.items():
+ if k in ["messages", "tools"]:
+ log_kwargs[k] = len(v) if v is not None else 0
+ elif k == "api_key":
+ # Mask API key in logs for security
+ log_kwargs[k] = "***" if v else None
+ else:
+ log_kwargs[k] = v
+ logger.info(f"llm_kwargs={log_kwargs}")
+
+ return llm_kwargs
+
+ async def _stream_chat(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ stream_kwargs: Optional[dict] = None,
+ ) -> AsyncGenerator[StreamChunk, None]:
+ """Execute async streaming chat requests and yield processed response chunks."""
+ # Create streaming completion request using LiteLLM asynchronously
+ stream_kwargs = stream_kwargs or {}
+ completion = await litellm.acompletion(**stream_kwargs)
+
+ # Track accumulated tool calls across chunks
+ ret_tools: List[ToolCall] = []
+ # Flag to track if we've started receiving answer content
+ is_answering: bool = False
+
+ async for chunk in completion:
+ # Handle usage information (typically the last chunk)
+ if not chunk.choices:
+ if hasattr(chunk, "usage") and chunk.usage:
+ yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump())
+
+ else:
+ delta = chunk.choices[0].delta
+
+ # Check for reasoning content (models that support thinking)
+ if hasattr(delta, "reasoning_content") and delta.reasoning_content is not None:
+ yield StreamChunk(chunk_type=ChunkEnum.THINK, chunk=delta.reasoning_content)
+
+ else:
+ if not is_answering:
+ is_answering = True
+
+ # Yield regular text content
+ if delta.content is not None:
+ yield StreamChunk(chunk_type=ChunkEnum.ANSWER, chunk=delta.content)
+
+ # Process tool calls - LiteLLM streams them incrementally
+ if hasattr(delta, "tool_calls") and delta.tool_calls is not None:
+ for tool_call in delta.tool_calls:
+ self._accumulate_tool_call_chunk(tool_call, ret_tools)
+
+ # After streaming completes, validate and yield complete tool calls
+ for tool_data in self._validate_and_serialize_tools(ret_tools, tools):
+ yield StreamChunk(chunk_type=ChunkEnum.TOOL, chunk=tool_data)
diff --git a/reme_ai/core/llm/lite_llm_sync.py b/reme_ai/core/llm/lite_llm_sync.py
new file mode 100644
index 00000000..5f814a2e
--- /dev/null
+++ b/reme_ai/core/llm/lite_llm_sync.py
@@ -0,0 +1,101 @@
+"""Synchronous LiteLLM-based LLM implementation for the ReMe framework.
+
+This module provides a unified synchronous interface for 100+ LLM providers via LiteLLM,
+supporting streaming completions, tool calling, and reasoning content. For
+asynchronous operations, refer to the LiteLLM class in the lite_llm module.
+"""
+
+from typing import List, Generator, Optional
+
+import litellm
+
+from .lite_llm import LiteLLM
+from ..context import C
+from ..enumeration import ChunkEnum
+from ..schema import Message
+from ..schema import StreamChunk
+from ..schema import ToolCall
+
+
+@C.register_llm("litellm_sync")
+class LiteLLMSync(LiteLLM):
+ """
+ Synchronous LiteLLM client for executing chat completions and streaming responses.
+
+ This class extends the base LiteLLM implementation to provide synchronous
+ execution of streaming methods, inheriting initialization and configuration
+ logic from the parent class.
+
+ Example:
+ >>> llm = LiteLLMSync(
+ ... model_name="qwen3-max",
+ ... api_key="sk-...",
+ ... temperature=0.7
+ ... )
+ >>> messages = [Message(role=Role.USER, content="Hello!")]
+ >>> for chunk in llm.chat(messages):
+ ... print(chunk)
+ """
+
+ def _stream_chat_sync(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ stream_kwargs: Optional[dict] = None,
+ ) -> Generator[StreamChunk, None, None]:
+ """
+ Internal synchronous generator for processing streaming chat completion chunks.
+
+ This method orchestrates the LiteLLM completion lifecycle by categorizing
+ raw API chunks into usage data, reasoning content (thinking), regular
+ text responses, and aggregated tool calls.
+
+ Args:
+ messages: List of conversation messages to send to the model.
+ tools: Optional list of tool definitions available for the model to call.
+ stream_kwargs: Dictionary of pre-built parameters for the LiteLLM API.
+
+ Yields:
+ StreamChunk: Wrapped response fragments categorized by ChunkEnum.
+
+ Raises:
+ ValueError: If tool call arguments fail validation or serialization.
+ """
+ # Create streaming completion request using LiteLLM
+ stream_kwargs = stream_kwargs or {}
+ completion = litellm.completion(**stream_kwargs)
+
+ # Track accumulated tool calls across chunks
+ ret_tools: List[ToolCall] = []
+ # Flag to track if we've started receiving answer content
+ is_answering: bool = False
+
+ for chunk in completion:
+ # Handle usage information (typically the last chunk)
+ if not chunk.choices:
+ if hasattr(chunk, "usage") and chunk.usage:
+ yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump())
+
+ else:
+ delta = chunk.choices[0].delta
+
+ # Check for reasoning content (models that support thinking)
+ if hasattr(delta, "reasoning_content") and delta.reasoning_content is not None:
+ yield StreamChunk(chunk_type=ChunkEnum.THINK, chunk=delta.reasoning_content)
+
+ else:
+ if not is_answering:
+ is_answering = True
+
+ # Yield regular text content
+ if delta.content is not None:
+ yield StreamChunk(chunk_type=ChunkEnum.ANSWER, chunk=delta.content)
+
+ # Process tool calls - LiteLLM streams them incrementally
+ if hasattr(delta, "tool_calls") and delta.tool_calls is not None:
+ for tool_call in delta.tool_calls:
+ self._accumulate_tool_call_chunk(tool_call, ret_tools)
+
+ # After streaming completes, validate and yield complete tool calls
+ for tool_data in self._validate_and_serialize_tools(ret_tools, tools):
+ yield StreamChunk(chunk_type=ChunkEnum.TOOL, chunk=tool_data)
diff --git a/reme_ai/core/llm/openai_llm.py b/reme_ai/core/llm/openai_llm.py
new file mode 100644
index 00000000..2c593d6f
--- /dev/null
+++ b/reme_ai/core/llm/openai_llm.py
@@ -0,0 +1,117 @@
+"""Asynchronous OpenAI-compatible LLM implementation supporting streaming, tool calls, and reasoning content."""
+
+import os
+from typing import List, AsyncGenerator, Optional
+
+from loguru import logger
+from openai import AsyncOpenAI
+
+from .base_llm import BaseLLM
+from ..context import C
+from ..enumeration import ChunkEnum
+from ..schema import Message
+from ..schema import StreamChunk
+from ..schema import ToolCall
+
+
+@C.register_llm("openai")
+class OpenAILLM(BaseLLM):
+ """Asynchronous LLM client for OpenAI-compatible APIs supporting streaming completions and tool execution."""
+
+ def __init__(
+ self,
+ api_key: Optional[str] = None,
+ base_url: Optional[str] = None,
+ **kwargs,
+ ):
+ """Initialize the OpenAI async client with API credentials and model configuration."""
+ super().__init__(**kwargs)
+ self.api_key: str = api_key or os.getenv("REME_LLM_API_KEY", "")
+ self.base_url: str = base_url or os.getenv("REME_LLM_BASE_URL", "")
+
+ # Create client using factory method
+ self._client = self._create_client()
+
+ def _create_client(self):
+ """Create and return an instance of the AsyncOpenAI client."""
+ return AsyncOpenAI(api_key=self.api_key, base_url=self.base_url)
+
+ def _build_stream_kwargs(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ log_params: bool = True,
+ **kwargs,
+ ) -> dict:
+ """Construct the parameter dictionary for the OpenAI Chat Completions API call."""
+ # Construct the API parameters by merging multiple sources
+ llm_kwargs = {
+ "model": self.model_name,
+ "messages": [x.simple_dump() for x in messages],
+ "tools": [x.simple_input_dump() for x in tools] if tools else None,
+ "stream": True,
+ **self.kwargs,
+ **kwargs,
+ }
+
+ # Log parameters for debugging, with message/tool counts instead of full content
+ if log_params:
+ log_kwargs: dict = {}
+ for k, v in llm_kwargs.items():
+ if k in ["messages", "tools"]:
+ log_kwargs[k] = len(v) if v is not None else 0
+ else:
+ log_kwargs[k] = v
+ logger.info(f"llm_kwargs={log_kwargs}")
+
+ return llm_kwargs
+
+ async def _stream_chat(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ stream_kwargs: Optional[dict] = None,
+ ) -> AsyncGenerator[StreamChunk, None]:
+ """Generate a stream of chat completion chunks including text, reasoning content, and tool calls."""
+ # Create streaming completion request to OpenAI API asynchronously
+ stream_kwargs = stream_kwargs or {}
+ completion = await self._client.chat.completions.create(**stream_kwargs)
+
+ # Track accumulated tool calls across chunks
+ ret_tools: List[ToolCall] = []
+ # Flag to track if we've started receiving answer content
+ is_answering: bool = False
+
+ async for chunk in completion:
+ # Handle usage information (typically the last chunk)
+ if not chunk.choices:
+ if hasattr(chunk, "usage") and chunk.usage:
+ yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump())
+
+ else:
+ delta = chunk.choices[0].delta
+
+ # Check for reasoning content (o1-preview, o1-mini models)
+ if hasattr(delta, "reasoning_content") and delta.reasoning_content is not None:
+ yield StreamChunk(chunk_type=ChunkEnum.THINK, chunk=delta.reasoning_content)
+
+ else:
+ if not is_answering:
+ is_answering = True
+
+ # Yield regular text content
+ if delta.content is not None:
+ yield StreamChunk(chunk_type=ChunkEnum.ANSWER, chunk=delta.content)
+
+ # Process tool calls - OpenAI streams them incrementally
+ if delta.tool_calls is not None:
+ for tool_call in delta.tool_calls:
+ self._accumulate_tool_call_chunk(tool_call, ret_tools)
+
+ # After streaming completes, validate and yield complete tool calls
+ for tool_data in self._validate_and_serialize_tools(ret_tools, tools):
+ yield StreamChunk(chunk_type=ChunkEnum.TOOL, chunk=tool_data)
+
+ async def close(self):
+ """Asynchronously close the OpenAI client and release network resources."""
+ await self._client.close()
diff --git a/reme_ai/core/llm/openai_llm_sync.py b/reme_ai/core/llm/openai_llm_sync.py
new file mode 100644
index 00000000..3aadab1a
--- /dev/null
+++ b/reme_ai/core/llm/openai_llm_sync.py
@@ -0,0 +1,71 @@
+"""Synchronous OpenAI-compatible LLM implementation supporting streaming, tool calls, and reasoning content."""
+
+from typing import List, Generator, Optional
+
+from openai import OpenAI
+
+from .openai_llm import OpenAILLM
+from ..context import C
+from ..enumeration import ChunkEnum
+from ..schema import Message
+from ..schema import StreamChunk
+from ..schema import ToolCall
+
+
+@C.register_llm("openai_sync")
+class OpenAILLMSync(OpenAILLM):
+ """Synchronous LLM client for OpenAI-compatible APIs, inheriting from OpenAILLM."""
+
+ def _create_client(self):
+ """Create and return an instance of the synchronous OpenAI client."""
+ return OpenAI(api_key=self.api_key, base_url=self.base_url)
+
+ def _stream_chat_sync(
+ self,
+ messages: List[Message],
+ tools: Optional[List[ToolCall]] = None,
+ stream_kwargs: Optional[dict] = None,
+ ) -> Generator[StreamChunk, None, None]:
+ """Synchronously generate a stream of chat completion chunks including text, reasoning, and tool calls."""
+ # Create streaming completion request to OpenAI API
+ stream_kwargs = stream_kwargs or {}
+ completion = self._client.chat.completions.create(**stream_kwargs)
+
+ # Track accumulated tool calls across chunks
+ ret_tools: List[ToolCall] = []
+ # Flag to track if we've started receiving answer content
+ is_answering: bool = False
+
+ for chunk in completion:
+ # Handle usage information (typically the last chunk)
+ if not chunk.choices:
+ if hasattr(chunk, "usage") and chunk.usage:
+ yield StreamChunk(chunk_type=ChunkEnum.USAGE, chunk=chunk.usage.model_dump())
+
+ else:
+ delta = chunk.choices[0].delta
+
+ # Check for reasoning content (o1-preview, o1-mini models)
+ if hasattr(delta, "reasoning_content") and delta.reasoning_content is not None:
+ yield StreamChunk(chunk_type=ChunkEnum.THINK, chunk=delta.reasoning_content)
+
+ else:
+ if not is_answering:
+ is_answering = True
+
+ # Yield regular text content
+ if delta.content is not None:
+ yield StreamChunk(chunk_type=ChunkEnum.ANSWER, chunk=delta.content)
+
+ # Process tool calls - OpenAI streams them incrementally
+ if delta.tool_calls is not None:
+ for tool_call in delta.tool_calls:
+ self._accumulate_tool_call_chunk(tool_call, ret_tools)
+
+ # After streaming completes, validate and yield complete tool calls
+ for tool_data in self._validate_and_serialize_tools(ret_tools, tools):
+ yield StreamChunk(chunk_type=ChunkEnum.TOOL, chunk=tool_data)
+
+ def close_sync(self):
+ """Close the synchronous OpenAI client and release network resources."""
+ self._client.close()
diff --git a/reme_ai/core/schema/__init__.py b/reme_ai/core/schema/__init__.py
new file mode 100644
index 00000000..f55f054e
--- /dev/null
+++ b/reme_ai/core/schema/__init__.py
@@ -0,0 +1,40 @@
+"""schema"""
+
+from .message import ContentBlock, Message, Trajectory
+from .request import Request
+from .response import Response
+from .service_config import (
+ CmdConfig,
+ EmbeddingModelConfig,
+ FlowConfig,
+ HttpConfig,
+ LLMConfig,
+ MCPConfig,
+ ServiceConfig,
+ TokenCounterConfig,
+ VectorStoreConfig,
+)
+from .stream_chunk import StreamChunk
+from .tool_call import ToolAttr, ToolCall
+from .vector_node import VectorNode
+
+__all__ = [
+ "ContentBlock",
+ "EmbeddingModelConfig",
+ "FlowConfig",
+ "HttpConfig",
+ "LLMConfig",
+ "MCPConfig",
+ "Message",
+ "Request",
+ "Response",
+ "ServiceConfig",
+ "StreamChunk",
+ "TokenCounterConfig",
+ "Trajectory",
+ "ToolAttr",
+ "ToolCall",
+ "VectorNode",
+ "VectorStoreConfig",
+ "CmdConfig",
+]
diff --git a/reme_ai/core/schema/message.py b/reme_ai/core/schema/message.py
new file mode 100644
index 00000000..496ef7ec
--- /dev/null
+++ b/reme_ai/core/schema/message.py
@@ -0,0 +1,136 @@
+"""Data models for multi-modal conversation history and LLM interaction trajectories."""
+
+import datetime
+import json
+from typing import Any
+
+from pydantic import BaseModel, ConfigDict, Field, model_validator
+
+from .tool_call import ToolCall
+from ..enumeration import Role
+
+
+class ContentBlock(BaseModel):
+ """
+ Individual unit of multi-modal content like text, images, or video.
+ examples:
+ {
+ "type": "image_url",
+ "image_url": {
+ "url": "https://img.alicdn.com/imgextra/i1/O1CN01gDEY8M1W114Hi3XcN_!!6000000002727-0-tps-1024-406.jpg"
+ },
+ }
+
+ {
+ "type": "video",
+ "video": [
+ "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/xzsgiz/football1.jpg",
+ "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/tdescd/football2.jpg",
+ "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/zefdja/football3.jpg",
+ "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/aedbqh/football4.jpg",
+ ],
+ }
+
+ {
+ "type": "text",
+ "text": "How do you solve this problem?"
+ }
+ """
+
+ model_config = ConfigDict(extra="allow")
+
+ type: str = Field(default="")
+ content: str | dict | list = Field(default="")
+
+ @model_validator(mode="before")
+ @classmethod
+ def init_block(cls, data: dict[str, Any]) -> dict[str, Any]:
+ """Dynamically maps the type-specific key to the content field."""
+ content_type = data.get("type", "")
+ if content_type and content_type in data:
+ data["content"] = data[content_type]
+ return data
+
+ def simple_dump(self) -> dict[str, Any]:
+ """Serializes the block into an API-compatible dictionary format."""
+ return {
+ "type": self.type,
+ self.type: self.content,
+ **self.model_extra,
+ }
+
+
+class Message(BaseModel):
+ """Data model for a single dialogue entry including roles and tool interactions."""
+
+ name: str | None = Field(default=None)
+ role: Role = Field(default=Role.USER)
+ content: str | list[ContentBlock] = Field(default="")
+ reasoning_content: str = Field(default="")
+ tool_calls: list[ToolCall] = Field(default_factory=list)
+ tool_call_id: str = Field(default="")
+ time_created: str = Field(
+ default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
+ )
+ metadata: dict[str, Any] = Field(default_factory=dict)
+
+ def dump_content(self) -> str | list[dict[str, Any]]:
+ """Returns content as a raw string or a list of serialized blocks."""
+ if isinstance(self.content, str):
+ return self.content
+ return [block.simple_dump() for block in self.content]
+
+ def simple_dump(self, add_reasoning: bool = True) -> dict[str, Any]:
+ """Transforms the message into a simplified dictionary for standard APIs."""
+ result = {"role": self.role.value, "content": self.dump_content()}
+
+ if add_reasoning and self.reasoning_content:
+ result["reasoning_content"] = self.reasoning_content
+ if self.tool_calls:
+ result["tool_calls"] = [tc.simple_output_dump() for tc in self.tool_calls]
+ if self.tool_call_id:
+ result["tool_call_id"] = self.tool_call_id
+
+ return result
+
+ def format_message(
+ self,
+ index: int | None = None,
+ add_time: bool = False,
+ use_name: bool = False,
+ add_reasoning: bool = True,
+ add_tools: bool = True,
+ ) -> str:
+ """Generates a human-readable string representation of the message."""
+ prefix = f"round{index} " if index is not None else ""
+ time_str = f"[{self.time_created}] " if add_time else ""
+ header = f"{self.name or self.role.value if use_name else self.role.value}:\n"
+
+ lines = [f"{prefix}{time_str}{header}"]
+
+ if add_reasoning and self.reasoning_content:
+ lines.append(f"{self.reasoning_content}\n")
+
+ if isinstance(self.content, str):
+ lines.append(self.content)
+ elif isinstance(self.content, list):
+ for block in self.content:
+ text = (
+ block.content if isinstance(block.content, str) else json.dumps(block.content, ensure_ascii=False)
+ )
+ lines.append(str(text))
+
+ if add_tools and self.tool_calls:
+ for tc in self.tool_calls:
+ lines.append(f" - tool_call={tc.name} params={tc.arguments}")
+
+ return "\n".join(lines).strip()
+
+
+class Trajectory(BaseModel):
+ """Sequence of messages representing a full conversation session and its evaluation."""
+
+ task_id: str = Field(default="")
+ messages: list[Message] = Field(default_factory=list)
+ score: float = Field(default=0.0)
+ metadata: dict[str, Any] = Field(default_factory=dict)
diff --git a/reme_ai/core/schema/request.py b/reme_ai/core/schema/request.py
new file mode 100644
index 00000000..aa054219
--- /dev/null
+++ b/reme_ai/core/schema/request.py
@@ -0,0 +1,17 @@
+"""Defines the data structure for processing incoming user requests and message history."""
+
+from typing import List
+
+from pydantic import Field, BaseModel, ConfigDict
+
+from .message import Message
+
+
+class Request(BaseModel):
+ """Represents a structured request payload containing a query, message list, and metadata."""
+
+ model_config = ConfigDict(extra="allow")
+
+ query: str = Field(default="")
+ messages: List[Message] = Field(default_factory=list)
+ metadata: dict = Field(default_factory=dict)
diff --git a/reme_ai/core/schema/response.py b/reme_ai/core/schema/response.py
new file mode 100644
index 00000000..3104bc6e
--- /dev/null
+++ b/reme_ai/core/schema/response.py
@@ -0,0 +1,11 @@
+"""Defines the standardized data structure for model output responses."""
+
+from pydantic import Field, BaseModel
+
+
+class Response(BaseModel):
+ """Represents a structured response containing the execution result, status, and metadata."""
+
+ answer: str | dict | list = Field(default="")
+ success: bool = Field(default=True)
+ metadata: dict = Field(default_factory=dict)
diff --git a/reme_ai/core/schema/service_config.py b/reme_ai/core/schema/service_config.py
new file mode 100644
index 00000000..67cac68f
--- /dev/null
+++ b/reme_ai/core/schema/service_config.py
@@ -0,0 +1,113 @@
+"""Configuration schemas for service components using Pydantic models."""
+
+from typing import Dict, List
+
+from pydantic import BaseModel, Field, ConfigDict
+
+from .tool_call import ToolCall
+
+
+class MCPConfig(BaseModel):
+ """Configuration for Model Context Protocol transport and network settings."""
+
+ model_config = ConfigDict(extra="allow")
+
+ transport: str = Field(default="")
+ host: str = Field(default="0.0.0.0")
+ port: int = Field(default=8001)
+
+
+class HttpConfig(BaseModel):
+ """Configuration for the HTTP server interface and connection lifecycle."""
+
+ model_config = ConfigDict(extra="allow")
+
+ host: str = Field(default="0.0.0.0")
+ port: int = Field(default=8001)
+ timeout_keep_alive: int = Field(default=3600)
+ limit_concurrency: int = Field(default=1000)
+
+
+class CmdConfig(BaseModel):
+ """Configuration for command-line flow execution parameters."""
+
+ model_config = ConfigDict(extra="allow")
+
+ flow: str = Field(default="")
+
+
+class FlowConfig(ToolCall):
+ """Configuration for workflow execution, caching, and error handling."""
+
+ model_config = ConfigDict(extra="allow")
+
+ flow_content: str = Field(default="")
+ stream: bool = Field(default=False)
+ raise_exception: bool = Field(default=True)
+ enable_cache: bool = Field(default=False)
+ cache_path: str = Field(default="cache/flow")
+ cache_expire_hours: float = Field(default=0.1)
+
+
+class LLMConfig(BaseModel):
+ """Configuration for Large Language Model backend and model identification."""
+
+ model_config = ConfigDict(extra="allow")
+
+ backend: str = Field(default="")
+ model_name: str = Field(default="")
+
+
+class EmbeddingModelConfig(BaseModel):
+ """Configuration for embedding model backends and identity."""
+
+ model_config = ConfigDict(extra="allow")
+
+ backend: str = Field(default="")
+ model_name: str = Field(default="")
+
+
+class VectorStoreConfig(BaseModel):
+ """Configuration for vector database storage and associated embeddings."""
+
+ model_config = ConfigDict(extra="allow")
+
+ backend: str = Field(default="local")
+ collection_name: str = Field(default="remy")
+ embedding_model: str = Field(default="default")
+
+
+class TokenCounterConfig(BaseModel):
+ """Configuration for token counting services and model mapping."""
+
+ model_config = ConfigDict(extra="allow")
+
+ backend: str = Field(default="base")
+ model_name: str = Field(default="")
+
+
+class ServiceConfig(BaseModel):
+ """Root configuration schema aggregating all service-level settings and components."""
+
+ model_config = ConfigDict(extra="allow")
+
+ backend: str = Field(default="")
+ enable_logo: bool = Field(default=True)
+ language: str = Field(default="")
+ thread_pool_max_workers: int = Field(default=16)
+ ray_max_workers: int = Field(default=-1)
+ disabled_flows: List[str] = Field(default_factory=list)
+ enabled_flows: List[str] = Field(default_factory=list)
+ external_mcp: Dict[str, dict] = Field(
+ default_factory=dict,
+ description="External MCP Server configuration",
+ )
+
+ mcp: MCPConfig = Field(default_factory=MCPConfig)
+ http: HttpConfig = Field(default_factory=HttpConfig)
+ cmd: CmdConfig = Field(default_factory=CmdConfig)
+ flow: Dict[str, FlowConfig] = Field(default_factory=dict)
+ llm: Dict[str, LLMConfig] = Field(default_factory=dict)
+ embedding_model: Dict[str, EmbeddingModelConfig] = Field(default_factory=dict)
+ vector_store: Dict[str, VectorStoreConfig] = Field(default_factory=dict)
+ token_counter: Dict[str, TokenCounterConfig] = Field(default_factory=dict)
diff --git a/reme_ai/core/schema/stream_chunk.py b/reme_ai/core/schema/stream_chunk.py
new file mode 100644
index 00000000..764981fd
--- /dev/null
+++ b/reme_ai/core/schema/stream_chunk.py
@@ -0,0 +1,14 @@
+"""Defines the data structure for individual data packets in a streaming response."""
+
+from pydantic import Field, BaseModel
+
+from ..enumeration import ChunkEnum
+
+
+class StreamChunk(BaseModel):
+ """Represents a single chunk of streamed data including its type, content, and completion status."""
+
+ chunk_type: ChunkEnum = Field(default=ChunkEnum.ANSWER)
+ chunk: str | dict | list = Field(default="")
+ done: bool = Field(default=False)
+ metadata: dict = Field(default_factory=dict)
diff --git a/reme_ai/core/schema/tool_call.py b/reme_ai/core/schema/tool_call.py
new file mode 100644
index 00000000..a94dfb1d
--- /dev/null
+++ b/reme_ai/core/schema/tool_call.py
@@ -0,0 +1,191 @@
+"""
+MCP Tool Schema definitions for recursive JSON Schema representation.
+"""
+
+import json
+from typing import Any, Dict, List, Optional, Union
+
+from mcp.types import Tool
+from pydantic import BaseModel, ConfigDict, Field, model_validator, field_validator
+
+from ..enumeration.json_schema_enum import JsonSchemaEnum
+
+
+class ToolAttr(BaseModel):
+ """Recursive model representing JSON Schema attributes for tool parameters."""
+
+ model_config = ConfigDict(extra="allow")
+
+ type: str = Field(default=JsonSchemaEnum.STRING.value, description="The data type of the attribute")
+ description: Optional[str] = Field(default=None, description="Description of the attribute")
+ required: Optional[List[str]] = Field(default=None, description="Required property names for object types")
+ properties: Optional[Dict[str, "ToolAttr"]] = Field(default=None, description="Child properties for objects")
+ items: Optional[Union[Dict[str, Any], "ToolAttr"]] = Field(default=None, description="Schema for array items")
+ enum: Optional[List[str]] = Field(default=None, description="Allowed values for the attribute")
+
+ @field_validator("type")
+ @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]
+
+ if v not in valid_types:
+ raise ValueError(f"Invalid type: '{v}'. Must be one of {valid_types}")
+ return v
+
+ def simple_input_dump(self) -> dict:
+ """Serializes the attribute into a standard JSON Schema dictionary."""
+ res: dict = {"type": self.type}
+ if self.description:
+ res["description"] = self.description
+ if self.enum:
+ res["enum"] = self.enum
+
+ if self.type == "object" and self.properties:
+ res["properties"] = {
+ k: v.simple_input_dump() if isinstance(v, ToolAttr) else v for k, v in self.properties.items()
+ }
+ if self.required:
+ res["required"] = self.required
+
+ if self.type == "array" and self.items:
+ res["items"] = self.items.simple_input_dump() if isinstance(self.items, ToolAttr) else self.items
+
+ return res
+
+
+# Enable recursive type resolution
+ToolAttr.model_rebuild()
+
+
+class ToolCall(BaseModel):
+ """
+ Model representing a tool definition and its call structure.
+ Supports parsing from standard JSON Schema formats and converting to MCP Tool objects.
+ input:
+ {
+ "type": "function",
+ "function": {
+ "name": "get_current_weather",
+ "description": "It is very useful when you want to check the weather of a specified city.",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "location": {
+ "type": "string",
+ "description": "Cities or counties, such as Beijing, Hangzhou, Yuhang District, etc.",
+ }
+ },
+ "required": ["location"]
+ }
+ }
+ }
+ output:
+ {
+ "index": 0,
+ "id": "call_6596dafa2a6a46f7a217da",
+ "function": {
+ "arguments": "{\"location\": \"Beijing\"}",
+ "name": "get_current_weather"
+ },
+ "type": "function",
+ }
+ """
+
+ index: int = 0
+ id: str = ""
+ type: str = "function"
+ name: str = ""
+ description: str = ""
+
+ arguments: str = Field(default="", description="JSON string of tool execution arguments")
+
+ parameters: ToolAttr = Field(
+ default_factory=lambda: ToolAttr(type="object", properties={}, required=[]),
+ description="Specification for input parameters",
+ )
+
+ output: ToolAttr = Field(
+ default_factory=lambda: ToolAttr(type="object", properties={}),
+ description="Specification for the execution result (Schema)",
+ )
+
+ @model_validator(mode="before")
+ @classmethod
+ def init_tool_call(cls, data: dict) -> dict:
+ """Initializes the model by parsing tool-specific body data."""
+ data = data.copy()
+ t_type = data.get("type", "function")
+ body = data.get(t_type, {})
+
+ # Extract basic metadata
+ data["name"] = body.get("name", data.get("name", ""))
+ data["arguments"] = body.get("arguments", data.get("arguments", ""))
+ data["description"] = body.get("description", data.get("description", ""))
+
+ # Handle parameters mapping
+ if "parameters" in body:
+ params = body["parameters"]
+ # If parameters is already a dict, ensure it matches ToolAttr structure
+ if isinstance(params, dict):
+ data["parameters"] = ToolAttr(**params)
+
+ # Handle output mapping (if provided in source)
+ if "output" in body and isinstance(body["output"], dict):
+ data["output"] = ToolAttr(**body["output"])
+
+ return data
+
+ def simple_input_dump(self) -> dict:
+ """Returns a standardized tool definition dictionary."""
+ return {
+ "type": self.type,
+ self.type: {
+ "name": self.name,
+ "description": self.description,
+ "parameters": self.parameters.simple_input_dump(),
+ },
+ }
+
+ @classmethod
+ def from_mcp_tool(cls, tool: Tool) -> "ToolCall":
+ """Creates a ToolCall instance from an MCP Tool object."""
+ # MCP Tool inputSchema maps directly to our parameters ToolAttr
+ return cls(
+ name=tool.name,
+ description=tool.description or "",
+ parameters=ToolAttr(**tool.inputSchema),
+ )
+
+ def to_mcp_tool(self) -> Tool:
+ """Converts the instance back into an MCP Tool object."""
+ return Tool(
+ name=self.name,
+ description=self.description,
+ inputSchema=self.parameters.simple_input_dump(),
+ )
+
+ @property
+ def argument_dict(self) -> dict:
+ """Parse and return arguments as a dictionary."""
+ return json.loads(self.arguments)
+
+ def check_argument(self) -> bool:
+ """Check if arguments can be parsed as valid JSON."""
+ try:
+ _ = self.argument_dict
+ return True
+ except Exception:
+ return False
+
+ def simple_output_dump(self) -> dict:
+ """Convert ToolCall to output format dictionary for API responses."""
+ return {
+ "index": self.index,
+ "id": self.id,
+ self.type: {
+ "arguments": self.arguments,
+ "name": self.name,
+ },
+ "type": self.type,
+ }
diff --git a/reme_ai/core/schema/vector_node.py b/reme_ai/core/schema/vector_node.py
new file mode 100644
index 00000000..937ef4be
--- /dev/null
+++ b/reme_ai/core/schema/vector_node.py
@@ -0,0 +1,15 @@
+"""Defines the data structure for individual vector embedding nodes within a retrieval system."""
+
+from typing import List, Dict
+from uuid import uuid4
+
+from pydantic import BaseModel, Field
+
+
+class VectorNode(BaseModel):
+ """Represents a discrete unit of text content paired with its corresponding vector embedding and metadata."""
+
+ vector_id: str = Field(default_factory=lambda: uuid4().hex)
+ content: str = Field(default="")
+ vector: List[float] | None = Field(default=None)
+ metadata: Dict[str, str | bool | int | float] = Field(default_factory=dict)
diff --git a/reme_ai/core/utils/__init__.py b/reme_ai/core/utils/__init__.py
new file mode 100644
index 00000000..ed1916b2
--- /dev/null
+++ b/reme_ai/core/utils/__init__.py
@@ -0,0 +1,14 @@
+"""utils"""
+
+from .case_converter import snake_to_camel, camel_to_snake
+from .env_utils import load_env
+from .singleton import singleton
+from .timer import timer
+
+__all__ = [
+ "snake_to_camel",
+ "camel_to_snake",
+ "load_env",
+ "singleton",
+ "timer",
+]
diff --git a/reme_ai/core/utils/case_converter.py b/reme_ai/core/utils/case_converter.py
new file mode 100644
index 00000000..28815282
--- /dev/null
+++ b/reme_ai/core/utils/case_converter.py
@@ -0,0 +1,28 @@
+"""Case conversion utility for PascalCase, camelCase, and snake_case."""
+
+import re
+
+# Acronyms that should remain uppercase in Pascal/camelCase
+_ACRONYMS = {"LLM", "API", "URL", "HTTP", "JSON", "XML", "AI", "MCP"}
+_ACRONYM_MAP = {word.lower(): word for word in _ACRONYMS}
+
+
+def camel_to_snake(content: str) -> str:
+ """Convert PascalCase or camelCase to snake_case."""
+ # Normalize acronyms to title case (e.g., LLM -> Llm) to assist regex splitting
+ for word in _ACRONYMS:
+ content = content.replace(word, word.capitalize())
+
+ # Insert underscores between case transitions and convert to lowercase
+ return re.sub(r"(? str:
+ """Convert snake_case to PascalCase (preserving defined acronyms)."""
+ return "".join(_ACRONYM_MAP.get(part.lower(), part.capitalize()) for part in content.split("_") if part)
+
+
+if __name__ == "__main__":
+ # Quick verification
+ print(camel_to_snake("OpenAILLMClient")) # open_ai_llm_client
+ print(snake_to_camel("open_ai_llm_client")) # OpenAILLMClient
diff --git a/reme_ai/core/utils/env_utils.py b/reme_ai/core/utils/env_utils.py
new file mode 100644
index 00000000..433039cd
--- /dev/null
+++ b/reme_ai/core/utils/env_utils.py
@@ -0,0 +1,63 @@
+"""Environment variable loader utility for managing .env files."""
+
+import os
+from pathlib import Path
+
+from loguru import logger
+
+# Global flag to ensure environment is loaded only once
+_ENV_LOADED = False
+
+
+def _parse_env_file(path: Path) -> None:
+ """Parse and inject key-value pairs from a .env file into os.environ."""
+ try:
+ with path.open(encoding="utf-8") as file:
+ for line in file:
+ line = line.strip()
+ if not line or line.startswith("#"):
+ continue
+
+ if "=" in line:
+ key, value = line.split("=", 1)
+ # Strip whitespace and common quotes
+ os.environ[key.strip()] = value.strip().strip("'\"")
+ except PermissionError as err:
+ logger.warning(f"Permission denied for {path}: {err}")
+ except Exception as err:
+ logger.error(f"Failed to load {path}: {err}")
+ raise
+
+
+def load_env(path: str | Path | None = None, enable_log: bool = True) -> None:
+ """Search and load the .env file into the system environment."""
+ global _ENV_LOADED # pylint: disable=global-statement
+ if _ENV_LOADED:
+ return
+
+ if path:
+ path = Path(path)
+ if path.exists():
+ _parse_env_file(path)
+ _ENV_LOADED = True
+ else:
+ logger.warning(f".env not found at: {path}")
+ return
+
+ # Search current directory and up to 5 levels of parents
+ for directory in [Path.cwd(), *Path.cwd().parents[:5]]:
+ env_path = directory / ".env"
+ if env_path.exists():
+ if enable_log:
+ logger.info(f"Loading environment from: {env_path}")
+ _parse_env_file(env_path)
+ _ENV_LOADED = True
+ return
+
+ logger.warning(".env file not found in search path")
+
+
+def reset_env_flag() -> None:
+ """Reset the internal load state flag."""
+ global _ENV_LOADED # pylint: disable=global-statement
+ _ENV_LOADED = False
diff --git a/reme_ai/core/utils/singleton.py b/reme_ai/core/utils/singleton.py
new file mode 100644
index 00000000..0f6c5907
--- /dev/null
+++ b/reme_ai/core/utils/singleton.py
@@ -0,0 +1,17 @@
+"""Module providing a decorator to implement the Singleton design pattern."""
+
+
+def singleton(cls):
+ """A class decorator that ensures only one instance of a class exists."""
+
+ # Dictionary to cache the single instance of the class
+ _instance = {}
+
+ def _singleton(*args, **kwargs):
+ """Return the existing instance or create a new one if it doesn't exist."""
+ if cls not in _instance:
+ # Create and store the instance if it's the first call
+ _instance[cls] = cls(*args, **kwargs)
+ return _instance[cls]
+
+ return _singleton
diff --git a/reme_ai/core/utils/timer.py b/reme_ai/core/utils/timer.py
new file mode 100644
index 00000000..f03225fa
--- /dev/null
+++ b/reme_ai/core/utils/timer.py
@@ -0,0 +1,66 @@
+"""
+Utility module for timing function execution with log metadata preservation.
+"""
+
+import functools
+import inspect
+import time
+from typing import Any, Callable, TypeVar, cast
+
+from loguru import logger
+
+# Type variable to preserve the signature of the decorated callable
+F = TypeVar("F", bound=Callable[..., Any])
+
+
+def timer(func: F) -> F:
+ """
+ Decorator that logs execution time and patches log records with original function metadata.
+ """
+ # Extract original function metadata to ensure logs point to the correct source
+ func_name = func.__name__
+ try:
+ # Retrieve the source file path and the starting line number
+ file_path = inspect.getsourcefile(func) or "unknown"
+ _, line_no = inspect.getsourcelines(func)
+ except Exception:
+ file_path = "unknown"
+ line_no = 0
+
+ def patcher(record):
+ """Modifies the log record to reflect the decorated function's location."""
+ record["function"] = func_name
+ record["file"].name = file_path.split("/")[-1]
+ record["file"].path = file_path
+ record["line"] = line_no
+
+ @functools.wraps(func)
+ async def async_wrapper(*args: Any, **kwargs: Any) -> Any:
+ """Timer wrapper for asynchronous functions."""
+ start_time = time.perf_counter()
+ try:
+ return await func(*args, **kwargs)
+ finally:
+ duration = time.perf_counter() - start_time
+ # Use patch to inject metadata instead of relying on stack depth
+ logger.patch(patcher).info(
+ "========== cost={:.6f}s ==========",
+ duration,
+ )
+
+ @functools.wraps(func)
+ def sync_wrapper(*args: Any, **kwargs: Any) -> Any:
+ """Timer wrapper for synchronous functions."""
+ start_time = time.perf_counter()
+ try:
+ return func(*args, **kwargs)
+ finally:
+ duration = time.perf_counter() - start_time
+ logger.patch(patcher).info(
+ "========== cost={:.6f}s ==========",
+ duration,
+ )
+
+ if inspect.iscoroutinefunction(func):
+ return cast(F, async_wrapper)
+ return cast(F, sync_wrapper)
diff --git a/tests/test_llm.py b/tests/test_llm.py
new file mode 100644
index 00000000..308a69a9
--- /dev/null
+++ b/tests/test_llm.py
@@ -0,0 +1,422 @@
+"""
+Async unit tests for LLM classes (OpenAILLM and LiteLLM) covering:
+- Async non-streaming chat
+- Async chat with stream print
+- Async streaming chat
+- Async chat with tools
+
+Usage:
+ python test_llm.py --openai # Test OpenAILLM only
+ python test_llm.py --litellm # Test LiteLLM only
+ python test_llm.py --all # Test both LLMs
+"""
+
+# flake8: noqa: E402
+# pylint: disable=C0413
+
+import asyncio
+import argparse
+from typing import Type
+
+from reme_ai.core.utils import load_env
+
+load_env()
+
+from reme_ai.core.llm import OpenAILLM, LiteLLM, BaseLLM
+from reme_ai.core.schema import Message, ToolCall
+from reme_ai.core.enumeration import Role, ChunkEnum
+
+
+def get_llm(llm_class: Type[BaseLLM]) -> BaseLLM:
+ """Create and return an LLM instance."""
+ return llm_class(
+ model_name="qwen3-30b-a3b-instruct-2507",
+ max_retries=2,
+ raise_exception=True,
+ )
+
+
+def get_multi_turn_messages() -> list[Message]:
+ """Create multi-turn conversation messages for testing."""
+ return [
+ Message(
+ role=Role.SYSTEM,
+ content="You are a helpful AI assistant with expertise in mathematics, science, and general knowledge.",
+ ),
+ Message(
+ role=Role.USER,
+ content="Hello! I'm working on a science project about renewable energy. "
+ "Can you help me understand the basics?",
+ ),
+ Message(
+ role=Role.ASSISTANT,
+ content="Of course! I'd be happy to help. Renewable energy comes from sources that naturally replenish, "
+ "like solar, wind, hydro, geothermal, and biomass. What specific aspect would you like to explore?",
+ ),
+ Message(
+ role=Role.USER,
+ content="I'm particularly interested in solar energy. Can you explain how solar panels work and calculate "
+ "how much energy a typical home solar system might produce?",
+ ),
+ Message(
+ role=Role.ASSISTANT,
+ content="Solar panels work through photovoltaic cells that convert sunlight into electricity. "
+ "When photons hit the silicon cells, they knock electrons loose, creating an electric current."
+ "\n\nFor energy calculation: A typical home solar system is 5-10kW. With average 4-5 peak sun "
+ "hours per day, a 6kW system would produce approximately 24-30 kWh daily, or 720-900 kWh monthly.",
+ ),
+ Message(
+ role=Role.USER,
+ content="That's helpful! Now, given that calculation, if electricity costs $0.12 per kWh, "
+ "estimate the annual savings. Also, briefly mention what factors might affect this.",
+ ),
+ ]
+
+
+def get_test_tools() -> list[ToolCall]:
+ """Create comprehensive test tools for tool calling."""
+ return [
+ ToolCall(
+ **{
+ "type": "function",
+ "function": {
+ "name": "calculate_energy_savings",
+ "description": "Calculate annual energy savings based on solar production and electricity rates",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "monthly_kwh": {
+ "type": "number",
+ "description": "Monthly energy production in kWh",
+ },
+ "electricity_rate": {
+ "type": "number",
+ "description": "Electricity cost per kWh in dollars",
+ },
+ "system_efficiency": {
+ "type": "number",
+ "description": "System efficiency factor (0-1), defaults to 0.85",
+ },
+ },
+ "required": ["monthly_kwh", "electricity_rate"],
+ },
+ },
+ },
+ ),
+ ToolCall(
+ **{
+ "type": "function",
+ "function": {
+ "name": "get_weather_data",
+ "description": "Get current weather and solar irradiance data for a location",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "location": {
+ "type": "string",
+ "description": "City name or coordinates, e.g., 'San Francisco' or '37.7749,-122.4194'",
+ },
+ "include_forecast": {
+ "type": "boolean",
+ "description": "Whether to include 7-day forecast",
+ },
+ "unit": {
+ "type": "string",
+ "description": "Temperature unit",
+ "enum": ["celsius", "fahrenheit"],
+ },
+ },
+ "required": ["location"],
+ },
+ },
+ },
+ ),
+ ToolCall(
+ **{
+ "type": "function",
+ "function": {
+ "name": "analyze_panel_efficiency",
+ "description": "Analyze solar panel efficiency based on various environmental factors",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "panel_type": {
+ "type": "string",
+ "description": "Type of solar panel",
+ "enum": ["monocrystalline", "polycrystalline", "thin-film"],
+ },
+ "temperature": {
+ "type": "number",
+ "description": "Ambient temperature in Celsius",
+ },
+ "age_years": {
+ "type": "number",
+ "description": "Age of the panel in years",
+ },
+ },
+ "required": ["panel_type", "temperature"],
+ },
+ },
+ },
+ ),
+ ]
+
+
+def get_tool_test_messages() -> list[Message]:
+ """Create multi-turn messages that should trigger tool calling."""
+ return [
+ Message(
+ role=Role.SYSTEM,
+ content="You are a helpful assistant with access to weather and energy calculation tools. Use them when "
+ "appropriate.",
+ ),
+ Message(
+ role=Role.USER,
+ content="I'm planning to install solar panels in San Francisco. Can you help me understand the weather "
+ "patterns there?",
+ ),
+ Message(
+ role=Role.ASSISTANT,
+ content="I'd be happy to help! San Francisco has a Mediterranean climate with "
+ "mild temperatures year-round. Let me get the current weather data for you.",
+ ),
+ Message(
+ role=Role.USER,
+ content="Great! Also, I'm considering monocrystalline panels. If my system produces 800 kWh monthly "
+ "and electricity costs $0.15 per kWh, what would be my annual savings?",
+ ),
+ ]
+
+
+async def test_async_chat(llm_class: Type[BaseLLM], llm_name: str):
+ """Test asynchronous non-streaming chat with multi-turn conversation."""
+ print(f"\n{'='*60}")
+ print(f"Testing {llm_name}: Async Non-Streaming Chat")
+ print(f"{'='*60}")
+
+ llm = get_llm(llm_class)
+ messages = get_multi_turn_messages()
+
+ print(f"Input: {len(messages)} messages in conversation")
+ print(f"Last user message: {messages[-1].content[:100]}...")
+
+ response = await llm.chat(messages=messages)
+
+ assert response is not None, f"{llm_name}: Response is None"
+ assert response.role == Role.ASSISTANT, f"{llm_name}: Wrong role"
+ assert isinstance(response.content, str), f"{llm_name}: Content is not string"
+ assert len(response.content) > 0, f"{llm_name}: Empty response"
+
+ print(f"\nResponse preview: {response.content[:200]}...")
+ print(f"Full response length: {len(response.content)} characters")
+ print(f"\nFull message:\n{response.simple_dump()}")
+
+ await llm.close()
+ print(f"✓ PASSED: {llm_name} async chat")
+
+
+async def test_async_chat_with_stream_print(llm_class: Type[BaseLLM], llm_name: str):
+ """Test asynchronous chat with stream print enabled."""
+ print(f"\n{'='*60}")
+ print(f"Testing {llm_name}: Async Chat with Stream Print")
+ print(f"{'='*60}")
+
+ llm = get_llm(llm_class)
+ messages = get_multi_turn_messages()
+
+ print(f"Input: {len(messages)} messages in conversation")
+ print(f"Last user message: {messages[-1].content[:100]}...")
+ print("\nStreaming output:")
+ print("-" * 60)
+
+ response = await llm.chat(messages=messages, enable_stream_print=True)
+
+ print("\n" + "-" * 60)
+
+ assert response is not None, f"{llm_name}: Response is None"
+ assert response.role == Role.ASSISTANT, f"{llm_name}: Wrong role"
+ assert isinstance(response.content, str), f"{llm_name}: Content is not string"
+ assert len(response.content) > 0, f"{llm_name}: Empty response"
+
+ print(f"\nFull message:\n{response.simple_dump()}")
+
+ await llm.close()
+ print(f"✓ PASSED: {llm_name} async chat with stream print")
+
+
+async def test_async_stream_chat(llm_class: Type[BaseLLM], llm_name: str):
+ """Test asynchronous streaming chat."""
+ print(f"\n{'='*60}")
+ print(f"Testing {llm_name}: Async Streaming Chat")
+ print(f"{'='*60}")
+
+ llm = get_llm(llm_class)
+ messages = get_multi_turn_messages()
+
+ print(f"Input: {len(messages)} messages in conversation")
+ print(f"Last user message: {messages[-1].content[:100]}...")
+ print("\nStreaming chunks:")
+ print("-" * 60)
+
+ chunks = []
+ answer_content = ""
+
+ async for chunk in llm.stream_chat(messages=messages):
+ chunks.append(chunk)
+ if chunk.chunk_type == ChunkEnum.ANSWER:
+ answer_content += chunk.chunk
+ print(chunk.chunk, end="", flush=True)
+
+ print("\n" + "-" * 60)
+
+ assert len(chunks) > 0, f"{llm_name}: No chunks received"
+ assert len(answer_content) > 0, f"{llm_name}: Empty answer content"
+
+ # Check that we received at least one ANSWER or USAGE chunk
+ chunk_types = [c.chunk_type for c in chunks]
+ assert ChunkEnum.ANSWER in chunk_types or ChunkEnum.USAGE in chunk_types, f"{llm_name}: No ANSWER or USAGE chunks"
+
+ # Print the final assembled message
+ if chunks and hasattr(chunks[-1], "message") and chunks[-1].message:
+ print(f"\nFull message:\n{chunks[-1].message.simple_dump()}")
+
+ print(f"\nTotal chunks: {len(chunks)}")
+ print(f"Answer length: {len(answer_content)} characters")
+
+ await llm.close()
+ print(f"✓ PASSED: {llm_name} async streaming chat")
+
+
+async def test_async_chat_with_tools(llm_class: Type[BaseLLM], llm_name: str):
+ """Test asynchronous chat with tool calling."""
+ print(f"\n{'='*60}")
+ print(f"Testing {llm_name}: Async Chat with Tools")
+ print(f"{'='*60}")
+
+ llm = get_llm(llm_class)
+ messages = get_tool_test_messages()
+ tools = get_test_tools()
+
+ print(f"Input: {len(messages)} messages, {len(tools)} tools available")
+ print(f"Tools: {[tool.name for tool in tools]}")
+ print(f"Last user message: {messages[-1].content[:100]}...")
+
+ response = await llm.chat(messages=messages, tools=tools)
+
+ assert response is not None, f"{llm_name}: Response is None"
+ assert response.role == Role.ASSISTANT, f"{llm_name}: Wrong role"
+ # Response should contain either content or tool_calls
+ assert response.content or response.tool_calls, f"{llm_name}: No content or tool_calls"
+
+ if response.tool_calls:
+ print(f"\n✓ Tool calls detected: {len(response.tool_calls)}")
+ for i, tool_call in enumerate(response.tool_calls, 1):
+ print(f"\n Tool call #{i}:")
+ print(f" - Name: {tool_call.name}")
+ print(f" - Arguments: {tool_call.arguments}")
+ # Validate that arguments are valid JSON
+ assert tool_call.check_argument(), f"{llm_name}: Invalid tool arguments"
+ print(" - ✓ Arguments validated")
+ else:
+ print("\n⚠ No tool calls (response with text instead)")
+ print(f"Response preview: {response.content[:200]}...")
+
+ print(f"\nFull message:\n{response.simple_dump()}")
+
+ await llm.close()
+ print(f"✓ PASSED: {llm_name} async chat with tools")
+
+
+async def run_all_tests_for_llm(llm_class: Type[BaseLLM], llm_name: str):
+ """Run all tests for a specific LLM class."""
+ print(f"\n\n{'#'*60}")
+ print(f"# Running all tests for: {llm_name}")
+ print(f"{'#'*60}")
+
+ await test_async_chat(llm_class, llm_name)
+ await test_async_chat_with_stream_print(llm_class, llm_name)
+ await test_async_stream_chat(llm_class, llm_name)
+ await test_async_chat_with_tools(llm_class, llm_name)
+
+ print(f"\n{'='*60}")
+ print(f"✓ All tests passed for {llm_name}!")
+ print(f"{'='*60}")
+
+
+async def main():
+ """Main entry point for running tests."""
+ parser = argparse.ArgumentParser(
+ description="Run async LLM tests",
+ formatter_class=argparse.RawDescriptionHelpFormatter,
+ epilog="""
+Examples:
+ python test_llm.py --openai # Test OpenAILLM only
+ python test_llm.py --litellm # Test LiteLLM only
+ python test_llm.py --all # Test both LLMs
+ """,
+ )
+ parser.add_argument(
+ "--openai",
+ action="store_true",
+ help="Test OpenAILLM",
+ )
+ parser.add_argument(
+ "--litellm",
+ action="store_true",
+ help="Test LiteLLM",
+ )
+ parser.add_argument(
+ "--all",
+ action="store_true",
+ help="Run tests for all available LLMs",
+ )
+
+ args = parser.parse_args()
+
+ # Determine which LLMs to test
+ llms_to_test = []
+
+ if args.all:
+ llms_to_test = [
+ (OpenAILLM, "OpenAILLM"),
+ (LiteLLM, "LiteLLM"),
+ ]
+ elif args.openai and args.litellm:
+ llms_to_test = [
+ (OpenAILLM, "OpenAILLM"),
+ (LiteLLM, "LiteLLM"),
+ ]
+ elif args.openai:
+ llms_to_test = [(OpenAILLM, "OpenAILLM")]
+ elif args.litellm:
+ llms_to_test = [(LiteLLM, "LiteLLM")]
+ else:
+ # Default to all LLMs if no argument provided
+ llms_to_test = [
+ (OpenAILLM, "OpenAILLM"),
+ (LiteLLM, "LiteLLM"),
+ ]
+ print("No LLM specified, defaulting to --all (testing all LLMs)")
+ print("Use --openai or --litellm to test a specific one\n")
+
+ # Run tests for each LLM
+ for llm_class, llm_name in llms_to_test:
+ try:
+ await run_all_tests_for_llm(llm_class, llm_name)
+ except Exception as e:
+ print(f"\n✗ FAILED: {llm_name} tests failed with error:")
+ print(f" {type(e).__name__}: {e}")
+ raise
+
+ # Final summary
+ print(f"\n\n{'#'*60}")
+ print("# TEST SUMMARY")
+ print(f"{'#'*60}")
+ print(f"✓ All tests passed for {len(llms_to_test)} LLM(s):")
+ for _, llm_name in llms_to_test:
+ print(f" - {llm_name}")
+ print(f"{'#'*60}\n")
+
+
+if __name__ == "__main__":
+ asyncio.run(main())
diff --git a/tests/test_llm_sync.py b/tests/test_llm_sync.py
new file mode 100644
index 00000000..98751f87
--- /dev/null
+++ b/tests/test_llm_sync.py
@@ -0,0 +1,421 @@
+"""
+Sync unit tests for LLM classes (OpenAILLM and LiteLLM) covering:
+- Sync non-streaming chat
+- Sync chat with stream print
+- Sync streaming chat
+- Sync chat with tools
+
+Usage:
+ python test_llm_sync.py --openai # Test OpenAILLM only
+ python test_llm_sync.py --litellm # Test LiteLLM only
+ python test_llm_sync.py --all # Test both LLMs
+"""
+
+# flake8: noqa: E402
+# pylint: disable=C0413
+
+import argparse
+from typing import Type
+
+from reme_ai.core.utils import load_env
+
+load_env()
+
+from reme_ai.core.llm import OpenAILLMSync, LiteLLMSync, BaseLLM
+from reme_ai.core.schema import Message, ToolCall
+from reme_ai.core.enumeration import Role, ChunkEnum
+
+
+def get_llm(llm_class: Type[BaseLLM]) -> BaseLLM:
+ """Create and return an LLM instance."""
+ return llm_class(
+ model_name="qwen3-30b-a3b-instruct-2507",
+ max_retries=2,
+ raise_exception=True,
+ )
+
+
+def get_multi_turn_messages() -> list[Message]:
+ """Create multi-turn conversation messages for testing."""
+ return [
+ Message(
+ role=Role.SYSTEM,
+ content="You are a helpful AI assistant with expertise in mathematics, science, and general knowledge.",
+ ),
+ Message(
+ role=Role.USER,
+ content="Hello! I'm working on a science project about renewable energy. Can you help me understand "
+ "the basics?",
+ ),
+ Message(
+ role=Role.ASSISTANT,
+ content="Of course! I'd be happy to help. Renewable energy comes from sources that naturally replenish, "
+ "like solar, wind, hydro, geothermal, and biomass. What specific aspect would you like to explore?",
+ ),
+ Message(
+ role=Role.USER,
+ content="I'm particularly interested in solar energy. Can you explain how solar panels work "
+ "and calculate how much energy a typical home solar system might produce?",
+ ),
+ Message(
+ role=Role.ASSISTANT,
+ content="Solar panels work through photovoltaic cells that convert sunlight into electricity. "
+ "When photons hit the silicon cells, they knock electrons loose, creating an electric current."
+ "\n\nFor energy calculation: A typical home solar system is 5-10kW. With average 4-5 peak sun "
+ "hours per day, a 6kW system would produce approximately 24-30 kWh daily, or 720-900 kWh monthly.",
+ ),
+ Message(
+ role=Role.USER,
+ content="That's helpful! Now, given that calculation, if electricity costs $0.12 per kWh, "
+ "estimate the annual savings. Also, briefly mention what factors might affect this.",
+ ),
+ ]
+
+
+def get_test_tools() -> list[ToolCall]:
+ """Create comprehensive test tools for tool calling."""
+ return [
+ ToolCall(
+ **{
+ "type": "function",
+ "function": {
+ "name": "calculate_energy_savings",
+ "description": "Calculate annual energy savings based on solar production and electricity rates",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "monthly_kwh": {
+ "type": "number",
+ "description": "Monthly energy production in kWh",
+ },
+ "electricity_rate": {
+ "type": "number",
+ "description": "Electricity cost per kWh in dollars",
+ },
+ "system_efficiency": {
+ "type": "number",
+ "description": "System efficiency factor (0-1), defaults to 0.85",
+ },
+ },
+ "required": ["monthly_kwh", "electricity_rate"],
+ },
+ },
+ },
+ ),
+ ToolCall(
+ **{
+ "type": "function",
+ "function": {
+ "name": "get_weather_data",
+ "description": "Get current weather and solar irradiance data for a location",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "location": {
+ "type": "string",
+ "description": "City name or coordinates, e.g., 'San Francisco' or '37.7749,-122.4194'",
+ },
+ "include_forecast": {
+ "type": "boolean",
+ "description": "Whether to include 7-day forecast",
+ },
+ "unit": {
+ "type": "string",
+ "description": "Temperature unit",
+ "enum": ["celsius", "fahrenheit"],
+ },
+ },
+ "required": ["location"],
+ },
+ },
+ },
+ ),
+ ToolCall(
+ **{
+ "type": "function",
+ "function": {
+ "name": "analyze_panel_efficiency",
+ "description": "Analyze solar panel efficiency based on various environmental factors",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "panel_type": {
+ "type": "string",
+ "description": "Type of solar panel",
+ "enum": ["monocrystalline", "polycrystalline", "thin-film"],
+ },
+ "temperature": {
+ "type": "number",
+ "description": "Ambient temperature in Celsius",
+ },
+ "age_years": {
+ "type": "number",
+ "description": "Age of the panel in years",
+ },
+ },
+ "required": ["panel_type", "temperature"],
+ },
+ },
+ },
+ ),
+ ]
+
+
+def get_tool_test_messages() -> list[Message]:
+ """Create multi-turn messages that should trigger tool calling."""
+ return [
+ Message(
+ role=Role.SYSTEM,
+ content="You are a helpful assistant with access to weather and energy calculation tools. "
+ "Use them when appropriate.",
+ ),
+ Message(
+ role=Role.USER,
+ content="I'm planning to install solar panels in San Francisco. Can you help me understand the "
+ "weather patterns there?",
+ ),
+ Message(
+ role=Role.ASSISTANT,
+ content="I'd be happy to help! San Francisco has a Mediterranean climate with mild temperatures "
+ "year-round. Let me get the current weather data for you.",
+ ),
+ Message(
+ role=Role.USER,
+ content="Great! Also, I'm considering monocrystalline panels. If my system produces 800 kWh monthly "
+ "and electricity costs $0.15 per kWh, what would be my annual savings?",
+ ),
+ ]
+
+
+def test_sync_chat(llm_class: Type[BaseLLM], llm_name: str):
+ """Test synchronous non-streaming chat with multi-turn conversation."""
+ print(f"\n{'='*60}")
+ print(f"Testing {llm_name}: Sync Non-Streaming Chat")
+ print(f"{'='*60}")
+
+ llm = get_llm(llm_class)
+ messages = get_multi_turn_messages()
+
+ print(f"Input: {len(messages)} messages in conversation")
+ print(f"Last user message: {messages[-1].content[:100]}...")
+
+ response = llm.chat_sync(messages=messages)
+
+ assert response is not None, f"{llm_name}: Response is None"
+ assert response.role == Role.ASSISTANT, f"{llm_name}: Wrong role"
+ assert isinstance(response.content, str), f"{llm_name}: Content is not string"
+ assert len(response.content) > 0, f"{llm_name}: Empty response"
+
+ print(f"\nResponse preview: {response.content[:200]}...")
+ print(f"Full response length: {len(response.content)} characters")
+ print(f"\nFull message:\n{response.simple_dump()}")
+
+ llm.close_sync()
+ print(f"✓ PASSED: {llm_name} sync chat")
+
+
+def test_sync_chat_with_stream_print(llm_class: Type[BaseLLM], llm_name: str):
+ """Test synchronous chat with stream print enabled."""
+ print(f"\n{'='*60}")
+ print(f"Testing {llm_name}: Sync Chat with Stream Print")
+ print(f"{'='*60}")
+
+ llm = get_llm(llm_class)
+ messages = get_multi_turn_messages()
+
+ print(f"Input: {len(messages)} messages in conversation")
+ print(f"Last user message: {messages[-1].content[:100]}...")
+ print("\nStreaming output:")
+ print("-" * 60)
+
+ response = llm.chat_sync(messages=messages, enable_stream_print=True)
+
+ print("\n" + "-" * 60)
+
+ assert response is not None, f"{llm_name}: Response is None"
+ assert response.role == Role.ASSISTANT, f"{llm_name}: Wrong role"
+ assert isinstance(response.content, str), f"{llm_name}: Content is not string"
+ assert len(response.content) > 0, f"{llm_name}: Empty response"
+
+ print(f"\nFull message:\n{response.simple_dump()}")
+
+ llm.close_sync()
+ print(f"✓ PASSED: {llm_name} sync chat with stream print")
+
+
+def test_sync_stream_chat(llm_class: Type[BaseLLM], llm_name: str):
+ """Test synchronous streaming chat."""
+ print(f"\n{'='*60}")
+ print(f"Testing {llm_name}: Sync Streaming Chat")
+ print(f"{'='*60}")
+
+ llm = get_llm(llm_class)
+ messages = get_multi_turn_messages()
+
+ print(f"Input: {len(messages)} messages in conversation")
+ print(f"Last user message: {messages[-1].content[:100]}...")
+ print("\nStreaming chunks:")
+ print("-" * 60)
+
+ chunks = []
+ answer_content = ""
+
+ for chunk in llm.stream_chat_sync(messages=messages):
+ chunks.append(chunk)
+ if chunk.chunk_type == ChunkEnum.ANSWER:
+ answer_content += chunk.chunk
+ print(chunk.chunk, end="", flush=True)
+
+ print("\n" + "-" * 60)
+
+ assert len(chunks) > 0, f"{llm_name}: No chunks received"
+ assert len(answer_content) > 0, f"{llm_name}: Empty answer content"
+
+ # Check that we received at least one ANSWER or USAGE chunk
+ chunk_types = [c.chunk_type for c in chunks]
+ assert ChunkEnum.ANSWER in chunk_types or ChunkEnum.USAGE in chunk_types, f"{llm_name}: No ANSWER or USAGE chunks"
+
+ # Print the final assembled message
+ if chunks and hasattr(chunks[-1], "message") and chunks[-1].message:
+ print(f"\nFull message:\n{chunks[-1].message.simple_dump()}")
+
+ print(f"\nTotal chunks: {len(chunks)}")
+ print(f"Answer length: {len(answer_content)} characters")
+
+ llm.close_sync()
+ print(f"✓ PASSED: {llm_name} sync streaming chat")
+
+
+def test_sync_chat_with_tools(llm_class: Type[BaseLLM], llm_name: str):
+ """Test synchronous chat with tool calling."""
+ print(f"\n{'='*60}")
+ print(f"Testing {llm_name}: Sync Chat with Tools")
+ print(f"{'='*60}")
+
+ llm = get_llm(llm_class)
+ messages = get_tool_test_messages()
+ tools = get_test_tools()
+
+ print(f"Input: {len(messages)} messages, {len(tools)} tools available")
+ print(f"Tools: {[tool.name for tool in tools]}")
+ print(f"Last user message: {messages[-1].content[:100]}...")
+
+ response = llm.chat_sync(messages=messages, tools=tools)
+
+ assert response is not None, f"{llm_name}: Response is None"
+ assert response.role == Role.ASSISTANT, f"{llm_name}: Wrong role"
+ # Response should contain either content or tool_calls
+ assert response.content or response.tool_calls, f"{llm_name}: No content or tool_calls"
+
+ if response.tool_calls:
+ print(f"\n✓ Tool calls detected: {len(response.tool_calls)}")
+ for i, tool_call in enumerate(response.tool_calls, 1):
+ print(f"\n Tool call #{i}:")
+ print(f" - Name: {tool_call.name}")
+ print(f" - Arguments: {tool_call.arguments}")
+ # Validate that arguments are valid JSON
+ assert tool_call.check_argument(), f"{llm_name}: Invalid tool arguments"
+ print(" - ✓ Arguments validated")
+ else:
+ print("\n⚠ No tool calls (response with text instead)")
+ print(f"Response preview: {response.content[:200]}...")
+
+ print(f"\nFull message:\n{response.simple_dump()}")
+
+ llm.close_sync()
+ print(f"✓ PASSED: {llm_name} sync chat with tools")
+
+
+def run_all_tests_for_llm(llm_class: Type[BaseLLM], llm_name: str):
+ """Run all tests for a specific LLM class."""
+ print(f"\n\n{'#'*60}")
+ print(f"# Running all tests for: {llm_name}")
+ print(f"{'#'*60}")
+
+ test_sync_chat(llm_class, llm_name)
+ test_sync_chat_with_stream_print(llm_class, llm_name)
+ test_sync_stream_chat(llm_class, llm_name)
+ test_sync_chat_with_tools(llm_class, llm_name)
+
+ print(f"\n{'='*60}")
+ print(f"✓ All tests passed for {llm_name}!")
+ print(f"{'='*60}")
+
+
+def main():
+ """Main entry point for running tests."""
+ parser = argparse.ArgumentParser(
+ description="Run sync LLM tests",
+ formatter_class=argparse.RawDescriptionHelpFormatter,
+ epilog="""
+Examples:
+ python test_llm_sync.py --openai # Test OpenAILLM only
+ python test_llm_sync.py --litellm # Test LiteLLM only
+ python test_llm_sync.py --all # Test both LLMs
+ """,
+ )
+ parser.add_argument(
+ "--openai",
+ action="store_true",
+ help="Test OpenAILLM",
+ )
+ parser.add_argument(
+ "--litellm",
+ action="store_true",
+ help="Test LiteLLM",
+ )
+ parser.add_argument(
+ "--all",
+ action="store_true",
+ help="Run tests for all available LLMs",
+ )
+
+ args = parser.parse_args()
+
+ # Determine which LLMs to test
+ llms_to_test = []
+
+ if args.all:
+ llms_to_test = [
+ (OpenAILLMSync, "OpenAILLMSync"),
+ (LiteLLMSync, "LiteLLMSync"),
+ ]
+ elif args.openai and args.litellm:
+ llms_to_test = [
+ (OpenAILLMSync, "OpenAILLMSync"),
+ (LiteLLMSync, "LiteLLMSync"),
+ ]
+ elif args.openai:
+ llms_to_test = [(OpenAILLMSync, "OpenAILLMSync")]
+ elif args.litellm:
+ llms_to_test = [(LiteLLMSync, "LiteLLMSync")]
+ else:
+ # Default to all LLMs if no argument provided
+ llms_to_test = [
+ (OpenAILLMSync, "OpenAILLMSync"),
+ (LiteLLMSync, "LiteLLMSync"),
+ ]
+ print("No LLM specified, defaulting to --all (testing all LLMs)")
+ print("Use --openai or --litellm to test a specific one\n")
+
+ # Run tests for each LLM
+ for llm_class, llm_name in llms_to_test:
+ try:
+ run_all_tests_for_llm(llm_class, llm_name)
+ except Exception as e:
+ print(f"\n✗ FAILED: {llm_name} tests failed with error:")
+ print(f" {type(e).__name__}: {e}")
+ raise
+
+ # Final summary
+ print(f"\n\n{'#'*60}")
+ print("# TEST SUMMARY")
+ print(f"{'#'*60}")
+ print(f"✓ All tests passed for {len(llms_to_test)} LLM(s):")
+ for _, llm_name in llms_to_test:
+ print(f" - {llm_name}")
+ print(f"{'#'*60}\n")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/tests/test_message.py b/tests/test_message.py
new file mode 100644
index 00000000..141174c5
--- /dev/null
+++ b/tests/test_message.py
@@ -0,0 +1,195 @@
+"""Test cases for message schema and serialization."""
+
+import unittest
+
+from mcp.types import Tool
+
+from reme_ai.core.enumeration import Role
+from reme_ai.core.schema import ToolAttr, ToolCall, ContentBlock, Message
+
+
+class TestModelDefinitions(unittest.TestCase):
+ """Test suite for validating message models and their serialization methods."""
+
+ def test_tool_attr_serialization(self):
+ """Test if ToolAttr correctly dumps to JSON schema format."""
+ # Test simple string attribute with enum
+ attr = ToolAttr(
+ type="string",
+ description="The city name",
+ enum=["Beijing", "London"],
+ )
+ dump = attr.simple_input_dump()
+
+ print("\n=== ToolAttr.simple_input_dump() (string with enum) ===")
+ print(dump)
+
+ self.assertEqual(dump["type"], "string")
+ self.assertEqual(dump["enum"], ["Beijing", "London"])
+ self.assertIn("description", dump)
+
+ # Test object attribute with required child properties
+ obj_attr = ToolAttr(
+ type="object",
+ description="User information",
+ properties={
+ "name": ToolAttr(type="string", description="User name"),
+ "age": ToolAttr(type="number", description="User age"),
+ },
+ required=["name"], # 'name' is required, 'age' is optional
+ )
+ obj_dump = obj_attr.simple_input_dump()
+
+ print("\n=== ToolAttr.simple_input_dump() (object with required) ===")
+ print(obj_dump)
+
+ self.assertEqual(obj_dump["type"], "object")
+ self.assertIn("properties", obj_dump)
+ self.assertEqual(obj_dump["required"], ["name"])
+ self.assertIn("name", obj_dump["properties"])
+ self.assertIn("age", obj_dump["properties"])
+
+ def test_tool_call_initialization(self):
+ """Test if ToolCall correctly parses raw OpenAI-style tool definitions."""
+ raw_input = {
+ "type": "function",
+ "function": {
+ "name": "get_weather",
+ "description": "Check weather info",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "location": {"type": "string", "description": "City name"},
+ "unit": {"type": "string", "description": "Temperature unit"},
+ },
+ "required": ["location"],
+ },
+ },
+ }
+
+ tc = ToolCall(**raw_input)
+
+ print("\n=== ToolCall.simple_input_dump() ===")
+ print(tc.simple_input_dump())
+ print("\n=== ToolCall.simple_output_dump() ===")
+ print(tc.simple_output_dump())
+
+ self.assertEqual(tc.name, "get_weather")
+ self.assertIn("location", tc.parameters.properties)
+ self.assertIn("unit", tc.parameters.properties)
+ # Check that 'location' is in the required list at ToolCall level
+ self.assertIn("location", tc.parameters.required)
+ self.assertNotIn("unit", tc.parameters.required)
+
+ def test_tool_call_argument_parsing(self):
+ """Test JSON argument parsing and validation."""
+ tc = ToolCall(name="test", arguments='{"key": "value"}')
+
+ self.assertTrue(tc.check_argument())
+ self.assertEqual(tc.argument_dict["key"], "value")
+
+ # Test invalid JSON
+ tc.arguments = "{invalid_json}"
+ self.assertFalse(tc.check_argument())
+
+ def test_content_block_dynamic_mapping(self):
+ """Test if ContentBlock correctly identifies content based on type key."""
+ # Test Image Block
+ img_data = {"type": "image_url", "image_url": {"url": "http://test.com/a.jpg"}}
+ block = ContentBlock(**img_data)
+ self.assertEqual(block.type, "image_url")
+ self.assertEqual(block.content["url"], "http://test.com/a.jpg")
+
+ # Test Text Block
+ text_data = {"type": "text", "text": "Hello World"}
+ block = ContentBlock(**text_data)
+ self.assertEqual(block.content, "Hello World")
+
+ def test_message_simple_dump(self):
+ """Test the transformation of Message to standard API dictionary."""
+ msg = Message(
+ role=Role.ASSISTANT,
+ content="Thinking...",
+ reasoning_content="I should check the weather first.",
+ tool_calls=[ToolCall(name="get_weather", arguments='{"city": "NY"}', id="call_123")],
+ )
+
+ dump = msg.simple_dump(add_reasoning=True)
+
+ print("\n=== Message.simple_dump(add_reasoning=True) ===")
+ print(dump)
+
+ dump_no_reasoning = msg.simple_dump(add_reasoning=False)
+ print("\n=== Message.simple_dump(add_reasoning=False) ===")
+ print(dump_no_reasoning)
+
+ self.assertEqual(dump["role"], "assistant")
+ self.assertEqual(dump["reasoning_content"], "I should check the weather first.")
+ self.assertEqual(len(dump["tool_calls"]), 1)
+ self.assertEqual(dump["tool_calls"][0]["id"], "call_123")
+
+ def test_message_format_human_readable(self):
+ """Test the string representation of messages for logging/UI."""
+ msg = Message(
+ role=Role.USER,
+ content=[
+ ContentBlock(type="text", text="Look at this:"),
+ ContentBlock(type="image_url", image_url={"url": "img.png"}),
+ ],
+ )
+
+ formatted = msg.format_message(index=1, use_name=False)
+
+ self.assertIn("round1", formatted)
+ self.assertIn("user:", formatted)
+ self.assertIn("Look at this:", formatted)
+ self.assertIn("img.png", formatted)
+
+ def test_mcp_conversion(self):
+ """Test the interoperability with MCP Tool format."""
+ # Create a mock MCP Tool
+ mcp_tool = Tool(
+ name="calculator",
+ description="adds numbers",
+ inputSchema={
+ "type": "object",
+ "properties": {
+ "a": {"type": "number", "description": "First number"},
+ "b": {"type": "number", "description": "Second number"},
+ },
+ "required": ["a"],
+ },
+ )
+
+ # From MCP to ToolCall
+ tc = ToolCall.from_mcp_tool(mcp_tool)
+
+ print("\n=== ToolCall from MCP - simple_input_dump() ===")
+ print(tc.simple_input_dump())
+ print("\n=== ToolCall from MCP - simple_output_dump() ===")
+ print(tc.simple_output_dump())
+
+ self.assertEqual(tc.name, "calculator")
+ self.assertIn("a", tc.parameters.properties)
+ self.assertIn("b", tc.parameters.properties)
+ # Check that 'a' is in the required list
+ self.assertIn("a", tc.parameters.required)
+ self.assertNotIn("b", tc.parameters.required)
+
+ # From ToolCall back to MCP structure (via to_mcp_tool)
+ # Note: This checks the logic of constructing the dict for Tool(...)
+ mcp_compatible = tc.to_mcp_tool()
+
+ print("\n=== MCP Tool converted back ===")
+ print(f"Name: {mcp_compatible.name}")
+ print(f"Description: {mcp_compatible.description}")
+ print(f"InputSchema: {mcp_compatible.inputSchema}")
+
+ self.assertEqual(mcp_compatible.name, "calculator")
+ self.assertIn("a", mcp_compatible.inputSchema["properties"])
+ self.assertIn("b", mcp_compatible.inputSchema["properties"])
+ self.assertEqual(mcp_compatible.inputSchema["required"], ["a"])
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/tests/test_timer.py b/tests/test_timer.py
new file mode 100644
index 00000000..c9714e38
--- /dev/null
+++ b/tests/test_timer.py
@@ -0,0 +1,64 @@
+"""
+This module provides a suite of tests to verify universal timer decorator functionality using loguru.
+"""
+
+import asyncio
+import time
+
+from loguru import logger
+
+from reme_ai.core.utils import timer
+
+
+@timer
+def test_sync_function(seconds: float) -> str:
+ """Tests timing of a standard synchronous function."""
+ time.sleep(seconds)
+ return "sync done"
+
+
+@timer
+async def test_async_function(seconds: float) -> str:
+ """Tests timing of an asynchronous function."""
+ await asyncio.sleep(seconds)
+ return "async done"
+
+
+class TestMemberMethods:
+ """Container class to test class method decoration."""
+
+ @timer
+ def test_sync_method(self, seconds: float) -> None:
+ """Tests a synchronous instance method."""
+ time.sleep(seconds)
+
+ @timer
+ async def test_async_method(self, seconds: float) -> None:
+ """Tests an asynchronous instance method."""
+ await asyncio.sleep(seconds)
+
+
+def run_all_tests() -> None:
+ """
+ Manual test runner.
+ Notice that the logs will now point to the line numbers below
+ (where the function is actually called).
+ """
+ logger.info("Starting tests and verifying stack trace...")
+
+ # 1. Test Sync Function
+ test_sync_function(0.1)
+
+ # 2. Test Async Function
+ asyncio.run(test_async_function(0.1))
+
+ # 3. Test Class Methods
+ tester = TestMemberMethods()
+ tester.test_sync_method(0.05)
+ asyncio.run(tester.test_async_method(0.05))
+
+ logger.success("All tests completed.")
+
+
+if __name__ == "__main__":
+ run_all_tests()
diff --git a/tests/test_tool_call.py b/tests/test_tool_call.py
new file mode 100644
index 00000000..30c0a37e
--- /dev/null
+++ b/tests/test_tool_call.py
@@ -0,0 +1,360 @@
+"""simple tool call test"""
+
+import json
+
+from reme_ai.core.schema.tool_call import ToolCall
+
+
+def test_simple_schema():
+ """测试简单的工具定义:只有基本类型参数"""
+ print("\n========== 测试简单 Schema ==========")
+
+ raw_definition = {
+ "type": "function",
+ "function": {
+ "name": "get_weather",
+ "description": "获取指定城市的天气信息",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "city": {"type": "string", "description": "城市名称"},
+ "unit": {"type": "string", "description": "温度单位", "enum": ["celsius", "fahrenheit"]},
+ },
+ "required": ["city"],
+ },
+ },
+ }
+
+ # 解析
+ tool_call = ToolCall.model_validate(raw_definition)
+ print(f"工具名称: {tool_call.name}")
+ print(f"必填参数: {tool_call.parameters.required}")
+
+ # 导出并验证相等性
+ dumped_data = tool_call.simple_input_dump()
+ print(f"\n原始定义:\n{json.dumps(raw_definition, indent=2, ensure_ascii=False)}")
+ print(f"\n导出结果:\n{json.dumps(dumped_data, indent=2, ensure_ascii=False)}")
+
+ # 验证相等
+ assert dumped_data == raw_definition, "简单 Schema 导出结果与原始定义不一致"
+ print("\n✅ 简单 Schema 测试通过:raw_definition == simple_input_dump()")
+
+
+def test_medium_nested_schema():
+ """测试中等复杂度:包含一层对象嵌套"""
+ print("\n========== 测试中等复杂 Schema ==========")
+
+ raw_definition = {
+ "type": "function",
+ "function": {
+ "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"],
+ },
+ },
+ }
+
+ # 解析
+ tool_call = ToolCall.model_validate(raw_definition)
+ print(f"工具名称: {tool_call.name}")
+ print(f"根级必填项: {tool_call.parameters.required}")
+
+ customer_attr = tool_call.parameters.properties["customer"]
+ print(f"Customer 子属性: {list(customer_attr.properties.keys())}")
+ print(f"Customer 必填项: {customer_attr.required}")
+
+ # 导出并验证相等性
+ dumped_data = tool_call.simple_input_dump()
+ print(f"\n原始定义:\n{json.dumps(raw_definition, indent=2, ensure_ascii=False)}")
+ print(f"\n导出结果:\n{json.dumps(dumped_data, indent=2, ensure_ascii=False)}")
+
+ # 验证相等
+ assert dumped_data == raw_definition, "中等复杂 Schema 导出结果与原始定义不一致"
+ print("\n✅ 中等复杂 Schema 测试通过:raw_definition == simple_input_dump()")
+
+
+def test_nested_schema():
+ """测试复杂嵌套:包含对象嵌套和数组嵌套"""
+ print("\n========== 测试复杂嵌套 Schema ==========")
+
+ # 1. 模拟一个来自 LLM 或 MCP 的复杂嵌套定义
+ raw_definition = {
+ "type": "function",
+ "function": {
+ "name": "register_user",
+ "description": "注册新用户,包含复杂的元数据和标签",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "username": {"type": "string", "description": "用户名"},
+ "metadata": {
+ "type": "object",
+ "description": "用户元数据",
+ "properties": {
+ "age": {"type": "integer"},
+ "location": {"type": "string"},
+ },
+ "required": ["age"],
+ },
+ "tags": {
+ "type": "array",
+ "description": "用户标签列表",
+ "items": {
+ "type": "object",
+ "properties": {
+ "tag_id": {"type": "string"},
+ "level": {"type": "number"},
+ },
+ "required": ["tag_id"],
+ },
+ },
+ },
+ "required": ["username", "metadata"],
+ },
+ },
+ }
+
+ # 2. 解析:将原始字典转化为 ToolCall 实例
+ tool_call = ToolCall.model_validate(raw_definition)
+
+ print(f"工具名称: {tool_call.name}")
+ print(f"根级必填项: {tool_call.parameters.required}")
+
+ # 验证嵌套深度
+ metadata_attr = tool_call.parameters.properties["metadata"]
+ print(f"Metadata 子属性: {list(metadata_attr.properties.keys())}")
+ print(f"Metadata 必填项: {metadata_attr.required}")
+
+ # 3. 导出:验证 simple_input_dump 是否生成了正确的 JSON Schema
+ dumped_data = tool_call.simple_input_dump()
+
+ print(f"\n原始定义:\n{json.dumps(raw_definition, indent=2, ensure_ascii=False)}")
+ print(f"\n导出结果:\n{json.dumps(dumped_data, indent=2, ensure_ascii=False)}")
+
+ # 验证相等
+ assert dumped_data == raw_definition, "复杂嵌套 Schema 导出结果与原始定义不一致"
+ print("\n✅ 复杂嵌套 Schema 测试通过:raw_definition == simple_input_dump()")
+
+ # 4. 转换验证:测试 to_mcp_tool
+ mcp_tool = tool_call.to_mcp_tool()
+ assert mcp_tool.name == "register_user"
+ assert "properties" in mcp_tool.inputSchema["properties"]["tags"]["items"]
+ print("✅ 嵌套结构在 MCP Tool 转换中成功保留")
+
+
+def test_array_of_primitives():
+ """测试数组嵌套:数组元素为基本类型"""
+ print("\n========== 测试基本类型数组 Schema ==========")
+
+ raw_definition = {
+ "type": "function",
+ "function": {
+ "name": "batch_process",
+ "description": "批量处理文件",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "file_paths": {
+ "type": "array",
+ "description": "文件路径列表",
+ "items": {"type": "string"},
+ },
+ "priorities": {
+ "type": "array",
+ "description": "优先级列表",
+ "items": {"type": "integer"},
+ },
+ },
+ "required": ["file_paths"],
+ },
+ },
+ }
+
+ # 解析
+ tool_call = ToolCall.model_validate(raw_definition)
+ print(f"工具名称: {tool_call.name}")
+ print(f"必填参数: {tool_call.parameters.required}")
+
+ file_paths_attr = tool_call.parameters.properties["file_paths"]
+ print(f"file_paths 类型: {file_paths_attr.type}")
+ t_items_type = file_paths_attr.items.type if hasattr(file_paths_attr.items, "type") else file_paths_attr.items
+ print(f"file_paths items 类型: {t_items_type}")
+
+ # 导出并验证相等性
+ dumped_data = tool_call.simple_input_dump()
+ print(f"\n原始定义:\n{json.dumps(raw_definition, indent=2, ensure_ascii=False)}")
+ print(f"\n导出结果:\n{json.dumps(dumped_data, indent=2, ensure_ascii=False)}")
+
+ # 验证相等
+ assert dumped_data == raw_definition, "基本类型数组 Schema 导出结果与原始定义不一致"
+ print("\n✅ 基本类型数组 Schema 测试通过:raw_definition == simple_input_dump()")
+
+
+def test_deep_nested_schema():
+ """测试深层嵌套:三层以上的嵌套结构"""
+ print("\n========== 测试深层嵌套 Schema ==========")
+
+ raw_definition = {
+ "type": "function",
+ "function": {
+ "name": "create_project",
+ "description": "创建项目,包含复杂的团队和任务结构",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "project_name": {"type": "string", "description": "项目名称"},
+ "team": {
+ "type": "object",
+ "description": "团队信息",
+ "properties": {
+ "leader": {
+ "type": "object",
+ "description": "团队负责人",
+ "properties": {
+ "name": {"type": "string"},
+ "contact": {
+ "type": "object",
+ "properties": {
+ "email": {"type": "string"},
+ "phone": {"type": "string"},
+ },
+ "required": ["email"],
+ },
+ },
+ "required": ["name", "contact"],
+ },
+ "members": {
+ "type": "array",
+ "description": "团队成员列表",
+ "items": {
+ "type": "object",
+ "properties": {
+ "name": {"type": "string"},
+ "role": {"type": "string"},
+ "skills": {
+ "type": "array",
+ "items": {"type": "string"},
+ },
+ },
+ "required": ["name", "role"],
+ },
+ },
+ },
+ "required": ["leader"],
+ },
+ },
+ "required": ["project_name", "team"],
+ },
+ },
+ }
+
+ # 解析
+ tool_call = ToolCall.model_validate(raw_definition)
+ print(f"工具名称: {tool_call.name}")
+ print(f"根级必填项: {tool_call.parameters.required}")
+
+ team_attr = tool_call.parameters.properties["team"]
+ leader_attr = team_attr.properties["leader"]
+ contact_attr = leader_attr.properties["contact"]
+ print(f"Team 必填项: {team_attr.required}")
+ print(f"Leader 必填项: {leader_attr.required}")
+ print(f"Contact 必填项: {contact_attr.required}")
+
+ # 导出并验证相等性
+ dumped_data = tool_call.simple_input_dump()
+ print(f"\n原始定义:\n{json.dumps(raw_definition, indent=2, ensure_ascii=False)}")
+ print(f"\n导出结果:\n{json.dumps(dumped_data, indent=2, ensure_ascii=False)}")
+
+ # 验证相等
+ assert dumped_data == raw_definition, "深层嵌套 Schema 导出结果与原始定义不一致"
+ print("\n✅ 深层嵌套 Schema 测试通过:raw_definition == simple_input_dump()")
+
+
+def test_mixed_types_schema():
+ """测试混合类型:包含所有基本类型和嵌套类型"""
+ print("\n========== 测试混合类型 Schema ==========")
+
+ raw_definition = {
+ "type": "function",
+ "function": {
+ "name": "configure_system",
+ "description": "配置系统参数,包含各种类型",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "enabled": {"type": "boolean", "description": "是否启用"},
+ "max_connections": {"type": "integer", "description": "最大连接数"},
+ "timeout": {"type": "number", "description": "超时时间(秒)"},
+ "mode": {
+ "type": "string",
+ "description": "运行模式",
+ "enum": ["development", "production", "testing"],
+ },
+ "allowed_ips": {
+ "type": "array",
+ "description": "允许的IP地址列表",
+ "items": {"type": "string"},
+ },
+ "database": {
+ "type": "object",
+ "description": "数据库配置",
+ "properties": {
+ "host": {"type": "string"},
+ "port": {"type": "integer"},
+ "ssl_enabled": {"type": "boolean"},
+ },
+ "required": ["host", "port"],
+ },
+ },
+ "required": ["enabled", "mode"],
+ },
+ },
+ }
+
+ # 解析
+ tool_call = ToolCall.model_validate(raw_definition)
+ print(f"工具名称: {tool_call.name}")
+ print(f"根级必填项: {tool_call.parameters.required}")
+
+ # 验证各种类型
+ print(f"enabled 类型: {tool_call.parameters.properties['enabled'].type}")
+ print(f"max_connections 类型: {tool_call.parameters.properties['max_connections'].type}")
+ print(f"timeout 类型: {tool_call.parameters.properties['timeout'].type}")
+ print(f"mode 枚举值: {tool_call.parameters.properties['mode'].enum}")
+
+ # 导出并验证相等性
+ dumped_data = tool_call.simple_input_dump()
+ print(f"\n原始定义:\n{json.dumps(raw_definition, indent=2, ensure_ascii=False)}")
+ print(f"\n导出结果:\n{json.dumps(dumped_data, indent=2, ensure_ascii=False)}")
+
+ # 验证相等
+ assert dumped_data == raw_definition, "混合类型 Schema 导出结果与原始定义不一致"
+ print("\n✅ 混合类型 Schema 测试通过:raw_definition == simple_input_dump()")
+
+
+if __name__ == "__main__":
+ test_simple_schema()
+ test_medium_nested_schema()
+ test_nested_schema()
+ test_array_of_primitives()
+ test_deep_nested_schema()
+ test_mixed_types_schema()
+ print("\n" + "=" * 50)
+ print("🎉 所有测试用例通过!")
+ print("=" * 50)