mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
refactor(memory): update memory system with service context and vector store enhancements
This commit is contained in:
parent
d3fa645832
commit
e77d228fae
16 changed files with 85 additions and 49 deletions
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -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/*
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
38
reme/reme.py
38
reme/reme.py
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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: 测试记忆检索 - 验证个人信息")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue