mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
refactor(memory): refactor memory tools and handlers for improved structure
This commit is contained in:
parent
c174aade76
commit
c06fd78763
23 changed files with 390 additions and 423 deletions
|
|
@ -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 = ""
|
||||
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ system_prompt: |
|
|||
|
||||
## User Question
|
||||
{context}
|
||||
|
||||
|
||||
## Retrieval Strategy
|
||||
### Phase 1 `retrieve_memory`
|
||||
- Purpose: Search for relevant memories using semantic similarity
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
"""Personal memory summarizer agent for two-phase personal memory processing."""
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from ..base_memory_agent import BaseMemoryAgent
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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 = ""
|
||||
|
|
|
|||
492
reme/reme.py
492
reme/reme.py
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -39,4 +39,4 @@ __all__ = [
|
|||
|
||||
for name in __all__:
|
||||
tool_class = globals()[name]
|
||||
R.op.register()(tool_class)
|
||||
R.op.register()(tool_class)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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] = []
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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] = []
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue