refactor(memory): refactor memory tools and handlers for improved structure

This commit is contained in:
jinli.yl 2026-01-28 20:10:38 +08:00
parent c174aade76
commit c06fd78763
23 changed files with 390 additions and 423 deletions

View file

@ -22,8 +22,10 @@ class PersonalRetriever(BaseMemoryAgent):
read_all_profiles_tool: BaseTool | None = self.pop_tool("read_all_profiles")
if read_all_profiles_tool is not None:
all_profiles = await read_all_profiles_tool.call(memory_target=self.memory_target,
service_context=self.service_context)
all_profiles = await read_all_profiles_tool.call(
memory_target=self.memory_target,
service_context=self.service_context,
)
else:
all_profiles = ""

View file

@ -6,7 +6,7 @@ system_prompt: |
## User Question
{context}
## Retrieval Strategy
### Phase 1 `retrieve_memory`
- Purpose: Search for relevant memories using semantic similarity

View file

@ -1,4 +1,5 @@
"""Personal memory summarizer agent for two-phase personal memory processing."""
from loguru import logger
from ..base_memory_agent import BaseMemoryAgent

View file

@ -1,15 +1,15 @@
system_prompt_s1_zh: |
你是一个记忆Agent负责管理关于 {memory_target} 的 {memory_type} 类型记忆。
## 最新对话
Format: round<index> [<timestamp>] <role/name>: <content>
{context}
## 任务
### 步骤1
根据`最新对话`的内容,在 `add_draft_and_retrieve_similar_memory` 中创建记忆草稿 `memory_draft`。
工具会根据memory_draft的内容行向量检索返回历史相似记忆确保在第二步的时候更好的管理记忆库记忆。
### 步骤2
使用`update_memory`更新向量库记忆。
通过`memory_ids_to_delete`删除历史记忆,`memories_to_add`添加新记忆包括message_time和memory_content。
@ -23,16 +23,16 @@ user_message_s1_zh: |
system_prompt_s2_zh: |
你是一个Profile Agent负责管理关于 {memory_target} 的 Profile。
## 最新对话
Format: round<index> [<timestamp>] <role/name>: <content>
{context}
## 任务
### 步骤1
根据`最新对话`的内容,在 `add_draft_and_read_all_profiles` 中创建记忆草稿 `profile_draft`。
工具会直接返回所有的Profile确保在第二步的时候更好的管理Profile。
### 步骤2
使用`update_profile`更新profile库。
通过`profile_ids_to_delete`删除历史Profile`profiles_to_add`添加新Profile包括message_time、profile_key和profile_value。
@ -46,16 +46,16 @@ user_message_s2_zh: |
system_prompt_s1: |
You are a Memory Agent responsible for managing {memory_type} type memories about {memory_target}.
## Latest Conversation
Format: round<index> [<timestamp>] <role/name>: <content>
{context}
## Task
### Step 1
Based on the content of `Latest Conversation`, create a memory draft `memory_draft` in `add_draft_and_retrieve_similar_memory`.
The tool will perform vector retrieval based on the content of memory_draft and return historically similar memories to better manage the memory store in Step 2.
### Step 2
Use `update_memory` to update the vector store memories.
Delete historical memories through `memory_ids_to_delete`, add new memories through `memories_to_add`, including message_time and memory_content.
@ -69,16 +69,16 @@ user_message_s1: |
system_prompt_s2: |
You are a Profile Agent responsible for managing the Profile about {memory_target}.
## Latest Conversation
Format: round<index> [<timestamp>] <role/name>: <content>
{context}
## Task
### Step 1
Based on the content of `Latest Conversation`, create a profile draft `profile_draft` in `add_draft_and_read_all_profiles`.
The tool will directly return all Profiles to better manage the Profile store in Step 2.
### Step 2
Use `update_profile` to update the profile store.
Delete historical Profiles through `profile_ids_to_delete`, add new Profiles through `profiles_to_add`, including message_time, profile_key, and profile_value.

View file

@ -1,6 +1,10 @@
"""Procedural memory retriever agent implementation."""
from ..base_memory_agent import BaseMemoryAgent
from ....core.enumeration import MemoryType
class ProceduralRetriever(BaseMemoryAgent):
"""Agent responsible for retrieving procedural memories."""
memory_type: MemoryType = MemoryType.PROCEDURAL

View file

@ -1,6 +1,10 @@
"""Procedural memory summarizer agent implementation."""
from ..base_memory_agent import BaseMemoryAgent
from ....core.enumeration import MemoryType
class ProceduralSummarizer(BaseMemoryAgent):
"""Agent responsible for summarizing procedural memories."""
memory_type: MemoryType = MemoryType.PROCEDURAL

View file

@ -1,6 +1,10 @@
"""Tool memory retriever agent implementation."""
from ..base_memory_agent import BaseMemoryAgent
from ....core.enumeration import MemoryType
class ToolRetriever(BaseMemoryAgent):
"""Agent responsible for retrieving tool-related memories."""
memory_type: MemoryType = MemoryType.TOOL

View file

@ -1,6 +1,10 @@
"""Tool memory summarizer agent implementation."""
from ..base_memory_agent import BaseMemoryAgent
from ....core.enumeration import MemoryType
class ToolSummarizer(BaseMemoryAgent):
"""Agent responsible for summarizing tool-related memories."""
memory_type: MemoryType = MemoryType.TOOL

View file

@ -1,3 +1,5 @@
"""High-level entry point for configuring and running ReMe services and flows."""
import asyncio
from .context import PromptHandler, ServiceContext
@ -11,21 +13,22 @@ from .vector_store import BaseVectorStore
class Application:
"""Application wrapper that wires together service context, flows, and runtimes."""
def __init__(
self,
*args,
llm_api_key: str | None = None,
llm_api_base: str | None = None,
embedding_api_key: str | None = None,
embedding_api_base: str | None = None,
enable_logo: bool = True,
parser: type[PydanticConfigParser] | None = None,
llm: dict | None = None,
embedding_model: dict | None = None,
vector_store: dict | None = None,
token_counter: dict | None = None,
**kwargs,
self,
*args,
llm_api_key: str | None = None,
llm_api_base: str | None = None,
embedding_api_key: str | None = None,
embedding_api_base: str | None = None,
enable_logo: bool = True,
parser: type[PydanticConfigParser] | None = None,
llm: dict | None = None,
embedding_model: dict | None = None,
vector_store: dict | None = None,
token_counter: dict | None = None,
**kwargs,
):
# ServiceContext
self.service_context = ServiceContext(
@ -94,10 +97,10 @@ class Application:
stream_queue = asyncio.Queue()
task = asyncio.create_task(flow.call(stream_queue=stream_queue, **kwargs))
async for chunk in execute_stream_task(
stream_queue=stream_queue,
task=task,
task_name=name,
as_bytes=False,
stream_queue=stream_queue,
task=task,
task_name=name,
as_bytes=False,
):
yield chunk

View file

@ -149,12 +149,12 @@ class MemoryNode(BaseModel):
)
def format(
self,
include_memory_id: bool = True,
include_when_to_use: bool = True,
include_content: bool = True,
include_message_time: bool = True,
ref_memory_id_key: str = "",
self,
include_memory_id: bool = True,
include_when_to_use: bool = True,
include_content: bool = True,
include_message_time: bool = True,
ref_memory_id_key: str = "",
) -> str:
"""Format memory node as string with configurable fields."""
line = ""

View file

@ -1,25 +1,36 @@
"""ReMe classes for simplified configuration and execution."""
import asyncio
import sys
from pathlib import Path
from loguru import logger
from .core import Application
from .agent.memory.default import ReMeSummarizer, PersonalSummarizer, PersonalRetriever, ReMeRetriever
from reme.agent.memory import BaseMemoryAgent
from .agent.memory.default import (
ReMeSummarizer,
PersonalSummarizer,
PersonalRetriever,
ReMeRetriever,
ProceduralSummarizer,
ToolSummarizer,
ProceduralRetriever,
ToolRetriever,
)
from .config import ReMeConfigParser
from .core.context import PromptHandler, ServiceContext
from .core.embedding import BaseEmbeddingModel
from .core import Application
from .core.enumeration import MemoryType
from .core.flow import BaseFlow
from .core.llm import BaseLLM
from .core.schema import Response, Message, MemoryNode, VectorNode
from .core.token_counter import BaseTokenCounter
from .core.utils import execute_stream_task, get_now_time
from .core.vector_store import BaseVectorStore
from .tool.memory import UpdateUserProfile, RetrieveMemory, AddMemory, DelegateTask, ReadHistory, ReadUserProfile, \
ProfileHandler
from .core.schema import Message
from .tool.memory import (
RetrieveMemory,
DelegateTask,
ReadHistory,
ProfileHandler,
MemoryHandler,
AddDraftAndRetrieveSimilarMemory,
UpdateMemoryV2,
AddDraftAndReadAllProfiles,
UpdateProfile,
AddHistory,
ReadAllProfiles,
)
class ReMe(Application):
@ -37,18 +48,10 @@ class ReMe(Application):
embedding_model: dict | None = None,
vector_store: dict | None = None,
token_counter: dict | None = None,
personal_memory_target: list[str] | None = None,
procedural_memory_target: list[str] | None = None,
tool_memory_target: list[str] | None = None,
profile_path: str = "reme_profile",
main_summary_version: str = "default",
personal_summary_version: str = "default",
procedural_summary_version: str = "default",
tool_summary_version: str = "default",
main_retrieve_version: str = "default",
personal_retrieve_version: str = "default",
procedural_retrieve_version: str = "default",
tool_retrieve_version: str = "default",
personal_memory_target: list[str] | None = None,
procedural_memory_target: list[str] | None = None,
tool_memory_target: list[str] | None = None,
profile_dir: str = "reme_profile",
**kwargs,
):
super().__init__(
@ -80,300 +83,239 @@ class ReMe(Application):
for name in tool_memory_target:
assert name not in memory_target_type_mapping, f"Memory target name {name} is already used."
memory_target_type_mapping[name] = MemoryType.TOOL
self.service_context.memory_target_type_mapping = memory_target_type_mapping
self.profile_path: str = profile_path
@property
def memory_target_type_mapping(self) -> dict[str, MemoryType]:
mapping = {}
if self.service_context.personal_memory_target:
for name in self.service_context.personal_memory_target:
assert name not in mapping, f"Memory target name {name} is already used."
mapping[name] = MemoryType.PERSONAL
if self.service_context.procedural_memory_target:
for name in self.service_context.procedural_memory_target:
assert name not in mapping, f"Memory target name {name} is already used."
mapping[name] = MemoryType.PROCEDURAL
if self.service_context.tool_memory_target:
for name in self.service_context.tool_memory_target:
assert name not in mapping, f"Memory target name {name} is already used."
mapping[name] = MemoryType.TOOL
return mapping
self.profile_dir: str = profile_dir
def add_meta_memory(self, memory_type: str | MemoryType, memory_target: str):
memory_type = MemoryType(memory_type)
if memory_type is MemoryType.PERSONAL:
personal_memory_target = self.service_context.personal_memory_target
if memory_target not in personal_memory_target:
personal_memory_target.append(memory_target)
else:
logger.warning(f"Memory target {memory_target} is already added.")
elif memory_type is MemoryType.PROCEDURAL:
procedural_memory_target = self.service_context.procedural_memory_target
if memory_target not in procedural_memory_target:
procedural_memory_target.append(memory_target)
else:
logger.warning(f"Memory target {memory_target} is already added.")
elif memory_type is MemoryType.TOOL:
tool_memory_target = self.service_context.tool_memory_target
if memory_target not in tool_memory_target:
tool_memory_target.append(memory_target)
else:
logger.warning(f"Memory target {memory_target} is already added.")
"""Register or validate a memory target with the given memory type."""
if memory_target in self.service_context.memory_target_type_mapping:
assert self.service_context.memory_target_type_mapping[memory_target] is memory_type
else:
self.service_context.memory_target_type_mapping[memory_target] = MemoryType(memory_type)
async def summary_memory(
self,
messages: list[Message | dict],
description: str = "",
user_name: str = "",
task_name: str = "",
tool_name: str = "",
user_name: str | list[str] = "",
task_name: str | list[str] = "",
tool_name: str | list[str] = "",
enable_thinking_params: bool = False,
version: str = "default",
retrieve_top_k: int = 20,
return_dict: bool = False,
**kwargs,
) -> str | dict:
"""Summarize messages and store them in memory for the specified user(s)."""
if 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]
"""Summarize personal, procedural and tool memories for the given context."""
format_messages: list[Message] = []
for message in messages:
if isinstance(message, dict):
assert message.get("time_created"), "message must have time_created field."
message = Message(**message)
format_messages.append(message)
if not meta_memories:
meta_memories = [
{
"memory_type": "personal",
"memory_target": name,
}
for name in user_name
]
if version == "default":
reme_summarizer = ReMeSummarizer(
meta_memories=meta_memories,
tools=[
DelegateTask(
memory_agents=[
PersonalSummarizer(
tools=[
RetrieveMemory(enable_thinking_params=enable_thinking_params),
AddMemory(enable_thinking_params=enable_thinking_params),
UpdateUserProfile(enable_thinking_params=enable_thinking_params),
],
),
],
),
],
)
else:
raise NotImplementedError
result = await reme_summarizer.call(
messages=messages,
description=description,
service_context=self.service_context,
**kwargs,
personal_summarizer: BaseMemoryAgent
if version:
personal_summarizer = PersonalSummarizer(
tools=[
AddDraftAndRetrieveSimilarMemory(
enable_thinking_params=enable_thinking_params,
top_k=retrieve_top_k,
),
UpdateMemoryV2(enable_thinking_params=enable_thinking_params),
AddDraftAndReadAllProfiles(
enable_thinking_params=enable_thinking_params,
profile_dir=self.profile_dir,
),
UpdateProfile(
enable_thinking_params=enable_thinking_params,
profile_dir=self.profile_dir,
),
],
)
if return_dict:
return result
else:
return result["answer"]
else:
raise NotImplementedError
procedural_summarizer: BaseMemoryAgent
if version == "default":
procedural_summarizer = ProceduralSummarizer(tools=[])
else:
raise NotImplementedError
tool_summarizer: BaseMemoryAgent
if version == "default":
tool_summarizer = ToolSummarizer(tools=[])
else:
raise NotImplementedError
memory_agents = []
if user_name:
if isinstance(user_name, str):
for message in format_messages:
message.name = user_name
self.add_meta_memory(MemoryType.PERSONAL, user_name)
elif isinstance(user_name, list):
for name in user_name:
self.add_meta_memory(MemoryType.PERSONAL, name)
else:
raise RuntimeError("user_name must be str or list[str]")
memory_agents.append(personal_summarizer)
if task_name:
if isinstance(task_name, str):
self.add_meta_memory(MemoryType.PROCEDURAL, task_name)
elif isinstance(task_name, list):
for name in task_name:
self.add_meta_memory(MemoryType.PROCEDURAL, name)
else:
raise RuntimeError("task_name must be str or list[str]")
memory_agents.append(procedural_summarizer)
if tool_name:
if isinstance(tool_name, str):
self.add_meta_memory(MemoryType.TOOL, tool_name)
elif isinstance(tool_name, list):
for name in tool_name:
self.add_meta_memory(MemoryType.TOOL, name)
else:
raise RuntimeError("tool_name must be str or list[str]")
memory_agents.append(tool_summarizer)
if not memory_agents:
memory_agents = [personal_summarizer, procedural_summarizer, tool_summarizer]
reme_summarizer: BaseMemoryAgent
if version == "default":
reme_summarizer = ReMeSummarizer(tools=[AddHistory(), DelegateTask(memory_agents=memory_agents)])
else:
raise NotImplementedError
result = await reme_summarizer.call(
messages=messages,
description=description,
service_context=self.service_context,
**kwargs,
)
if return_dict:
return result
else:
return result["answer"]
async def retrieve_memory(
self,
query: str = "",
top_k: int = 20,
description: str = "",
messages: list[dict] | None = None,
user_name: str | list[str] = "",
task_name: str | list[str] = "",
tool_name: str | list[str] = "",
enable_thinking_params: bool = False,
meta_memories: list[dict] = None,
version: str = "default",
retrieve_top_k: int = 20,
enable_memory_target: bool = True,
return_dict: bool = False,
**kwargs,
) -> str | dict:
"""Retrieve relevant memories for the specified user(s) based on query or messages."""
if 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]
"""Retrieve relevant personal, procedural and tool memories for a query."""
if not meta_memories:
meta_memories = [
{
"memory_type": "personal",
"memory_target": name,
}
for name in user_name
]
if version == "default":
reme_retriever = ReMeRetriever(
meta_memories=meta_memories,
tools=[
DelegateTask(
memory_agents=[
PersonalRetriever(
tools=[
RetrieveMemory(enable_thinking_params=enable_thinking_params, top_k=top_k),
ReadHistory(enable_thinking_params=enable_thinking_params),
],
),
],
),
],
)
else:
raise NotImplementedError
result = await reme_retriever.call(
query=query,
messages=messages,
description=description,
service_context=self.service_context,
**kwargs,
personal_retriever: BaseMemoryAgent
if version:
personal_retriever = PersonalRetriever(
tools=[
ReadAllProfiles(
enable_thinking_params=enable_thinking_params,
profile_dir=self.profile_dir,
),
RetrieveMemory(
enable_thinking_params=enable_thinking_params,
top_k=retrieve_top_k,
enable_memory_target=enable_memory_target,
),
ReadHistory(enable_thinking_params=enable_thinking_params),
],
)
if return_dict:
return result
else:
return result["answer"]
else:
raise NotImplementedError
async def add_memory(
self,
memory_content: str,
user_name: str,
memory_type: str | MemoryType | None = None,
memory_target: str = "",
when_to_use: str = "",
ref_memory_id: str = "",
author: str = "",
score: float = 0,
conversation_time: str = "",
**kwargs,
) -> MemoryNode:
"""Add a new memory to the vector store for the specified user."""
procedural_retriever: BaseMemoryAgent
if version == "default":
procedural_retriever = ProceduralRetriever(tools=[])
else:
raise NotImplementedError
tool_retriever: BaseMemoryAgent
if version == "default":
tool_retriever = ToolRetriever(tools=[])
else:
raise NotImplementedError
memory_agents = []
if user_name:
memory_type = MemoryType.PERSONAL
memory_target = user_name
if isinstance(user_name, str):
self.add_meta_memory(MemoryType.PERSONAL, user_name)
elif isinstance(user_name, list):
for name in user_name:
self.add_meta_memory(MemoryType.PERSONAL, name)
else:
raise RuntimeError("user_name must be str or list[str]")
memory_agents.append(personal_retriever)
if task_name:
if isinstance(task_name, str):
self.add_meta_memory(MemoryType.PROCEDURAL, task_name)
elif isinstance(task_name, list):
for name in task_name:
self.add_meta_memory(MemoryType.PROCEDURAL, name)
else:
raise RuntimeError("task_name must be str or list[str]")
memory_agents.append(procedural_retriever)
if tool_name:
if isinstance(tool_name, str):
self.add_meta_memory(MemoryType.TOOL, tool_name)
elif isinstance(tool_name, list):
for name in tool_name:
self.add_meta_memory(MemoryType.TOOL, name)
else:
raise RuntimeError("tool_name must be str or list[str]")
memory_agents.append(tool_retriever)
if not memory_agents:
memory_agents = [personal_retriever, procedural_retriever, tool_retriever]
reme_retriever: BaseMemoryAgent
if version == "default":
reme_retriever = ReMeRetriever(tools=[DelegateTask(memory_agents=memory_agents)])
else:
memory_type = MemoryType(memory_type)
assert memory_target, "memory_target is required"
raise NotImplementedError
metadata = kwargs.copy()
if conversation_time:
metadata["conversation_time"] = conversation_time
memory_node = MemoryNode(
memory_type=memory_type,
memory_target=memory_target,
when_to_use=when_to_use,
content=memory_content,
ref_memory_id=ref_memory_id,
author=author,
score=score,
metadata=metadata,
result = await reme_retriever.call(
query=query,
messages=messages,
description=description,
service_context=self.service_context,
**kwargs,
)
vector_node = memory_node.to_vector_node()
await self.vector_store.delete([vector_node.vector_id])
await self.vector_store.insert([vector_node])
return memory_node
async def update_memory(
self,
memory_id: str,
memory_content: str,
user_name: str,
memory_type: str | MemoryType | None = None,
memory_target: str = "",
when_to_use: str = "",
ref_memory_id: str = "",
author: str = "",
score: float = 0,
conversation_time: str = "",
**kwargs,
) -> MemoryNode:
"""Update an existing memory in the vector store by its ID."""
if user_name:
memory_type = MemoryType.PERSONAL
memory_target = user_name
if return_dict:
return result
else:
memory_type = MemoryType(memory_type)
assert memory_target, "memory_target is required"
return result["answer"]
metadata = kwargs.copy()
if conversation_time:
metadata["conversation_time"] = conversation_time
@property
def profile_path(self) -> Path:
"""Get the path to the profile directory."""
return Path(self.profile_dir) / self.vector_store.collection_name
memory_node = MemoryNode(
memory_type=memory_type,
memory_target=memory_target,
when_to_use=when_to_use,
content=memory_content,
ref_memory_id=ref_memory_id,
author=author,
score=score,
metadata=metadata,
)
vector_node = memory_node.to_vector_node()
await self.vector_store.delete(list(set([memory_id, vector_node.vector_id])))
await self.vector_store.insert([vector_node])
return memory_node
async def delete_memory(self, memory_id: str | list[str]):
"""Delete one or more memories from the vector store by their IDs."""
vector_ids = [memory_id] if isinstance(memory_id, str) else memory_id
await self.vector_store.delete(list(set(vector_ids)))
async def delete_all_memories(self):
"""Delete all memories from the vector store."""
await self.vector_store.delete_all()
async def get_memory(self, memory_id: str | list[str]) -> MemoryNode | list[MemoryNode]:
"""Retrieve one or more memories from the vector store by their IDs."""
vector_ids = [memory_id] if isinstance(memory_id, str) else memory_id
vector_nodes = await self.vector_store.get(vector_ids)
if isinstance(vector_nodes, VectorNode):
return vector_nodes.to_memory_node()
else:
return [node.to_memory_node() for node in vector_nodes]
async def get_all_memories(self) -> list[MemoryNode]:
"""Retrieve all memories from the vector store."""
return [node.to_memory_node() for node in await self.vector_store.list()]
def get_memory_handler(self, memory_target: str) -> MemoryHandler:
"""Get the memory handler for the specified memory target."""
return MemoryHandler(memory_target=memory_target, service_context=self.service_context)
def get_profile_handler(self, user_name: str) -> ProfileHandler:
"""Get the profile handler for the specified user."""
profile_path = Path(self.profile_path) / self.vector_store.collection_name
return ProfileHandler(memory_target=user_name, profile_path=profile_path)
return ProfileHandler(memory_target=user_name, profile_path=self.profile_path)
async def context_offload(self):
"""working memory summary"""

View file

@ -39,4 +39,4 @@ __all__ = [
for name in __all__:
tool_class = globals()[name]
R.op.register()(tool_class)
R.op.register()(tool_class)

View file

@ -1,5 +1,4 @@
"""Add draft profile and read all profiles from local storage"""
from pathlib import Path
from loguru import logger
@ -11,9 +10,8 @@ from ...core.schema import ToolCall
class AddDraftAndReadAllProfiles(BaseMemoryTool):
"""Tool to add draft profile and read all profiles"""
def __init__(self, profile_path: str, enable_memory_target: bool = False, **kwargs):
def __init__(self, enable_memory_target: bool = False, **kwargs):
super().__init__(**kwargs)
self.profile_path: str = profile_path
self.enable_memory_target: bool = enable_memory_target
def _build_query_parameters(self) -> dict:
@ -86,10 +84,7 @@ class AddDraftAndReadAllProfiles(BaseMemoryTool):
continue
targets_processed.add(target)
profile_handler = ProfileHandler(
profile_path=Path(self.profile_path) / self.vector_store.collection_name,
memory_target=target,
)
profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=target)
profiles_str = profile_handler.read_all(add_profile_id=True)
if profiles_str:

View file

@ -80,11 +80,13 @@ class AddDraftAndRetrieveSimilarMemory(BaseMemoryTool):
if target not in queries_by_target:
queries_by_target[target] = []
queries_by_target[target].append({
"query": item["memory_draft"],
"limit": self.top_k,
"filters": {},
})
queries_by_target[target].append(
{
"query": item["memory_draft"],
"limit": self.top_k,
"filters": {},
},
)
# Execute batch searches for each target
memory_nodes: list[MemoryNode] = []

View file

@ -11,10 +11,10 @@ class AddMemory(BaseMemoryTool):
"""Tool to add memories to vector store"""
def __init__(
self,
enable_memory_target: bool = False,
enable_when_to_use: bool = False,
**kwargs,
self,
enable_memory_target: bool = False,
enable_when_to_use: bool = False,
**kwargs,
):
super().__init__(**kwargs)
self.enable_memory_target: bool = enable_memory_target
@ -112,14 +112,16 @@ class AddMemory(BaseMemoryTool):
except Exception:
logger.warning(f"Invalid message time format: {message_time}")
memory_dicts.append({
"content": memory_content,
"when_to_use": when_to_use,
"message_time": message_time,
"ref_memory_id": self.history_id,
"author": self.author,
"metadata": metadata,
})
memory_dicts.append(
{
"content": memory_content,
"when_to_use": when_to_use,
"message_time": message_time,
"ref_memory_id": self.history_id,
"author": self.author,
"metadata": metadata,
},
)
if memory_dicts:
handler = MemoryHandler(target, self.service_context)

View file

@ -1,6 +1,7 @@
"""Base class for memory tool"""
from abc import ABCMeta
from pathlib import Path
from ...core.enumeration import MemoryType
from ...core.op import BaseTool
@ -14,11 +15,13 @@ class BaseMemoryTool(BaseTool, metaclass=ABCMeta):
self,
enable_multiple: bool = True,
enable_thinking_params: bool = False,
profile_dir: str = "",
**kwargs,
):
super().__init__(**kwargs)
self.enable_multiple: bool = enable_multiple
self.enable_thinking_params: bool = enable_thinking_params
self.profile_dir: str = profile_dir
def _build_tool_call(self) -> ToolCall:
"""Build and return the tool call schema"""
@ -98,3 +101,8 @@ class BaseMemoryTool(BaseTool, metaclass=ABCMeta):
def memory_target_type_mapping(self) -> dict[str, MemoryType]:
"""Get the memory target type mapping from context."""
return self.context.memory_target_type_mapping
@property
def profile_path(self) -> Path:
"""Get the path to the profile directory for the current collection."""
return Path(self.profile_dir) / self.vector_store.collection_name

View file

@ -1,3 +1,5 @@
"""Memory handler"""
from ...core.context import ServiceContext
from ...core.enumeration import MemoryType
from ...core.schema import MemoryNode
@ -45,14 +47,14 @@ class MemoryHandler:
return memory_nodes
async def add(
self,
content: str,
when_to_use: str = "",
message_time: str = "",
ref_memory_id: str = "",
author: str = "",
score: float = 0.0,
**kwargs,
self,
content: str,
when_to_use: str = "",
message_time: str = "",
ref_memory_id: str = "",
author: str = "",
score: float = 0.0,
**kwargs,
) -> MemoryNode:
"""Add a single memory node and return its memory_id."""
memory_dict = {
@ -123,15 +125,15 @@ class MemoryHandler:
return updated_nodes
async def update(
self,
memory_id: str,
content: str | None = None,
when_to_use: str | None = None,
message_time: str | None = None,
ref_memory_id: str | None = None,
author: str | None = None,
score: float | None = None,
**kwargs,
self,
memory_id: str,
content: str | None = None,
when_to_use: str | None = None,
message_time: str | None = None,
ref_memory_id: str | None = None,
author: str | None = None,
score: float | None = None,
**kwargs,
) -> MemoryNode:
"""Update a memory node's content, when_to_use, or other fields."""
update_dict: dict = {"memory_id": memory_id}
@ -154,11 +156,11 @@ class MemoryHandler:
return memory_nodes[0]
async def search(
self,
query: str | list[str],
limit: int = 5,
filters: dict | None = None,
**kwargs,
self,
query: str | list[str],
limit: int = 5,
filters: dict | None = None,
**kwargs,
) -> list[MemoryNode]:
"""Search for similar memory nodes based on query text."""
filters = filters or {}
@ -195,11 +197,11 @@ class MemoryHandler:
return list(seen_ids.values())
async def list(
self,
filters: dict | None = None,
limit: int | None = None,
sort_key: str | None = None,
reverse: bool = True,
self,
filters: dict | None = None,
limit: int | None = None,
sort_key: str | None = None,
reverse: bool = True,
) -> list[MemoryNode]:
"""List memory nodes with optional filtering and sorting."""
filters = filters or {}

View file

@ -1,4 +1,5 @@
"""Profile Handler for managing user profiles in local memory"""
from pathlib import Path
from loguru import logger
@ -36,7 +37,9 @@ class ProfileHandler:
removed_count = len(sorted_nodes) - self.max_capacity
nodes = sorted_nodes[removed_count:]
logger.info(
f"Capacity limit reached: removed {removed_count} oldest profiles (kept {len(nodes)}/{self.max_capacity})")
f"Capacity limit reached: removed {removed_count} oldest profiles "
f"(kept {len(nodes)}/{self.max_capacity})",
)
nodes_data = [node.model_dump(exclude_none=True) for node in nodes]
self.cache_handler.save(self.cache_key, nodes_data)
@ -191,9 +194,6 @@ class ProfileHandler:
def read_all(self, add_profile_id: bool = False, add_history_id: bool = False) -> str:
"""Read all profiles and return formatted string"""
nodes = self.get_all()
formatted_profiles = [
self._format_node(node, add_profile_id, add_history_id)
for node in nodes
]
formatted_profiles = [self._format_node(node, add_profile_id, add_history_id) for node in nodes]
logger.info(f"Read {len(formatted_profiles)} profiles from {self.cache_key}")
return "\n".join(formatted_profiles).strip()

View file

@ -1,5 +1,4 @@
"""Read user profile tool"""
from pathlib import Path
from loguru import logger
@ -11,10 +10,9 @@ from ...core.schema import ToolCall
class ReadAllProfiles(BaseMemoryTool):
"""Tool to read all user profiles"""
def __init__(self, profile_path: str, **kwargs):
def __init__(self, **kwargs):
kwargs["enable_multiple"] = False
super().__init__(**kwargs)
self.profile_path: str = profile_path
def _build_tool_call(self) -> ToolCall:
"""Build and return the tool call schema"""
@ -30,16 +28,12 @@ class ReadAllProfiles(BaseMemoryTool):
)
async def execute(self):
profile_handler = ProfileHandler(
profile_path=Path(self.profile_path) / self.vector_store.collection_name,
memory_target=self.memory_target,
)
profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target)
profiles_str = profile_handler.read_all(add_profile_id=True)
if not profiles_str:
output = "No profiles found."
logger.info(output)
return output
logger.info(f"Successfully read profiles")
logger.info("Successfully read profiles")
return profiles_str

