feat(memory): extend memory system with procedural and tool memory support

This commit is contained in:
jinli.yl 2026-01-28 16:54:28 +08:00
parent afa4eb9114
commit c174aade76
19 changed files with 709 additions and 182 deletions

View file

@ -1,7 +1,9 @@
"""A simple chatbot."""
from . import chat
from . import memory
__all__ = [
"chat",
"memory",
]

View file

@ -1,13 +1,26 @@
"""Default memory agents for personal and ReMe memory operations."""
"""Default memory agents for personal, procedural, tool and ReMe memory operations."""
from .personal_retriever import PersonalRetriever
from .personal_summarizer import PersonalSummarizer
from .procedural_retriever import ProceduralRetriever
from .procedural_summarizer import ProceduralSummarizer
from .reme_retriever import ReMeRetriever
from .reme_summarizer import ReMeSummarizer
from .tool_retriever import ToolRetriever
from .tool_summarizer import ToolSummarizer
from ....core import R
__all__ = [
"PersonalRetriever",
"PersonalSummarizer",
"ProceduralRetriever",
"ProceduralSummarizer",
"ReMeRetriever",
"ReMeSummarizer",
"ToolRetriever",
"ToolSummarizer",
]
for name in __all__:
tool_class = globals()[name]
R.op.register()(tool_class)

View file

@ -20,6 +20,13 @@ class PersonalRetriever(BaseMemoryAgent):
else:
raise ValueError("input must have either `query` or `messages`")
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)
else:
all_profiles = ""
return [
Message(
role=Role.SYSTEM,
@ -27,7 +34,7 @@ class PersonalRetriever(BaseMemoryAgent):
prompt_name="system_prompt",
memory_type=self.memory_type.value,
memory_target=self.memory_target,
user_profile=await self.read_user_profile(show_id="history"),
user_profile=all_profiles,
context=context.strip(),
),
),
@ -59,5 +66,4 @@ class PersonalRetriever(BaseMemoryAgent):
async def execute(self):
result = await super().execute()
result["retrieved_nodes"] = self.retrieved_nodes
return result

View file

@ -1,45 +1,44 @@
system_prompt: |
You are a memory agent managing **{memory_type}** memories about **{memory_target}**.
You are a memory retrieval Agent responsible for retrieving {memory_type} memories about {memory_target}.
## User Profile
{user_profile}
## Question
## User Question
{context}
## Retrieval Strategy
**Tool 1: Vector Search (`retrieve_memory`)**
### Phase 1 `retrieve_memory`
- Purpose: Search for relevant memories using semantic similarity
- Try at least 3-5 different queries before moving to next tool:
- Try at least 3-5 different queries before moving to next phase:
* Direct question
* Direct question reformulation
* Different phrasings and perspectives
* Entity-focused queries (names, places, events)
* Various keyword combinations
- Time range filtering (optional):
- Time filter (optional):
* Format: single date '20200101' or range '20200101,20200102'
* Example: '20200101,20200102' for 20200101 <= time <= 20200102
* Single-sided: '0,20200102' (before date) or '20200101,99999999' (after date)
- If no results: retry with different time ranges or remove time constraints
**Tool 2: Read History (`read_history`) - ONLY AFTER Tool 1**
### Phase 2 `read_history`
- Purpose: Read full original conversation context
- Use this ONLY after completing multiple retrieve_memory attempts
- Extract history_id from retrieved memory results
- Extract history_id from context
- Prioritize most relevant or recent history entries
- Read multiple histories if needed for complete understanding
## Response Requirements
- Answer ONLY based on retrieved memories and user profile - NO hallucination or inference
- Answer ONLY based on retrieved memories / user profile / history - NO hallucination or inference
- Always cite the source: reference specific memories with their timestamps
- If information conflicts, present all versions with their respective times
- Try multiple search angles before concluding no information exists
## Output Format
When answering, structure your response as follows:
- [timestamp][Relevant history/memory/profile from context]
If no relevant information found after thorough search (5+ queries), state:
### Output Format
1. When answering, structure your response as follows:
- [timestamp][Relevant retrieved memories / user profile / history from context]
2. If no relevant information found after thorough search (5+ queries), state:
"No relevant information found after thorough search using multiple query strategies."
user_message: |

View file

@ -1,5 +1,4 @@
"""Personal memory summarizer agent for two-phase personal memory processing."""
from loguru import logger
from ..base_memory_agent import BaseMemoryAgent
@ -13,13 +12,12 @@ class PersonalSummarizer(BaseMemoryAgent):
memory_type: MemoryType = MemoryType.PERSONAL
async def _build_phase1_messages(self) -> list[Message]:
"""Build messages for phase 1: retrieve and add memory."""
async def _build_s1_messages(self) -> list[Message]:
return [
Message(
role=Role.SYSTEM,
content=self.prompt_format(
prompt_name="system_prompt_phase1",
prompt_name="system_prompt_s1",
context=self.context.history_node.content,
memory_type=self.memory_type.value,
memory_target=self.memory_target,
@ -27,26 +25,24 @@ class PersonalSummarizer(BaseMemoryAgent):
),
Message(
role=Role.USER,
content=self.get_prompt("user_message_phase1"),
content=self.get_prompt("user_message_s1"),
),
]
async def _build_phase2_messages(self) -> list[Message]:
"""Build messages for phase 2: update user profile."""
async def _build_s2_messages(self) -> list[Message]:
return [
Message(
role=Role.SYSTEM,
content=self.prompt_format(
prompt_name="system_prompt_phase2",
prompt_name="system_prompt_s2",
context=self.context.history_node.content,
memory_type=self.memory_type.value,
memory_target=self.memory_target,
user_profile=await self.read_user_profile(show_id="profile"),
),
),
Message(
role=Role.USER,
content=self.get_prompt("user_message_phase2"),
content=self.get_prompt("user_message_s2"),
),
]
@ -73,34 +69,52 @@ class PersonalSummarizer(BaseMemoryAgent):
)
async def execute(self):
"""Execute two-phase memory processing: retrieve/add -> update profile."""
tools = self.tools
for i, tool in enumerate(tools):
memory_tools = []
profile_tools = []
for i, tool in enumerate(self.tools):
tool_name = tool.tool_call.name
if "_memory" in tool_name:
memory_tools.append(tool)
elif "_profile" in tool_name:
profile_tools.append(tool)
else:
raise ValueError(f"[{self.__class__.__name__}] unknown tool_name={tool_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):
stage = "s1-memory"
messages_s1 = await self._build_s1_messages()
for i, message in enumerate(messages_s1):
role = message.name or message.role
logger.info(f"[{self.__class__.__name__} S1] role={role} {message.simple_dump(as_dict=False)}")
tools_phase1, messages_phase1, success_phase1 = await self.react(messages_phase1, tools[:-1], stage="S1")
logger.info(f"[{self.__class__.__name__} {stage}] role={role} {message.simple_dump(as_dict=False)}")
tools_s1, messages_s1, success_s1 = await self.react(messages_s1, memory_tools, stage=stage)
messages_phase2 = await self._build_phase2_messages()
for i, message in enumerate(messages_phase2):
role = message.name or message.role
logger.info(f"[{self.__class__.__name__} S2] role={role} {message.simple_dump(as_dict=False)}")
tools_phase2, messages_phase2, success_phase2 = await self.react(messages_phase2, tools[-1:], stage="S2")
if profile_tools:
stage = "s2-profile"
messages_s2 = await self._build_s2_messages()
for i, message in enumerate(messages_s2):
role = message.name or message.role
logger.info(f"[{self.__class__.__name__} {stage}] role={role} {message.simple_dump(as_dict=False)}")
tools_s2, messages_s2, success_s2 = await self.react(messages_s2, profile_tools, stage=stage)
else:
tools_s2, messages_s2, success_s2 = [], [], True
success = success_phase1 and success_phase2
messages = messages_phase1 + messages_phase2
tools = tools_phase1 + tools_phase2
success = success_s1 and success_s2
messages = messages_s1 + messages_s2
tools = tools_s1 + tools_s2
memory_nodes = []
for tool in tools:
if tool.memory_nodes:
memory_nodes.extend(tool.memory_nodes)
profile_nodes = []
for tool in tools:
if tool.profile_nodes:
profile_nodes.extend(tool.profile_nodes)
return {
"answer": memory_nodes,
"success": success,
"messages": messages,
"tools": tools,
"profile_nodes": profile_nodes,
}

View file

@ -1,45 +1,91 @@
system_prompt_phase1: |
You are a memory agent managing **{memory_type}** memories about **{memory_target}**.
## Latest Conversation:
Message format: `round<index> [<timestamp>] <role/name>: <content>` (timestamp: YYYY-MM-DD HH:MM:SS).
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。
要求:
- 原样提取最新对话中的内容,不得推断、假设或编造。
- 最后记忆库包含所有的历史记忆和新的记忆,例如记录在同一个主题下用户不同时间的变化。
- 最后记忆库有比较好的组织,同一主题的记忆放到同一条中,不要有重复/多余的记忆。
## Task: Retrieve Similar Memories and Add New Memories
**CRITICAL**: Extract ONLY explicitly stated information. DO NOT infer, assume, or fabricate.
user_message_s1_zh: |
严格按照步骤1和步骤2完成任务
### Step 1: Retrieve Similar Memories
Use `retrieve_memory` to search for existing similar memories about **{memory_target}**.
- Use appropriate queries to find relevant existing memories
- Check if new information already exists in the memory store
### Step 2: Add New Memories
Use `add_memory` to add new memories:
- Extract and summarize important information about **{memory_target}**
- Set `update_time` (format: 2020-01-01 00:00:00; use 0000-00-00 00:00:00 if unavailable)
- If the information is completely identical to existing memory, skip adding
user_message_phase1: |
First retrieve similar memories, then extract and add new personal memories from the conversation.
system_prompt_phase2: |
You are a memory agent managing **{memory_type}** memories about **{memory_target}**.
## Latest Conversation:
Message format: `round<index> [<timestamp>] <role/name>: <content>` (timestamp: YYYY-MM-DD HH:MM:SS).
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。
要求:
- 原样提取最新对话中的内容,不得推断、假设或编造。
- 最后Profile库只保留用户最新的状态。例如用户开始喜欢吃苹果后来只吃喜欢香蕉可以记录水果偏好香蕉
- 最后Profile库有比较好的组织同一主题的Profile放到同一条中不要有重复/多余的Profile。
## Current User Profile:
UserProfile format: `profile_id=<id> update_time=<timestamp> <content>`.
{user_profile}
user_message_s2_zh: |
严格按照步骤1和步骤2完成任务
## Task: Update Profile with `UpdateUserProfile`
**CRITICAL**: Extract ONLY explicitly stated information. DO NOT infer, assume, or fabricate.
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.
Requirements:
- Extract content from the latest conversation as-is, without inference, assumption, or fabrication.
- The final memory store should contain all historical memories and new memories, for example, recording user changes at different times under the same topic.
- The final memory store should be well-organized, with memories on the same topic placed in one entry, without duplicate/redundant memories.
Synchronize profile/memories with new information from the conversation, including **{memory_target}**' current status:
- `profile_ids_to_delete`: Remove conflicting, or redundant entries.
- `profiles_to_add`: Add new profiles/memories with `update_time`, e.g. `YYYY-MM-DD HH:MM:SS`, {memory_target} did something.
- Maintain profiles that are concise, mutually exclusive, and collectively comprehensive with no information loss.
user_message_s1: |
Strictly complete the task following Step 1 and Step 2
user_message_phase2: |
Update user profile using `UpdateUserProfile` based on the conversation and current profile.
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.
Requirements:
- Extract content from the latest conversation as-is, without inference, assumption, or fabrication.
- The final Profile store should only keep the user's latest state. For example, if the user initially liked apples but later only likes bananas, record: Fruit preference: banana
- The final Profile store should be well-organized, with Profiles on the same topic placed in one entry, without duplicate/redundant Profiles.
user_message_s2: |
Strictly complete the task following Step 1 and Step 2

View file

@ -0,0 +1,6 @@
from ..base_memory_agent import BaseMemoryAgent
from ....core.enumeration import MemoryType
class ProceduralRetriever(BaseMemoryAgent):
memory_type: MemoryType = MemoryType.PROCEDURAL

View file

@ -0,0 +1,6 @@
from ..base_memory_agent import BaseMemoryAgent
from ....core.enumeration import MemoryType
class ProceduralSummarizer(BaseMemoryAgent):
memory_type: MemoryType = MemoryType.PROCEDURAL

View file

@ -0,0 +1,6 @@
from ..base_memory_agent import BaseMemoryAgent
from ....core.enumeration import MemoryType
class ToolRetriever(BaseMemoryAgent):
memory_type: MemoryType = MemoryType.TOOL

View file

@ -0,0 +1,6 @@
from ..base_memory_agent import BaseMemoryAgent
from ....core.enumeration import MemoryType
class ToolSummarizer(BaseMemoryAgent):
memory_type: MemoryType = MemoryType.TOOL

View file

@ -11,6 +11,7 @@ from . import service
from . import token_counter
from . import utils
from . import vector_store
from .application import Application
from .context import R
__all__ = [
@ -25,5 +26,6 @@ __all__ = [
"token_counter",
"utils",
"vector_store",
"Application",
"R",
]

106
reme/core/application.py Normal file
View file

@ -0,0 +1,106 @@
import asyncio
from .context import PromptHandler, ServiceContext
from .embedding import BaseEmbeddingModel
from .flow import BaseFlow
from .llm import BaseLLM
from .schema import Response
from .token_counter import BaseTokenCounter
from .utils import execute_stream_task, PydanticConfigParser
from .vector_store import BaseVectorStore
class Application:
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,
):
# ServiceContext
self.service_context = ServiceContext(
*args,
llm_api_key=llm_api_key,
llm_api_base=llm_api_base,
embedding_api_key=embedding_api_key,
embedding_api_base=embedding_api_base,
service_config=None,
parser=parser,
config_path=None,
enable_logo=enable_logo,
llm=llm,
embedding_model=embedding_model,
vector_store=vector_store,
token_counter=token_counter,
**kwargs,
)
# PromptHandler
self.prompt_handler = PromptHandler(language=self.service_context.language)
# LLM & EmbeddingModel & VectorStore & TokenCounter
self.llm: BaseLLM | None = self.service_context.llms.get("default", None)
self.embedding_model: BaseEmbeddingModel | None = self.service_context.embedding_models.get("default", None)
self.vector_store: BaseVectorStore | None = self.service_context.vector_stores.get("default", None)
self.token_counter: BaseTokenCounter | None = self.service_context.token_counters.get("default", None)
async def __aenter__(self):
"""Async context manager entry."""
return self
def __enter__(self):
"""Context manager entry."""
return self
async def close(self):
"""Close"""
return await self.service_context.close()
def close_sync(self):
"""Close synchronously"""
self.service_context.close_sync()
async def __aexit__(self, exc_type=None, exc_val=None, exc_tb=None):
"""Async context manager exit."""
await self.close()
return False
def __exit__(self, exc_type=None, exc_val=None, exc_tb=None):
"""Context manager exit."""
self.close_sync()
return False
async def execute_flow(self, name: str, **kwargs) -> Response:
"""Execute a flow with the given name and parameters."""
assert name in self.service_context.flows, f"Flow {name} not found"
flow: BaseFlow = self.service_context.flows[name]
return await flow.call(**kwargs)
async def execute_stream_flow(self, name: str, **kwargs):
"""Execute a stream flow with the given name and parameters."""
assert name in self.service_context.flows, f"Flow {name} not found"
flow: BaseFlow = self.service_context.flows[name]
assert flow.stream is True, "non-stream flow is not supported in execute_stream_flow!"
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,
):
yield chunk
def run_service(self):
"""Run the configured service (HTTP, MCP, or CMD)."""
self.service_context.service.run()

View file

@ -6,6 +6,7 @@ from pathlib import Path
from loguru import logger
from .core import Application
from .agent.memory.default import ReMeSummarizer, PersonalSummarizer, PersonalRetriever, ReMeRetriever
from .config import ReMeConfigParser
from .core.context import PromptHandler, ServiceContext
@ -21,7 +22,7 @@ from .tool.memory import UpdateUserProfile, RetrieveMemory, AddMemory, DelegateT
ProfileHandler
class ReMe:
class ReMe(Application):
"""ReMe with config file support and flow execution methods."""
def __init__(
@ -50,7 +51,20 @@ class ReMe:
tool_retrieve_version: str = "default",
**kwargs,
):
# MemoryTarget -> MemoryType
super().__init__(
*args,
llm_api_key=llm_api_key,
llm_api_base=llm_api_base,
embedding_api_key=embedding_api_key,
embedding_api_base=embedding_api_base,
enable_logo=enable_logo,
parser=ReMeConfigParser,
llm=llm,
embedding_model=embedding_model,
vector_store=vector_store,
token_counter=token_counter,
**kwargs,
)
memory_target_type_mapping: dict[str, MemoryType] = {}
if personal_memory_target:
for name in personal_memory_target:
@ -66,37 +80,9 @@ class ReMe:
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
# ServiceContext
self.service_context = ServiceContext(
*args,
llm_api_key=llm_api_key,
llm_api_base=llm_api_base,
embedding_api_key=embedding_api_key,
embedding_api_base=embedding_api_base,
service_config=None,
parser=ReMeConfigParser,
config_path=None,
enable_logo=enable_logo,
llm=llm,
embedding_model=embedding_model,
vector_store=vector_store,
token_counter=token_counter,
memory_target_type_mapping=memory_target_type_mapping,
**kwargs,
)
self.service_context.memory_target_type_mapping = memory_target_type_mapping
self.profile_path: str = profile_path
# PromptHandler
self.prompt_handler = PromptHandler(language=self.service_context.language)
# LLM & EmbeddingModel & VectorStore & TokenCounter
self.llm: BaseLLM | None = self.service_context.llms.get("default", None)
self.embedding_model: BaseEmbeddingModel | None = self.service_context.embedding_models.get("default", None)
self.vector_store: BaseVectorStore | None = self.service_context.vector_stores.get("default", None)
self.token_counter: BaseTokenCounter | None = self.service_context.token_counters.get("default", None)
@property
def memory_target_type_mapping(self) -> dict[str, MemoryType]:
mapping = {}
@ -395,57 +381,6 @@ class ReMe:
async def context_reload(self):
"""working memory retrieve"""
async def execute_flow(self, name: str, **kwargs) -> Response:
"""Execute a flow with the given name and parameters."""
assert name in self.service_context.flows, f"Flow {name} not found"
flow: BaseFlow = self.service_context.flows[name]
return await flow.call(**kwargs)
async def execute_stream_flow(self, name: str, **kwargs):
"""Execute a stream flow with the given name and parameters."""
assert name in self.service_context.flows, f"Flow {name} not found"
flow: BaseFlow = self.service_context.flows[name]
assert flow.stream is True, "non-stream flow is not supported in execute_stream_flow!"
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,
):
yield chunk
def run_service(self):
"""Run the configured service (HTTP, MCP, or CMD)."""
self.service_context.service.run()
async def __aenter__(self):
"""Async context manager entry."""
return self
def __enter__(self):
"""Context manager entry."""
return self
async def close(self):
"""Close"""
return await self.service_context.close()
def close_sync(self):
"""Close synchronously"""
self.service_context.close_sync()
async def __aexit__(self, exc_type=None, exc_val=None, exc_tb=None):
"""Async context manager exit."""
await self.close()
return False
def __exit__(self, exc_type=None, exc_val=None, exc_tb=None):
"""Context manager exit."""
self.close_sync()
return False
def main():
"""Main entry point for running ReMe from command line."""

