[dev] add add_memory op

This commit is contained in:
jinli.yl 2024-07-13 19:45:35 +08:00
parent 4bb001a810
commit 44bf421487
14 changed files with 47 additions and 45 deletions

View file

@ -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

View file

@ -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,

View file

@ -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:

View file

@ -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 -----"]

View file

@ -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.

View file

@ -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

View file

@ -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

View file

@ -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)

View file

@ -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)

View file

@ -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.")

View file

@ -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)

View file

@ -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)

View file

@ -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`.

View file

@ -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)