refactor(memory): update memory tool implementations and parameters

This commit is contained in:
jinli.yl 2026-01-06 23:46:51 +08:00
parent 5306ff7d40
commit d29aff4e6c
11 changed files with 93 additions and 174 deletions

View file

@ -32,27 +32,29 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta):
return {}
def _build_tool_call(self) -> ToolCall:
tool_call_params: dict = {
"description": self.get_prompt("tool" + ("_multiple" if self.enable_multiple else "")),
}
if self.enable_multiple:
parameters = self._build_multiple_parameters()
else:
parameters = self._build_parameters()
if self.enable_thinking_params and "thinking" not in parameters["properties"]:
parameters["properties"] = {
"thinking": {
"type": "string",
"description": "Your thinking and reasoning about how to fill in the parameters",
},
**parameters["properties"],
}
parameters["required"] = ["thinking", *parameters["required"]]
if parameters:
tool_call_params["parameters"] = parameters
return ToolCall(
**{
"description": self.get_prompt("tool" + ("_multiple" if self.enable_multiple else "")),
"parameters": parameters,
},
)
if self.enable_thinking_params and "thinking" not in parameters["properties"]:
parameters["properties"] = {
"thinking": {
"type": "string",
"description": "Your thinking and reasoning about how to fill in the parameters",
},
**parameters["properties"],
}
parameters["required"] = ["thinking", *parameters["required"]]
return ToolCall(**tool_call_params)
@property
def meta_memory(self) -> CacheHandler:
@ -84,19 +86,20 @@ class BaseMemoryTool(BaseOp, metaclass=ABCMeta):
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.
This is a shared utility method for subclasses that need to create MemoryNode instances.
"""
"""Build MemoryNode from content, when_to_use, and metadata."""
return MemoryNode(
memory_type=self.memory_type,
memory_target=self.memory_target,
memory_type=memory_type or self.memory_type,
memory_target=memory_target or self.memory_target,
when_to_use=when_to_use or "",
content=memory_content,
ref_memory_id=self.ref_memory_id,
author=self.author,
ref_memory_id=ref_memory_id or self.ref_memory_id,
author=author or self.author,
metadata=metadata or {},
)

View file

@ -4,108 +4,47 @@ from loguru import logger
from ...base_memory_tool import BaseMemoryTool
from ....core.context import C
from ....core.schema import MemoryNode
from ....core.enumeration import MemoryType
from ....core.schema import ToolCall, Message
from ....core.utils import format_messages
@C.register_op()
class AddHistoryMemory(BaseMemoryTool):
"""Add history memory from conversation messages."""
def __init__(self, add_metadata: bool = True, **kwargs):
super().__init__(**kwargs)
self.add_metadata: bool = add_metadata
def _build_item_schema(self) -> tuple[dict, list[str]]:
properties = {
"messages": {
"type": "array",
"description": self.get_prompt("messages"),
"items": {"type": "object"},
},
}
required = ["messages"]
if self.add_metadata:
properties["metadata"] = {
"type": "object",
"description": self.get_prompt("metadata"),
}
return properties, required
def _build_parameters(self) -> dict:
properties, required = self._build_item_schema()
return {
"type": "object",
"properties": properties,
"required": required,
}
def _build_multiple_parameters(self) -> dict:
item_properties, required_fields = self._build_item_schema()
return {
"type": "object",
"properties": {
"histories": {
"type": "array",
"description": self.get_prompt("histories"),
"items": {
"type": "object",
"properties": item_properties,
"required": required_fields,
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": self.get_prompt("tool"),
"parameters": {
"type": "object",
"properties": {
"messages": {
"type": "array",
"description": self.get_prompt("messages"),
"items": {"type": "object"},
},
},
"required": ["messages"],
},
},
"required": ["histories"],
}
def _format_messages(self, messages: list) -> str:
return "\n".join([f"{msg.get('role', 'unknown')}: {msg.get('content', '')}" for msg in messages])
def _extract_history_data(self, hist_dict: dict) -> tuple[list, dict]:
messages = hist_dict.get("messages", [])
metadata = hist_dict.get("metadata", {}) if self.add_metadata else {}
return messages, metadata
)
async def execute(self):
memory_nodes: list[MemoryNode] = []
if self.enable_multiple:
histories: list[dict] = self.context.get("histories", [])
if not histories:
self.output = "No histories provided for addition."
return
for hist in histories:
messages, metadata = self._extract_history_data(hist)
if not messages:
logger.warning("Skipping history with empty messages")
continue
memory_content = self._format_messages(messages)
memory_nodes.append(
self._build_memory_node(memory_content, when_to_use="", metadata=metadata),
)
else:
messages, metadata = self._extract_history_data(self.context)
if not messages:
self.output = "No messages provided for addition."
return
memory_content = self._format_messages(messages)
memory_nodes.append(
self._build_memory_node(memory_content, when_to_use="", metadata=metadata),
)
if not memory_nodes:
self.output = "No valid histories provided for addition."
messages: list[Message | dict] = self.context.get("messages", [])
if not messages:
self.output = "No messages provided for addition."
return
vector_nodes = [node.to_vector_node() for node in memory_nodes]
vector_ids: list[str] = [node.vector_id for node in vector_nodes]
messages = [Message(**m) if isinstance(m, dict) else m for m in messages]
await self.vector_store.delete(vector_ids=vector_ids)
await self.vector_store.insert(nodes=vector_nodes)
memory_content = format_messages(messages)
memory_node = self._build_memory_node(memory_content=memory_content, memory_type=MemoryType.HISTORY)
self.output = f"Successfully added {len(memory_nodes)} history memories to vector_store."
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.output = "Successfully added history memory to vector_store."
logger.info(self.output)

View file

@ -50,11 +50,7 @@ class ReadHistoryMemory(BaseMemoryTool):
logger.warning(self.output)
return
nodes = await self.vector_store.search(
query="",
top_k=len(memory_ids),
filter_dict={"vector_id": memory_ids},
)
nodes = await self.vector_store.get(vector_ids=memory_ids)
if not nodes:
self.output = "No history memories found with the provided IDs."
@ -62,14 +58,5 @@ class ReadHistoryMemory(BaseMemoryTool):
return
memories: list[MemoryNode] = [MemoryNode.from_vector_node(n) for n in nodes]
output_lines = []
for memory in memories:
output_lines.append(f"Memory ID: {memory.vector_id}")
output_lines.append(f"Content:\n{memory.content}")
if memory.metadata:
output_lines.append(f"Metadata: {memory.metadata}")
output_lines.append("---")
self.output = "\n".join(output_lines)
self.output = "---\n".join([m.content for m in memories])
logger.info(f"Successfully read {len(memories)} history memories.")

View file

@ -22,12 +22,6 @@ class ReadIdentityMemory(BaseMemoryTool):
}
async def execute(self):
result = self.meta_memory.load("identity_memory")
identity_memory = result if result is not None else ""
if identity_memory:
self.output = f"Identity memory:\n{identity_memory}"
logger.info("Retrieved identity memory")
else:
self.output = "No identity memory found."
logger.info(self.output)
identity_memory = self.meta_memory.load("identity_memory") or ""
self.output = identity_memory or "No identity memory found."
logger.info(self.output)

View file

@ -66,13 +66,23 @@ class AddMetaMemory(BaseMemoryTool):
def _load_meta_memories(self) -> list[dict]:
"""Load existing meta memories from cache."""
result = self.meta_memory.load("meta_memories")
return result if result is not None else []
return self.meta_memory.load("meta_memories") or []
def _save_meta_memories(self, memories: list[dict]) -> bool:
"""Save meta memories to cache."""
return self.meta_memory.save("meta_memories", memories)
@staticmethod
def _filter_memory_type_target(memory_type: str, memory_target: str, existing_set: set) -> bool:
result = (
memory_type in [MemoryType.PERSONAL.value, MemoryType.PROCEDURAL.value]
and memory_target
and (memory_type, memory_target) not in existing_set
)
if result:
existing_set.add((memory_type, memory_target))
return result
async def execute(self):
"""Execute addition: load existing, merge with new, and save.
@ -88,24 +98,14 @@ class AddMetaMemory(BaseMemoryTool):
for mem in meta_memories:
memory_type = mem.get("memory_type", "")
memory_target = mem.get("memory_target", "")
if memory_type and (memory_type, memory_target) not in existing_set:
new_memories.append(
{
"memory_type": memory_type,
"memory_target": memory_target,
},
)
existing_set.add((memory_type, memory_target))
if self._filter_memory_type_target(memory_type, memory_target, existing_set):
new_memories.append({"memory_type": memory_type, "memory_target": memory_target})
else:
memory_type = self.context.get("memory_type", "")
memory_target = self.context.get("memory_target", "")
if memory_type and (memory_type, memory_target) not in existing_set:
new_memories.append(
{
"memory_type": memory_type,
"memory_target": memory_target,
},
)
if self._filter_memory_type_target(memory_type, memory_target, existing_set):
new_memories.append({"memory_type": memory_type, "memory_target": memory_target})
if not new_memories:
self.output = "No new meta memories to add (all entries already exist or invalid)."

View file

@ -21,4 +21,4 @@ memory_target: |
The target identifier for this memory category.
Examples:
- For personal memory: person's name (e.g., "John", "Alice")
- For procedural memory: process name (e.g., "deployment", "code_review")
- For procedural memory: domain or topic name (e.g., "deployment", "code_review")

View file

@ -51,8 +51,7 @@ class ReadMetaMemory(BaseMemoryTool):
filtered_memories = []
for m in all_memories:
memory_type = MemoryType(m.get("memory_type"))
if memory_type in (MemoryType.PERSONAL, MemoryType.PROCEDURAL):
if m.get("memory_type") in [MemoryType.PERSONAL.value, MemoryType.PROCEDURAL.value]:
filtered_memories.append(m)
if self.enable_tool_memory:
@ -102,8 +101,7 @@ class ReadMetaMemory(BaseMemoryTool):
memories = self._load_meta_memories()
if memories:
formatted = self._format_memory_metadata(memories)
self.output = formatted
self.output = self._format_memory_metadata(memories)
logger.info(f"Retrieved {len(memories)} meta memory entries")
else:
self.output = "No memory metadata found."

View file

@ -114,9 +114,7 @@ class AddMemory(BaseMemoryTool):
logger.warning("Skipping memory with empty content")
continue
memory_nodes.append(
self._build_memory_node(memory_content, when_to_use, metadata),
)
memory_nodes.append(self._build_memory_node(memory_content, when_to_use=when_to_use, metadata=metadata))
else:
memory_content, when_to_use, metadata = self._extract_memory_data(self.context)
@ -124,9 +122,7 @@ class AddMemory(BaseMemoryTool):
self.output = "No memory content provided for addition."
return
memory_nodes.append(
self._build_memory_node(memory_content, when_to_use, metadata),
)
memory_nodes.append(self._build_memory_node(memory_content, when_to_use=when_to_use, metadata=metadata))
if not memory_nodes:
self.output = "No valid memories provided for addition."

View file

@ -16,11 +16,7 @@ class AddSummaryMemory(AddMemory):
- No when_to_use field (add_when_to_use=False)
"""
def __init__(
self,
add_metadata: bool = True,
**kwargs,
):
def __init__(self, add_metadata: bool = True, **kwargs):
"""Initialize AddSummaryMemory.
Args:

View file

@ -45,11 +45,11 @@ class DeleteMemory(BaseMemoryTool):
if self.enable_multiple:
memory_ids = self.context.get("memory_ids", [])
else:
single_id = self.context.get("memory_id", "")
memory_ids = [single_id] if single_id else []
memory_id = self.context.get("memory_id", "")
memory_ids = [memory_id] if memory_id else []
# Filter out empty IDs
memory_ids = [mid for mid in memory_ids if mid]
memory_ids = [m for m in memory_ids if m]
if not memory_ids:
self.output = "No valid memory IDs provided for deletion."

View file

@ -126,7 +126,13 @@ class UpdateMemory(BaseMemoryTool):
logger.warning(f"Skipping memory with missing id or content: {mem}")
continue
old_memory_ids.append(memory_id)
new_memory_nodes.append(self._build_memory_node(memory_content, when_to_use, metadata))
new_memory_nodes.append(
self._build_memory_node(
memory_content,
when_to_use=when_to_use,
metadata=metadata,
),
)
else:
memory_id, memory_content, when_to_use, metadata = self._extract_memory_data(self.context)
@ -134,7 +140,7 @@ class UpdateMemory(BaseMemoryTool):
self.output = "No memory ID or content provided for update."
return
old_memory_ids.append(memory_id)
new_memory_nodes.append(self._build_memory_node(memory_content, when_to_use, metadata))
new_memory_nodes.append(self._build_memory_node(memory_content, when_to_use=when_to_use, metadata=metadata))
if not old_memory_ids or not new_memory_nodes:
self.output = "No valid memories provided for update."