From c1957af6af16713a7a8c0b788dd1a5b882cce315 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sat, 13 Jul 2024 23:02:29 +0800 Subject: [PATCH] [dev] modify op name & service base func --- config/demo_config.yaml | 4 +- memory_scope/chat/cli_memory_chat.py | 19 +++-- .../memory/operation/base_operation.py | 4 +- .../memory/operation/base_workflow.py | 3 +- .../memory/operation/write_memory_op.py | 3 + .../memory/service/base_memory_service.py | 41 ++-------- .../memory/service/chat_memory_service.py | 79 ++++++++----------- memory_scope/memory/worker/base_worker.py | 4 + .../worker/frontend/print_memory_worker.py | 4 +- .../worker/frontend/retrieve_memory_worker.py | 8 +- .../memory/worker/memory_base_worker.py | 17 +++- .../summary/long_contra_repeat_worker.py | 5 +- .../worker/write/contra_repeat_worker.py | 2 +- .../memory/worker/write/info_filter_worker.py | 1 - .../memory/worker/write/load_memory_worker.py | 32 ++++---- .../worker/write/store_memory_worker.py | 2 +- memory_scope/utils/memory_handler.py | 36 ++++----- 17 files changed, 122 insertions(+), 142 deletions(-) diff --git a/config/demo_config.yaml b/config/demo_config.yaml index 8a5e2589..6bfbcd54 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -39,13 +39,13 @@ memory_service: description: "delete all long-term memory" write_memory: - class: memory.operation.write_memory + class: memory.operation.write_memory_op 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 + class: memory.operation.backend_operation workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory description: "summary observation memory of the user" interval_time: 30 diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 0fe1aab6..72073686 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -110,7 +110,8 @@ class CliMemoryChat(BaseMemoryChat): if self._memory_service not in G_CONTEXT.memory_service_dict: raise ValueError("Missing declaration of memory_service in yaml configuration: " + self._memory_service) self._memory_service = G_CONTEXT.memory_service_dict[self._memory_service] - self._memory_service.start_service() # ⭐ Initialize and start the memory service + self._memory_service.init_service() + self._memory_service.start_backend_service() return self._memory_service @property @@ -230,7 +231,7 @@ class CliMemoryChat(BaseMemoryChat): command, kwargs = self.parse_query_command(query) if command == "exit": - self.memory_service.stop_service() + self.memory_service.stop_backend_service() continue_run = False elif command == "clear": @@ -250,16 +251,23 @@ class CliMemoryChat(BaseMemoryChat): refresh_time = kwargs.pop("refresh_time", "") if refresh_time and refresh_time.isdigit(): refresh_time = int(refresh_time) + self.memory_service.stop_backend_service() while True: result = self.memory_service.do_operation(op_name=command, **kwargs) os.system("clear") self.print_logo() - questionary.print(result) + if result: + questionary.print(result) + else: + questionary.print(f"command={command} result is empty! kwargs={kwargs}") time.sleep(refresh_time) else: result = self.memory_service.do_operation(op_name=command, **kwargs) - questionary.print(result) + if result: + questionary.print(result) + else: + questionary.print(f"command={command} result is empty! kwargs={kwargs}") else: questionary.print(f"Unknown command={command} received.") @@ -297,6 +305,7 @@ class CliMemoryChat(BaseMemoryChat): questionary.print(f"{self.assistant_name}: ", end="", style="bold") # Fetch and display AI's response, with support for streaming + self.memory_service.start_backend_service() if self.stream: model_response = None for model_response in self.chat_with_memory(query=query): @@ -316,7 +325,7 @@ class CliMemoryChat(BaseMemoryChat): questionary.print("User interrupt occurred.") is_exit = questionary.confirm("Continue exit?").unsafe_ask() if is_exit: - self.memory_service.stop_service() + self.memory_service.stop_backend_service() break except Exception as e: diff --git a/memory_scope/memory/operation/base_operation.py b/memory_scope/memory/operation/base_operation.py index 5c7df589..8bce8c57 100644 --- a/memory_scope/memory/operation/base_operation.py +++ b/memory_scope/memory/operation/base_operation.py @@ -17,18 +17,16 @@ class BaseOperation(metaclass=ABCMeta): operation_type: OPERATION_TYPE = "frontend" - def __init__(self, name: str, description: str = "", **kwargs): + def __init__(self, name: str, description: str = ""): """ Initializes a new instance of the BaseOperation. Args: name (str): The name identifying the operation. description (str): An optional description detailing the operation's purpose or behavior. - **kwargs: Arbitrary keyword arguments for custom settings or parameters. """ self.name: str = name self.description: str = description - self.kwargs: dict = kwargs def init_workflow(self, **kwargs): """ diff --git a/memory_scope/memory/operation/base_workflow.py b/memory_scope/memory/operation/base_workflow.py index beb4b5a2..5ee5b765 100644 --- a/memory_scope/memory/operation/base_workflow.py +++ b/memory_scope/memory/operation/base_workflow.py @@ -146,6 +146,7 @@ class BaseWorkflow(object): worker = self.worker_dict[name] worker.run() if not worker.continue_run: + self.logger.warning(f"worker={worker.name} stop workflow!") return False return True @@ -159,7 +160,7 @@ class BaseWorkflow(object): they are submitted for parallel execution using a thread pool. The workflow will stop if any sub-workflow returns False. """ - with Timer(self.name, time_log_type="wrap"): + with Timer(f"workflow.{self.name}", time_log_type="wrap"): self.context[WORKFLOW_NAME] = self.name # Iterate over each part of the workflow diff --git a/memory_scope/memory/operation/write_memory_op.py b/memory_scope/memory/operation/write_memory_op.py index e4334266..8d481841 100644 --- a/memory_scope/memory/operation/write_memory_op.py +++ b/memory_scope/memory/operation/write_memory_op.py @@ -30,6 +30,9 @@ class WriteMemoryOp(BackendOperation): Any: The result obtained from running the workflow. """ + if not self.chat_messages: + return + # Use shallow copy to prevent adding new messages. chat_messages = self.chat_messages.copy() diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index c37c46c2..f627ad86 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -41,38 +41,10 @@ class BaseMemoryService(metaclass=ABCMeta): self.logger = Logger.get_logger() self.kwargs = kwargs - @abstractmethod - def _init_operation(self, memory_operations: Dict[str, dict]): - """ - Initializes the memory operations with a given dictionary of operations. - - This method is to be implemented by subclasses to set up or configure - the memory operations based on the provided dictionary. - - Args: - memory_operations (Dict[str, dict]): A dictionary containing configuration - details for each memory operation. - - Raises: - NotImplementedError: This exception is raised to indicate that the method - needs to be overridden in the subclass. - """ - raise NotImplementedError - @abstractmethod def add_messages(self, messages: List[Message] | Message): raise NotImplementedError - def start_service(self, **kwargs): - """ - This method is intended to initiate the service with provided keyword arguments, - preparing the necessary resources for executing operations. - - Args: - **kwargs: Additional keyword arguments used to configure the service upon startup. - """ - pass - @abstractmethod def do_operation(self, op_name: str, **kwargs): """ @@ -123,9 +95,12 @@ class BaseMemoryService(metaclass=ABCMeta): assert self.read_message_key in self._operation_dict, f"op={self.read_message_key} is not inited!" return self.do_operation(self.read_message_key) - def stop_service(self): - """ - Placeholder method to stop the service. - Intended to be overridden by subclasses to define specific shutdown logic. - """ + @abstractmethod + def init_service(self): + raise NotImplementedError + + def start_backend_service(self): + pass + + def stop_backend_service(self): pass diff --git a/memory_scope/memory/service/chat_memory_service.py b/memory_scope/memory/service/chat_memory_service.py index 2bd96758..49731f82 100644 --- a/memory_scope/memory/service/chat_memory_service.py +++ b/memory_scope/memory/service/chat_memory_service.py @@ -1,5 +1,6 @@ -from typing import List, Dict +from typing import List +from memory_scope.memory.operation.base_operation import BaseOperation from memory_scope.memory.service.base_memory_service import BaseMemoryService from memory_scope.scheme.message import Message from memory_scope.utils.tool_functions import init_instance_by_config @@ -12,35 +13,6 @@ class ChatMemoryService(BaseMemoryService): self.contextual_msg_count: int = contextual_msg_count assert self.history_msg_count >= self.contextual_msg_count - self._init_operation(self.memory_operations) - - def _init_operation(self, memory_operations: Dict[str, dict]): - """ - Initializes memory operations based on the provided configuration dictionary. - Ensures that each operation is only initialized once by checking for duplicates. - - Args: - memory_operations (Dict[str, dict]): A dictionary where keys are operation names - and values are configuration dictionaries for each operation. - - Note: - Logs a warning if an attempt is made to initialize an operation with a name that already exists. - Logs an info message upon successful initialization of each operation. - """ - for name, operation_config in memory_operations.items(): - if name in self._operation_dict: - self.logger.warning(f"memory operation={name} is repeated!") - continue - - # ⭐ Initialize operation instance by its config - self._operation_dict[name] = init_instance_by_config( - config=operation_config, - name=name, - chat_messages=self.chat_messages, - message_lock=self.message_lock, - contextual_msg_count=self.contextual_msg_count) - self.logger.info(f"service={self.__class__.__name__} init operation={name}") - def add_messages(self, messages: List[Message] | Message): """ Adds a single message or a list of messages to the chat history, ensuring the message list @@ -56,28 +28,16 @@ class ChatMemoryService(BaseMemoryService): # Sort the messages by their creation time to maintain chronological order messages = sorted(messages, key=lambda x: x.time_created) - + # Append the sorted messages to the chat history self.chat_messages.extend(messages) - + # If the chat history exceeds the allowed message count, remove the oldest messages if len(self.chat_messages) > self.history_msg_count: gap_size = len(self.chat_messages) - self.history_msg_count for _ in range(gap_size): self.chat_messages.pop(0) - def start_service(self, **kwargs): - """ - Initializes and starts backend operations defined in `_operation_dict`. - - Args: - **kwargs: Additional keyword arguments passed to each operation's initialization. - """ - for _, operation in self._operation_dict.items(): - operation.init_workflow(**kwargs) # Initialize workflow for each operation - if operation.operation_type == "backend": - operation.run_operation_backend() # Run backend operations - def do_operation(self, op_name: str, **kwargs): """ Executes a specific operation by its name with provided keyword arguments. @@ -97,10 +57,37 @@ class ChatMemoryService(BaseMemoryService): return return self._operation_dict[op_name].run_operation(**kwargs) # Execute the operation - def stop_service(self): + def init_service(self): + for name, operation_config in self.memory_operations.items(): + if name in self._operation_dict: + self.logger.warning(f"memory operation={name} is repeated!") + continue + + # ⭐ Initialize operation instance by its config + operation: BaseOperation = init_instance_by_config( + config=operation_config, + name=name, + chat_messages=self.chat_messages, + message_lock=self.message_lock, + contextual_msg_count=self.contextual_msg_count) + operation.init_workflow() # Initialize workflow for each operation + + self._operation_dict[name] = operation + self.logger.info(f"service={self.__class__.__name__} init operation={name}") + + def start_backend_service(self): + for _, operation in self._operation_dict.items(): + if operation.operation_type == "backend": + # Run backend operations + operation.run_operation_backend() + self.logger.info(f"start operation={operation.name}...") + + def stop_backend_service(self): """ Stops all backend operations that are currently running. """ for _, operation in self._operation_dict.items(): if operation.operation_type == "backend": - operation.stop_operation_backend() # Stop backend operations + # Stop backend operations + operation.stop_operation_backend() + self.logger.info(f"stop operation={operation.name}...") diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index 5a11b6db..047eb3d5 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -62,6 +62,10 @@ class BaseWorker(metaclass=ABCMeta): def run(self): with Timer(f"worker.{self.name}", time_log_type="wrap"): + self.continue_run = True + self.async_task_list.clear() + self.thread_task_list.clear() + if self.raise_exception: self._run() else: diff --git a/memory_scope/memory/worker/frontend/print_memory_worker.py b/memory_scope/memory/worker/frontend/print_memory_worker.py index 1d470bbb..d53e01ac 100644 --- a/memory_scope/memory/worker/frontend/print_memory_worker.py +++ b/memory_scope/memory/worker/frontend/print_memory_worker.py @@ -1,8 +1,8 @@ from typing import List from memory_scope.constants.common_constants import RETRIEVE_MEMORY_NODES, RESULT -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.utils.datetime_handler import DatetimeHandler @@ -25,7 +25,7 @@ class PrintMemoryWorker(MemoryBaseWorker): dt_handler = DatetimeHandler(node.timestamp) dt = dt_handler.datetime_format("%Y%m%d-%H:%M:%S") line = f"{dt} {node.content}" - if MemoryNodeStatus(node.status) is MemoryNodeStatus.EXPIRED: + if StoreStatusEnum(node.store_status) is StoreStatusEnum.EXPIRED: if node.content in expired_content_set: continue else: diff --git a/memory_scope/memory/worker/frontend/retrieve_memory_worker.py b/memory_scope/memory/worker/frontend/retrieve_memory_worker.py index 838f0fd2..02eb3060 100644 --- a/memory_scope/memory/worker/frontend/retrieve_memory_worker.py +++ b/memory_scope/memory/worker/frontend/retrieve_memory_worker.py @@ -114,12 +114,12 @@ class RetrieveMemoryWorker(MemoryBaseWorker): self.logger.info(f"memory_node_list.size={len(memory_node_list)}") if not memory_node_list: - self.continue_run = False return memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True) for node in memory_node_list: node.action_status = ActionStatusEnum.NONE - self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} " - f"type={node.memory_type} status={node.status}") - self.memory_handler.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list) + self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} type={node.memory_type} " + f"store_status={node.store_status} action_status={node.action_status}") + + self.memory_handler.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list) diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py index 361681d2..a7d2865d 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -5,6 +5,7 @@ from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, from memory_scope.memory.worker.base_worker import BaseWorker from memory_scope.models.base_model import BaseModel from memory_scope.scheme.message import Message +from memory_scope.storage.base_memory_store import BaseMemoryStore from memory_scope.storage.base_monitor import BaseMonitor from memory_scope.utils.global_context import G_CONTEXT from memory_scope.utils.memory_handler import MemoryHandler @@ -36,6 +37,8 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): self._embedding_model: BaseModel | str = embedding_model self._generation_model: BaseModel | str = generation_model self._rank_model: BaseModel | str = rank_model + + self._memory_store: BaseMemoryStore | None = None self._monitor: BaseMonitor | None = None self._user_name: str | None = None @@ -105,10 +108,10 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): return self._rank_model @property - def memory_handler(self) -> MemoryHandler: - if not self.has_content(MEMORY_HANDLER): - self.set_context(MEMORY_HANDLER, MemoryHandler()) # Initialize the memory handler if not present - return self.get_context(MEMORY_HANDLER) + def memory_store(self) -> BaseMemoryStore: + if self._memory_store is None: + self._memory_store = G_CONTEXT.memory_store + return self._memory_store @property def monitor(self) -> BaseMonitor: @@ -160,6 +163,12 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): self._prompt_handler = PromptHandler(self.FILE_PATH, **self.kwargs) return self._prompt_handler + @property + def memory_handler(self) -> MemoryHandler: + if not self.has_content(MEMORY_HANDLER): + self.set_context(MEMORY_HANDLER, MemoryHandler()) # Initialize the memory handler if not present + return self.get_context(MEMORY_HANDLER) + def __getattr__(self, key: str): """ Custom attribute access to directly retrieve values from kwargs. 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 91979d76..ede7f284 100644 --- a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py +++ b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py @@ -22,7 +22,7 @@ class LongContraRepeatWorker(MemoryBaseWorker): def retrieve_similar_content(self, node: MemoryNode) -> (MemoryNode, List[MemoryNode]): """ - Retrieves memory nodes with content similar to the given node, filtering by user, target, status, and memory type. + Retrieves memory nodes with content similar to the given node, filtering by user/target/status/memory_type. Only returns nodes whose similarity score meets or exceeds the predefined threshold. Args: @@ -139,7 +139,8 @@ class LongContraRepeatWorker(MemoryBaseWorker): node.store_status = StoreStatusEnum.EXPIRED.value merge_obs_nodes.append(node) - self.logger.info(f"after_long_contra_repeat: {node.content} {node.status}") + self.logger.info(f"after_long_contra_repeat: {node.content} store_status={node.store_status} " + f"action_status={node.action_status}") # save context self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes) diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py index d441d927..478834e6 100644 --- a/memory_scope/memory/worker/write/contra_repeat_worker.py +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -74,7 +74,7 @@ class ContraRepeatWorker(MemoryBaseWorker): # parse text idx_merge_obs_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__) if len(idx_merge_obs_list) <= 0: - self.add_run_info("idx_merge_obs_list is empty!") + self.logger.warning("idx_merge_obs_list is empty!") return # add merged obs diff --git a/memory_scope/memory/worker/write/info_filter_worker.py b/memory_scope/memory/worker/write/info_filter_worker.py index 81b90a85..51d64b93 100644 --- a/memory_scope/memory/worker/write/info_filter_worker.py +++ b/memory_scope/memory/worker/write/info_filter_worker.py @@ -76,7 +76,6 @@ class InfoFilterWorker(MemoryBaseWorker): info_score_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__) if len(info_score_list) != len(info_messages): self.logger.warning(f"score_size != messages_size, {len(info_score_list)} vs {len(info_messages)}") - self.continue_run = False return # filter messages diff --git a/memory_scope/memory/worker/write/load_memory_worker.py b/memory_scope/memory/worker/write/load_memory_worker.py index 51425fd3..14d97f66 100644 --- a/memory_scope/memory/worker/write/load_memory_worker.py +++ b/memory_scope/memory/worker/write/load_memory_worker.py @@ -5,7 +5,6 @@ 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 from memory_scope.utils.datetime_handler import DatetimeHandler from memory_scope.utils.timer import timer @@ -63,24 +62,18 @@ class LoadMemoryWorker(MemoryBaseWorker): self.memory_handler.set_memories(INSIGHT_NODES, nodes) @timer - def retrieve_today_memory(self): - if not self.today_obs_top_k: + def retrieve_today_memory(self, query: str, dt: str): + if not self.retrieve_today_top_k: return - if not self.chat_messages: - self.logger.warning("chat_messages is empty!") - return - - message: Message = self.chat_messages[-1] - dt_handler = DatetimeHandler(message.time_created) filter_dict = { "user_name": self.user_name, "target_name": self.target_name, "store_status": StoreStatusEnum.VALID.value, "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], - "dt": dt_handler.datetime_format(), + "dt": dt, } - nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=message.content, + nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query, top_k=self.retrieve_today_top_k, filter_dict=filter_dict) @@ -95,9 +88,14 @@ class LoadMemoryWorker(MemoryBaseWorker): This method serves as the controller for data retrieval operations, enhancing efficiency by handling tasks concurrently. """ - mock_query = "-" # Placeholder query - self.submit_thread_task(self.retrieve_not_reflected_memory, query=mock_query) - self.submit_thread_task(self.retrieve_not_updated_memory, query=mock_query) - self.submit_thread_task(self.retrieve_insight_memory, query=mock_query) - self.submit_thread_task(self.retrieve_today_memory) - self.gather_thread_result() # Waits for all submitted tasks to complete + + # Placeholder query + query = "-" + dt = DatetimeHandler().datetime_format() + self.submit_thread_task(self.retrieve_not_reflected_memory, query=query) + self.submit_thread_task(self.retrieve_not_updated_memory, query=query) + self.submit_thread_task(self.retrieve_insight_memory, query=query) + self.submit_thread_task(self.retrieve_today_memory, query=query, dt=dt) + + # Waits for all submitted tasks to complete + self.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 fc8f7edf..02e0357e 100644 --- a/memory_scope/memory/worker/write/store_memory_worker.py +++ b/memory_scope/memory/worker/write/store_memory_worker.py @@ -23,4 +23,4 @@ class StoreMemoryWorker(MemoryBaseWorker): timestamp=dt_handler.timestamp) self.memory_handler.update_memories(nodes=node) else: - self.memory_handler.update_memories(self.store_key) + self.memory_handler.update_memories(key=self.store_key) diff --git a/memory_scope/utils/memory_handler.py b/memory_scope/utils/memory_handler.py index 0a4fca3f..780ee1e3 100644 --- a/memory_scope/utils/memory_handler.py +++ b/memory_scope/utils/memory_handler.py @@ -37,7 +37,7 @@ class MemoryHandler(object): self._id_memory_dict.clear() self._key_id_dict.clear() - def add_memories(self, nodes: MemoryNode | List[MemoryNode], log_repeat: bool = True): + def set_memories(self, key: str, nodes: MemoryNode | List[MemoryNode], log_repeat: bool = True): if nodes is None: nodes = [] elif isinstance(nodes, MemoryNode): @@ -47,15 +47,13 @@ class MemoryHandler(object): if node.memory_id in self._id_memory_dict: if log_repeat: self.logger.warning(f"repeated_id memory id={node.memory_id} content={node.content} " - f"status={node.status}") + f"store_status={node.store_status} action_status={node.action_status}") continue self._id_memory_dict[node.memory_id] = node self.logger.info(f"add to memory context memory id={node.memory_id} content={node.content} " - f"status={node.status}") + f"store_status={node.store_status} action_status={node.action_status}") - def set_memories(self, key: str, nodes: MemoryNode | List[MemoryNode], log_repeat: bool = True): - self.add_memories(nodes=nodes, log_repeat=log_repeat) self._key_id_dict[key] = [n.memory_id for n in nodes] def get_memories(self, keys: str | List[str]) -> List[MemoryNode]: @@ -78,25 +76,23 @@ class MemoryHandler(object): keys = [keys] for key in keys: + if key not in self._key_id_dict: + continue memory_ids: List[str] = self._key_id_dict[key] if memory_ids: memories.update({x: self._id_memory_dict[x] for x in memory_ids}) return list(memories.values()) - def update_memories(self, keys: str | List[str] = None, nodes: MemoryNode | List[MemoryNode] = None): - if keys is None: - keys = [] - elif keys == "all": - keys = list(self._id_memory_dict.keys()) - elif isinstance(keys, str): - keys = [keys] - - # combine keys + def update_memories(self, key: str = "", nodes: MemoryNode | List[MemoryNode] = None): ids: Set[str] = set() - for key in keys: - t_ids: List[str] = self._key_id_dict[key] - if t_ids: - ids.update(t_ids) + + if key == "all": + ids.update(self._id_memory_dict.keys()) + elif key: + for k in key.split(","): + t_ids: List[str] = self._key_id_dict.get(k.strip()) + if t_ids: + ids.update(t_ids) # Remove and collect nodes by IDs update_nodes: List[MemoryNode] = [self._id_memory_dict.pop(_) for _ in ids] @@ -148,9 +144,9 @@ class MemoryHandler(object): if modified_memories: for n in modified_memories: n.action_status = ActionStatusEnum.NONE.value - self.memory_store.batch_update(c_modified_memories, update_embedding=False) + self.memory_store.batch_update(modified_memories, update_embedding=False) # set memories expired delete_memories = [n for n in nodes if n.action_status == ActionStatusEnum.DELETE] if delete_memories: - self.memory_store.batch_delete(nodes) + self.memory_store.batch_delete(delete_memories)