diff --git a/reme_ai/bench/__init__.py b/reme_ai/bench/__init__.py
new file mode 100644
index 00000000..e69de29b
diff --git a/reme_ai/core/config/default.yaml b/reme_ai/core/config/default.yaml
index 8209216d..08dfe23d 100644
--- a/reme_ai/core/config/default.yaml
+++ b/reme_ai/core/config/default.yaml
@@ -32,6 +32,7 @@ vector_store:
default:
backend: local
embedding_model: default
+ collection_name: reme
token_counter:
default:
diff --git a/reme_ai/core/embedding/base_embedding_model.py b/reme_ai/core/embedding/base_embedding_model.py
index 3725f96d..fff59c83 100644
--- a/reme_ai/core/embedding/base_embedding_model.py
+++ b/reme_ai/core/embedding/base_embedding_model.py
@@ -36,17 +36,17 @@ class BaseEmbeddingModel(ABC):
self.raise_exception = raise_exception
self.kwargs = kwargs
- async def _get_embeddings(self, input_text: list[str]) -> list[list[float]]:
+ async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Internal async implementation for calling the embedding API with batch input."""
- def _get_embeddings_sync(self, input_text: list[str]) -> list[list[float]]:
+ def _get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Internal synchronous implementation for calling the embedding API with batch input."""
- async def get_embedding(self, input_text: str) -> list[float]:
+ async def get_embedding(self, input_text: str, **kwargs) -> list[float]:
"""Async get embedding for a single text with exponential backoff retries."""
for i in range(self.max_retries):
try:
- result = await self._get_embeddings([input_text])
+ result = await self._get_embeddings([input_text], **kwargs)
return result[0]
except Exception as e:
logger.error(f"Model {self.model_name} failed: {e}")
@@ -57,7 +57,7 @@ class BaseEmbeddingModel(ABC):
await asyncio.sleep(i + 1)
return []
- async def get_embeddings(self, input_text: list[str]) -> list[list[float]]:
+ async def get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Async get embeddings with automatic batching and exponential backoff retries."""
# Split into batches and process sequentially to respect rate limits
results = []
@@ -66,7 +66,7 @@ class BaseEmbeddingModel(ABC):
# Process each batch with retry logic
for retry in range(self.max_retries):
try:
- batch_res = await self._get_embeddings(batch)
+ batch_res = await self._get_embeddings(batch, **kwargs)
if batch_res:
results.extend(batch_res)
break
@@ -79,11 +79,11 @@ class BaseEmbeddingModel(ABC):
await asyncio.sleep(retry + 1)
return results
- def get_embedding_sync(self, input_text: str) -> list[float]:
+ def get_embedding_sync(self, input_text: str, **kwargs) -> list[float]:
"""Synchronous get embedding for a single text with retry logic."""
for i in range(self.max_retries):
try:
- result = self._get_embeddings_sync([input_text])
+ result = self._get_embeddings_sync([input_text], **kwargs)
return result[0]
except Exception as exc:
logger.error(f"Model {self.model_name} failed: {exc}")
@@ -94,7 +94,7 @@ class BaseEmbeddingModel(ABC):
time.sleep(i + 1)
return []
- def get_embeddings_sync(self, input_text: list[str]) -> list[list[float]]:
+ def get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Synchronous get embeddings with automatic batching and retry logic."""
results = []
for i in range(0, len(input_text), self.max_batch_size):
@@ -102,7 +102,7 @@ class BaseEmbeddingModel(ABC):
# Process each batch with retry logic
for retry in range(self.max_retries):
try:
- batch_res = self._get_embeddings_sync(batch)
+ batch_res = self._get_embeddings_sync(batch, **kwargs)
if batch_res:
results.extend(batch_res)
break
@@ -115,15 +115,15 @@ class BaseEmbeddingModel(ABC):
time.sleep(retry + 1)
return results
- async def get_node_embedding(self, node: VectorNode) -> VectorNode:
+ async def get_node_embedding(self, node: VectorNode, **kwargs) -> VectorNode:
"""Async generate and populate vector field for a single VectorNode object."""
- node.vector = await self.get_embedding(node.content)
+ node.vector = await self.get_embedding(node.content, **kwargs)
return node
- async def get_node_embeddings(self, nodes: list[VectorNode]) -> list[VectorNode]:
+ async def get_node_embeddings(self, nodes: list[VectorNode], **kwargs) -> list[VectorNode]:
"""Async generate and populate vector fields for a batch of VectorNode objects."""
contents = [node.content for node in nodes]
- embeddings: list[list[float]] = await self.get_embeddings(contents)
+ embeddings: list[list[float]] = await self.get_embeddings(contents, **kwargs)
if len(embeddings) == len(nodes):
for node, vec in zip(nodes, embeddings):
@@ -132,15 +132,15 @@ class BaseEmbeddingModel(ABC):
logger.warning(f"Mismatch: got {len(embeddings)} vectors for {len(nodes)} nodes")
return nodes
- def get_node_embedding_sync(self, node: VectorNode) -> VectorNode:
+ def get_node_embedding_sync(self, node: VectorNode, **kwargs) -> VectorNode:
"""Synchronously generate and populate vector field for a single VectorNode object."""
- node.vector = self.get_embedding_sync(node.content)
+ node.vector = self.get_embedding_sync(node.content, **kwargs)
return node
- def get_node_embeddings_sync(self, nodes: list[VectorNode]) -> list[VectorNode]:
+ def get_node_embeddings_sync(self, nodes: list[VectorNode], **kwargs) -> list[VectorNode]:
"""Synchronously generate and populate vector fields for a batch of VectorNode objects."""
contents = [node.content for node in nodes]
- embeddings: list[list[float]] = self.get_embeddings_sync(contents)
+ embeddings: list[list[float]] = self.get_embeddings_sync(contents, **kwargs)
if len(embeddings) == len(nodes):
for node, vec in zip(nodes, embeddings):
diff --git a/reme_ai/core/embedding/openai_embedding_model.py b/reme_ai/core/embedding/openai_embedding_model.py
index e10c7e66..a7e9f0c8 100644
--- a/reme_ai/core/embedding/openai_embedding_model.py
+++ b/reme_ai/core/embedding/openai_embedding_model.py
@@ -33,13 +33,15 @@ class OpenAIEmbeddingModel(BaseEmbeddingModel):
"""Create and return an internal AsyncOpenAI client instance."""
return AsyncOpenAI(api_key=self.api_key, base_url=self.base_url)
- async def _get_embeddings(self, input_text: list[str]) -> list[list[float]]:
+ async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Fetch embeddings from the API for a batch of strings."""
completion = await self._client.embeddings.create(
model=self.model_name,
input=input_text,
dimensions=self.dimensions,
encoding_format=self.encoding_format,
+ **self.kwargs,
+ **kwargs,
)
result_emb = [[] for _ in range(len(input_text))]
diff --git a/reme_ai/core/embedding/openai_embedding_model_sync.py b/reme_ai/core/embedding/openai_embedding_model_sync.py
index b777e6d6..760732cd 100644
--- a/reme_ai/core/embedding/openai_embedding_model_sync.py
+++ b/reme_ai/core/embedding/openai_embedding_model_sync.py
@@ -14,13 +14,15 @@ class OpenAIEmbeddingModelSync(OpenAIEmbeddingModel):
"""Create and return an internal synchronous OpenAI client instance."""
return OpenAI(api_key=self.api_key, base_url=self.base_url)
- def _get_embeddings_sync(self, input_text: list[str]) -> list[list[float]]:
+ def _get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Fetch embeddings synchronously from the API for a batch of strings."""
completion = self._client.embeddings.create(
model=self.model_name,
input=input_text,
dimensions=self.dimensions,
encoding_format=self.encoding_format,
+ **self.kwargs,
+ **kwargs,
)
result_emb = [[] for _ in range(len(input_text))]
diff --git a/reme_ai/core/flow/__init__.py b/reme_ai/core/flow/__init__.py
index 6d5a053b..e74a2b5c 100644
--- a/reme_ai/core/flow/__init__.py
+++ b/reme_ai/core/flow/__init__.py
@@ -3,9 +3,11 @@
from .base_flow import BaseFlow
from .cmd_flow import CmdFlow
from .expression_flow import ExpressionFlow
+from .simple_flow import SimpleFlow
__all__ = [
"BaseFlow",
"CmdFlow",
"ExpressionFlow",
+ "SimpleFlow",
]
diff --git a/reme_ai/core/flow/base_flow.py b/reme_ai/core/flow/base_flow.py
index 1802af78..5fb3d00e 100644
--- a/reme_ai/core/flow/base_flow.py
+++ b/reme_ai/core/flow/base_flow.py
@@ -25,6 +25,7 @@ class BaseFlow(ABC):
def __init__(
self,
name: str = "",
+ flow_op: BaseOp | None = None,
stream: bool = False,
raise_exception: bool = True,
enable_cache: bool = False,
@@ -43,7 +44,7 @@ class BaseFlow(ABC):
self.cache_expire_hours: float = cache_expire_hours
self.flow_params: dict = kwargs
- self._flow_op: BaseOp | None = None
+ self._flow_op: BaseOp | None = flow_op
self._cache: CacheHandler | None = None
self._flow_printed: bool = False
self._tool_call: ToolCall | None = None
@@ -129,6 +130,12 @@ class BaseFlow(ABC):
self._flow_op = self._build_flow()
return self._flow_op
+ @flow_op.setter
+ def flow_op(self, op: BaseOp):
+ """Set the root operation of the flow."""
+ self._flow_op = op
+ self._flow_printed = False
+
@property
def async_mode(self) -> bool:
"""Check if the current flow operation tree is asynchronous."""
diff --git a/reme_ai/core/flow/simple_flow.py b/reme_ai/core/flow/simple_flow.py
new file mode 100644
index 00000000..51a1381d
--- /dev/null
+++ b/reme_ai/core/flow/simple_flow.py
@@ -0,0 +1,17 @@
+"""Simple flow implementation that directly uses a predefined flow operation."""
+
+from .base_flow import BaseFlow
+from ..op import BaseOp
+from ..schema import ToolCall
+
+
+class SimpleFlow(BaseFlow):
+ """Simple flow that directly uses a predefined flow operation."""
+
+ def _build_flow(self) -> BaseOp:
+ assert self._flow_op is not None
+ return self._flow_op.copy()
+
+ def _build_tool_call(self) -> ToolCall:
+ assert self._flow_op is not None
+ return self._flow_op.tool_call
diff --git a/reme_ai/core/llm/base_llm.py b/reme_ai/core/llm/base_llm.py
index 5235198c..04370cc3 100644
--- a/reme_ai/core/llm/base_llm.py
+++ b/reme_ai/core/llm/base_llm.py
@@ -24,53 +24,6 @@ class BaseLLM(ABC):
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]):
"""Assemble incremental tool call fragments into complete ToolCall objects."""
@@ -137,14 +90,15 @@ class BaseLLM(ABC):
"""Internal synchronous generator for streaming raw response chunks."""
raise NotImplementedError
- async def _stream_with_retry(
+ async def stream_chat(
self,
- operation_name: str,
messages: list[Message],
- tools: list[ToolCall] | None,
- stream_kwargs: dict,
+ tools: list[ToolCall] | None = None,
+ **kwargs,
) -> AsyncGenerator[StreamChunk, None]:
- """Execute the async streaming operation with retry logic and error recovery."""
+ """Public async interface for streaming chat completions with retries."""
+ stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs)
+
for i in range(self.max_retries):
try:
async for chunk in self._stream_chat(messages=messages, tools=tools, stream_kwargs=stream_kwargs):
@@ -152,7 +106,7 @@ class BaseLLM(ABC):
return
except Exception as e:
- logger.exception(f"{operation_name} with model={self.model_name} encounter error with e={e.args}")
+ logger.exception(f"stream chat with model={self.model_name} encounter error with e={e.args}")
if i == self.max_retries - 1:
if self.raise_exception:
@@ -163,21 +117,22 @@ class BaseLLM(ABC):
yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e))
await asyncio.sleep(i + 1)
- def _stream_with_retry_sync(
+ def stream_chat_sync(
self,
- operation_name: str,
messages: list[Message],
- tools: list[ToolCall] | None,
- stream_kwargs: dict,
+ tools: list[ToolCall] | None = None,
+ **kwargs,
) -> Generator[StreamChunk, None, None]:
- """Execute the synchronous streaming operation with retry logic and error recovery."""
+ """Public synchronous interface for streaming chat completions with retries."""
+ stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs)
+
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}")
+ logger.exception(f"stream chat sync with model={self.model_name} encounter error with e={e.args}")
if i == self.max_retries - 1:
if self.raise_exception:
@@ -188,27 +143,6 @@ class BaseLLM(ABC):
yield StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e))
time.sleep(i + 1)
- async def stream_chat(
- self,
- messages: list[Message],
- tools: list[ToolCall] | None = 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: list[ToolCall] | None = 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],
@@ -227,9 +161,46 @@ class BaseLLM(ABC):
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)
+ # Process 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,
+ )
- return self._create_message_from_state(state)
+ 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)
+
+ return Message(
+ role=Role.ASSISTANT,
+ reasoning_content=state["reasoning_content"],
+ content=state["answer_content"],
+ tool_calls=state["tool_calls"],
+ )
def _chat_sync(
self,
@@ -249,57 +220,46 @@ class BaseLLM(ABC):
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)
+ # Process 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,
+ )
- return self._create_message_from_state(state)
+ 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
- async def _execute_with_retry(
- self,
- operation_name: str,
- operation_fn: Callable[[], Any],
- callback_fn: Callable[[Message], Any] | None = 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
+ 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
- except Exception as e:
- logger.exception(f"{operation_name} with model={self.model_name} encounter error with e={e.args}")
+ 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)
- if i == self.max_retries - 1:
- if self.raise_exception:
- raise e
- return default_value
+ elif stream_chunk.chunk_type is ChunkEnum.ERROR:
+ if enable_stream_print:
+ print(f"\n{stream_chunk.chunk}", flush=True)
- await asyncio.sleep(1 + i)
- return default_value
-
- def _execute_with_retry_sync(
- self,
- operation_name: str,
- operation_fn: Callable[[], Message],
- callback_fn: Callable[[Message], Any] | None = 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
+ return Message(
+ role=Role.ASSISTANT,
+ reasoning_content=state["reasoning_content"],
+ content=state["answer_content"],
+ tool_calls=state["tool_calls"],
+ )
async def chat(
self,
@@ -311,17 +271,26 @@ class BaseLLM(ABC):
**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,
- )
+ for i in range(self.max_retries):
+ try:
+ result = await self._chat(
+ messages=messages,
+ tools=tools,
+ enable_stream_print=enable_stream_print,
+ **kwargs,
+ )
+ return callback_fn(result) if callback_fn else result
+
+ except Exception as e:
+ logger.exception(f"chat 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 chat_sync(
self,
@@ -333,17 +302,26 @@ class BaseLLM(ABC):
**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,
- )
+ for i in range(self.max_retries):
+ try:
+ result = self._chat_sync(
+ messages=messages,
+ tools=tools,
+ enable_stream_print=enable_stream_print,
+ **kwargs,
+ )
+ return callback_fn(result) if callback_fn else result
+
+ except Exception as e:
+ logger.exception(f"chat sync 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 close(self):
"""Release any asynchronous resources or connections held by the client."""
diff --git a/reme_ai/core/op/base_op.py b/reme_ai/core/op/base_op.py
index 4f96ec32..24c416d4 100644
--- a/reme_ai/core/op/base_op.py
+++ b/reme_ai/core/op/base_op.py
@@ -377,5 +377,6 @@ class BaseOp:
copy_op = self.__class__(*self._init_args, **{**self._init_kwargs, **kwargs})
if self.sub_ops:
copy_op.sub_ops.clear()
- copy_op.add_sub_ops(self.sub_ops)
+ for op in self.sub_ops:
+ copy_op.add_sub_op(op.copy())
return copy_op
diff --git a/reme_ai/core/reme.py b/reme_ai/core/reme.py
deleted file mode 100644
index 552a526b..00000000
--- a/reme_ai/core/reme.py
+++ /dev/null
@@ -1,48 +0,0 @@
-"""ReMe classes for simplified configuration and execution."""
-
-from .application import Application
-from .config import ReMeConfigParser
-from .context import C
-
-
-class ReMe(Application):
- """Simplified ReMe application that auto-initializes the service context."""
-
- def __init__(
- self,
- *args,
- llm_api_key: str | None = None,
- llm_api_base: str | None = None,
- embedding_api_key: str | None = None,
- embedding_api_base: str | None = None,
- enable_logo: bool = True,
- llm: dict | None = None,
- embedding_model: dict | None = None,
- vector_store: dict | None = None,
- token_counter: dict | None = None,
- **kwargs,
- ):
- super().__init__(
- *args,
- llm_api_key=llm_api_key,
- llm_api_base=llm_api_base,
- embedding_api_key=embedding_api_key,
- embedding_api_base=embedding_api_base,
- service_config=None,
- parser=ReMeConfigParser,
- config_path=None,
- enable_logo=enable_logo,
- llm=llm,
- embedding_model=embedding_model,
- vector_store=vector_store,
- token_counter=token_counter,
- **kwargs,
- )
-
- C.initialize_service_context()
-
- async def summary(self):
- """Execute summary operations."""
-
- async def retrieve(self):
- """Execute retrieve operations."""
diff --git a/reme_ai/core/schema/memory_node.py b/reme_ai/core/schema/memory_node.py
index 95d58f67..fa9bd25b 100644
--- a/reme_ai/core/schema/memory_node.py
+++ b/reme_ai/core/schema/memory_node.py
@@ -153,7 +153,7 @@ class MemoryNode(BaseModel):
str: Formatted string with when_to_use, content, and ref_memory_id.
"""
parts: list[str] = [
- f"memory_id={self.memory_id}" f"modified_time={self.time_modified}",
+ f"memory_id={self.memory_id} modified_time={self.time_modified}",
]
if self.when_to_use:
@@ -166,7 +166,7 @@ class MemoryNode(BaseModel):
parts.append(f"metadata={json.dumps(self.metadata, ensure_ascii=False)}")
if self.ref_memory_id:
- parts.append(f"history_memory.ref_memory_id={self.ref_memory_id}")
+ parts.append(f"ref_memory_id={self.ref_memory_id}")
return " ".join(parts)
diff --git a/reme_ai/core/schema/message.py b/reme_ai/core/schema/message.py
index 510bb7c4..ca53f2fc 100644
--- a/reme_ai/core/schema/message.py
+++ b/reme_ai/core/schema/message.py
@@ -83,7 +83,8 @@ class Message(BaseModel):
add_reasoning: bool = True,
add_time_created: bool = False,
add_metadata: bool = False,
- ) -> dict:
+ enable_json_dump: bool = False,
+ ) -> dict | str:
"""Transforms the message into a simplified dictionary for standard APIs."""
result = {}
if add_name and self.name:
@@ -107,7 +108,10 @@ class Message(BaseModel):
if add_metadata:
result["metadata"] = self.metadata
- return result
+ if enable_json_dump:
+ return json.dumps(result, ensure_ascii=False)
+ else:
+ return result
def format_message(
self,
diff --git a/reme_ai/core/schema/service_config.py b/reme_ai/core/schema/service_config.py
index 6471e982..240f4b1a 100644
--- a/reme_ai/core/schema/service_config.py
+++ b/reme_ai/core/schema/service_config.py
@@ -74,7 +74,7 @@ class VectorStoreConfig(BaseModel):
model_config = ConfigDict(extra="allow")
backend: str = Field(default="local")
- collection_name: str = Field(default="remy")
+ collection_name: str = Field(default="reme")
embedding_model: str = Field(default="default")
diff --git a/reme_ai/core/service/cmd_service.py b/reme_ai/core/service/cmd_service.py
index bed8ec86..c02efa20 100644
--- a/reme_ai/core/service/cmd_service.py
+++ b/reme_ai/core/service/cmd_service.py
@@ -26,9 +26,7 @@ class CmdService(BaseService):
super().run()
if self._cmd_flow.async_mode:
- response = run_coro_safely(
- self._cmd_flow.call(**C.service_config.cmd.model_extra),
- )
+ response = run_coro_safely(self._cmd_flow.call(**C.service_config.cmd.model_extra))
else:
response = self._cmd_flow.call_sync(**C.service_config.cmd.model_extra)
diff --git a/reme_ai/mem_agent/__init__.py b/reme_ai/mem_agent/__init__.py
index 69495a12..ffb5f332 100644
--- a/reme_ai/mem_agent/__init__.py
+++ b/reme_ai/mem_agent/__init__.py
@@ -1,15 +1,13 @@
"""memory agent"""
+from . import chat
from . import retriever
from . import summarizer
from .base_memory_agent import BaseMemoryAgent
-from .simple_chat import SimpleChat
-from .stream_chat import StreamChat
__all__ = [
+ "chat",
"retriever",
"summarizer",
"BaseMemoryAgent",
- "StreamChat",
- "SimpleChat",
]
diff --git a/reme_ai/mem_agent/base_memory_agent.py b/reme_ai/mem_agent/base_memory_agent.py
index 7b6bdba0..f5b33329 100644
--- a/reme_ai/mem_agent/base_memory_agent.py
+++ b/reme_ai/mem_agent/base_memory_agent.py
@@ -20,18 +20,16 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
self,
tools: list[BaseMemoryTool],
add_think_tool: bool = False, # only for instruct model
- force_tool_language: bool = True,
tool_call_interval: float = 0,
max_steps: int = 20,
**kwargs,
):
- super().__init__(**kwargs)
- self.tools: list[BaseMemoryTool] = tools or []
+ tools = tools or []
if add_think_tool:
- self.tools.append(ThinkTool())
- if force_tool_language and self.language:
- for tool in self.tools:
- tool.language = self.language
+ tools.append(ThinkTool())
+ kwargs["sub_ops"] = tools
+ super().__init__(**kwargs)
+ self.sub_ops: list[BaseMemoryTool] = [t for t in self.sub_ops if isinstance(t, BaseMemoryTool)]
self.tool_call_interval: float = tool_call_interval
self.max_steps: int = max_steps
@@ -72,6 +70,15 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
},
)
+ @property
+ def tools(self):
+ """Returns the list of memory tools available to this agent."""
+ return self.sub_ops
+
+ @tools.setter
+ def tools(self, tools: list[BaseMemoryTool]):
+ self.sub_ops = tools
+
def get_messages(self) -> list[Message]:
"""Extracts and returns messages from the context query or messages."""
if self.context.get("query"):
@@ -93,7 +100,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
**kwargs,
)
messages.append(assistant_message)
- logger.info(f"step{step + 1}.assistant={assistant_message.model_dump_json()}")
+ logger.info(f"step{step + 1}.assistant={assistant_message.simple_dump(enable_json_dump=True)}")
should_act = bool(assistant_message.tool_calls)
return assistant_message, should_act
@@ -110,7 +117,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
logger.warning(f"unknown tool_call.name={tool_call.name}")
continue
- logger.info(f"step{step + 1}.{j} submit tool_calls={tool_call.name} argument={tool_call.argument_dict}")
+ logger.info(f"step{step + 1}.{j} submit tool_calls={tool_call.name} argument={tool_call.arguments}")
tool_copy: BaseMemoryTool = tool_dict[tool_call.name].copy()
tool_copy.tool_call.id = tool_call.id
tool_list.append(tool_copy)
@@ -150,7 +157,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
async def execute(self):
messages = await self.build_messages()
for i, message in enumerate(messages):
- logger.info(f"step0.{i} {message.role} {message.name or ''} {message.simple_dump()}")
+ logger.info(f"step0.{i} {message.role} {message.name or ''} {message.simple_dump(enable_json_dump=True)}")
self.messages, self.success = await self.react(messages)
if self.success and self.messages:
@@ -163,6 +170,11 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
"""Returns the target memory identifier from context."""
return self.context.get("memory_target", "")
+ @property
+ def description(self) -> str:
+ """Returns the description of the messages."""
+ return self.context.get("description", "")
+
@property
def ref_memory_id(self) -> str:
"""Returns the reference memory ID from context."""
diff --git a/reme_ai/mem_agent/chat/__init__.py b/reme_ai/mem_agent/chat/__init__.py
new file mode 100644
index 00000000..be2fc055
--- /dev/null
+++ b/reme_ai/mem_agent/chat/__init__.py
@@ -0,0 +1,11 @@
+"""chat agent"""
+
+from .remy_agent import ReMyAgent
+from .simple_chat import SimpleChat
+from .stream_chat import StreamChat
+
+__all__ = [
+ "ReMyAgent",
+ "StreamChat",
+ "SimpleChat",
+]
diff --git a/reme_ai/mem_agent/retriever/remy_agent.py b/reme_ai/mem_agent/chat/remy_agent.py
similarity index 100%
rename from reme_ai/mem_agent/retriever/remy_agent.py
rename to reme_ai/mem_agent/chat/remy_agent.py
diff --git a/reme_ai/mem_agent/retriever/remy_agent.yaml b/reme_ai/mem_agent/chat/remy_agent.yaml
similarity index 100%
rename from reme_ai/mem_agent/retriever/remy_agent.yaml
rename to reme_ai/mem_agent/chat/remy_agent.yaml
diff --git a/reme_ai/mem_agent/simple_chat.py b/reme_ai/mem_agent/chat/simple_chat.py
similarity index 93%
rename from reme_ai/mem_agent/simple_chat.py
rename to reme_ai/mem_agent/chat/simple_chat.py
index fc97b7a4..8a71c7a8 100644
--- a/reme_ai/mem_agent/simple_chat.py
+++ b/reme_ai/mem_agent/chat/simple_chat.py
@@ -2,10 +2,10 @@
from loguru import logger
-from ..core.context import C
-from ..core.enumeration import Role
-from ..core.op import BaseOp
-from ..core.schema import Message, ToolCall
+from ...core.context import C
+from ...core.enumeration import Role
+from ...core.op import BaseOp
+from ...core.schema import Message, ToolCall
@C.register_op()
diff --git a/reme_ai/mem_agent/stream_chat.py b/reme_ai/mem_agent/chat/stream_chat.py
similarity index 93%
rename from reme_ai/mem_agent/stream_chat.py
rename to reme_ai/mem_agent/chat/stream_chat.py
index bb410cf6..470e4647 100644
--- a/reme_ai/mem_agent/stream_chat.py
+++ b/reme_ai/mem_agent/chat/stream_chat.py
@@ -2,10 +2,10 @@
from loguru import logger
-from ..core.context import C
-from ..core.enumeration import Role, ChunkEnum
-from ..core.op import BaseOp
-from ..core.schema import Message, ToolCall
+from ...core.context import C
+from ...core.enumeration import Role, ChunkEnum
+from ...core.op import BaseOp
+from ...core.schema import Message, ToolCall
@C.register_op()
diff --git a/reme_ai/mem_agent/retriever/__init__.py b/reme_ai/mem_agent/retriever/__init__.py
index 83f1266f..0d34a3fd 100644
--- a/reme_ai/mem_agent/retriever/__init__.py
+++ b/reme_ai/mem_agent/retriever/__init__.py
@@ -1,9 +1,7 @@
"""memory retriever"""
from .reme_retriever import ReMeRetriever
-from .remy_agent import ReMyAgent
__all__ = [
"ReMeRetriever",
- "ReMyAgent",
]
diff --git a/reme_ai/mem_agent/retriever/reme_retriever.py b/reme_ai/mem_agent/retriever/reme_retriever.py
index 4a7a3fa9..9ba15d66 100644
--- a/reme_ai/mem_agent/retriever/reme_retriever.py
+++ b/reme_ai/mem_agent/retriever/reme_retriever.py
@@ -13,25 +13,32 @@ from ...core.utils import get_now_time, format_messages
class ReMeRetriever(BaseMemoryAgent):
"""Memory agent that retrieves and builds messages with meta memory context."""
- def __init__(self, enable_tool_memory: bool = True, **kwargs):
- """Initialize retriever with tool memory option."""
+ def __init__(self, meta_memories: list[dict] = None, **kwargs):
super().__init__(**kwargs)
- self.enable_tool_memory = enable_tool_memory
+ self.meta_memories: list[dict] = meta_memories
- async def _read_meta_memories(self) -> str:
+ @staticmethod
+ async def _read_meta_memories() -> str:
"""Read and return meta memories as string."""
from ...mem_tool import ReadMetaMemory
- op = ReadMetaMemory(enable_tool_memory=self.enable_tool_memory, enable_identity_memory=False)
+ op = ReadMetaMemory(enable_identity_memory=False)
await op.call()
return str(op.output)
async def build_messages(self) -> List[Message]:
"""Build messages with system prompt and user message."""
+ from ...mem_tool import ReadMetaMemory
+
+ if self.meta_memories:
+ meta_memory_info = ReadMetaMemory().format_memory_metadata(self.meta_memories)
+ else:
+ meta_memory_info = await self._read_meta_memories()
+
system_prompt = self.prompt_format(
prompt_name="system_prompt",
now_time=get_now_time(),
- meta_memory_info=await self._read_meta_memories(),
+ meta_memory_info=meta_memory_info,
context=format_messages(self.get_messages()),
)
diff --git a/reme_ai/mem_agent/summarizer/personal_summarizer.py b/reme_ai/mem_agent/summarizer/personal_summarizer.py
index 357e7b9f..424eec39 100644
--- a/reme_ai/mem_agent/summarizer/personal_summarizer.py
+++ b/reme_ai/mem_agent/summarizer/personal_summarizer.py
@@ -18,7 +18,7 @@ class PersonalSummarizer(BaseMemoryAgent):
system_prompt = self.prompt_format(
prompt_name="system_prompt",
now_time=get_now_time(),
- context=format_messages(self.get_messages()),
+ context=self.description + "\n" + format_messages(self.get_messages()),
memory_type=self.memory_type.value,
memory_target=self.memory_target,
)
@@ -37,5 +37,6 @@ class PersonalSummarizer(BaseMemoryAgent):
memory_target=self.memory_target,
memory_type=self.memory_type.value,
author=self.author,
+ ref_memory_id=self.ref_memory_id,
**kwargs,
)
diff --git a/reme_ai/mem_agent/summarizer/personal_summarizer.yaml b/reme_ai/mem_agent/summarizer/personal_summarizer.yaml
index eafe1da5..7ed0b091 100644
--- a/reme_ai/mem_agent/summarizer/personal_summarizer.yaml
+++ b/reme_ai/mem_agent/summarizer/personal_summarizer.yaml
@@ -6,7 +6,7 @@ tool: |
or conflicts with existing memories, and perform add, update, or delete operations as needed.
system_prompt: |
- You are a professional memory agent specializing in the domain of **{memory_target}**. Your task is to update the main agent's {memory_type} memory regarding {memory_target} based on the context.
+ You are a professional memory agent. Your task is to update the main agent's {memory_type} memory regarding {memory_target} based on the context.
## Context:
{context}
diff --git a/reme_ai/mem_agent/summarizer/procedural_summarizer.yaml b/reme_ai/mem_agent/summarizer/procedural_summarizer.yaml
index 83a8867d..c80cdfe1 100644
--- a/reme_ai/mem_agent/summarizer/procedural_summarizer.yaml
+++ b/reme_ai/mem_agent/summarizer/procedural_summarizer.yaml
@@ -7,7 +7,7 @@ tool: |
and successful strategies from successes to improve future performance.
system_prompt: |
- You are a professional memory Agent specializing in the domain of **{memory_target}**. Your task is to update the main agent's {memory_type} memory regarding {memory_target} based on the context.
+ You are a professional memory Agent specializing. Your task is to update the main agent's {memory_type} memory regarding {memory_target} based on the context.
## Context:
{context}
diff --git a/reme_ai/mem_tool/__init__.py b/reme_ai/mem_tool/__init__.py
index 15392685..9b8820dc 100644
--- a/reme_ai/mem_tool/__init__.py
+++ b/reme_ai/mem_tool/__init__.py
@@ -1,6 +1,7 @@
"""Memory tool operations."""
from .base_memory_tool import BaseMemoryTool
+from .hands_off_tool import HandsOffTool
from .history.add_history_memory import AddHistoryMemory
from .history.read_history_memory import ReadHistoryMemory
from .identity.read_identity_memory import ReadIdentityMemory
@@ -16,6 +17,7 @@ from .vector.vector_retrieve_memory import VectorRetrieveMemory
__all__ = [
"BaseMemoryTool",
+ "HandsOffTool",
"AddHistoryMemory",
"ReadHistoryMemory",
"ReadIdentityMemory",
diff --git a/reme_ai/mem_tool/hands_off_tool.py b/reme_ai/mem_tool/hands_off_tool.py
index e0b63ecb..a3924ba1 100644
--- a/reme_ai/mem_tool/hands_off_tool.py
+++ b/reme_ai/mem_tool/hands_off_tool.py
@@ -17,17 +17,17 @@ if TYPE_CHECKING:
class HandsOffTool(BaseMemoryTool):
"""Distribute memory tasks to appropriate agents based on memory_type."""
- def __init__(self, memory_agents: list["BaseMemoryAgent"], force_agent_language: bool = True, **kwargs):
+ def __init__(self, memory_agents: list["BaseMemoryAgent"], **kwargs):
+ kwargs["sub_ops"] = memory_agents or []
super().__init__(**kwargs)
- self.memory_agent_dict: dict[MemoryType, "BaseMemoryAgent"] = {}
- if memory_agents:
- for agent in memory_agents:
- if agent.memory_type is None:
- continue
+ from ..mem_agent import BaseMemoryAgent
- self.memory_agent_dict[agent.memory_type] = agent
- if force_agent_language and self.language:
- agent.language = self.language
+ self.sub_ops: list[BaseMemoryAgent] = [a for a in self.sub_ops if isinstance(a, BaseMemoryAgent)]
+
+ @property
+ def memory_agent_dict(self) -> dict[MemoryType, "BaseMemoryAgent"]:
+ """Returns a dictionary mapping memory types to their corresponding agents."""
+ return {a.memory_type: a for a in self.sub_ops}
def _build_item_schema(self) -> tuple[dict, list[str]]:
"""Build shared schema properties and required fields for memory tasks."""
@@ -117,11 +117,13 @@ class HandsOffTool(BaseMemoryTool):
continue
agent_copy = self.memory_agent_dict[memory_type].copy()
- agent_list.append({
- "agent": agent_copy,
- "memory_type": memory_type,
- "memory_target": memory_target,
- })
+ agent_list.append(
+ {
+ "agent": agent_copy,
+ "memory_type": memory_type,
+ "memory_target": memory_target,
+ },
+ )
logger.info(f"Task {i}: Submitting {memory_type.value} agent for target={memory_target}")
self.submit_async_task(
@@ -138,12 +140,14 @@ class HandsOffTool(BaseMemoryTool):
results = []
for i, (agent, memory_type, memory_target) in enumerate(agent_list):
result_str = str(agent.output)
- results.append({
- "memory_type": memory_type.value,
- "memory_target": memory_target,
- "result": result_str[:200] + ("..." if len(result_str) > 200 else ""),
- })
+ results.append(
+ {
+ "memory_type": memory_type.value,
+ "memory_target": memory_target,
+ "result": result_str[:200] + ("..." if len(result_str) > 200 else ""),
+ },
+ )
logger.info(f"Task {i}: Completed {memory_type.value} agent for target={memory_target}")
results_str = json.dumps(results, ensure_ascii=False, indent=2)
- self.set_output(f"Successfully executed {len(results)} memory tasks:\n{results_str}")
+ self.output = f"Successfully executed {len(results)} memory tasks:\n{results_str}"
diff --git a/reme_ai/mem_tool/meta/read_meta_memory.py b/reme_ai/mem_tool/meta/read_meta_memory.py
index b763ec1a..07ad1ecf 100644
--- a/reme_ai/mem_tool/meta/read_meta_memory.py
+++ b/reme_ai/mem_tool/meta/read_meta_memory.py
@@ -17,20 +17,17 @@ class ReadMetaMemory(BaseMemoryTool):
def __init__(
self,
- enable_tool_memory: bool = False,
enable_identity_memory: bool = False,
**kwargs,
):
"""Initialize ReadMetaMemory.
Args:
- enable_tool_memory: Include TOOL type meta memory. Defaults to False.
enable_identity_memory: Include IDENTITY type meta memory. Defaults to False.
**kwargs: Additional arguments for BaseMemoryTool.
"""
kwargs["enable_multiple"] = False
super().__init__(**kwargs)
- self.enable_tool_memory = enable_tool_memory
self.enable_identity_memory = enable_identity_memory
def _build_parameters(self) -> dict:
@@ -54,14 +51,6 @@ class ReadMetaMemory(BaseMemoryTool):
if m.get("memory_type") in [MemoryType.PERSONAL.value, MemoryType.PROCEDURAL.value]:
filtered_memories.append(m)
- if self.enable_tool_memory:
- filtered_memories.append(
- {
- "memory_type": MemoryType.TOOL.value,
- "memory_target": "tool_guidelines",
- },
- )
-
if self.enable_identity_memory:
filtered_memories.append(
{
@@ -72,7 +61,7 @@ class ReadMetaMemory(BaseMemoryTool):
return filtered_memories
- def _format_memory_metadata(self, memories: list[dict[str, str]]) -> str:
+ def format_memory_metadata(self, memories: list[dict[str, str]]) -> str:
"""Format memory metadata into a readable string.
Args:
@@ -101,7 +90,7 @@ class ReadMetaMemory(BaseMemoryTool):
memories = self._load_meta_memories()
if memories:
- self.output = self._format_memory_metadata(memories)
+ self.output = self.format_memory_metadata(memories)
logger.info(f"Retrieved {len(memories)} meta memory entries")
else:
self.output = "No memory metadata found."
diff --git a/reme_ai/reme.py b/reme_ai/reme.py
new file mode 100644
index 00000000..93b00ede
--- /dev/null
+++ b/reme_ai/reme.py
@@ -0,0 +1,168 @@
+"""ReMe classes for simplified configuration and execution."""
+
+from typing import Literal
+
+from .core.application import Application
+from .core.config import ReMeConfigParser
+from .core.context import C
+from .core.enumeration import Role
+from .core.vector_store import BaseVectorStore
+from .mem_agent.summarizer import (
+ ReMeSummarizer,
+ # ToolSummarizer,
+ PersonalSummarizer,
+ ProceduralSummarizer,
+ # IdentitySummarizer,
+)
+from .mem_agent.retriever import ReMeRetriever
+
+# from .mem_agent.chat import ReMyAgent
+from .mem_tool import (
+ HandsOffTool,
+ ReadHistoryMemory,
+ # ReadIdentityMemory,
+ # UpdateIdentityMemory,
+ AddMetaMemory,
+ AddMemory,
+ AddSummaryMemory,
+ DeleteMemory,
+ UpdateMemory,
+ VectorRetrieveMemory,
+)
+from .core.schema import Message
+
+
+class ReMe(Application):
+ """Simplified ReMe application that auto-initializes the service context."""
+
+ def __init__(
+ self,
+ *args,
+ llm_api_key: str | None = None,
+ llm_api_base: str | None = None,
+ embedding_api_key: str | None = None,
+ embedding_api_base: str | None = None,
+ enable_logo: bool = True,
+ llm: dict | None = None,
+ embedding_model: dict | None = None,
+ vector_store: dict | None = None,
+ token_counter: dict | None = None,
+ enable_identity_memory: bool = True,
+ enable_tool_memory: bool = True,
+ force_tool_language: bool = True,
+ add_think_tool: bool = False,
+ **kwargs,
+ ):
+ super().__init__(
+ *args,
+ llm_api_key=llm_api_key,
+ llm_api_base=llm_api_base,
+ embedding_api_key=embedding_api_key,
+ embedding_api_base=embedding_api_base,
+ service_config=None,
+ parser=ReMeConfigParser,
+ config_path=None,
+ enable_logo=enable_logo,
+ llm=llm,
+ embedding_model=embedding_model,
+ vector_store=vector_store,
+ token_counter=token_counter,
+ **kwargs,
+ )
+
+ C.initialize_service_context()
+ self.enable_identity_memory = enable_identity_memory
+ self.enable_tool_memory = enable_tool_memory
+ self.force_tool_language = force_tool_language
+ self.add_think_tool = add_think_tool
+
+ self._personal_summarizer = PersonalSummarizer(
+ tools=[VectorRetrieveMemory(), AddMemory(), DeleteMemory(), UpdateMemory()],
+ )
+ self._procedural_summarizer = ProceduralSummarizer(
+ tools=[VectorRetrieveMemory(), AddMemory(), DeleteMemory(), UpdateMemory()],
+ )
+ hands_off_tool = HandsOffTool(memory_agents=[self._personal_summarizer, self._procedural_summarizer])
+ self._reme_summarizer = ReMeSummarizer(
+ tools=[AddMetaMemory(), AddSummaryMemory(), hands_off_tool],
+ enable_identity_memory=self.enable_identity_memory,
+ enable_tool_memory=self.enable_tool_memory,
+ force_tool_language=self.force_tool_language,
+ add_think_tool=self.add_think_tool,
+ )
+ self._reme_retriever = ReMeRetriever(
+ tools=[VectorRetrieveMemory(add_memory_type_target=True), ReadHistoryMemory()],
+ )
+
+ self.vector_store: BaseVectorStore = C.get_vector_store("default")
+
+ @staticmethod
+ def _prepare_messages(messages: list[dict | Message], user_id: str, assistant_id: str):
+ if not messages:
+ return []
+
+ messages = [Message(**m) if isinstance(m, dict) else m for m in messages]
+ for message in messages:
+ if message.role is Role.USER and user_id:
+ message.name = user_id
+ elif message.role is Role.ASSISTANT and assistant_id:
+ message.name = assistant_id
+ return messages
+
+ async def summary(
+ self,
+ messages: list[dict],
+ description: str = "",
+ user_id: str = "",
+ assistant_id: str = "",
+ memory_mode: Literal["personal", "procedural", "auto"] = "personal",
+ **kwargs,
+ ):
+ """Summarizes messages and stores them as memory based on the specified memory mode."""
+ messages = self._prepare_messages(messages, user_id, assistant_id)
+
+ if memory_mode == "personal":
+ return await self._personal_summarizer.call(
+ messages=messages,
+ description=description,
+ memory_target=user_id,
+ **kwargs,
+ )
+ elif memory_mode == "procedural":
+ return await self._procedural_summarizer.call(
+ messages=messages,
+ description=description,
+ memory_target=user_id,
+ **kwargs,
+ )
+ else:
+ return await self._reme_summarizer.call(
+ messages=messages,
+ description=description,
+ memory_target=user_id,
+ **kwargs,
+ )
+
+ async def retrieve(
+ self,
+ query: str = "",
+ messages: list[dict] | None = None,
+ description: str = "",
+ user_id: str = "",
+ assistant_id: str = "",
+ memory_mode: Literal["personal", "procedural", "auto"] = "personal",
+ **kwargs,
+ ):
+ """Retrieves relevant memories based on the query and specified memory mode."""
+ messages = self._prepare_messages(messages, user_id, assistant_id)
+
+ if memory_mode == "personal":
+ self._reme_retriever.meta_memories = [{"memory_type": "personal", "memory_target": user_id}]
+ return await self._reme_retriever.call(query=query, messages=messages, description=description, **kwargs)
+
+ elif memory_mode == "procedural":
+ self._reme_retriever.meta_memories = [{"memory_type": "procedural", "memory_target": user_id}]
+ return await self._reme_retriever.call(query=query, messages=messages, description=description, **kwargs)
+
+ else:
+ return await self._reme_retriever.call(query=query, messages=messages, description=description, **kwargs)
diff --git a/tests/test_reme.py b/tests/test_reme.py
new file mode 100644
index 00000000..2e57e8a7
--- /dev/null
+++ b/tests/test_reme.py
@@ -0,0 +1,111 @@
+"""Test module for ReMe memory system functionality."""
+
+import asyncio
+
+from reme_ai.core.schema import VectorNode, MemoryNode
+from reme_ai.reme import ReMe
+
+reme = ReMe(
+ vector_store={"collection_name": "reme"},
+)
+
+
+async def test_reme():
+ """Tests ReMe memory system with personal information storage and retrieval."""
+ # 构建一段包含个人信息的对话
+ await reme.vector_store.delete_collection("reme")
+
+ messages = [
+ {
+ "role": "user",
+ "content": "你好,我是张伟,今年28岁,是一名软件工程师。",
+ },
+ {
+ "role": "assistant",
+ "content": "你好张伟!很高兴认识你。作为一名软件工程师,你主要从事什么方向的开发工作呢?",
+ },
+ {
+ "role": "user",
+ "content": "我主要做后端开发,擅长Python和Go语言。最近在研究AI Agent相关的技术。",
+ },
+ {
+ "role": "assistant",
+ "content": "很棒!Python和Go都是非常实用的语言。AI Agent是当前很热门的方向,你在这方面有什么具体的研究重点吗?",
+ },
+ {
+ "role": "user",
+ "content": "我特别关注记忆系统的设计,希望能让AI Agent具有长期记忆能力。我的工作地点在北京,平时喜欢看技术博客和参加技术分享会。",
+ },
+ {
+ "role": "assistant",
+ "content": "记忆系统确实是AI Agent的核心能力之一。北京有很多优秀的技术社区和活动,相信你能找到很多志同道合的朋友。",
+ },
+ {
+ "role": "user",
+ "content": "是的,我每周末都会去参加一些技术沙龙。对了,我的邮箱是zhangwei@example.com,如果有好的技术资料可以发给我。",
+ },
+ {
+ "role": "assistant",
+ "content": "好的,我记下了。保持学习的热情很重要,祝你在AI Agent领域的研究越来越深入!",
+ },
+ ]
+
+ print("=" * 60)
+ print("步骤1: 开始总结对话并生成记忆")
+ print("=" * 60)
+
+ # 对对话进行总结,生成记忆
+ await reme.summary(
+ messages=messages,
+ user_id="zhangwei",
+ description="用户自我介绍和技术兴趣分享",
+ ref_memory_id="ref_123",
+ )
+
+ print("\n✓ 记忆总结完成")
+
+ print("\n" + "=" * 60)
+ print("步骤2: 查看已存储的记忆节点")
+ print("=" * 60)
+
+ # 列出所有存储的记忆节点
+ nodes: list[VectorNode] = await reme.vector_store.list()
+ for i, node in enumerate(nodes, 1):
+ memory_node = MemoryNode.from_vector_node(node)
+ print(f"{i} {memory_node.memory_type} {memory_node.memory_target} {memory_node.format_memory()}")
+
+ print("\n" + "=" * 60)
+ print("步骤3: 测试记忆检索 - 验证个人信息")
+ print("=" * 60)
+
+ # 测试问题1: 检索用户姓名
+ query1 = "用户叫什么名字?"
+ print(f"\n问题1: {query1}")
+ result1 = await reme.retrieve(query=query1, user_id="zhangwei")
+ print(f"检索结果:\n{result1}")
+
+ # 测试问题2: 检索技术背景
+ query2 = "用户擅长什么编程语言和技术方向?"
+ print(f"\n问题2: {query2}")
+ result2 = await reme.retrieve(query=query2, user_id="zhangwei")
+ print(f"检索结果:\n{result2}")
+
+ # 测试问题3: 检索个人信息
+ query3 = "用户的工作地点和联系方式是什么?"
+ print(f"\n问题3: {query3}")
+ result3 = await reme.retrieve(query=query3, user_id="zhangwei")
+ print(f"检索结果:\n{result3}")
+
+ # 测试问题4: 检索兴趣爱好
+ query4 = "用户平时有什么爱好或活动?"
+ print(f"\n问题4: {query4}")
+ result4 = await reme.retrieve(query=query4, user_id="zhangwei")
+ print(f"检索结果:\n{result4}")
+
+ print("\n" + "=" * 60)
+ print("测试完成!")
+ print("=" * 60)
+
+
+if __name__ == "__main__":
+ asyncio.run(test_reme())