mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
feat(memory): extend memory system with procedural and tool memory support
This commit is contained in:
parent
afa4eb9114
commit
c174aade76
19 changed files with 709 additions and 182 deletions
|
|
@ -1,7 +1,9 @@
|
|||
"""A simple chatbot."""
|
||||
|
||||
from . import chat
|
||||
from . import memory
|
||||
|
||||
__all__ = [
|
||||
"chat",
|
||||
"memory",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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: |
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
6
reme/agent/memory/default/procedural_retriever.py
Normal file
6
reme/agent/memory/default/procedural_retriever.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from ..base_memory_agent import BaseMemoryAgent
|
||||
from ....core.enumeration import MemoryType
|
||||
|
||||
|
||||
class ProceduralRetriever(BaseMemoryAgent):
|
||||
memory_type: MemoryType = MemoryType.PROCEDURAL
|
||||
6
reme/agent/memory/default/procedural_summarizer.py
Normal file
6
reme/agent/memory/default/procedural_summarizer.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from ..base_memory_agent import BaseMemoryAgent
|
||||
from ....core.enumeration import MemoryType
|
||||
|
||||
|
||||
class ProceduralSummarizer(BaseMemoryAgent):
|
||||
memory_type: MemoryType = MemoryType.PROCEDURAL
|
||||
6
reme/agent/memory/default/tool_retriever.py
Normal file
6
reme/agent/memory/default/tool_retriever.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from ..base_memory_agent import BaseMemoryAgent
|
||||
from ....core.enumeration import MemoryType
|
||||
|
||||
|
||||
class ToolRetriever(BaseMemoryAgent):
|
||||
memory_type: MemoryType = MemoryType.TOOL
|
||||
6
reme/agent/memory/default/tool_summarizer.py
Normal file
6
reme/agent/memory/default/tool_summarizer.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from ..base_memory_agent import BaseMemoryAgent
|
||||
from ....core.enumeration import MemoryType
|
||||
|
||||
|
||||
class ToolSummarizer(BaseMemoryAgent):
|
||||
memory_type: MemoryType = MemoryType.TOOL
|
||||
|
|
@ -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
106
reme/core/application.py
Normal 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()
|
||||
99
reme/reme.py
99
reme/reme.py
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
||||
|
|
|
|||
105
reme/tool/memory/add_draft_and_read_all_profiles.py
Normal file
105
reme/tool/memory/add_draft_and_read_all_profiles.py
Normal 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
|
||||
107
reme/tool/memory/add_draft_and_retrieve_similar_memory.py
Normal file
107
reme/tool/memory/add_draft_and_retrieve_similar_memory.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
|
@ -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"],
|
||||
|
|
|
|||
155
reme/tool/memory/update_memory_v2.py
Normal file
155
reme/tool/memory/update_memory_v2.py
Normal 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)
|
||||
Loading…
Add table
Reference in a new issue