View file

@ -1,31 +1,39 @@
"""memory tools"""
from .add_draft_and_read_all_profiles import AddDraftAndReadAllProfiles
from .add_draft_and_retrieve_similar_memory import AddDraftAndRetrieveSimilarMemory
from .add_history import AddHistory
from .add_memory import AddMemory
from .base_memory_tool import BaseMemoryTool
from .delegate_task import DelegateTask
from .delete_memory import DeleteMemory
from .memory_handler import MemoryHandler
from .profile_handler import ProfileHandler
from .read_all_profiles import ReadAllProfiles
from .read_history import ReadHistory
from .read_profile import ReadProfile
from .retrieve_memory import RetrieveMemory
from .retrieve_recent_memory import RetrieveRecentMemory
from .update_memory import UpdateMemory
from .update_memory_v2 import UpdateMemoryV2
from .update_profile import UpdateProfile
from ...core import R
__all__ = [
"AddDraftAndReadAllProfiles",
"AddDraftAndRetrieveSimilarMemory",
"AddHistory",
"AddMemory",
"BaseMemoryTool",
"DelegateTask",
"DeleteMemory",
"MemoryHandler",
"ProfileHandler",
"ReadAllProfiles",
"ReadHistory",
"ReadProfile",
"RetrieveMemory",
"RetrieveRecentMemory",
"UpdateMemory",
"UpdateMemoryV2",
"UpdateProfile",
]

View file

@ -0,0 +1,105 @@
"""Add draft profile and read all profiles from local storage"""
from pathlib import Path
from loguru import logger
from .base_memory_tool import BaseMemoryTool
from .profile_handler import ProfileHandler
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):
super().__init__(**kwargs)
self.profile_path: str = profile_path
self.enable_memory_target: bool = enable_memory_target
def _build_query_parameters(self) -> dict:
"""Build the query parameters schema"""
properties = {
"profile_draft": {
"type": "string",
"description": "profile_draft",
},
}
required = ["profile_draft"]
if self.enable_memory_target:
properties["memory_target"] = {
"type": "string",
"description": "memory_target",
}
required.append("memory_target")
return {
"type": "object",
"properties": properties,
"required": required,
}
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": "Add draft profile and read all profiles from local storage.",
"parameters": self._build_query_parameters(),
},
)
def _build_multiple_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": "Add draft profile and read all profiles from local storage.",
"parameters": {
"type": "object",
"properties": {
"draft_items": {
"type": "array",
"description": "List of draft profile items.",
"items": self._build_query_parameters(),
},
},
"required": ["draft_items"],
},
},
)
async def execute(self):
if self.enable_multiple:
draft_items = self.context.get("draft_items", [])
else:
draft_items = [self.context]
# Collect all profiles from all targets
all_profiles = []
targets_processed = set()
for item in draft_items:
if self.enable_memory_target:
target = item["memory_target"]
else:
target = self.memory_target
# Skip if already processed this target
if target in targets_processed:
continue
targets_processed.add(target)
profile_handler = ProfileHandler(
profile_path=Path(self.profile_path) / self.vector_store.collection_name,
memory_target=target,
)
profiles_str = profile_handler.read_all(add_profile_id=True)
if profiles_str:
all_profiles.append(f"## Profiles for {target}:\n{profiles_str}")
if not all_profiles:
output = "No profiles found."
logger.info(output)
return output
output = "\n\n".join(all_profiles)
logger.info(f"Successfully read profiles for {len(targets_processed)} target(s)")
return output

View file

@ -0,0 +1,107 @@
"""Add draft memory and retrieve similar memories from vector store"""
from loguru import logger
from .base_memory_tool import BaseMemoryTool
from .memory_handler import MemoryHandler
from ...core.schema import ToolCall, MemoryNode
from ...core.utils import deduplicate_memories
class AddDraftAndRetrieveSimilarMemory(BaseMemoryTool):
"""Tool to add draft memory and retrieve similar memories"""
def __init__(self, top_k: int = 20, enable_memory_target: bool = False, **kwargs):
super().__init__(**kwargs)
self.top_k: int = top_k
self.enable_memory_target: bool = enable_memory_target
def _build_query_parameters(self) -> dict:
"""Build the query parameters schema"""
properties = {
"memory_draft": {
"type": "string",
"description": "memory_draft",
},
}
required = ["memory_draft"]
if self.enable_memory_target:
properties["memory_target"] = {
"type": "string",
"description": "memory_target",
}
required.append("memory_target")
return {
"type": "object",
"properties": properties,
"required": required,
}
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": "Add draft memory and retrieve similar memories from the vector store.",
"parameters": self._build_query_parameters(),
},
)
def _build_multiple_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": "Add draft memory and retrieve similar memories from the vector store.",
"parameters": {
"type": "object",
"properties": {
"draft_items": {
"type": "array",
"description": "List of draft memory items.",
"items": self._build_query_parameters(),
},
},
"required": ["draft_items"],
},
},
)
async def execute(self):
if self.enable_multiple:
draft_items = self.context.get("draft_items", [])
else:
draft_items = [self.context]
queries_by_target: dict[str, list[dict]] = {}
for item in draft_items:
if self.enable_memory_target:
target = item["memory_target"]
else:
target = self.memory_target
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": {},
})
# Execute batch searches for each target
memory_nodes: list[MemoryNode] = []
for target, searches in queries_by_target.items():
handler = MemoryHandler(target, self.service_context)
nodes = await handler.batch_search(searches)
memory_nodes.extend(nodes)
memory_nodes = deduplicate_memories(memory_nodes)
retrieved_ids = {n.memory_id for n in self.retrieved_nodes if n.memory_id}
new_nodes = [n for n in memory_nodes if n.memory_id not in retrieved_ids]
self.retrieved_nodes.extend(new_nodes)
if not new_nodes:
output = "No similar memories found."
else:
output = "\n".join([n.format(ref_memory_id_key="history_id") for n in new_nodes])
logger.info(f"Retrieved {len(memory_nodes)} similar memories, {len(new_nodes)} new after deduplication")
return output

View file

@ -8,7 +8,7 @@ from .profile_handler import ProfileHandler
from ...core.schema import ToolCall
class ReadProfile(BaseMemoryTool):
class ReadAllProfiles(BaseMemoryTool):
"""Tool to read all user profiles"""
def __init__(self, profile_path: str, **kwargs):
@ -35,7 +35,7 @@ class ReadProfile(BaseMemoryTool):
memory_target=self.memory_target,
)
profiles_str = profile_handler.read_all()
profiles_str = profile_handler.read_all(add_profile_id=True)
if not profiles_str:
output = "No profiles found."
logger.info(output)

View file

@ -11,29 +11,34 @@ from ...core.utils import deduplicate_memories
class RetrieveMemory(BaseMemoryTool):
"""Tool to retrieve memories using similarity search"""
def __init__(self, top_k: int = 20, enable_memory_target: bool = False, **kwargs):
def __init__(self, top_k: int = 20, enable_memory_target: bool = False, enable_time_filter: bool = False, **kwargs):
super().__init__(**kwargs)
self.top_k: int = top_k
self.enable_memory_target: bool = enable_memory_target
self.enable_time_filter: bool = enable_time_filter
def _build_query_parameters(self) -> dict:
"""Build the query parameters schema based on enabled features."""
properties = {
"query": {
"type": "string",
"description": "query text for vector similarity search.",
},
"time_range": {
"type": "string",
"description": "optional time range filter. Format: '20200101' or '20200101,20200102'",
"description": "query",
},
}
required = ["query"]
if self.enable_time_filter:
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.",
}
if self.enable_memory_target:
properties["memory_target"] = {
"type": "string",
"description": "target memory type to search in.",
"description": "memory_target",
}
required.append("memory_target")
@ -46,7 +51,7 @@ class RetrieveMemory(BaseMemoryTool):
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": "retrieve memories using vector similarity search.",
"description": "Retrieve relevant memories from the vector store using semantic similarity search.",
"parameters": self._build_query_parameters(),
},
)
@ -54,13 +59,13 @@ class RetrieveMemory(BaseMemoryTool):
def _build_multiple_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": "retrieve memories using multiple queries with vector similarity search.",
"description": "Retrieve relevant memories from the vector store using semantic similarity search.",
"parameters": {
"type": "object",
"properties": {
"query_items": {
"type": "array",
"description": "list of query items for vector similarity search.",
"description": "List of query items.",
"items": self._build_query_parameters(),
},
},
@ -85,14 +90,14 @@ class RetrieveMemory(BaseMemoryTool):
queries_by_target[target] = []
filters = {}
time_range = item.get("time_range")
if time_range:
time_range = time_range.strip()
if "," in time_range:
start, end = time_range.split(",")
time_filter = item.get("time_filter")
if time_filter:
time_filter = time_filter.strip()
if "," in time_filter:
start, end = time_filter.split(",")
filters = {"time_int": [int(start.strip()), int(end.strip())]}
else:
filters = {"time_int": [int(time_range), int(time_range)]}
filters = {"time_int": [int(time_filter), int(time_filter)]}
queries_by_target[target].append({
"query": item["query"],

View file

@ -0,0 +1,155 @@
"""Update memory in vector store"""
from loguru import logger
from .base_memory_tool import BaseMemoryTool
from .memory_handler import MemoryHandler
from ...core.schema import ToolCall
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,
):
kwargs["enable_multiple"] = True
super().__init__(name=name, **kwargs)
self.enable_memory_target: bool = enable_memory_target
self.enable_when_to_use: bool = enable_when_to_use
def _build_add_memory_parameters(self) -> dict:
"""Build the add memory parameters schema based on enabled features."""
properties = {
"message_time": {
"type": "string",
"description": "message time, e.g. '2020-01-01 00:00:00'",
},
"memory_content": {
"type": "string",
"description": "content of the memory.",
},
}
required = ["message_time", "memory_content"]
if self.enable_when_to_use:
properties["when_to_use"] = {
"type": "string",
"description": "description of when to use this memory.",
}
required.append("when_to_use")
if self.enable_memory_target:
properties["memory_target"] = {
"type": "string",
"description": "target memory type for this memory.",
}
required.append("memory_target")
return {
"type": "object",
"properties": properties,
"required": required,
}
def _build_multiple_tool_call(self) -> ToolCall:
"""Build and return the multiple tool call schema"""
return ToolCall(
**{
"description": "update memories by removing and adding memory entries.",
"parameters": {
"type": "object",
"properties": {
"memory_ids_to_delete": {
"type": "array",
"description": "List of memory IDs to delete",
"items": {
"type": "string"
},
},
"memories_to_add": {
"type": "array",
"description": "List of memories to add",
"items": self._build_add_memory_parameters(),
},
},
"required": ["memory_ids_to_delete", "memories_to_add"],
},
},
)
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]))
memories_to_add = self.context.get("memories_to_add", [])
if not memory_ids_to_delete and not memories_to_add:
return "No memories to remove or add, operation completed."
# Group memories by memory_target if enabled
if self.enable_memory_target:
memories_by_target = {}
for mem in memories_to_add:
target = mem.get("memory_target", self.memory_target)
if target not in memories_by_target:
memories_by_target[target] = []
memories_by_target[target].append(mem)
else:
memories_by_target = {self.memory_target: memories_to_add}
# Delete memories (all at once, regardless of target)
removed_count = 0
if memory_ids_to_delete:
# Use the default memory_target handler for deletion
handler = MemoryHandler(self.memory_target, self.service_context)
await handler.delete(memory_ids_to_delete)
removed_count = len(memory_ids_to_delete)
# Add new memories by target
added_count = 0
all_memory_nodes = []
for target, target_memories in memories_by_target.items():
# Parse and prepare add data
add_dicts = []
for mem in target_memories:
memory_content = mem.get("memory_content", "")
message_time = mem.get("message_time", "")
when_to_use = mem.get("when_to_use", "") if self.enable_when_to_use else ""
metadata = {}
try:
metadata["time_int"] = int(message_time.split(" ")[0].replace("-", ""))
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,
})
if add_dicts:
handler = MemoryHandler(target, self.service_context)
memory_nodes = await handler.add_batch(add_dicts)
all_memory_nodes.extend(memory_nodes)
added_count += len(memory_nodes)
# Extend memory_nodes for tracking
self.memory_nodes.extend(all_memory_nodes)
# Build output message
operations = []
if removed_count > 0:
operations.append(f"removed {removed_count} old memories.")
if added_count > 0:
operations.append(f"added {added_count} new memories.")
operations.append("Operation completed.")
logger.info("\n".join(operations))
return "\n".join(operations)