From 0edaaa425663a4434bb87abe51753442d5745430 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sun, 14 Jul 2024 15:22:37 +0800 Subject: [PATCH] [dev] rename worker name --- config/demo_config.yaml | 64 +++++++------------ memory_scope/chat/cli_memory_chat.py | 8 +++ .../memory/service/base_memory_service.py | 2 +- .../memory/service/chat_memory_service.py | 4 +- memory_scope/memory/worker/base_worker.py | 6 +- memory_scope/memory/worker/dummy_worker.py | 4 +- .../worker/frontend/extract_time_worker.py | 3 + .../worker/frontend/fuse_rerank_worker.py | 8 +++ .../worker/frontend/retrieve_memory_worker.py | 13 ++-- .../worker/frontend/set_query_worker.py | 16 +++-- .../worker/frontend/update_status_worker.py | 24 ------- .../memory/worker/memory_base_worker.py | 12 ---- .../summary/get_reflection_subject_worker.py | 5 ++ .../summary/long_contra_repeat_worker.py | 5 ++ .../worker/summary/update_insight_worker.py | 7 +- .../worker/write/contra_repeat_worker.py | 5 ++ .../worker/write/get_observation_worker.py | 3 + .../memory/worker/write/info_filter_worker.py | 5 ++ .../memory/worker/write/load_memory_worker.py | 5 ++ .../worker/write/store_memory_worker.py | 26 -------- .../worker/write/update_memory_worker.py | 60 +++++++++++++++++ memory_scope/scheme/message.py | 2 +- memory_scope/utils/memory_handler.py | 45 ++++--------- 23 files changed, 179 insertions(+), 153 deletions(-) delete mode 100644 memory_scope/memory/worker/frontend/update_status_worker.py delete mode 100644 memory_scope/memory/worker/write/store_memory_worker.py create mode 100644 memory_scope/memory/worker/write/update_memory_worker.py diff --git a/config/demo_config.yaml b/config/demo_config.yaml index 6bfbcd54..61480e64 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -20,33 +20,33 @@ memory_service: read_memory: class: memory.operation.frontend_operation - workflow: set_query,[extract_time|retrieve_memory1,semantic_rank],fuse_rerank + workflow: set_query,[extract_time|retrieve_obs_ins,semantic_rank],fuse_rerank description: "read long-term memory" list_memory: class: memory.operation.frontend_operation - workflow: set_query,retrieve_memory2,print_memory + workflow: set_query,retrieve_top_memory,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 + workflow: set_query,retrieve_all_memory,delete_memory description: "delete all long-term memory" add_memory: class: memory.operation.frontend_operation - workflow: store_memory - description: "delete all long-term memory" + workflow: add_memory + description: "add a single observation" write_memory: class: memory.operation.write_memory_op - workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory + workflow: info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat,store_memory description: "write observation memory of the user" interval_time: 5 summary_memory: class: memory.operation.backend_operation - workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory + workflow: load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat,store_memory description: "summary observation memory of the user" interval_time: 30 @@ -60,15 +60,13 @@ worker: class: memory.worker.frontend.read_message_worker set_query: class: memory.worker.frontend.set_query_worker - retrieve_memory1: + retrieve_obs_ins: class: memory.worker.frontend.retrieve_memory_worker retrieve_obs_top_k: 100 - retrieve_ins_pf_top_k: 100 - retrieve_expired_top_k: 0 + retrieve_ins_top_k: 100 extract_time: class: memory.worker.frontend.extract_time_worker generation_model: dashscope_generation - generation_model_top_k: 1 semantic_rank: class: memory.worker.frontend.semantic_rank_worker rank_model: dashscope_rank @@ -82,75 +80,59 @@ worker: insight: 2.0 fuse_time_ratio: 2.0 fuse_rerank_top_k: 10 - retrieve_memory2: + retrieve_top_memory: class: memory.worker.frontend.retrieve_memory_worker retrieve_obs_top_k: 100 - retrieve_ins_pf_top_k: 100 + retrieve_ins_top_k: 100 retrieve_expired_top_k: 100 print_memory: class: memory.worker.frontend.print_memory_worker - retrieve_memory3: + retrieve_all_memory: class: memory.worker.frontend.retrieve_memory_worker retrieve_obs_top_k: 10000 - retrieve_ins_pf_top_k: 10000 + retrieve_ins_top_k: 10000 retrieve_expired_top_k: 10000 delete_memory: - class: memory.worker.frontend.update_status_worker + class: memory.worker.frontend.update_memory_worker + method: modify_action_status expired_action: delete valid_action_dict: obs_customized: delete insight: delete observation: delete + add_memory: + class: memory.worker.frontend.update_memory_worker + method: from_query info_filter: class: memory.worker.write.info_filter_worker generation_model: dashscope_generation - preserved_scores: 2,3 - info_filter_msg_max_size: 200 - generation_model_top_k: 1 - load_memory1: + load_today_memory: class: memory.worker.write.load_memory_worker - retrieve_not_reflected_top_k: 0 - retrieve_not_updated_top_k: 0 - retrieve_insight_top_k: 0 retrieve_today_top_k: 100 get_observation: class: memory.worker.write.get_observation_worker generation_model: dashscope_generation - generation_model_top_k: 1 get_observation_with_time: class: memory.worker.write.get_observation_with_time_worker generation_model: dashscope_generation - generation_model_top_k: 1 contra_repeat: class: memory.worker.write.contra_repeat_worker generation_model: dashscope_generation - generation_model_top_k: 1 - retrieve_top_k: 30 - contra_repeat_max_count: 50 store_memory: - class: memory.worker.write.store_memory_worker - store_key: all - load_memory2: + class: memory.worker.write.update_memory_worker + method: from_memory_key + memory_key: all + load_obs_and_insight: class: memory.worker.write.load_memory_worker retrieve_not_reflected_top_k: 100 retrieve_not_updated_top_k: 100 retrieve_insight_top_k: 100 - retrieve_today_top_k: 0 get_reflection_subject: class: memory.worker.summary.get_reflection_subject_worker - retrieve_top_k: 100 - reflect_obs_cnt_threshold: 10 - generation_model_top_k: 1 update_insight: class: memory.worker.summary.update_insight_worker - update_insight_threshold: 0.1 - generation_model_top_k: 1 - update_insight_max_thread: 10 long_contra_repeat: class: memory.worker.summary.long_contra_repeat_worker - long_contra_repeat_top_k: 2 - long_contra_repeat_threshold: 0.1 - generation_model_top_k: 1 models: dashscope_generation: diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 72073686..76dfa7f5 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -230,6 +230,10 @@ class CliMemoryChat(BaseMemoryChat): continue_run = True command, kwargs = self.parse_query_command(query) + # Print prompt for AI's response + questionary.print("> ", end="", style="fg:yellow") + questionary.print(f"{self.assistant_name}: ", end="", style="bold") + if command == "exit": self.memory_service.stop_backend_service() continue_run = False @@ -257,6 +261,8 @@ class CliMemoryChat(BaseMemoryChat): os.system("clear") self.print_logo() if result: + if isinstance(result, list): + result = "\n".join([str(x) for x in result]) questionary.print(result) else: questionary.print(f"command={command} result is empty! kwargs={kwargs}") @@ -265,6 +271,8 @@ class CliMemoryChat(BaseMemoryChat): else: result = self.memory_service.do_operation(op_name=command, **kwargs) if result: + if isinstance(result, list): + result = "\n".join([str(x) for x in result]) questionary.print(result) else: questionary.print(f"command={command} result is empty! kwargs={kwargs}") diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index f627ad86..18c4d81a 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -96,7 +96,7 @@ class BaseMemoryService(metaclass=ABCMeta): return self.do_operation(self.read_message_key) @abstractmethod - def init_service(self): + def init_service(self, **kwargs): raise NotImplementedError def start_backend_service(self): diff --git a/memory_scope/memory/service/chat_memory_service.py b/memory_scope/memory/service/chat_memory_service.py index 49731f82..1da26091 100644 --- a/memory_scope/memory/service/chat_memory_service.py +++ b/memory_scope/memory/service/chat_memory_service.py @@ -57,7 +57,7 @@ class ChatMemoryService(BaseMemoryService): return return self._operation_dict[op_name].run_operation(**kwargs) # Execute the operation - def init_service(self): + def init_service(self, **kwargs): for name, operation_config in self.memory_operations.items(): if name in self._operation_dict: self.logger.warning(f"memory operation={name} is repeated!") @@ -70,7 +70,7 @@ class ChatMemoryService(BaseMemoryService): 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 + operation.init_workflow(**kwargs) # Initialize workflow for each operation self._operation_dict[name] = operation self.logger.info(f"service={self.__class__.__name__} init operation={name}") diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index 047eb3d5..32bac93f 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -24,13 +24,17 @@ class BaseWorker(metaclass=ABCMeta): self.raise_exception: bool = raise_exception self.is_multi_thread: bool = is_multi_thread self.thread_pool: ThreadPoolExecutor = thread_pool - self.kwargs: dict = kwargs self.continue_run: bool = True self.async_task_list: list = [] self.thread_task_list: list = [] self.logger: Logger = Logger.get_logger() + self._parse_params(**kwargs) + + def _parse_params(self, **kwargs): + pass + def submit_async_task(self, fn, *args, **kwargs): if self.is_multi_thread: raise RuntimeError(f"async_task is not allowed in multi_thread condition") diff --git a/memory_scope/memory/worker/dummy_worker.py b/memory_scope/memory/worker/dummy_worker.py index cabda23d..e4a4b1c8 100644 --- a/memory_scope/memory/worker/dummy_worker.py +++ b/memory_scope/memory/worker/dummy_worker.py @@ -1,10 +1,10 @@ import datetime from memory_scope.constants.common_constants import RESULT, WORKFLOW_NAME, CHAT_KWARGS -from memory_scope.memory.worker.base_worker import BaseWorker +from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker -class DummyWorker(BaseWorker): +class DummyWorker(MemoryBaseWorker): def _run(self): """ Executes the dummy worker's run logic by logging workflow entry, capturing the current timestamp, diff --git a/memory_scope/memory/worker/frontend/extract_time_worker.py b/memory_scope/memory/worker/frontend/extract_time_worker.py index 60ee0f41..2fb0e7fa 100644 --- a/memory_scope/memory/worker/frontend/extract_time_worker.py +++ b/memory_scope/memory/worker/frontend/extract_time_worker.py @@ -18,6 +18,9 @@ class ExtractTimeWorker(MemoryBaseWorker): EXTRACT_TIME_PATTERN = r'-\s*(\S+):(\d+)' FILE_PATH: str = __file__ + def _parse_params(self, **kwargs): + self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1) + def _run(self): """ Executes the primary logic of identifying and extracting time data from an LLM's response. diff --git a/memory_scope/memory/worker/frontend/fuse_rerank_worker.py b/memory_scope/memory/worker/frontend/fuse_rerank_worker.py index 121bc588..8dcaf5a9 100644 --- a/memory_scope/memory/worker/frontend/fuse_rerank_worker.py +++ b/memory_scope/memory/worker/frontend/fuse_rerank_worker.py @@ -8,6 +8,12 @@ from memory_scope.utils.datetime_handler import DatetimeHandler class FuseRerankWorker(MemoryBaseWorker): + def _parse_params(self, **kwargs): + self.fuse_score_threshold: float = kwargs.get("fuse_score_threshold", 0.1) + self.fuse_ratio_dict: Dict[str, float] = kwargs.get("fuse_ratio_dict", {}) + self.fuse_time_ratio: float = kwargs.get("fuse_time_ratio", 2.0) + self.fuse_rerank_top_k: int = kwargs.get("fuse_rerank_top_k", 10) + @staticmethod def match_node_time(extract_time_dict: Dict[str, str], node: MemoryNode): if extract_time_dict: @@ -65,6 +71,8 @@ class FuseRerankWorker(MemoryBaseWorker): continue # Calculate type-based adjustment factor + if node.memory_type not in self.fuse_ratio_dict: + self.logger.warning(f"{node.memory_type} 'factor is not configured!") type_ratio: float = self.fuse_ratio_dict.get(node.memory_type, 0.1) # Determine time relevance adjustment factor diff --git a/memory_scope/memory/worker/frontend/retrieve_memory_worker.py b/memory_scope/memory/worker/frontend/retrieve_memory_worker.py index 02eb3060..47107275 100644 --- a/memory_scope/memory/worker/frontend/retrieve_memory_worker.py +++ b/memory_scope/memory/worker/frontend/retrieve_memory_worker.py @@ -16,6 +16,11 @@ class RetrieveMemoryWorker(MemoryBaseWorker): facilitating efficient memory retrieval operations within a given scope. """ + def _parse_params(self, **kwargs): + self.retrieve_obs_top_k: int = kwargs.get("retrieve_obs_top_k", 0) + self.retrieve_ins_top_k: int = kwargs.get("retrieve_ins_top_k", 0) + self.retrieve_expired_top_k: int = kwargs.get("retrieve_expired_top_k", 0) + @timer def retrieve_from_observation(self, query: str) -> List[MemoryNode]: """ @@ -45,7 +50,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker): filter_dict=filter_dict) @timer - def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]: + def retrieve_from_insight(self, query: str) -> List[MemoryNode]: """ Retrieves memories marked as insights from the database based on a query, filtered by user, target, and set to active status. @@ -58,7 +63,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker): limited by 'retrieve_ins_pf_top_k'. Returns an empty list if 'retrieve_ins_pf_top_k' is not set. """ - if not self.retrieve_ins_pf_top_k: + if not self.retrieve_ins_top_k: return [] filter_dict = { @@ -69,7 +74,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker): } # ⭐ Retrieve insights matching the query, filtered, and limited by top_k return self.memory_store.retrieve_memories(query=query, - top_k=self.retrieve_ins_pf_top_k, + top_k=self.retrieve_ins_top_k, filter_dict=filter_dict) @timer @@ -104,7 +109,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker): """ query, _ = self.get_context(QUERY_WITH_TS) self.submit_thread_task(self.retrieve_from_observation, query=query) - self.submit_thread_task(self.retrieve_from_insight_and_profile, query=query) + self.submit_thread_task(self.retrieve_from_insight, query=query) self.submit_thread_task(self.retrieve_expired_memory, query=query) memory_node_list: List[MemoryNode] = [] diff --git a/memory_scope/memory/worker/frontend/set_query_worker.py b/memory_scope/memory/worker/frontend/set_query_worker.py index 455ba560..80d10937 100644 --- a/memory_scope/memory/worker/frontend/set_query_worker.py +++ b/memory_scope/memory/worker/frontend/set_query_worker.py @@ -24,16 +24,20 @@ class SetQueryWorker(MemoryBaseWorker): query = "_" # Default query value query_timestamp = int(datetime.datetime.now().timestamp()) # Current timestamp as default - # Check if a specific 'query' has been provided via chat kwargs if "query" in self.chat_kwargs: + # Check if a specific 'query' has been provided via chat kwargs query = self.chat_kwargs["query"] + if not query: + query = "" + query = query.strip() - # If no explicit query is given, use the content of the latest chat message elif self.chat_messages: - message = self.chat_messages[-1] - assert message.role == MessageRoleEnum.USER.value - query = message.content - query_timestamp = message.time_created + # If no explicit query is given, use the content of the latest chat message + chat_messages = [msg for msg in self.chat_messages if msg.role == MessageRoleEnum.USER.value] + if chat_messages: + message = chat_messages[-1] + query = message.content + query_timestamp = message.time_created # Store the determined query and its timestamp in the context self.set_context(QUERY_WITH_TS, (query, query_timestamp)) diff --git a/memory_scope/memory/worker/frontend/update_status_worker.py b/memory_scope/memory/worker/frontend/update_status_worker.py deleted file mode 100644 index 0ea5c5d2..00000000 --- a/memory_scope/memory/worker/frontend/update_status_worker.py +++ /dev/null @@ -1,24 +0,0 @@ -from typing import List - -from memory_scope.constants.common_constants import RETRIEVE_MEMORY_NODES -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 - - -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.memory_handler.get_memories(RETRIEVE_MEMORY_NODES) - if not memory_node_list: - return - - for node in memory_node_list: - if node.store_status == StoreStatusEnum.EXPIRED.value: - node.action_status = expired_action - - elif node.memory_type in valid_action_dict: - node.action_status = valid_action_dict[node.memory_type] - - self.memory_handler.update_memories(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 a7d2865d..4c3be5e0 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -169,18 +169,6 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): 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. - - Args: - key (str): The attribute key to look up in kwargs. - - Returns: - Any: The value associated with the key in kwargs. - """ - return self.kwargs[key] - @staticmethod def get_language_value(languages: dict | list[dict]) -> Any | list[Any]: """ 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 9532a05e..59ec9cee 100644 --- a/memory_scope/memory/worker/summary/get_reflection_subject_worker.py +++ b/memory_scope/memory/worker/summary/get_reflection_subject_worker.py @@ -19,6 +19,11 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): """ FILE_PATH: str = __file__ + def _parse_params(self, **kwargs): + self.reflect_obs_cnt_threshold: int = kwargs.get("reflect_obs_cnt_threshold", 10) + self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1) + self.reflect_num_questions: int = kwargs.get("reflect_num_questions", 5) + def new_insight_node(self, insight_key: str) -> MemoryNode: """ Creates a new MemoryNode for an insight with the given key, enriched with current datetime metadata. 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 ede7f284..919834bd 100644 --- a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py +++ b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py @@ -20,6 +20,11 @@ class LongContraRepeatWorker(MemoryBaseWorker): """ FILE_PATH: str = __file__ + def _parse_params(self, **kwargs): + self.long_contra_repeat_top_k: int = kwargs.get("long_contra_repeat_top_k", 2) + self.long_contra_repeat_threshold: float = kwargs.get("long_contra_repeat_threshold", 0.1) + self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1) + 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/memory_type. diff --git a/memory_scope/memory/worker/summary/update_insight_worker.py b/memory_scope/memory/worker/summary/update_insight_worker.py index 919a28dc..6c758582 100644 --- a/memory_scope/memory/worker/summary/update_insight_worker.py +++ b/memory_scope/memory/worker/summary/update_insight_worker.py @@ -19,6 +19,11 @@ class UpdateInsightWorker(MemoryBaseWorker): """ FILE_PATH: str = __file__ + def _parse_params(self, **kwargs): + self.update_insight_threshold: float = kwargs.get("update_insight_threshold", 0.1) + self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1) + self.update_insight_max_count: int = kwargs.get("update_insight_max_count", 10) + def filter_obs_nodes(self, insight_node: MemoryNode, obs_nodes: List[MemoryNode]) -> (MemoryNode, List[MemoryNode], float): @@ -178,7 +183,7 @@ class UpdateInsightWorker(MemoryBaseWorker): if not filtered_nodes: continue result_list.append(result) - result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)[: self.update_insight_max_thread] + result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)[: self.update_insight_max_count] # Submit tasks to update insights for the top nodes for insight_node, filtered_nodes, _ in result_sorted: diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py index 478834e6..f216557b 100644 --- a/memory_scope/memory/worker/write/contra_repeat_worker.py +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -23,6 +23,11 @@ class ContraRepeatWorker(MemoryBaseWorker): """ FILE_PATH: str = __file__ + def _parse_params(self, **kwargs): + self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1) + self.retrieve_top_k: int = kwargs.get("retrieve_top_k", 30) + self.contra_repeat_max_count: int = kwargs.get("contra_repeat_max_count", 50) + def _run(self): """ Executes the primary routine of the ContraRepeatWorker which involves fetching memory nodes, diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index 2a4ed53f..f04738c9 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -16,6 +16,9 @@ class GetObservationWorker(MemoryBaseWorker): FILE_PATH: str = __file__ OBS_STORE_KEY: str = NEW_OBS_NODES + def _parse_params(self, **kwargs): + self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1) + def add_observation(self, message: Message, time_infer: str, obs_content: str, keywords: str): dt_handler = DatetimeHandler(dt=message.time_created) diff --git a/memory_scope/memory/worker/write/info_filter_worker.py b/memory_scope/memory/worker/write/info_filter_worker.py index 51d64b93..8ea21c5d 100644 --- a/memory_scope/memory/worker/write/info_filter_worker.py +++ b/memory_scope/memory/worker/write/info_filter_worker.py @@ -17,6 +17,11 @@ class InfoFilterWorker(MemoryBaseWorker): """ FILE_PATH: str = __file__ + def _parse_params(self, **kwargs): + self.preserved_scores: str = kwargs.get("preserved_scores", "2,3") + self.info_filter_msg_max_size: int = kwargs.get("info_filter_msg_max_size", 200) + self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1) + def _run(self): """ Filters user messages in the chat, generates a prompt incorporating these messages, diff --git a/memory_scope/memory/worker/write/load_memory_worker.py b/memory_scope/memory/worker/write/load_memory_worker.py index 14d97f66..14796396 100644 --- a/memory_scope/memory/worker/write/load_memory_worker.py +++ b/memory_scope/memory/worker/write/load_memory_worker.py @@ -10,6 +10,11 @@ from memory_scope.utils.timer import timer class LoadMemoryWorker(MemoryBaseWorker): + def _parse_params(self, **kwargs): + self.retrieve_not_reflected_top_k: int = kwargs.get("retrieve_not_reflected_top_k", 0) + self.retrieve_not_updated_top_k: int = kwargs.get("retrieve_not_updated_top_k", 0) + self.retrieve_insight_top_k: int = kwargs.get("retrieve_insight_top_k", 0) + self.retrieve_today_top_k: int = kwargs.get("retrieve_today_top_k", 0) @timer def retrieve_not_reflected_memory(self, query: str): diff --git a/memory_scope/memory/worker/write/store_memory_worker.py b/memory_scope/memory/worker/write/store_memory_worker.py deleted file mode 100644 index 02e0357e..00000000 --- a/memory_scope/memory/worker/write/store_memory_worker.py +++ /dev/null @@ -1,26 +0,0 @@ -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 -from memory_scope.utils.datetime_handler import DatetimeHandler - - -class StoreMemoryWorker(MemoryBaseWorker): - - def _run(self): - if "query" in self.chat_kwargs: - query = self.chat_kwargs["query"] - query = query.strip() - if not query: - return - - dt_handler = DatetimeHandler() - node = MemoryNode(user_name=self.user_name, - target_name=self.target_name, - content=query, - memory_type=MemoryTypeEnum.OBS_CUSTOMIZED.value, - action_status=ActionStatusEnum.NEW.value, - timestamp=dt_handler.timestamp) - self.memory_handler.update_memories(nodes=node) - else: - self.memory_handler.update_memories(key=self.store_key) diff --git a/memory_scope/memory/worker/write/update_memory_worker.py b/memory_scope/memory/worker/write/update_memory_worker.py new file mode 100644 index 00000000..a4ae6e0b --- /dev/null +++ b/memory_scope/memory/worker/write/update_memory_worker.py @@ -0,0 +1,60 @@ +from typing import Dict, List + +from memory_scope.enumeration.action_status_enum import ActionStatusEnum +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 + + +class UpdateMemoryWorker(MemoryBaseWorker): + + def _parse_params(self, **kwargs): + self.method: str = kwargs.get("method", "") + self.memory_key: str = kwargs.get("memory_key", "") + self.expired_action: str = kwargs.get("expired_action", "") + self.valid_action_dict: Dict[str, str] = kwargs.get("expired_action", {}) + + def from_query(self): + if "query" not in self.chat_kwargs: + return + + query = self.chat_kwargs["query"].strip() + if not query: + return + + dt_handler = DatetimeHandler() + node = MemoryNode(user_name=self.user_name, + target_name=self.target_name, + content=query, + memory_type=MemoryTypeEnum.OBS_CUSTOMIZED.value, + action_status=ActionStatusEnum.NEW.value, + timestamp=dt_handler.timestamp) + return [node] + + def from_memory_key(self): + if not self.memory_key: + return + + return self.memory_handler.get_memories(keys=self.memory_key) + + def modify_action_status(self): + nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all") + for node in nodes: + if self.expired_action and node.store_status == StoreStatusEnum.EXPIRED.value: + node.action_status = self.expired_action + + elif node.memory_type in self.valid_action_dict: + action_status = self.valid_action_dict[node.memory_type] + if action_status: + node.action_status = action_status + + return nodes + + def _run(self): + method = self.method.strip() + if not hasattr(self, method): + self.logger.info(f"method={method} is missing!") + return + self.memory_handler.update_memories(nodes=getattr(self, method)()) diff --git a/memory_scope/scheme/message.py b/memory_scope/scheme/message.py index e6196d4e..5575f4c1 100644 --- a/memory_scope/scheme/message.py +++ b/memory_scope/scheme/message.py @@ -22,7 +22,7 @@ class Message(BaseModel): content: str = Field(..., description="The primary content of the message") - time_created: int = Field(int(datetime.datetime.now().timestamp()), + time_created: int = Field(default_factory=lambda: int(datetime.datetime.now().timestamp()), description="Timestamp marking the message creation time") memorized: bool = Field(False, description="Indicates if the message is flagged for memory retention") diff --git a/memory_scope/utils/memory_handler.py b/memory_scope/utils/memory_handler.py index 780ee1e3..223f017b 100644 --- a/memory_scope/utils/memory_handler.py +++ b/memory_scope/utils/memory_handler.py @@ -1,4 +1,4 @@ -from typing import Dict, List, Set +from typing import Dict, List from memory_scope.enumeration.action_status_enum import ActionStatusEnum from memory_scope.enumeration.store_status_enum import StoreStatusEnum @@ -57,53 +57,34 @@ class MemoryHandler(object): self._key_id_dict[key] = [n.memory_id for n in nodes] def get_memories(self, keys: str | List[str]) -> List[MemoryNode]: - """ - Retrieves memory nodes associated with the given keys. - - This method accepts a single key or a list of keys. For each key, it fetches the - associated memory IDs from the context. If memory IDs are found, they are used to - collect the corresponding MemoryNode objects from the contex_memory_dict. The final - result is a list of unique MemoryNode values, avoiding duplicates. - - Args: - keys (str | List[str]): The key or list of keys to retrieve memories for. - - Returns: - List[MemoryNode]: A list of MemoryNode objects associated with the input keys. - """ memories: Dict[str, MemoryNode] = {} + if isinstance(keys, str): keys = [keys] for key in keys: - if key not in self._key_id_dict: + if key == "all": + memories.update(self._id_memory_dict) + break + elif key not in self._key_id_dict: continue - memory_ids: List[str] = self._key_id_dict[key] + + memory_ids: List[str] = self._key_id_dict.get(key.strip()) if memory_ids: memories.update({x: self._id_memory_dict[x] for x in memory_ids}) + return list(memories.values()) - def update_memories(self, key: str = "", nodes: MemoryNode | List[MemoryNode] = None): - ids: Set[str] = set() - - 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] + def update_memories(self, keys: str = "", nodes: MemoryNode | List[MemoryNode] = None): + update_memories: Dict[str, MemoryNode] = {n.memory_id: n for n in self.get_memories(keys=keys)} if nodes is not None: if isinstance(nodes, MemoryNode): nodes = [nodes] - update_nodes.extend(nodes) + update_memories.update({n.memory_id: n for n in nodes}) # Save collected nodes to memory store - self._update_memories(update_nodes) + self._update_memories(list(update_memories.values())) def _update_memories(self, nodes: List[MemoryNode]): """