View file

@ -31,8 +31,8 @@ class RetrieveMemory(BaseMemoryTool):
properties["time_filter"] = {
"type": "string",
"description": "Optional time filter to narrow down search results by date. "
"Format: single date '20200101' for exact date match, "
"or date range '20200101,20200102' for inclusive range filtering.",
"Format: single date '20200101' for exact date match, "
"or date range '20200101,20200102' for inclusive range filtering.",
}
if self.enable_memory_target:
@ -99,11 +99,13 @@ class RetrieveMemory(BaseMemoryTool):
else:
filters = {"time_int": [int(time_filter), int(time_filter)]}
queries_by_target[target].append({
"query": item["query"],
"limit": self.top_k,
"filters": filters,
})
queries_by_target[target].append(
{
"query": item["query"],
"limit": self.top_k,
"filters": filters,
},
)
# Execute batch searches for each target
memory_nodes: list[MemoryNode] = []

View file

@ -11,10 +11,10 @@ class UpdateMemory(BaseMemoryTool):
"""Tool to update memories in vector store"""
def __init__(
self,
enable_memory_target: bool = False,
enable_when_to_use: bool = False,
**kwargs,
self,
enable_memory_target: bool = False,
enable_when_to_use: bool = False,
**kwargs,
):
super().__init__(**kwargs)
self.enable_memory_target: bool = enable_memory_target
@ -116,14 +116,16 @@ class UpdateMemory(BaseMemoryTool):
except Exception:
logger.warning(f"Invalid message time format: {message_time}")
update_dicts.append({
"memory_id": mem.get("memory_id", ""),
"content": memory_content,
"when_to_use": when_to_use,
"message_time": message_time,
"author": self.author,
"metadata": metadata,
})
update_dicts.append(
{
"memory_id": mem.get("memory_id", ""),
"content": memory_content,
"when_to_use": when_to_use,
"message_time": message_time,
"author": self.author,
"metadata": metadata,
},
)
if update_dicts:
handler = MemoryHandler(target, self.service_context)

View file

@ -11,11 +11,11 @@ class UpdateMemoryV2(BaseMemoryTool):
"""Tool to update memories in vector store by deleting and adding memory entries"""
def __init__(
self,
name="update_memory",
enable_memory_target: bool = False,
enable_when_to_use: bool = False,
**kwargs,
self,
name="update_memory",
enable_memory_target: bool = False,
enable_when_to_use: bool = False,
**kwargs,
):
kwargs["enable_multiple"] = True
super().__init__(name=name, **kwargs)
@ -68,7 +68,7 @@ class UpdateMemoryV2(BaseMemoryTool):
"type": "array",
"description": "List of memory IDs to delete",
"items": {
"type": "string"
"type": "string",
},
},
"memories_to_add": {
@ -85,7 +85,7 @@ class UpdateMemoryV2(BaseMemoryTool):
async def execute(self):
# Get parameters
memory_ids_to_delete = self.context.get("memory_ids_to_delete", [])
memory_ids_to_delete = sorted(set([mid for mid in memory_ids_to_delete if mid]))
memory_ids_to_delete = sorted({mid for mid in memory_ids_to_delete if mid})
memories_to_add = self.context.get("memories_to_add", [])
if not memory_ids_to_delete and not memories_to_add:
@ -126,14 +126,16 @@ class UpdateMemoryV2(BaseMemoryTool):
except Exception:
logger.warning(f"Invalid message time format: {message_time}")
add_dicts.append({
"content": memory_content,
"when_to_use": when_to_use,
"message_time": message_time,
"ref_memory_id": self.history_id,
"author": self.author,
"metadata": metadata,
})
add_dicts.append(
{
"content": memory_content,
"when_to_use": when_to_use,
"message_time": message_time,
"ref_memory_id": self.history_id,
"author": self.author,
"metadata": metadata,
},
)
if add_dicts:
handler = MemoryHandler(target, self.service_context)

View file

@ -1,5 +1,4 @@
"""Update user profile tool"""
from pathlib import Path
from loguru import logger
@ -11,12 +10,10 @@ from ...core.schema import ToolCall
class UpdateProfile(BaseMemoryTool):
"""Tool to update user profile by adding or removing profile entries"""
def __init__(self, profile_path: str, **kwargs):
def __init__(self, **kwargs):
kwargs["enable_multiple"] = True
super().__init__(**kwargs)
self.profile_path: str = profile_path
def _build_multiple_tool_call(self) -> ToolCall:
"""Build and return the multiple tool call schema"""
return ToolCall(
@ -29,7 +26,7 @@ class UpdateProfile(BaseMemoryTool):
"type": "array",
"description": "List of profile IDs to delete",
"items": {
"type": "string"
"type": "string",
},
},
"profiles_to_add": {
@ -61,14 +58,11 @@ class UpdateProfile(BaseMemoryTool):
)
async def execute(self):
profile_handler = ProfileHandler(
profile_path=Path(self.profile_path) / self.vector_store.collection_name,
memory_target=self.memory_target,
)
profile_handler = ProfileHandler(profile_path=self.profile_path, memory_target=self.memory_target)
# Get parameters
profile_ids_to_delete = self.context.get("profile_ids_to_delete", [])
profile_ids_to_delete = sorted(set([pid for pid in profile_ids_to_delete if pid]))
profile_ids_to_delete = sorted({pid for pid in profile_ids_to_delete if pid})
profiles_to_add = self.context.get("profiles_to_add", [])
if not profile_ids_to_delete and not profiles_to_add: