mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
[dev] add add_memory op
This commit is contained in:
parent
4bb001a810
commit
44bf421487
14 changed files with 47 additions and 45 deletions
|
|
@ -17,23 +17,33 @@ memory_service:
|
|||
class: memory.operation.frontend_operation
|
||||
workflow: read_message
|
||||
description: "read short memory"
|
||||
|
||||
read_memory:
|
||||
class: memory.operation.frontend_operation
|
||||
workflow: set_query,[extract_time|retrieve_memory1,semantic_rank],fuse_rerank
|
||||
description: "read long-term memory"
|
||||
|
||||
list_memory:
|
||||
class: memory.operation.frontend_operation
|
||||
workflow: set_query,retrieve_memory2,print_memory
|
||||
description: "read all long-term memory of the user"
|
||||
|
||||
delete_memory:
|
||||
class: memory.operation.frontend_operation
|
||||
workflow: set_query,retrieve_memory3,delete_memory,store_memory
|
||||
workflow: set_query,retrieve_memory3,delete_memory
|
||||
description: "delete all long-term memory"
|
||||
|
||||
add_memory:
|
||||
class: memory.operation.frontend_operation
|
||||
workflow: store_memory
|
||||
description: "delete all long-term memory"
|
||||
|
||||
write_memory:
|
||||
class: memory.operation.write_memory
|
||||
workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory
|
||||
description: "write observation memory of the user"
|
||||
interval_time: 5
|
||||
|
||||
summary_memory:
|
||||
class: memory.operation.summary_memory
|
||||
workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
from memory_scope.constants.common_constants import CHAT_MESSAGES, RESULT, CHAT_KWARGS
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memory_scope.memory.operation.base_backend_operation import BaseBackendOperation
|
||||
from memory_scope.memory.operation.backend_operation import BackendOperation
|
||||
|
||||
|
||||
class WriteMemoryOp(BaseBackendOperation):
|
||||
class WriteMemoryOp(BackendOperation):
|
||||
|
||||
def __init__(self,
|
||||
message_lock=None,
|
||||
|
|
|
|||
|
|
@ -50,7 +50,7 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
"""
|
||||
# Parse input parameters from the worker's context
|
||||
extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT)
|
||||
memory_node_list: List[MemoryNode] = self.get_memories(RANKED_MEMORY_NODES)
|
||||
memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RANKED_MEMORY_NODES)
|
||||
|
||||
# Check if memory nodes are available; warn and return if not
|
||||
if not memory_node_list:
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ from memory_scope.utils.datetime_handler import DatetimeHandler
|
|||
class PrintMemoryWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True)
|
||||
|
||||
expired_content_list: List[str] = ["----- expired -----"]
|
||||
|
|
|
|||
|
|
@ -89,7 +89,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
|
|||
|
||||
def _run(self):
|
||||
"""
|
||||
Executes the main retrieval流程 for memories. It fetches the query from the context, initiates concurrent tasks
|
||||
Executes the main retrieval for memories. It fetches the query from the context, initiates concurrent tasks
|
||||
to retrieve memories from observations, insights, and expired sources, collects the results, sorts them by
|
||||
similarity score, logs the details, and finally sets the retrieved memory nodes.
|
||||
|
||||
|
|
|
|||
|
|
@ -29,7 +29,7 @@ class SemanticRankWorker(MemoryBaseWorker):
|
|||
"""
|
||||
# query
|
||||
query, _ = self.get_context(QUERY_WITH_TS)
|
||||
memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
if not memory_node_list:
|
||||
self.logger.warning("Retrieve memory nodes is empty!")
|
||||
return
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ class UpdateStatusWorker(MemoryBaseWorker):
|
|||
def _run(self):
|
||||
expired_action = self.expired_action
|
||||
valid_action_dict: dict = self.valid_action_dict
|
||||
memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
if not memory_node_list:
|
||||
return
|
||||
|
||||
|
|
|
|||
|
|
@ -52,8 +52,8 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
|
|||
- Parsing the model's responses for new insight keys.
|
||||
- Creating new insight nodes and updating the memory status accordingly.
|
||||
"""
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_memories(NOT_REFLECTED_NODES)
|
||||
insight_nodes: List[MemoryNode] = self.get_memories(INSIGHT_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_REFLECTED_NODES)
|
||||
insight_nodes: List[MemoryNode] = self.memory_handler.get_memories(INSIGHT_NODES)
|
||||
|
||||
# Count unaudited nodes
|
||||
not_reflected_count = len(not_reflected_nodes)
|
||||
|
|
|
|||
|
|
@ -56,7 +56,7 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
|
||||
The process helps in maintaining conversation coherence by resolving contradictions and redundancies.
|
||||
"""
|
||||
not_updated_nodes: List[MemoryNode] = self.get_memories(NOT_UPDATED_NODES)
|
||||
not_updated_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_UPDATED_NODES)
|
||||
for node in not_updated_nodes:
|
||||
self.submit_thread_task(fn=self.retrieve_similar_content, node=node)
|
||||
|
||||
|
|
@ -141,4 +141,4 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
self.logger.info(f"after_long_contra_repeat: {node.content} {node.status}")
|
||||
|
||||
# save context
|
||||
self.set_memories(MERGE_OBS_NODES, merge_obs_nodes)
|
||||
self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes)
|
||||
|
|
|
|||
|
|
@ -152,9 +152,9 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
5. Gather the results of all update tasks.
|
||||
6. Mark processed nodes as updated in memory.
|
||||
"""
|
||||
insight_nodes: List[MemoryNode] = self.get_memories(INSIGHT_NODES)
|
||||
not_updated_nodes: List[MemoryNode] = self.get_memories(NOT_UPDATED_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_memories(NOT_REFLECTED_NODES)
|
||||
insight_nodes: List[MemoryNode] = self.memory_handler.get_memories(INSIGHT_NODES)
|
||||
not_updated_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_UPDATED_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_REFLECTED_NODES)
|
||||
|
||||
if not insight_nodes:
|
||||
self.logger.warning("insight_nodes is empty, stopping processing.")
|
||||
|
|
|
|||
|
|
@ -37,13 +37,13 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
6. Updates the status of nodes accordingly.
|
||||
7. Persists the changes back to memory storage.
|
||||
"""
|
||||
all_obs_nodes: List[MemoryNode] = self.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES])
|
||||
all_obs_nodes: List[MemoryNode] = self.memory_handler.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES])
|
||||
if not all_obs_nodes:
|
||||
self.logger.info("all_obs_nodes is empty!")
|
||||
self.continue_run = False
|
||||
return
|
||||
|
||||
today_obs_nodes: List[MemoryNode] = self.get_memories(TODAY_NODES)
|
||||
today_obs_nodes: List[MemoryNode] = self.memory_handler.get_memories(TODAY_NODES)
|
||||
|
||||
if today_obs_nodes:
|
||||
all_obs_nodes.extend(today_obs_nodes)
|
||||
|
|
@ -111,4 +111,4 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
merge_obs_nodes.append(node)
|
||||
|
||||
# save context
|
||||
self.set_memories(MERGE_OBS_NODES, merge_obs_nodes, log_repeat=False)
|
||||
self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes, log_repeat=False)
|
||||
|
|
|
|||
|
|
@ -158,4 +158,4 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
keywords=keywords))
|
||||
|
||||
# Stores the extracted and structured observations in the conversation memory
|
||||
self.set_memories(self.OBS_STORE_KEY, new_obs_nodes)
|
||||
self.memory_handler.set_memories(self.OBS_STORE_KEY, new_obs_nodes)
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, NOT_UPDATED_NODES, INSIGHT_NODES, TODAY_NODES
|
||||
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memory_scope.enumeration.store_status_enum import StoreStatusEnum
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.scheme.message import Message
|
||||
|
|
@ -20,14 +20,14 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
filter_dict = {
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"store_status": StoreStatusEnum.VALID.value,
|
||||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_reflected": False,
|
||||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_not_reflected_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_memories(NOT_REFLECTED_NODES, nodes)
|
||||
self.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes)
|
||||
|
||||
@timer
|
||||
def retrieve_not_updated_memory(self, query: str):
|
||||
|
|
@ -37,14 +37,14 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
filter_dict = {
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"store_status": StoreStatusEnum.VALID.value,
|
||||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_updated": False,
|
||||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_not_updated_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_memories(NOT_UPDATED_NODES, nodes)
|
||||
self.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
|
||||
|
||||
@timer
|
||||
def retrieve_insight_memory(self, query: str):
|
||||
|
|
@ -54,13 +54,13 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
filter_dict = {
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"store_status": StoreStatusEnum.VALID.value,
|
||||
"memory_type": MemoryTypeEnum.INSIGHT.value,
|
||||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_insight_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_memories(INSIGHT_NODES, nodes)
|
||||
self.memory_handler.set_memories(INSIGHT_NODES, nodes)
|
||||
|
||||
@timer
|
||||
def retrieve_today_memory(self):
|
||||
|
|
@ -76,7 +76,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
filter_dict = {
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"store_status": StoreStatusEnum.VALID.value,
|
||||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"dt": dt_handler.datetime_format(),
|
||||
}
|
||||
|
|
@ -84,11 +84,11 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
top_k=self.today_obs_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
self.set_memories(TODAY_NODES, nodes)
|
||||
self.memory_handler.set_memories(TODAY_NODES, nodes)
|
||||
|
||||
def _run(self):
|
||||
"""
|
||||
Initiates asynchronous tasks to retrieve various types of memory data including
|
||||
Initiates asynchronous tasks to retrieve various types of memory data including
|
||||
not reflected, not updated, insights, and data from today. After submitting all tasks,
|
||||
it waits for their completion by calling `gather_thread_result`.
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from memory_scope.enumeration.action_status_enum import ActionStatusEnum
|
||||
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
|
|
@ -8,16 +8,8 @@ from memory_scope.utils.datetime_handler import DatetimeHandler
|
|||
class StoreMemoryWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
store_key: str = self.store_key
|
||||
|
||||
if store_key == "all":
|
||||
self.update_memories()
|
||||
|
||||
elif self.has_content(store_key):
|
||||
self.update_memories(store_key)
|
||||
|
||||
elif store_key in self.chat_kwargs:
|
||||
query = self.chat_kwargs[store_key]
|
||||
if "query" in self.chat_kwargs:
|
||||
query = self.chat_kwargs["query"]
|
||||
query = query.strip()
|
||||
if not query:
|
||||
return
|
||||
|
|
@ -27,8 +19,8 @@ class StoreMemoryWorker(MemoryBaseWorker):
|
|||
target_name=self.target_name,
|
||||
content=query,
|
||||
memory_type=MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
status=MemoryNodeStatus.NEW.value,
|
||||
timestamp=dt_handler.timestamp,
|
||||
obs_reflected=False,
|
||||
obs_updated=False)
|
||||
self.memory_store.update_memories(node)
|
||||
action_status=ActionStatusEnum.NEW.value,
|
||||
timestamp=dt_handler.timestamp)
|
||||
self.memory_handler.update_memories(nodes=node)
|
||||
else:
|
||||
self.memory_handler.update_memories(self.store_key)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue