From 44bf4214877108eace2d4d2971223f0641c0c7ed Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sat, 13 Jul 2024 19:45:35 +0800 Subject: [PATCH] [dev] add add_memory op --- config/demo_config.yaml | 12 +++++++++- .../memory/operation/write_memory_op.py | 4 ++-- .../worker/frontend/fuse_rerank_worker.py | 2 +- .../worker/frontend/print_memory_worker.py | 2 +- .../worker/frontend/retrieve_memory_worker.py | 2 +- .../worker/frontend/semantic_rank_worker.py | 2 +- .../worker/frontend/update_status_worker.py | 2 +- .../summary/get_reflection_subject_worker.py | 4 ++-- .../summary/long_contra_repeat_worker.py | 4 ++-- .../worker/summary/update_insight_worker.py | 6 ++--- .../worker/write/contra_repeat_worker.py | 6 ++--- .../worker/write/get_observation_worker.py | 2 +- .../memory/worker/write/load_memory_worker.py | 20 ++++++++-------- .../worker/write/store_memory_worker.py | 24 +++++++------------ 14 files changed, 47 insertions(+), 45 deletions(-) diff --git a/config/demo_config.yaml b/config/demo_config.yaml index ebf2dcf5..f3edf11d 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -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 diff --git a/memory_scope/memory/operation/write_memory_op.py b/memory_scope/memory/operation/write_memory_op.py index e5e9d1bf..e4334266 100644 --- a/memory_scope/memory/operation/write_memory_op.py +++ b/memory_scope/memory/operation/write_memory_op.py @@ -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, diff --git a/memory_scope/memory/worker/frontend/fuse_rerank_worker.py b/memory_scope/memory/worker/frontend/fuse_rerank_worker.py index 6e8d80ed..121bc588 100644 --- a/memory_scope/memory/worker/frontend/fuse_rerank_worker.py +++ b/memory_scope/memory/worker/frontend/fuse_rerank_worker.py @@ -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: diff --git a/memory_scope/memory/worker/frontend/print_memory_worker.py b/memory_scope/memory/worker/frontend/print_memory_worker.py index 286752bd..1d470bbb 100644 --- a/memory_scope/memory/worker/frontend/print_memory_worker.py +++ b/memory_scope/memory/worker/frontend/print_memory_worker.py @@ -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 -----"] diff --git a/memory_scope/memory/worker/frontend/retrieve_memory_worker.py b/memory_scope/memory/worker/frontend/retrieve_memory_worker.py index 6cd232d5..838f0fd2 100644 --- a/memory_scope/memory/worker/frontend/retrieve_memory_worker.py +++ b/memory_scope/memory/worker/frontend/retrieve_memory_worker.py @@ -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. diff --git a/memory_scope/memory/worker/frontend/semantic_rank_worker.py b/memory_scope/memory/worker/frontend/semantic_rank_worker.py index 596618f6..a7e4dcfd 100644 --- a/memory_scope/memory/worker/frontend/semantic_rank_worker.py +++ b/memory_scope/memory/worker/frontend/semantic_rank_worker.py @@ -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 diff --git a/memory_scope/memory/worker/frontend/update_status_worker.py b/memory_scope/memory/worker/frontend/update_status_worker.py index eabdf6cf..0ea5c5d2 100644 --- a/memory_scope/memory/worker/frontend/update_status_worker.py +++ b/memory_scope/memory/worker/frontend/update_status_worker.py @@ -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 diff --git a/memory_scope/memory/worker/summary/get_reflection_subject_worker.py b/memory_scope/memory/worker/summary/get_reflection_subject_worker.py index 4854ea31..304c8ab5 100644 --- a/memory_scope/memory/worker/summary/get_reflection_subject_worker.py +++ b/memory_scope/memory/worker/summary/get_reflection_subject_worker.py @@ -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) diff --git a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py index 26a370d4..59cfefb9 100644 --- a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py +++ b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py @@ -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) diff --git a/memory_scope/memory/worker/summary/update_insight_worker.py b/memory_scope/memory/worker/summary/update_insight_worker.py index ad7ef1e8..69f7e7cf 100644 --- a/memory_scope/memory/worker/summary/update_insight_worker.py +++ b/memory_scope/memory/worker/summary/update_insight_worker.py @@ -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.") diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py index 94d81601..e6354f4e 100644 --- a/memory_scope/memory/worker/write/contra_repeat_worker.py +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -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) diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index 4360e9d4..40334d95 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -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) diff --git a/memory_scope/memory/worker/write/load_memory_worker.py b/memory_scope/memory/worker/write/load_memory_worker.py index 5704677b..690a6deb 100644 --- a/memory_scope/memory/worker/write/load_memory_worker.py +++ b/memory_scope/memory/worker/write/load_memory_worker.py @@ -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`. diff --git a/memory_scope/memory/worker/write/store_memory_worker.py b/memory_scope/memory/worker/write/store_memory_worker.py index ea73a70f..fc8f7edf 100644 --- a/memory_scope/memory/worker/write/store_memory_worker.py +++ b/memory_scope/memory/worker/write/store_memory_worker.py @@ -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)