feat(core): implement memory node tracking and embedding text truncation

This commit is contained in:
jinli.yl 2026-01-09 17:51:38 +08:00
parent d91fc6a14c
commit 0b7843f557
31 changed files with 714 additions and 208 deletions

5
.gitignore vendored
View file

@ -32,4 +32,7 @@ site/*
docs/_build/*
test_compact_storage/*
test_working_memory/*
*.code-workspace
*.code-workspace
local_vector_store/*
bench_results/*
meta_memory/*

329
bench/eval_reme.py Normal file
View file

@ -0,0 +1,329 @@
"""ReMe evaluation script for HaluMem-like benchmarks."""
import asyncio
import copy
import json
import os
import re
import time
from datetime import datetime, timezone
from tqdm import tqdm
from reme_ai.core.enumeration import Role
from reme_ai.core.schema import Message, MemoryNode
from reme_ai.reme import ReMe
TEMPLATE_REME = """Memories for user {user_id}:
{memories}
"""
RETRY_TIMES = 3
WAIT_TIME = 2
# Default prompt for answering questions with memory context
PROMPT_REME = """You are a helpful AI assistant with access to the user's memories.
Use the following context to answer the user's question accurately.
Context:
{context}
Question: {question}
Please provide a detailed and accurate answer based on the available context.
If the context doesn't contain enough information to answer the question, say so clearly."""
async def add_memory_async(
reme: ReMe,
user_id: str,
messages: list[dict],
description: str = "",
):
"""Add memory to ReMe system asynchronously."""
start = time.time()
result = await reme.summary(
messages=messages,
user_id=user_id,
description=description,
memory_mode="personal",
)
duration_ms = (time.time() - start) * 1000
return result, duration_ms
async def search_memory_async(
reme: ReMe,
query: str,
user_id: str,
top_k: int = 20,
):
"""Search memory from ReMe system asynchronously."""
start = time.time()
result = await reme.retrieve(
query=query,
user_id=user_id,
memory_mode="personal",
top_k=top_k,
)
# Format the context
context = TEMPLATE_REME.format(
user_id=user_id,
memories=result if isinstance(result, str) else json.dumps(result, indent=4, ensure_ascii=False),
)
duration_ms = (time.time() - start) * 1000
return context, result, duration_ms
async def llm_request_async(reme: ReMe, prompt: str):
"""Make LLM request using ReMe's llm."""
messages = [
Message(role=Role.SYSTEM, content="You are a helpful assistant."),
Message(role=Role.USER, content=prompt),
]
response = await reme.llm.chat(messages=messages)
return response.content
def extract_user_name(persona_info: str):
"""Extract user name from persona info."""
match = re.search(r"Name:\s*(.*?); Gender:", persona_info)
if match:
username = match.group(1).strip()
return username
else:
raise ValueError("No name found.")
async def _process_session_questions(
session: dict,
new_session: dict,
reme: ReMe,
user_name: str,
top_k_value: int,
) -> None:
"""Process questions for a session."""
if "questions" not in session:
return
new_session["questions"] = []
for qa in session["questions"]:
context, _, duration_ms = await search_memory_async(
reme=reme,
query=qa["question"],
user_id=user_name,
top_k=top_k_value,
)
new_qa = copy.deepcopy(qa)
new_qa["context"] = context
new_qa["search_duration_ms"] = duration_ms
prompt = PROMPT_REME.format(
context=context,
question=qa["question"],
)
start_time = time.time()
response = await llm_request_async(reme, prompt)
new_qa["system_response"] = response
new_qa["response_duration_ms"] = (time.time() - start_time) * 1000
new_session["questions"].append(new_qa)
async def process_user_async(
user_data: dict,
top_k_value: int,
save_path: str,
reme: ReMe,
):
"""Process a single user's data asynchronously."""
user_name = extract_user_name(user_data["persona_info"])
sessions = user_data["sessions"]
tmp_dir = os.path.join(save_path, "tmp")
os.makedirs(tmp_dir, exist_ok=True)
tmp_file = os.path.join(tmp_dir, f"{user_data['uuid']}.json")
# Clear existing memories for this user
await reme.vector_store.delete_collection(f"reme_eval_{user_name}")
# Update collection name for this user
reme.vector_store.set_collection_name(f"reme_eval_{user_name}")
new_user_data = {
"uuid": user_data["uuid"],
"user_name": user_name,
"sessions": [],
}
for session in tqdm(sessions, total=len(sessions), desc=f"Processing user {user_name}"):
new_session = {
"memory_points": session["memory_points"],
"dialogue": session["dialogue"],
}
# Add messages to ReMe
dialogue = session["dialogue"]
# Parse timestamp and format as "YYYY-MM-DD HH:MM:SS"
date_format = "%b %d, %Y, %H:%M:%S"
# dt = datetime.strptime(session["start_time"], date_format).replace(tzinfo=timezone.utc)
# time_created = dt.strftime("%Y-%m-%d %H:%M:%S")
formatted_dialogue = [
{
"role": turn["role"],
"content": turn["content"],
"time_created": datetime.strptime(turn["timestamp"], date_format)
.replace(tzinfo=timezone.utc)
.strftime("%Y-%m-%d %H:%M:%S"),
}
for turn in dialogue
]
# Add memory
result, duration_ms = await add_memory_async(
reme=reme,
user_id=user_name,
messages=formatted_dialogue,
)
memories = []
for memory_modes in result:
for memory_mode in memory_modes:
if not isinstance(memory_mode, MemoryNode):
continue
memories.append(memory_mode.content)
print(memories)
if session.get("is_generated_qa_session", False):
new_session["add_dialogue_duration_ms"] = duration_ms
new_session["is_generated_qa_session"] = True
del new_session["dialogue"]
del new_session["memory_points"]
new_user_data["sessions"].append(new_session)
continue
# Store the result from summary
new_session["extracted_memories"] = memories
new_session["add_dialogue_duration_ms"] = duration_ms
# Search updated memories for memory points
# for memory in new_session["memory_points"]:
# if memory["is_update"] == "False" or not memory["original_memories"]:
# continue
#
# _, memories_from_system, duration_ms = await search_memory_async(
# reme=reme,
# query=memory["memory_content"],
# user_id=user_name,
# top_k=10,
# )
#
# memory["memories_from_system"] = str(memories_from_system)
# Process questions
await _process_session_questions(session, new_session, reme, user_name, top_k_value)
new_user_data["sessions"].append(new_session)
with open(tmp_file, "w", encoding="utf-8") as f:
json.dump(new_user_data, f, ensure_ascii=False, indent=2)
# raise NotImplementedError
# Save results
with open(tmp_file, "w", encoding="utf-8") as f:
json.dump(new_user_data, f, ensure_ascii=False, indent=2)
print(f"✅ Saved user {user_name} to {tmp_file}")
return {"uuid": user_data["uuid"], "status": "ok", "path": tmp_file}
def iter_jsonl(file_path: str):
"""Iterate over lines in a JSONL file."""
with open(file_path, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if line:
yield json.loads(line)
async def main_async(
data_path_arg: str,
version_arg: str = "default",
top_k_arg: int = 20,
):
"""Main evaluation function."""
frame = "reme"
save_path = f"bench_results/{frame}-{version_arg}/"
os.makedirs(save_path, exist_ok=True)
output_file = os.path.join(save_path, f"{frame}_eval_results.jsonl")
tmp_dir = os.path.join(save_path, "tmp")
os.makedirs(tmp_dir, exist_ok=True)
start_time = time.time()
# Initialize ReMe instance (will reuse for all users)
reme = ReMe()
# Load all user data
user_data_list = list(iter_jsonl(data_path_arg))
total_users = len(user_data_list)
print(f"Processing {total_users} users sequentially...")
# Sequential processing
for idx, user_data in enumerate(user_data_list, 1):
result = await process_user_async(user_data, top_k_arg, save_path, reme)
print(f"[{idx}/{total_users}] ✅ Finished {user_data['uuid']} ({result['status']})")
# Combine all results into final output
with open(output_file, "w", encoding="utf-8") as f_out:
for file in os.listdir(tmp_dir):
if file.endswith(".json"):
file_path = os.path.join(tmp_dir, file)
with open(file_path, "r", encoding="utf-8") as f_in:
data = json.load(f_in)
f_out.write(json.dumps(data, ensure_ascii=False) + "\n")
elapsed = time.time() - start_time
print(f"✅ All done in {elapsed:.2f}s")
print(f"✅ Final results saved to: {output_file}")
def main(
data_path_arg: str,
version_arg: str = "default",
top_k_arg: int = 20,
):
"""Synchronous entry point for main evaluation."""
asyncio.run(main_async(data_path_arg, version_arg, top_k_arg))
if __name__ == "__main__":
# Example usage - update these paths as needed
# Note: Don't use HaluMem-long.jsonl directly as each line is too large
# Instead, create a smaller test dataset or use a different data file
DEFAULT_DATA_PATH = "/Users/yuli/workspace/HaluMem/data/HaluMem-Long.jsonl"
DEFAULT_VERSION = "test"
DEFAULT_TOP_K = 20
main(
data_path_arg=DEFAULT_DATA_PATH,
version_arg=DEFAULT_VERSION,
top_k_arg=DEFAULT_TOP_K,
)

View file

@ -26,6 +26,7 @@ class BaseEmbeddingModel(ABC):
max_batch_size: int = 10,
max_retries: int = 3,
raise_exception: bool = True,
max_input_length: int = 8192,
**kwargs,
):
"""Initialize model configuration and parameters."""
@ -34,8 +35,22 @@ class BaseEmbeddingModel(ABC):
self.max_batch_size = max_batch_size
self.max_retries = max_retries
self.raise_exception = raise_exception
self.max_input_length = max_input_length
self.kwargs = kwargs
def _truncate_text(self, text: str) -> str:
"""Truncate text to max_input_length if it exceeds the limit."""
if len(text) > self.max_input_length:
logger.warning(
f"Text length {len(text)} exceeds max_input_length {self.max_input_length}, truncating"
)
return text[: self.max_input_length]
return text
def _truncate_texts(self, texts: list[str]) -> list[str]:
"""Truncate a list of texts to max_input_length."""
return [self._truncate_text(text) for text in texts]
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Internal async implementation for calling the embedding API with batch input."""
@ -44,9 +59,10 @@ class BaseEmbeddingModel(ABC):
async def get_embedding(self, input_text: str, **kwargs) -> list[float]:
"""Async get embedding for a single text with exponential backoff retries."""
truncated_text = self._truncate_text(input_text)
for i in range(self.max_retries):
try:
result = await self._get_embeddings([input_text], **kwargs)
result = await self._get_embeddings([truncated_text], **kwargs)
return result[0]
except Exception as e:
logger.error(f"Model {self.model_name} failed: {e}")
@ -59,10 +75,13 @@ class BaseEmbeddingModel(ABC):
async def get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Async get embeddings with automatic batching and exponential backoff retries."""
# Truncate all input texts first
truncated_texts = self._truncate_texts(input_text)
# Split into batches and process sequentially to respect rate limits
results = []
for i in range(0, len(input_text), self.max_batch_size):
batch = input_text[i : i + self.max_batch_size]
for i in range(0, len(truncated_texts), self.max_batch_size):
batch = truncated_texts[i : i + self.max_batch_size]
# Process each batch with retry logic
for retry in range(self.max_retries):
try:
@ -81,9 +100,10 @@ class BaseEmbeddingModel(ABC):
def get_embedding_sync(self, input_text: str, **kwargs) -> list[float]:
"""Synchronous get embedding for a single text with retry logic."""
truncated_text = self._truncate_text(input_text)
for i in range(self.max_retries):
try:
result = self._get_embeddings_sync([input_text], **kwargs)
result = self._get_embeddings_sync([truncated_text], **kwargs)
return result[0]
except Exception as exc:
logger.error(f"Model {self.model_name} failed: {exc}")
@ -96,9 +116,12 @@ class BaseEmbeddingModel(ABC):
def get_embeddings_sync(self, input_text: list[str], **kwargs) -> list[list[float]]:
"""Synchronous get embeddings with automatic batching and retry logic."""
# Truncate all input texts first
truncated_texts = self._truncate_texts(input_text)
results = []
for i in range(0, len(input_text), self.max_batch_size):
batch = input_text[i : i + self.max_batch_size]
for i in range(0, len(truncated_texts), self.max_batch_size):
batch = truncated_texts[i : i + self.max_batch_size]
# Process each batch with retry logic
for retry in range(self.max_retries):
try:

View file

@ -62,7 +62,7 @@ class BaseFlow(ABC):
payload = json.dumps(params, sort_keys=True, ensure_ascii=False, default=str)
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
except Exception as e:
logger.exception(f"{self.name} cache key serialization failed: {e}")
logger.exception(f"[{self.__class__.__name__}] {self.name} cache key serialization failed: {e}")
return None
def _maybe_load_cached(self, params: dict) -> Response | None:
@ -72,7 +72,7 @@ class BaseFlow(ABC):
if key := self._compute_cache_key(params):
if cached := self.cache.load(key):
logger.info(f"Loaded {self.name} response from cache.")
logger.info(f"[{self.__class__.__name__}] Loaded {self.name} response from cache.")
return Response(**cached)
return None
@ -92,7 +92,7 @@ class BaseFlow(ABC):
"""Recursively log the hierarchy of the flow's operation tree."""
prefix = " " * indent
op_type = "sequential" if isinstance(op, SequentialOp) else "parallel" if isinstance(op, ParallelOp) else name
logger.info(f"{prefix}{op_type} execution")
logger.info(f"[{self.__class__.__name__}] {prefix}{op_type} execution")
for sub_op in op.sub_ops or []:
self._print_operation_tree(sub_op.name, sub_op, indent + 2)
@ -160,15 +160,15 @@ class BaseFlow(ABC):
def print_flow(self):
"""Log the visual structure of the flow once."""
if not self._flow_printed:
logger.info(f"---------- [Flow Structure] {self.name} ----------")
logger.info(f"[{self.__class__.__name__}] ---------- [Flow Structure] {self.name} ----------")
self._print_operation_tree(self.name, self.flow_op, 0)
logger.info("-" * 50)
logger.info(f"[{self.__class__.__name__}] " + "-" * 50)
self._flow_printed = True
async def call(self, **kwargs) -> Response | asyncio.Queue:
"""Execute the flow asynchronously with parameter caching."""
kwargs["stream"] = self.stream
logger.info(f"{self.name} incoming params: {kwargs}")
logger.info(f"[{self.__class__.__name__}] {self.name} incoming params: {kwargs}")
if cached := self._maybe_load_cached(kwargs):
return cached
@ -187,7 +187,7 @@ class BaseFlow(ABC):
self._maybe_save_cache(kwargs, result)
return result
except Exception as e:
logger.exception(f"{self.name} async call failed: {e}")
logger.exception(f"[{self.__class__.__name__}] {self.name} async call failed: {e}")
if self.raise_exception:
raise e
if self.stream:
@ -199,7 +199,7 @@ class BaseFlow(ABC):
def call_sync(self, **kwargs) -> Response:
"""Execute the flow synchronously with parameter caching."""
logger.info(f"{self.name} incoming sync params: {kwargs}")
logger.info(f"[{self.__class__.__name__}] {self.name} incoming sync params: {kwargs}")
assert not self.stream, "Synchronous call cannot be used in stream mode."
if cached := self._maybe_load_cached(kwargs):
return cached
@ -214,7 +214,7 @@ class BaseFlow(ABC):
self._maybe_save_cache(kwargs, context.response)
return context.response
except Exception as e:
logger.exception(f"{self.name} sync call failed: {e}")
logger.exception(f"[{self.__class__.__name__}] {self.name} sync call failed: {e}")
if self.raise_exception:
raise e
context.add_response_error(e)

View file

@ -100,7 +100,7 @@ class BaseOp:
def _handle_failure(self, e: Exception, attempt: int):
"""Log failures and handle final retry logic."""
message = f"{self.name} failed (attempt {attempt + 1}): {e}"
message = f"[{self.__class__.__name__}] {self.name} failed (attempt {attempt + 1}): {e}"
if attempt == self.max_retries - 1:
logger.exception(message)
if self.raise_exception:
@ -305,7 +305,7 @@ class BaseOp:
results = []
for res in raw_results:
if isinstance(res, Exception):
logger.error(f"Async task failed: {res}")
logger.error(f"[{self.__class__.__name__}] Async task failed: {res}")
continue
if res:
results.extend(res if isinstance(res, list) else [res])

View file

@ -124,12 +124,12 @@ class Message(BaseModel):
"""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"
header = f"{self.name or self.role.value if use_name else self.role.value}:"
lines = [f"{prefix}{time_str}{header}"]
if add_reasoning and self.reasoning_content:
lines.append(f"{self.reasoning_content}\n")
lines.append(self.reasoning_content)
if isinstance(self.content, str):
lines.append(self.content)
@ -144,7 +144,7 @@ class Message(BaseModel):
for tc in self.tool_calls:
lines.append(f" - tool_call={tc.name} params={tc.arguments}")
return "\n".join(lines).strip()
return " ".join(lines).strip()
class Trajectory(BaseModel):

View file

@ -48,6 +48,10 @@ class BaseVectorStore(ABC):
"""Convert multiple text queries into vector embeddings using the configured model."""
return await self.embedding_model.get_embeddings(queries)
def set_collection_name(self, collection_name: str):
"""Change the name of the current collection."""
self.collection_name = collection_name
@abstractmethod
async def list_collections(self) -> list[str]:
"""Retrieve a list of all existing collection names in the store."""

View file

@ -103,7 +103,7 @@ class ChromaVectorStore(BaseVectorStore):
metadata = metadatas[i] if i < len(metadatas) and metadatas[i] else {}
if include_score and distances and i < len(distances):
metadata["_score"] = 1.0 - distances[i]
metadata["score"] = 1.0 - distances[i]
node = VectorNode(
vector_id=vector_id,
@ -300,7 +300,7 @@ class ChromaVectorStore(BaseVectorStore):
score_threshold = kwargs.get("score_threshold")
if score_threshold is not None:
nodes = [n for n in nodes if n.metadata.get("_score", 0) >= score_threshold]
nodes = [n for n in nodes if n.metadata.get("score", 0) >= score_threshold]
return nodes
async def delete(self, vector_ids: str | list[str], **kwargs):
@ -392,6 +392,15 @@ class ChromaVectorStore(BaseVectorStore):
await self._run_sync_in_executor(_recreate)
logger.info(f"Collection {self.collection_name} has been reset")
def set_collection_name(self, collection_name: str):
"""Set the collection name and reinitialize the collection object."""
super().set_collection_name(collection_name)
self.collection = self.client.get_or_create_collection(
name=collection_name,
metadata={"hnsw:space": "cosine"},
)
logger.info(f"Collection name set to {collection_name}, collection object reinitialized")
async def close(self):
"""Close the vector store and log the shutdown process."""
logger.info(f"ChromaDB vector store for collection {self.collection_name} closed")

View file

@ -279,7 +279,7 @@ class ESVectorStore(BaseVectorStore):
vector=source.get("vector"),
metadata=source.get("metadata", {}),
)
node.metadata["_score"] = hit["_score"]
node.metadata["score"] = hit["_score"]
results.append(node)
return results
@ -452,6 +452,12 @@ class ESVectorStore(BaseVectorStore):
return results
def set_collection_name(self, collection_name: str):
"""Set the collection name and ensure it's lowercase for Elasticsearch compatibility."""
collection_name = collection_name.lower()
super().set_collection_name(collection_name)
logger.info(f"Collection name set to {collection_name} (converted to lowercase)")
async def close(self):
"""Terminate the Elasticsearch client session and release resources."""
await self.client.close()

View file

@ -202,7 +202,7 @@ class LocalVectorStore(BaseVectorStore):
scored_nodes = scored_nodes[:limit]
results = []
for node, score in scored_nodes:
node.metadata["_score"] = score
node.metadata["score"] = score
results.append(node)
return results
@ -276,6 +276,12 @@ class LocalVectorStore(BaseVectorStore):
return filtered_nodes
def set_collection_name(self, collection_name: str):
"""Set the collection name and reinitialize the collection path."""
super().set_collection_name(collection_name)
self.collection_path = self.root_path / collection_name
logger.info(f"Collection name set to {collection_name}, path updated to {self.collection_path}")
async def close(self):
"""Close the vector store (no-op for local file system)."""
logger.info("Local vector store closed")

View file

@ -324,7 +324,7 @@ class PGVectorStore(BaseVectorStore):
if isinstance(metadata, str):
metadata = json.loads(metadata)
metadata["_score"] = 1 - distance
metadata["score"] = 1 - distance
metadata["_distance"] = distance
node = VectorNode(

View file

@ -304,7 +304,7 @@ class QdrantVectorStore(BaseVectorStore):
vector=point.vector if hasattr(point, "vector") else None,
metadata=payload.get("metadata", {}),
)
node.metadata["_score"] = point.score
node.metadata["score"] = point.score
nodes.append(node)
return nodes

View file

@ -1,13 +1,14 @@
"""Base memory agent for handling memory operations with tool-based reasoning."""
import asyncio
import json
from abc import ABCMeta
from loguru import logger
from ..core.enumeration import Role, MemoryType
from ..core.op import BaseOp
from ..core.schema import Message, ToolCall
from ..core.schema import Message, ToolCall, MemoryNode
from ..mem_tool import BaseMemoryTool, ThinkTool
@ -35,6 +36,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
self.messages: list[Message] = []
self.success: bool = True
self.memory_nodes: list[MemoryNode | str] = []
def _build_tool_call(self) -> ToolCall:
return ToolCall(
@ -71,7 +73,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
)
@property
def tools(self):
def tools(self) -> list[BaseMemoryTool]:
"""Returns the list of memory tools available to this agent."""
return self.sub_ops
@ -79,7 +81,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
def tools(self, tools: list[BaseMemoryTool]):
self.sub_ops = tools
def get_messages(self) -> list[Message]:
def get_messages(self) -> list[Message] | str:
"""Extracts and returns messages from the context query or messages."""
if self.context.get("query"):
messages = [Message(role=Role.USER, content=self.context.query)]
@ -100,7 +102,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
**kwargs,
)
messages.append(assistant_message)
logger.info(f"step{step + 1}.assistant={assistant_message.simple_dump(enable_json_dump=True)}")
logger.info(f"[{self.__class__.__name__}] step{step + 1}.assistant={assistant_message.simple_dump(enable_json_dump=True)}")
should_act = bool(assistant_message.tool_calls)
return assistant_message, should_act
@ -114,10 +116,10 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
for j, tool_call in enumerate(assistant_message.tool_calls):
if tool_call.name not in tool_dict:
logger.warning(f"unknown tool_call.name={tool_call.name}")
logger.warning(f"[{self.__class__.__name__}] unknown tool_call.name={tool_call.name}")
continue
logger.info(f"step{step + 1}.{j} submit tool_calls={tool_call.name} argument={tool_call.arguments}")
logger.info(f"[{self.__class__.__name__}] 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)
@ -129,6 +131,9 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
await self.join_async_tasks()
for j, op in enumerate(tool_list):
if op.memory_nodes:
self.memory_nodes.extend(op.memory_nodes)
tool_result = str(op.output)
tool_message = Message(
role=Role.TOOL,
@ -136,7 +141,7 @@ class BaseMemoryAgent(BaseOp, metaclass=ABCMeta):
tool_call_id=op.tool_call.id,
)
tool_result_messages.append(tool_message)
logger.info(f"step{step + 1}.{j} join tool_result={tool_result[:200]}...\n\n")
logger.info(f"[{self.__class__.__name__}] step{step + 1}.{j} join tool_result={tool_result[:500]}...\n\n")
return tool_result_messages
async def react(self, messages: list[Message]):
@ -157,13 +162,16 @@ 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(enable_json_dump=True)}")
logger.info(f"[{self.__class__.__name__}] step0.{i} {message.role} {message.name or ''} "
f"{message.simple_dump(enable_json_dump=True)}")
for i, tool in enumerate(self.tools):
logger.info(f"[{self.__class__.__name__}] step0.{i} tool_call={json.dumps(tool.tool_call.simple_input_dump(), ensure_ascii=False)}")
self.messages, self.success = await self.react(messages)
if self.success and self.messages:
self.output = self.messages[-1].content
else:
self.output = ""
self.output = "No relevant memories found."
@property
def memory_target(self) -> str:

View file

@ -13,33 +13,36 @@ 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, meta_memories: list[dict] = None, **kwargs):
def __init__(self, meta_memories: list[dict] | None = None, **kwargs):
super().__init__(**kwargs)
self.meta_memories: list[dict] = meta_memories
self.meta_memories: list[dict] = meta_memories or []
@staticmethod
async def _read_meta_memories() -> str:
"""Read and return meta memories as string."""
async def _read_meta_memories(self) -> str:
"""Fetch all meta-memory entries that define specialized memory agents."""
from ...mem_tool import ReadMetaMemory
op = ReadMetaMemory(enable_identity_memory=False)
await op.call()
return str(op.output)
if self.meta_memories:
return op.format_memory_metadata(self.meta_memories)
else:
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)
meta_memory_info = await self._read_meta_memories()
if self.context.get("query"):
context = self.context.query
elif self.context.get("messages"):
messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages]
context = self.description + format_messages(messages)
else:
meta_memory_info = await self._read_meta_memories()
raise ValueError("input must have either `query` or `messages`")
system_prompt = self.prompt_format(
prompt_name="system_prompt",
now_time=get_now_time(),
meta_memory_info=meta_memory_info,
context=format_messages(self.get_messages()),
context=context,
)
messages = [

View file

@ -1,12 +1,12 @@
tool: |
Retrieve relevant memories from the memory bank to assist in answering questions.
Retrieve relevant memories to assist in answering questions.
Use this tool when you need to search for historical information, user preferences,
procedural knowledge, or any other stored memories that may help answer the current query.
The agent will analyze the context, determine what information is needed, and perform
semantic searches across different memory types to find the most relevant memories.
system_prompt: |
You are a memory agent. Please analyze the context, retrieve relevant information from the memory bank when needed, and return a summary of the retrieved memories to assist in answering the user's question.
You are a memory agent. Please analyze the context, retrieve relevant memories when needed, and return a summary of the retrieved memories to assist in answering the user's question.
## Context
{context}
@ -47,4 +47,4 @@ system_prompt: |
- If multiple attempts still yield no relevant memory, output `<NO_RELEVANT_MEMORY>`.
user_message: |
Please analyze the context, retrieve relevant information from the memory bank when needed, and return a summary of the retrieved memories to assist in answering the user's question.
Please analyze the context, retrieve relevant memories when needed, and return a summary of the retrieved memories to assist in answering the user's question.

View file

@ -3,7 +3,7 @@
from ..base_memory_agent import BaseMemoryAgent
from ...core.context import C
from ...core.enumeration import Role, MemoryType
from ...core.schema import Message
from ...core.schema import Message, ToolCall
from ...core.utils import get_now_time, format_messages
@ -13,6 +13,36 @@ class PersonalSummarizer(BaseMemoryAgent):
memory_type: MemoryType = MemoryType.PERSONAL
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": self.get_prompt("tool"),
"parameters": {
"type": "object",
"properties": {
"messages": {
"type": "array",
"items": {
"type": "object",
"properties": {
"role": {
"type": "string",
"description": "role",
},
"content": {
"type": "string",
"description": "content",
},
},
"required": ["role", "content"],
},
},
},
"required": ["messages"],
},
},
)
async def build_messages(self) -> list[Message]:
"""Construct messages with context, memory_target, and memory_type information."""
system_prompt = self.prompt_format(
@ -34,9 +64,9 @@ class PersonalSummarizer(BaseMemoryAgent):
return await super()._acting_step(
assistant_message,
step,
memory_target=self.memory_target,
memory_type=self.memory_type.value,
author=self.author,
memory_target=self.memory_target,
ref_memory_id=self.ref_memory_id,
author=self.author,
**kwargs,
)

View file

@ -8,7 +8,7 @@ from loguru import logger
from ..base_memory_agent import BaseMemoryAgent
from ...core.context import C
from ...core.enumeration import Role
from ...core.schema import Message, MemoryNode
from ...core.schema import Message, MemoryNode, ToolCall
from ...core.utils import get_now_time, format_messages
@ -16,11 +16,41 @@ from ...core.utils import get_now_time, format_messages
class ReMeSummarizer(BaseMemoryAgent):
"""Coordinates memory updates by delegating to specialized memory agents."""
def __init__(self, enable_tool_memory: bool = True, enable_identity_memory: bool = True, **kwargs):
"""Initialize with flags to enable/disable tool and identity memory processing."""
def __init__(self, meta_memories: list[dict] | None = None, enable_identity_memory: bool = False, **kwargs):
"""Initialize with flags to enable/disable identity memory processing."""
super().__init__(**kwargs)
self.enable_tool_memory = enable_tool_memory
self.enable_identity_memory = enable_identity_memory
self.meta_memories: list[dict] = meta_memories or []
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": self.get_prompt("tool"),
"parameters": {
"type": "object",
"properties": {
"messages": {
"type": "array",
"items": {
"type": "object",
"properties": {
"role": {
"type": "string",
"description": "role",
},
"content": {
"type": "string",
"description": "content",
},
},
"required": ["role", "content"],
},
},
},
"required": ["messages"],
},
},
)
async def _add_history_memory(self) -> MemoryNode:
"""Store conversation history and return the memory node."""
@ -28,7 +58,7 @@ class ReMeSummarizer(BaseMemoryAgent):
op = AddHistoryMemory()
await op.call(messages=self.get_messages())
return op.output
return op.memory_nodes[0]
@staticmethod
async def _read_identity_memory() -> str:
@ -43,27 +73,28 @@ class ReMeSummarizer(BaseMemoryAgent):
"""Fetch all meta-memory entries that define specialized memory agents."""
from ...mem_tool import ReadMetaMemory
op = ReadMetaMemory(
enable_tool_memory=self.enable_tool_memory,
enable_identity_memory=self.enable_identity_memory,
)
await op.call()
return str(op.output)
op = ReadMetaMemory(enable_identity_memory=self.enable_identity_memory)
if self.meta_memories:
return op.format_memory_metadata(self.meta_memories)
else:
await op.call()
return str(op.output)
async def build_messages(self) -> List[Message]:
"""Construct initial messages with context, identity, and meta-memory information."""
memory_node: MemoryNode = await self._add_history_memory()
self.context["ref_memory_id"] = memory_node.memory_id
now_time = get_now_time()
identity_memory = await self._read_identity_memory()
meta_memory_info = await self._read_meta_memories()
context = format_messages(self.get_messages())
context = self.description + "\n" + format_messages(self.get_messages())
logger.info(
f"now_time={now_time} "
f"memory_node={memory_node} "
f"memory_node={memory_node.content[:100]}... "
f"identity_memory={identity_memory} "
f"meta_memory_info={meta_memory_info} "
f"context={context}",
f"context={context[:100]}",
)
system_prompt = self.prompt_format(
@ -84,13 +115,21 @@ class ReMeSummarizer(BaseMemoryAgent):
async def _reasoning_step(self, messages: list[Message], step: int, **kwargs) -> tuple[Message, bool]:
"""Refresh meta-memory info in system prompt before each reasoning step."""
meta_memory_info = await self._read_meta_memories()
system_messages = [message for message in messages if message.role is Role.SYSTEM]
if system_messages:
system_message = system_messages[0]
pattern = r'("- <memory_type>\(<memory_target>\): <description>"\n)(.*?)(\n\n)'
replacement = rf"\g<1>{meta_memory_info}\g<3>"
system_message.content = re.sub(pattern, replacement, system_message.content, flags=re.DOTALL)
now_time = get_now_time()
identity_memory = await self._read_identity_memory()
meta_memory_info = await self._read_meta_memories()
context = self.description + "\n" + format_messages(self.get_messages())
system_message.content = self.prompt_format(
prompt_name="system_prompt",
now_time=now_time,
identity_memory=identity_memory,
meta_memory_info=meta_memory_info,
context=context,
)
return await super()._reasoning_step(messages, step, **kwargs)
@ -99,6 +138,8 @@ class ReMeSummarizer(BaseMemoryAgent):
return await super()._acting_step(
assistant_message,
step,
messages=self.context.get("messages", []),
description=self.context.get("description"),
ref_memory_id=self.context["ref_memory_id"],
author=self.author,
**kwargs,

View file

@ -6,11 +6,11 @@ tool: |
3. Delegating to specialized memory agents for detailed memory extraction and update
system_prompt: |
You are a Memory Agent responsible for performing necessary updates and summaries of the main Agent's memories based on the **context**.
# Context
{context}
You are a Memory Agent responsible for performing necessary updates and summaries of the main Agent's memories based on the **context**.
## Current Time
{now_time}
@ -25,11 +25,12 @@ system_prompt: |
## Your Tasks
### 1. Create New Meta Memory (if needed)
When the context contains significant personal or procedural information not yet covered by existing meta memories:
- Use `add_meta_memory` to create one or more new meta memory entries.
When the context contains significant new valuable information, first check if the Main Agent's Meta Memory already contains a corresponding `<memory_type>(<memory_target>)` entry:
- If the required `<memory_type>(<memory_target>)` does NOT exist in the Meta Memory, use `add_meta_memory` to create a new meta memory entry.
- For personal memories: specify `memory_type="personal"` and `memory_target=<person's name>`.
- For procedural memories: specify `memory_type="procedural"` and `memory_target=<topic or domain>`.
- Each meta memory entry will instantiate a dedicated specialized Memory Agent for that dimension.
- Only create new meta memory entries when necessary; avoid duplicating existing ones.
### 2. Add Summary Memory (if valuable)
When the context includes information worth remembering for quick future recall:

View file

@ -3,6 +3,8 @@
from abc import ABCMeta
from pathlib import Path
from loguru import logger
from ..core.enumeration import MemoryType
from ..core.op import BaseOp
from ..core.schema import ToolCall, MemoryNode
@ -23,7 +25,7 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta):
self.enable_multiple: bool = enable_multiple
self.enable_thinking_params: bool = enable_thinking_params
self.meta_memory_path: str = meta_memory_path
self._meta_memory: CacheHandler | None = None
self.memory_nodes: list[MemoryNode | str] = []
def _build_parameters(self) -> dict:
return {}
@ -58,10 +60,8 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta):
@property
def meta_memory(self) -> CacheHandler:
"""Get or create the meta memory cache handler."""
if self._meta_memory is None:
self._meta_memory = CacheHandler(Path(self.meta_memory_path) / self.vector_store.collection_name)
return self._meta_memory
"""Create the meta memory cache handler."""
return CacheHandler(Path(self.meta_memory_path) / self.vector_store.collection_name)
@property
def memory_type(self) -> MemoryType:
@ -94,7 +94,7 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta):
metadata: dict | None = None,
) -> MemoryNode:
"""Build MemoryNode from content, when_to_use, and metadata."""
return MemoryNode(
node = MemoryNode(
memory_type=memory_type or self.memory_type,
memory_target=memory_target or self.memory_target,
when_to_use=when_to_use or "",
@ -103,3 +103,6 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta):
author=author or self.author,
metadata=metadata or {},
)
logger.opt(depth=1).info(f"[{self.__class__.__name__}] build node={node.model_dump_json(indent=2, exclude_none=True)}")
return node

View file

@ -35,12 +35,7 @@ class HandsOffTool(BaseMemoryTool):
"memory_type": {
"type": "string",
"description": self.get_prompt("memory_type"),
"enum": [
MemoryType.IDENTITY.value,
MemoryType.PERSONAL.value,
MemoryType.PROCEDURAL.value,
MemoryType.TOOL.value,
],
"enum": [k.value for k in self.memory_agent_dict],
},
"memory_target": {
"type": "string",
@ -82,10 +77,7 @@ class HandsOffTool(BaseMemoryTool):
def _parse_memory_type_target(task: dict):
memory_type = task.get("memory_type", "")
memory_target = task.get("memory_target", "")
return {
"memory_type": MemoryType(memory_type),
"memory_target": memory_target,
}
return {"memory_type": MemoryType(memory_type), "memory_target": memory_target}
def _collect_tasks(self) -> list[dict]:
"""Collect memory tasks from context based on enable_multiple flag."""
@ -116,21 +108,17 @@ class HandsOffTool(BaseMemoryTool):
logger.warning(f"No agent found for memory_type={memory_type}")
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 = self.memory_agent_dict[memory_type].copy()
agent_list.append([agent, memory_type, memory_target])
logger.info(f"Task {i}: Submitting {memory_type.value} agent for target={memory_target}")
self.submit_async_task(
agent_copy.call,
agent.call,
query=self.context.get("query", ""),
messages=self.context.get("messages", []),
memory_type=memory_type,
memory_target=memory_target,
description=self.context.get("description"),
ref_memory_id=self.context.get("ref_memory_id", ""),
)
@ -140,6 +128,9 @@ class HandsOffTool(BaseMemoryTool):
results = []
for i, (agent, memory_type, memory_target) in enumerate(agent_list):
result_str = str(agent.output)
if agent.memory_nodes:
self.memory_nodes.extend(agent.memory_nodes)
results.append(
{
"memory_type": memory_type.value,

View file

@ -40,10 +40,11 @@ class AddHistoryMemory(BaseMemoryTool):
messages = [Message(**m) if isinstance(m, dict) else m for m in messages]
memory_content = format_messages(messages)
memory_node = self._build_memory_node(memory_content=memory_content, memory_type=MemoryType.HISTORY)
vector_node = memory_node.to_vector_node()
await self.vector_store.delete(vector_ids=[vector_node.vector_id])
await self.vector_store.insert(nodes=[vector_node])
self.memory_nodes.append(memory_node)
self.output = "Successfully added history memory to vector_store."
logger.info(self.output)

View file

@ -15,48 +15,49 @@ class ReadHistoryMemory(BaseMemoryTool):
return {
"type": "object",
"properties": {
"memory_id": {
"ref_memory_id": {
"type": "string",
"description": self.get_prompt("memory_id"),
"description": self.get_prompt("ref_memory_id"),
},
},
"required": ["memory_id"],
"required": ["ref_memory_id"],
}
def _build_multiple_parameters(self) -> dict:
return {
"type": "object",
"properties": {
"memory_ids": {
"ref_memory_ids": {
"type": "array",
"description": self.get_prompt("memory_ids"),
"description": self.get_prompt("ref_memory_ids"),
"items": {"type": "string"},
},
},
"required": ["memory_ids"],
"required": ["ref_memory_ids"],
}
async def execute(self):
if self.enable_multiple:
memory_ids: list[str] = self.context.get("memory_ids", [])
ref_memory_ids: list[str] = self.context.get("ref_memory_ids", [])
else:
memory_id = self.context.get("memory_id", "")
memory_ids: list[str] = [memory_id] if memory_id else []
ref_memory_id = self.context.get("ref_memory_id", "")
ref_memory_ids: list[str] = [ref_memory_id] if ref_memory_id else []
memory_ids = [mid for mid in memory_ids if mid]
ref_memory_ids = [mid for mid in ref_memory_ids if mid]
if not memory_ids:
self.output = "No valid history memory IDs provided for reading."
if not ref_memory_ids:
self.output = "No valid reference memory IDs provided for reading."
logger.warning(self.output)
return
nodes = await self.vector_store.get(vector_ids=memory_ids)
# Query original history dialogues by ref_memory_id
nodes = await self.vector_store.get(vector_ids=ref_memory_ids)
if not nodes:
self.output = "No history memories found with the provided IDs."
self.output = "No history memories found with the provided reference IDs."
logger.warning(self.output)
return
memories: list[MemoryNode] = [MemoryNode.from_vector_node(n) for n in nodes]
self.output = "---\n".join([m.content for m in memories])
logger.info(f"Successfully read {len(memories)} history memories.")
logger.info(f"Successfully read {len(memories)} history memories by reference IDs.")

View file

@ -1,11 +1,11 @@
tool: |
Read history memory by ID.
Read original history dialogue by reference memory ID.
tool_multiple: |
Read multiple history memories by IDs.
Read multiple original history dialogues by reference memory IDs.
memory_id: |
Unique identifier of the history memory.
ref_memory_id: |
Reference memory ID to query the original history dialogue.
memory_ids: |
List of unique identifiers of history memories.
ref_memory_ids: |
List of reference memory IDs to query the original history dialogues.

View file

@ -1,11 +1,13 @@
tool: |
Add a memory metadata entry to register a new memory type and target.
IMPORTANT: Before using this tool, verify that the Main Agent's Meta Memory does NOT already contain the same <memory_type>(<memory_target>) combination. Only create new entries if they don't exist.
Use this tool to define what types of memories should be tracked, such as:
- Personal memories: "John", "Alice" (person-specific preferences and context)
- Procedural memories: "deployment_process", "code_review_steps" (how-to knowledge)
tool_multiple: |
Add multiple memory metadata entries to register multiple memory types and targets at once.
Before using this tool, verify that the Main Agent's Meta Memory does NOT already contain the same <memory_type>(<memory_target>) combinations. Only create new entries for those that don't exist.
Use this tool to define multiple memory tracking categories in a single operation.
Each entry specifies a memory_type and memory_target for organizing different memory domains.

View file

@ -71,6 +71,7 @@ class AddMemory(BaseMemoryTool):
"description": metadata_description,
"properties": metadata_properties,
}
required.append("metadata")
return properties, required
@ -158,6 +159,7 @@ class AddMemory(BaseMemoryTool):
# Delete existing IDs (upsert behavior), then insert
await self.vector_store.delete(vector_ids=vector_ids)
await self.vector_store.insert(nodes=vector_nodes)
self.memory_nodes = memory_nodes
self.output = f"Successfully added {len(memory_nodes)} memories to vector_store."
logger.info(self.output)

View file

@ -4,6 +4,8 @@ from loguru import logger
from .add_memory import AddMemory
from ...core.context import C
from ...core.enumeration import MemoryType
from ...core.schema import MemoryNode
@C.register_op()
@ -53,6 +55,7 @@ class AddSummaryMemory(AddMemory):
"description": metadata_description,
"properties": metadata_properties,
}
required.append("metadata")
return {
"type": "object",
@ -60,6 +63,30 @@ class AddSummaryMemory(AddMemory):
"required": required,
}
def _build_memory_node(
self,
memory_content: str,
memory_type: MemoryType | None = None,
memory_target: str = "",
ref_memory_id: str = "",
when_to_use: str = "",
author: str = "",
metadata: dict | None = None,
) -> MemoryNode:
"""Build MemoryNode from content, when_to_use, and metadata."""
node = MemoryNode(
memory_type=MemoryType.SUMMARY,
memory_target="",
when_to_use="",
content=memory_content,
ref_memory_id=self.ref_memory_id,
author=self.author,
metadata=metadata or {},
)
logger.info(f"Adding summary memory: {node.model_dump_json(indent=2, exclude_none=True)}")
return node
async def execute(self):
"""Execute addition: map summary_memory to memory_content and call parent."""
# Map summary_memory to memory_content

View file

@ -56,5 +56,6 @@ class DeleteMemory(BaseMemoryTool):
return
await self.vector_store.delete(vector_ids=memory_ids)
self.memory_nodes = memory_ids
self.output = f"Successfully deleted {len(memory_ids)} memories from vector_store."
logger.info(self.output)

View file

@ -169,6 +169,7 @@ class UpdateMemory(BaseMemoryTool):
all_ids_to_delete = list(set(old_memory_ids + new_vector_ids))
await self.vector_store.delete(vector_ids=all_ids_to_delete)
await self.vector_store.insert(nodes=vector_nodes)
self.memory_nodes = new_memory_nodes
self.output = f"Update: deleted {len(old_memory_ids)} old memories, added {len(new_memory_nodes)} new memories."
logger.info(self.output)

View file

@ -23,7 +23,7 @@ class VectorRetrieveMemory(BaseMemoryTool):
enable_summary_memory: bool = False,
add_memory_type_target: bool = False,
metadata_desc: dict[str, str] | None = None,
top_k: int = 10,
top_k: int = 20,
**kwargs,
):
"""Initialize VectorRetrieveMemory.
@ -169,11 +169,7 @@ class VectorRetrieveMemory(BaseMemoryTool):
value = str(value).strip()
filter_dict[key] = [value] if not isinstance(value, list) else value
nodes: list[VectorNode] = await self.vector_store.search(
query=query,
top_k=self.top_k,
filter_dict=filter_dict,
)
nodes: list[VectorNode] = await self.vector_store.search(query=query, limit=self.top_k, filters=filter_dict)
memory_nodes: list[MemoryNode] = [MemoryNode.from_vector_node(n) for n in nodes]
@ -220,8 +216,8 @@ class VectorRetrieveMemory(BaseMemoryTool):
self.output = "No valid query texts provided for retrieval."
return
# Retrieve memories for all queries
memories: list[MemoryNode] = []
# Retrieve memory_nodes for all queries
memory_nodes: list[MemoryNode] = []
for item in query_items:
memory_type = item.get("memory_type") or default_memory_type
memory_target = item.get("memory_target") or default_memory_target
@ -237,14 +233,15 @@ class VectorRetrieveMemory(BaseMemoryTool):
query=item["query"],
metadata_filters=metadata_filters,
)
memories.extend(retrieved)
memory_nodes.extend(retrieved)
# Deduplicate and format output
memories = deduplicate_memories(memories)
memory_nodes = deduplicate_memories(memory_nodes)
self.memory_nodes = memory_nodes
if not memories:
self.output = "No memories found matching the query."
if not memory_nodes:
self.output = "No memory_nodes found matching the query."
else:
self.output = "\n".join([m.format_memory() for m in memories])
self.output = "\n".join([m.format_memory() for m in memory_nodes])
logger.info(f"Retrieved {len(memories)} memories")
logger.info(f"Retrieved {len(memory_nodes)} memory_nodes")

View file

@ -1,27 +1,18 @@
"""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.embedding import BaseEmbeddingModel
from .core.enumeration import Role
from .core.llm import BaseLLM
from .core.schema import Message
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_agent.summarizer import ReMeSummarizer, PersonalSummarizer
from .mem_tool import (
HandsOffTool,
ReadHistoryMemory,
# ReadIdentityMemory,
# UpdateIdentityMemory,
AddMetaMemory,
AddMemory,
AddSummaryMemory,
@ -29,7 +20,6 @@ from .mem_tool import (
UpdateMemory,
VectorRetrieveMemory,
)
from .core.schema import Message
class ReMe(Application):
@ -47,10 +37,6 @@ class ReMe(Application):
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__(
@ -71,30 +57,10 @@ class ReMe(Application):
)
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.llm: BaseLLM = C.get_llm("default")
self.vector_store: BaseVectorStore = C.get_vector_store("default")
self.embedding_model: BaseEmbeddingModel = C.get_embedding_model("default")
@staticmethod
def _prepare_messages(messages: list[dict | Message], user_id: str, assistant_id: str):
@ -115,33 +81,61 @@ class ReMe(Application):
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,
if user_id:
# halumem: user_id -> message.name
# locomo: add description
metadata_summary = {
"year": "The `year` information associated with the memory(Optional)",
"month": "The `month` information associated with the memory(Optional)",
"day": "The `day` information associated with the memory(Optional)",
"hour": "The `hour` information associated with the memory(Optional)",
# "year": "The year when the memory content occurred(Optional)",
# "month": "The month when the memory content occurred(Optional)",
# "day": "The day when the memory content occurred(Optional)",
# "hour": "The hour when the memory content occurred(Optional)",
}
meta_memories = [
{
"memory_type": "personal",
"memory_target": user_id,
},
]
messages = self._prepare_messages(messages, user_id, assistant_id)
personal_summarizer = PersonalSummarizer(
tools=[
VectorRetrieveMemory(
enable_summary_memory=False,
add_memory_type_target=False,
metadata_desc=None,
top_k=15,
),
AddMemory(add_when_to_use=False, metadata_desc=metadata_summary),
DeleteMemory(),
UpdateMemory(add_when_to_use=False, metadata_desc=metadata_summary),
],
)
elif memory_mode == "procedural":
return await self._procedural_summarizer.call(
messages=messages,
description=description,
memory_target=user_id,
**kwargs,
reme_summarizer = ReMeSummarizer(
meta_memories=meta_memories,
enable_identity_memory=False,
tools=[
AddMetaMemory(),
AddSummaryMemory(metadata_desc=metadata_summary),
HandsOffTool(memory_agents=[personal_summarizer]),
],
)
await reme_summarizer.call(messages=messages, description=description, **kwargs)
return reme_summarizer.memory_nodes
else:
return await self._reme_summarizer.call(
messages=messages,
description=description,
memory_target=user_id,
**kwargs,
)
raise NotImplementedError
async def retrieve(
self,
@ -150,19 +144,42 @@ class ReMe(Application):
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)
if user_id:
messages = self._prepare_messages(messages, user_id, assistant_id)
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)
metadata_retrieve = {
"year": "The year to filter memories(Optional)",
"month": "The month to filter memories(Optional)",
"day": "The day to filter memories(Optional)",
"hour": "The hour to filter memories(Optional)",
}
meta_memories = [
{
"memory_type": "personal",
"memory_target": user_id,
},
]
reme_retriever = ReMeRetriever(
meta_memories=meta_memories,
tools=[
VectorRetrieveMemory(
enable_summary_memory=True,
add_memory_type_target=True,
metadata_desc=metadata_retrieve,
top_k=20,
),
ReadHistoryMemory(),
],
)
await reme_retriever.call(query=query, messages=messages, description=description, **kwargs)
return reme_retriever.output
else:
return await self._reme_retriever.call(query=query, messages=messages, description=description, **kwargs)
raise NotImplementedError