refactor(memory): update memory system with service context and vector store enhancements

This commit is contained in:
jinli.yl 2026-01-26 17:33:07 +08:00
parent d3fa645832
commit e77d228fae
16 changed files with 85 additions and 49 deletions

1
.gitignore vendored
View file

@ -34,6 +34,7 @@ test_compact_storage/*
test_working_memory/*
*.code-workspace
local_vector_store/*
reme_local_memory/*
chroma_vector_store/*
bench_results/*
meta_memory/*

View file

@ -29,16 +29,19 @@ class BaseMemoryAgent(BaseReact, metaclass=ABCMeta):
from ...tool.memory import ReadUserProfile
read_tool = ReadUserProfile(show_id=show_id)
await read_tool.call(memory_target=self.memory_target)
await read_tool.call(memory_target=self.memory_target, service_context=self.service_context)
return str(read_tool.response.answer)
@staticmethod
async def read_history_node() -> MemoryNode:
"""Read and return the current history node from the context."""
async def add_history_node(self) -> MemoryNode:
"""Add history node"""
from ...tool.memory import AddHistory
add_history_tool = AddHistory()
await add_history_tool.call()
await add_history_tool.call(
messages=self.messages,
description=self.description,
service_context=self.service_context,
)
return add_history_tool.context.history_node
@property

View file

@ -3,7 +3,7 @@
from ..base_memory_agent import BaseMemoryAgent
from ....core.enumeration import Role, MemoryType
from ....core.op import BaseTool
from ....core.schema import Message
from ....core.schema import Message, MemoryNode
from ....core.utils import format_messages
@ -12,6 +12,10 @@ class PersonalRetriever(BaseMemoryAgent):
memory_type: MemoryType = MemoryType.PERSONAL
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.retrieved_nodes: list[MemoryNode] = []
async def build_messages(self) -> list[Message]:
if self.context.get("query"):
context = self.context.query
@ -52,5 +56,6 @@ class PersonalRetriever(BaseMemoryAgent):
step,
memory_type=self.memory_type.value,
memory_target=self.memory_target,
retrieved_nodes=self.retrieved_nodes,
**kwargs,
)

View file

@ -5,7 +5,7 @@ from loguru import logger
from ..base_memory_agent import BaseMemoryAgent
from ....core.enumeration import Role, MemoryType
from ....core.op import BaseTool
from ....core.schema import Message
from ....core.schema import Message, MemoryNode
class PersonalSummarizer(BaseMemoryAgent):
@ -13,6 +13,10 @@ class PersonalSummarizer(BaseMemoryAgent):
memory_type: MemoryType = MemoryType.PERSONAL
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.retrieved_nodes: list[MemoryNode] = []
async def _build_phase1_messages(self) -> list[Message]:
"""Build messages for phase 1: retrieve and add memory."""
return [
@ -68,6 +72,7 @@ class PersonalSummarizer(BaseMemoryAgent):
memory_target=self.memory_target,
history_node=self.history_node,
author=self.author,
retrieved_nodes=self.retrieved_nodes,
**kwargs,
)
@ -75,7 +80,7 @@ class PersonalSummarizer(BaseMemoryAgent):
"""Execute two-phase memory processing: retrieve/add -> update profile."""
tools = self.tools
for i, tool in enumerate(tools):
logger.info(f"[{self.__class__.__name__}] tool_call[{i}]={tool.tool_call.name}")
logger.info(f"[{self.__class__.__name__}] tool_call[{i}]={tool.tool_call.simple_input_dump(as_dict=False)}")
messages_phase1 = await self._build_phase1_messages()
for i, message in enumerate(messages_phase1):

View file

@ -52,15 +52,13 @@ class ReMeRetriever(BaseMemoryAgent):
description=self.description,
messages=self.messages,
query=self.query,
history_node=self.history_node,
author=self.author,
**kwargs,
)
async def execute(self):
await super().execute()
tools: list[BaseTool] = self.response.metadata["tools"]
result = await super().execute()
tools: list[BaseTool] = result["tools"]
hands_off_tool = tools[0]
agents: list[BaseMemoryAgent] = hands_off_tool.response.metadata["agents"]
@ -70,7 +68,7 @@ class ReMeRetriever(BaseMemoryAgent):
tools = []
for agent in agents:
answer += "\n" + agent.response.answer
success = success and agent.response.metadata["success"]
success = success and agent.response.success
messages += agent.response.metadata["messages"]
tools += agent.response.metadata["tools"]

View file

@ -14,7 +14,7 @@ class ReMeSummarizer(BaseMemoryAgent):
self.meta_memories: list[dict] = meta_memories or []
async def build_messages(self) -> list[Message]:
self.context.history_node = await self.read_history_node()
self.context.history_node = await self.add_history_node()
messages = [
Message(
@ -53,9 +53,8 @@ class ReMeSummarizer(BaseMemoryAgent):
)
async def execute(self):
await super().execute()
tools: list[BaseTool] = self.response.metadata["tools"]
result = await super().execute()
tools: list[BaseTool] = result["tools"]
hands_off_tool = tools[0]
agents: list[BaseMemoryAgent] = hands_off_tool.response.metadata["agents"]
@ -65,7 +64,7 @@ class ReMeSummarizer(BaseMemoryAgent):
tools = []
for agent in agents:
answer += "\n" + agent.response.answer
success = success and agent.response.metadata["success"]
success = success and agent.response.success
messages += agent.response.metadata["messages"]
tools += agent.response.metadata["tools"]

View file

@ -101,7 +101,6 @@ class PromptHandler(BaseContext):
prompt_file_path = Path(prompt_file_path)
if not prompt_file_path.exists():
logger.warning(f"Prompt file not found: {prompt_file_path}")
return self
suffix = prompt_file_path.suffix.lower()
@ -117,7 +116,6 @@ class PromptHandler(BaseContext):
f"Unsupported file format: {suffix}. " f"Supported formats: .yaml, .yml, .json",
)
logger.info(f"Loaded {len(prompt_dict or {})} prompts from {prompt_file_path}")
self.load_prompt_dict(prompt_dict, overwrite=overwrite)
except (yaml.YAMLError, json.JSONDecodeError) as e:

View file

@ -94,12 +94,12 @@ class BaseOp(metaclass=ABCMeta):
def _handle_failure(self, e: Exception, attempt: int) -> str | None:
"""Log failures and handle final retry logic."""
message = f"[{self.__class__.__name__}] {self.name} failed (attempt {attempt + 1}): {e}"
message = f"[{self.__class__.__name__}] failed (attempt {attempt + 1}): {e}"
if attempt == self.max_retries - 1:
logger.exception(message)
if self.raise_exception:
raise e
return f"{self.name} failed: {e}"
return f"[{self.__class__.__name__}] failed: {e}"
else:
logger.warning(message)
return None
@ -178,7 +178,7 @@ class BaseOp(metaclass=ABCMeta):
if k == "answer":
self.response.answer = v
elif k == "success":
self.response.success = v.lower() == "true"
self.response.success = v if isinstance(v, bool) else v.lower() == "true"
else:
self.response.metadata[k] = v
else:
@ -262,7 +262,7 @@ class BaseOp(metaclass=ABCMeta):
if isinstance(result, list):
results.extend(result)
else:
result.append(result)
results.append(result)
self._pending_tasks.clear()
return results

View file

@ -102,7 +102,7 @@ class BaseReact(BaseOp):
# Create isolated kwargs for each tool call to avoid parameter conflicts
tool_kwargs = {**kwargs, **tool_call.argument_dict}
self.submit_async_task(tool_copy.call, **tool_kwargs)
self.submit_async_task(tool_copy.call, service_context=self.service_context, **tool_kwargs)
if self.tool_call_interval > 0:
await asyncio.sleep(self.tool_call_interval)

View file

@ -91,6 +91,16 @@ class BaseVectorStore(ABC):
async def get(self, vector_ids: str | list[str]) -> VectorNode | list[VectorNode]:
"""Fetch specific vector nodes from the collection by their IDs."""
async def dump(self) -> list[VectorNode]:
"""Dump the vector store to a list of vector nodes."""
return await self.list()
async def load(self, nodes: list[VectorNode]):
"""Load the vector store from a list of vector nodes."""
await self.delete_all()
if nodes:
await self.insert(nodes)
@abstractmethod
async def list(
self,

View file

@ -108,13 +108,13 @@ class ReMe:
):
"""Summarize messages and store them in memory for the specified user(s)."""
if user_name:
for message in messages and isinstance(user_name, str):
if isinstance(message, dict) and not message.get("name"):
message["name"] = user_name
elif isinstance(message, Message) and not message.name:
message.name = user_name
if isinstance(user_name, str):
for message in messages:
if isinstance(message, dict) and not message.get("name"):
message["name"] = user_name
elif isinstance(message, Message) and not message.name:
message.name = user_name
user_name = [user_name]
if not meta_memories:
@ -147,7 +147,12 @@ class ReMe:
else:
raise NotImplementedError
return await reme_summarizer.call(messages=messages, description=description, **kwargs)
return await reme_summarizer.call(
messages=messages,
description=description,
service_context=self.service_context,
**kwargs,
)
else:
raise NotImplementedError
@ -166,14 +171,13 @@ class ReMe:
):
"""Retrieve relevant memories for the specified user(s) based on query or messages."""
if user_name:
if messages:
for message in messages and isinstance(user_name, str):
if isinstance(message, dict) and not message.get("name"):
message["name"] = user_name
elif isinstance(message, Message) and not message.name:
message.name = user_name
if isinstance(user_name, str):
if messages:
for message in messages:
if isinstance(message, dict) and not message.get("name"):
message["name"] = user_name
elif isinstance(message, Message) and not message.name:
message.name = user_name
user_name = [user_name]
if not meta_memories:
@ -206,7 +210,13 @@ class ReMe:
else:
raise NotImplementedError
return await reme_retriever.call(query=query, messages=messages, description=description, **kwargs)
return await reme_retriever.call(
query=query,
messages=messages,
description=description,
service_context=self.service_context,
**kwargs,
)
else:
raise NotImplementedError

View file

@ -90,7 +90,7 @@ class HandsOff(BaseMemoryTool):
for k in ["query", "messages", "description", "history_node"]:
if k in self.context:
task_kwargs[k] = self.context[k]
self.submit_async_task(agent.call, **task_kwargs)
self.submit_async_task(agent.call, service_context=self.service_context, **task_kwargs)
await self.join_async_tasks()

View file

@ -39,10 +39,10 @@ class AddHistory(BaseMemoryTool):
author=self.author,
)
self.context.history_node = history_node
logger.info(f"Adding history node: {history_node.model_dump_json(indent=2, exclude={'content'})}")
logger.info(f"Adding history node: {history_node.model_dump_json(indent=2)}")
vector_node = history_node.to_vector_node()
await self.vector_store.delete(vector_node.memory_id)
await self.vector_store.delete(vector_node.vector_id)
await self.vector_store.insert([vector_node])
return f"Successfully added history: {history_node.memory_id}"

View file

@ -106,4 +106,4 @@ class UpdateUserProfile(BaseMemoryTool):
operations.append(f"added {len(new_nodes)} new profiles.")
operations.append("Operation completed.")
logger.info("\n".join(operations))
return operations
return "\n".join(operations)

View file

@ -122,7 +122,16 @@ class RetrieveMemory(BaseMemoryTool):
if not new_memory_nodes:
output = "No new memory_nodes found matching the query (duplicates removed)."
else:
output = "\n".join([m.format_memory() for m in new_memory_nodes])
outputs = []
for node in new_memory_nodes:
line = ""
if "conversation_time" in node.metadata and node.metadata["conversation_time"]:
line += f"conversation_time={node.metadata['conversation_time']} "
line += node.content.strip() + " "
if node.ref_memory_id:
line += f"history_id={node.ref_memory_id} "
outputs.append(line.strip())
output = "\n".join(outputs)
logger.info(f"Retrieved {len(memory_nodes)} memory_nodes, {len(new_memory_nodes)} new after deduplication")
return output

View file

@ -2,12 +2,10 @@
import asyncio
from reme import ReMe
from reme.core.schema import VectorNode, MemoryNode
from reme.reme import ReMe
reme = ReMe(
vector_store={"collection_name": "reme"},
)
reme = ReMe(vector_store={"collection_name": "reme"})
async def test_reme():
@ -71,7 +69,7 @@ async def test_reme():
nodes: list[VectorNode] = await reme.default_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(f"{i} {memory_node.model_dump_json()}")
print("\n" + "=" * 60)
print("步骤3: 测试记忆检索 - 验证个人信息")