From 590cbf65c4fc4f9c100e1345e21f655bf361489f Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sat, 27 Jul 2024 01:49:20 +0800 Subject: [PATCH] [dev] rename MEMORY_HANDLER to MEMORY_MANAGER --- memoryscope/constants/common_constants.py | 4 +- .../memory/operation/backend_operation.py | 22 +++----- memoryscope/memory/operation/base_workflow.py | 21 +++++--- ...rvation_op.py => consolidate_operation.py} | 19 +++---- .../memory/operation/frontend_operation.py | 14 +++-- .../memory/service/memory_scope_service.py | 2 +- .../worker/backend/contra_repeat_worker.py | 6 +-- .../worker/backend/get_observation_worker.py | 8 +-- .../backend/get_reflection_subject_worker.py | 8 +-- .../worker/backend/load_memory_worker.py | 8 +-- .../backend/long_contra_repeat_worker.py | 4 +- .../worker/backend/update_insight_worker.py | 10 ++-- .../worker/backend/update_memory_worker.py | 10 ++-- .../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 | 4 +- .../memory/worker/memory_base_worker.py | 51 +++++++++++-------- .../worker/memory_manager.py} | 13 +++-- memoryscope/models/dummy_generation_model.py | 7 +-- memoryscope/utils/datetime_handler.py | 49 ++++++++++-------- tests/worker/test_workers_cn.py | 36 ++++++------- tests/worker/test_workers_en.py | 36 ++++++------- 23 files changed, 177 insertions(+), 161 deletions(-) rename memoryscope/memory/operation/{summary_observation_op.py => consolidate_operation.py} (83%) rename memoryscope/{utils/memory_handler.py => memory/worker/memory_manager.py} (96%) diff --git a/memoryscope/constants/common_constants.py b/memoryscope/constants/common_constants.py index adf9e548..4eba1d1d 100644 --- a/memoryscope/constants/common_constants.py +++ b/memoryscope/constants/common_constants.py @@ -5,11 +5,13 @@ WORKFLOW_NAME = "workflow_name" +MEMORYSCOPE_CONTEXT = "memoryscope_context" + RESULT = "result" CHAT_MESSAGES = "chat_messages" -MEMORY_HANDLER = "memory_handler" +MEMORY_MANAGER = "memory_manager" CHAT_KWARGS = "chat_kwargs" diff --git a/memoryscope/memory/operation/backend_operation.py b/memoryscope/memory/operation/backend_operation.py index bd5d3ac9..4a53a6b4 100644 --- a/memoryscope/memory/operation/backend_operation.py +++ b/memoryscope/memory/operation/backend_operation.py @@ -1,7 +1,6 @@ import time from typing import List - from memoryscope.constants.common_constants import CHAT_KWARGS, RESULT, CHAT_MESSAGES from memoryscope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE from memoryscope.memory.operation.base_workflow import BaseWorkflow @@ -54,17 +53,14 @@ class BackendOperation(BaseWorkflow, BaseOperation): Returns: Any: The result obtained after executing the workflow. """ - self.context.clear() - - # Add additional arguments to the context - kwargs.update(**self.kwargs) - self.context[CHAT_KWARGS] = kwargs - - # Include the most recent messages in the operation context - self.context[CHAT_MESSAGES] = self.chat_messages + # prepare kwargs + workflow_kwargs = { + CHAT_MESSAGES: self.chat_messages, + CHAT_KWARGS: {**kwargs, **self.kwargs}, + } # Execute the workflow with the prepared context - self.run_workflow() + self.run_workflow(**workflow_kwargs) # Retrieve the result from the context after workflow execution return self.context.get(RESULT) @@ -114,8 +110,8 @@ class BackendOperation(BaseWorkflow, BaseOperation): """ if not self._loop_switch: self._loop_switch = True - self._backend_task = G_CONTEXT.thread_pool.submit(self._loop_operation) - self.logger.info(f"start operation={operation.name}...") + self._backend_task = self.thread_pool.submit(self._loop_operation) + self.logger.info(f"start operation={self.name}...") def stop_operation_backend(self, wait_task_end: bool = False): """ @@ -128,5 +124,3 @@ class BackendOperation(BaseWorkflow, BaseOperation): self.logger.info(f"stop operation={self.name}...") else: self.logger.info(f"send stop signal to operation={self.name}...") - - diff --git a/memoryscope/memory/operation/base_workflow.py b/memoryscope/memory/operation/base_workflow.py index 1ec0a5f5..7cb75daf 100644 --- a/memoryscope/memory/operation/base_workflow.py +++ b/memoryscope/memory/operation/base_workflow.py @@ -4,7 +4,7 @@ from concurrent.futures import ThreadPoolExecutor, as_completed from itertools import zip_longest from typing import Dict, Any, List -from memoryscope.constants.common_constants import WORKFLOW_NAME +from memoryscope.constants.common_constants import WORKFLOW_NAME, MEMORYSCOPE_CONTEXT from memoryscope.memory.worker.base_worker import BaseWorker from memoryscope.memoryscope_context import MemoryscopeContext from memoryscope.utils.logger import Logger @@ -22,6 +22,7 @@ class BaseWorkflow(object): self.name: str = name self.memoryscope_context: MemoryscopeContext = memoryscope_context + self.thread_pool: ThreadPoolExecutor = self.memoryscope_context.thread_pool self.workflow: str = workflow self.kwargs = kwargs @@ -133,12 +134,11 @@ class BaseWorkflow(object): self.worker_dict[name] = init_instance_by_config( config=self.memoryscope_context.worker_conf_dict[name], - suffix_name="worker", name=name, is_multi_thread=is_backend or self.worker_dict[name], context=self.context, context_lock=self.context_lock, - thread_pool=self.memoryscope_context.thread_pool, + thread_pool=self.thread_pool, **kwargs) def _run_sub_workflow(self, worker_list: List[str]) -> bool: @@ -150,7 +150,7 @@ class BaseWorkflow(object): return False return True - def run_workflow(self): + def run_workflow(self, **kwargs): """ Executes the workflow by orchestrating the steps defined in `self.workflow_worker_list`. This method supports both sequential and parallel execution of sub-workflows based on the structure @@ -159,9 +159,18 @@ class BaseWorkflow(object): If a workflow part consists of a single item, it is executed sequentially. For parts with multiple items, they are submitted for parallel execution using a thread pool. The workflow will stop if any sub-workflow returns False. + + Args: + **kwargs: Additional keyword arguments to be passed to context. """ with Timer(f"workflow.{self.name}", time_log_type="wrap"): - self.context[WORKFLOW_NAME] = self.name + self.context.clear() + + self.context.update({ + WORKFLOW_NAME: self.name, + MEMORYSCOPE_CONTEXT: self.memoryscope_context, + **kwargs, + }) # Iterate over each part of the workflow for workflow_part in self.workflow_worker_list: @@ -174,7 +183,7 @@ class BaseWorkflow(object): t_list = [] # Submit tasks to the thread pool for sub_workflow in workflow_part: - t_list.append(G_CONTEXT.thread_pool.submit(self._run_sub_workflow, sub_workflow)) + t_list.append(self.thread_pool.submit(self._run_sub_workflow, sub_workflow)) # Check results; if any task returns False, stop the workflow flag = True diff --git a/memoryscope/memory/operation/summary_observation_op.py b/memoryscope/memory/operation/consolidate_operation.py similarity index 83% rename from memoryscope/memory/operation/summary_observation_op.py rename to memoryscope/memory/operation/consolidate_operation.py index 1da2e65c..967c3639 100644 --- a/memoryscope/memory/operation/summary_observation_op.py +++ b/memoryscope/memory/operation/consolidate_operation.py @@ -3,10 +3,10 @@ from memoryscope.enumeration.message_role_enum import MessageRoleEnum from memoryscope.memory.operation.backend_operation import BackendOperation -class SummaryObservationOp(BackendOperation): +class ConsolidateOperation(BackendOperation): def __init__(self, **kwargs): - super(SummaryObservationOp, self).__init__(**kwargs) + super(ConsolidateOperation, self).__init__(**kwargs) self.message_lock = kwargs.get("message_lock", None) self.contextual_msg_min_count: int = kwargs.get("contextual_msg_min_count", 0) @@ -43,17 +43,14 @@ class SummaryObservationOp(BackendOperation): f"contextual_msg_min_count({self.contextual_msg_min_count}), skip.") return - self.context.clear() - - # Add additional arguments to the context - kwargs.update(**self.kwargs) - self.context[CHAT_KWARGS] = kwargs - - # Include the most recent messages in the operation context - self.context[CHAT_MESSAGES] = chat_messages + # prepare kwargs + workflow_kwargs = { + CHAT_MESSAGES: chat_messages, + CHAT_KWARGS: {**kwargs, **self.kwargs}, + } # Execute the workflow with the prepared context - self.run_workflow() + self.run_workflow(**workflow_kwargs) # Retrieve the result from the context after workflow execution result = self.context.get(RESULT) diff --git a/memoryscope/memory/operation/frontend_operation.py b/memoryscope/memory/operation/frontend_operation.py index 8cde7977..abaca44a 100644 --- a/memoryscope/memory/operation/frontend_operation.py +++ b/memoryscope/memory/operation/frontend_operation.py @@ -39,17 +39,15 @@ class FrontendOperation(BaseWorkflow, BaseOperation): Returns: Any: The result obtained from executing the workflow. """ - self.context.clear() - # Include the most recent messages in the operation context - self.context[CHAT_MESSAGES] = self.chat_messages - - # Add additional arguments to the context - kwargs.update(**self.kwargs) - self.context[CHAT_KWARGS] = kwargs + # prepare kwargs + workflow_kwargs = { + CHAT_MESSAGES: self.chat_messages, + CHAT_KWARGS: {**kwargs, **self.kwargs}, + } # Execute the workflow with the prepared context - self.run_workflow() + self.run_workflow(**workflow_kwargs) # Retrieve the result from the context after workflow execution return self.context.get(RESULT) diff --git a/memoryscope/memory/service/memory_scope_service.py b/memoryscope/memory/service/memory_scope_service.py index eb9d14e8..bf1c9b94 100644 --- a/memoryscope/memory/service/memory_scope_service.py +++ b/memoryscope/memory/service/memory_scope_service.py @@ -76,7 +76,7 @@ class MemoryScopeService(BaseMemoryService): name=name, chat_messages=self.chat_messages, message_lock=self.message_lock, - context=self.context, + memoryscope_context=self.context, contextual_msg_max_count=self.contextual_msg_max_count, contextual_msg_min_count=self.contextual_msg_min_count) diff --git a/memoryscope/memory/worker/backend/contra_repeat_worker.py b/memoryscope/memory/worker/backend/contra_repeat_worker.py index 5146887f..85a18884 100644 --- a/memoryscope/memory/worker/backend/contra_repeat_worker.py +++ b/memoryscope/memory/worker/backend/contra_repeat_worker.py @@ -43,13 +43,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.memory_handler.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES]) + all_obs_nodes: List[MemoryNode] = self.memory_manager.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.memory_handler.get_memories(TODAY_NODES) + today_obs_nodes: List[MemoryNode] = self.memory_manager.get_memories(TODAY_NODES) if today_obs_nodes: all_obs_nodes.extend(today_obs_nodes) @@ -121,4 +121,4 @@ class ContraRepeatWorker(MemoryBaseWorker): merge_obs_nodes.append(node) # save context - self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes, log_repeat=False) + self.memory_manager.set_memories(MERGE_OBS_NODES, merge_obs_nodes, log_repeat=False) diff --git a/memoryscope/memory/worker/backend/get_observation_worker.py b/memoryscope/memory/worker/backend/get_observation_worker.py index b36d2b6f..385251d1 100644 --- a/memoryscope/memory/worker/backend/get_observation_worker.py +++ b/memoryscope/memory/worker/backend/get_observation_worker.py @@ -42,11 +42,11 @@ class GetObservationWorker(MemoryBaseWorker): MemoryTypeEnum.CONVERSATION.value: message.content, TIME_INFER: time_infer, "keywords": keywords, - **{k: str(v) for k, v in dt_handler.dt_info_dict.items()}, + **{k: str(v) for k, v in dt_handler.get_dt_info_dict(self.language).items()}, } if time_infer: - dt_info_dict = DatetimeHandler.extract_date_parts(input_string=time_infer) + dt_info_dict = DatetimeHandler.extract_date_parts(input_string=time_infer, language=self.language) meta_data.update({f"event_{k}": str(v) for k, v in dt_info_dict.items()}) obs_content = (f"{obs_content} ({self.get_language_value(TIME_INFER_WORD)}" f"{self.get_language_value(COLON_WORD)} {time_infer})") @@ -68,7 +68,7 @@ class GetObservationWorker(MemoryBaseWorker): """ filter_messages = [] for msg in self.chat_messages: - if not DatetimeHandler.has_time_word(query=msg.content): + if not DatetimeHandler.has_time_word(query=msg.content, language=self.language): filter_messages.append(msg) self.logger.info(f"after filter_messages.size from {len(self.chat_messages)} to {len(filter_messages)}") @@ -184,4 +184,4 @@ class GetObservationWorker(MemoryBaseWorker): keywords=keywords)) # Stores the extracted and structured observations in the conversation memory - self.memory_handler.set_memories(self.OBS_STORE_KEY, new_obs_nodes) + self.memory_manager.set_memories(self.OBS_STORE_KEY, new_obs_nodes) diff --git a/memoryscope/memory/worker/backend/get_reflection_subject_worker.py b/memoryscope/memory/worker/backend/get_reflection_subject_worker.py index 4c9d7b27..711206ae 100644 --- a/memoryscope/memory/worker/backend/get_reflection_subject_worker.py +++ b/memoryscope/memory/worker/backend/get_reflection_subject_worker.py @@ -36,7 +36,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): """ dt_handler = DatetimeHandler() # Prepare metadata with current datetime info - meta_data = {k: str(v) for k, v in dt_handler.dt_info_dict.items()} + meta_data = {k: str(v) for k, v in dt_handler.get_dt_info_dict.items()} return MemoryNode(user_name=self.user_name, target_name=self.target_name, @@ -58,8 +58,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.memory_handler.get_memories(NOT_REFLECTED_NODES) - insight_nodes: List[MemoryNode] = self.memory_handler.get_memories(INSIGHT_NODES) + not_reflected_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_REFLECTED_NODES) + insight_nodes: List[MemoryNode] = self.memory_manager.get_memories(INSIGHT_NODES) # Count unaudited nodes not_reflected_count = len(not_reflected_nodes) @@ -104,7 +104,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): new_insight_keys = ResponseTextParser(response.message.content, self.__class__.__name__).parse_v2() if new_insight_keys: for insight_key in new_insight_keys: - self.memory_handler.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key)) + self.memory_manager.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key)) # Mark unaudited nodes as reflected for node in not_reflected_nodes: diff --git a/memoryscope/memory/worker/backend/load_memory_worker.py b/memoryscope/memory/worker/backend/load_memory_worker.py index 1f5e4233..22c77792 100644 --- a/memoryscope/memory/worker/backend/load_memory_worker.py +++ b/memoryscope/memory/worker/backend/load_memory_worker.py @@ -33,7 +33,7 @@ class LoadMemoryWorker(MemoryBaseWorker): } nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_not_reflected_top_k, filter_dict=filter_dict) - self.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes) + self.memory_manager.set_memories(NOT_REFLECTED_NODES, nodes) @timer def retrieve_not_updated_memory(self): @@ -52,7 +52,7 @@ class LoadMemoryWorker(MemoryBaseWorker): } nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_not_updated_top_k, filter_dict=filter_dict) - self.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + self.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) @timer def retrieve_insight_memory(self): @@ -70,7 +70,7 @@ class LoadMemoryWorker(MemoryBaseWorker): } nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_insight_top_k, filter_dict=filter_dict) - self.memory_handler.set_memories(INSIGHT_NODES, nodes) + self.memory_manager.set_memories(INSIGHT_NODES, nodes) @timer def retrieve_today_memory(self, dt: str): @@ -93,7 +93,7 @@ class LoadMemoryWorker(MemoryBaseWorker): nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_today_top_k, filter_dict=filter_dict) - self.memory_handler.set_memories(TODAY_NODES, nodes) + self.memory_manager.set_memories(TODAY_NODES, nodes) def _run(self): """ diff --git a/memoryscope/memory/worker/backend/long_contra_repeat_worker.py b/memoryscope/memory/worker/backend/long_contra_repeat_worker.py index 94b306b4..056c63d7 100644 --- a/memoryscope/memory/worker/backend/long_contra_repeat_worker.py +++ b/memoryscope/memory/worker/backend/long_contra_repeat_worker.py @@ -63,7 +63,7 @@ class LongContraRepeatWorker(MemoryBaseWorker): The process helps in maintaining conversation coherence by resolving contradictions and redundancies. """ - not_updated_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_UPDATED_NODES) + not_updated_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_UPDATED_NODES) for node in not_updated_nodes: self.submit_thread_task(fn=self.retrieve_similar_content, node=node) @@ -157,4 +157,4 @@ class LongContraRepeatWorker(MemoryBaseWorker): f"action_status={node.action_status}") # save context - self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes) + self.memory_manager.set_memories(MERGE_OBS_NODES, merge_obs_nodes) diff --git a/memoryscope/memory/worker/backend/update_insight_worker.py b/memoryscope/memory/worker/backend/update_insight_worker.py index b56ed6f7..339d7871 100644 --- a/memoryscope/memory/worker/backend/update_insight_worker.py +++ b/memoryscope/memory/worker/backend/update_insight_worker.py @@ -95,7 +95,7 @@ class UpdateInsightWorker(MemoryBaseWorker): content = f"{key}{self.get_language_value(COLON_WORD)} {insight_value}" insight_node.content = content insight_node.value = insight_value - insight_node.meta_data.update({k: str(v) for k, v in dt_handler.dt_info_dict.items()}) + insight_node.meta_data.update({k: str(v) for k, v in dt_handler.get_dt_info_dict.items()}) insight_node.timestamp = dt_handler.timestamp insight_node.dt = dt_handler.datetime_format() if insight_node.action_status == ActionStatusEnum.NONE.value: @@ -175,9 +175,9 @@ class UpdateInsightWorker(MemoryBaseWorker): 6. Gather the results of all update tasks. 7. Mark processed nodes as updated in memory. """ - 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(keys=[NOT_REFLECTED_NODES, + insight_nodes: List[MemoryNode] = self.memory_manager.get_memories(INSIGHT_NODES) + not_updated_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_UPDATED_NODES) + not_reflected_nodes: List[MemoryNode] = self.memory_manager.get_memories(keys=[NOT_REFLECTED_NODES, NOT_UPDATED_NODES]) if not insight_nodes: @@ -216,7 +216,7 @@ class UpdateInsightWorker(MemoryBaseWorker): # delete empty nodes empty_nodes = [n for n in insight_nodes if not n.content.strip()] - self.memory_handler.delete_memories(empty_nodes) + self.memory_manager.delete_memories(empty_nodes) for node in not_updated_nodes: node.obs_updated = 1 diff --git a/memoryscope/memory/worker/backend/update_memory_worker.py b/memoryscope/memory/worker/backend/update_memory_worker.py index 87b182a7..9133de68 100644 --- a/memoryscope/memory/worker/backend/update_memory_worker.py +++ b/memoryscope/memory/worker/backend/update_memory_worker.py @@ -46,7 +46,7 @@ class UpdateMemoryWorker(MemoryBaseWorker): if not self.memory_key: return - return self.memory_handler.get_memories(keys=self.memory_key) + return self.memory_manager.get_memories(keys=self.memory_key) def delete_all(self): """ @@ -55,7 +55,7 @@ class UpdateMemoryWorker(MemoryBaseWorker): Returns: List[MemoryNode]: A list of all MemoryNode objects marked for deletion. """ - nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all") + nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all") for node in nodes: node.action_status = ActionStatusEnum.DELETE.value self.logger.info(f"delete_all.size={len(nodes)}") @@ -74,7 +74,7 @@ class UpdateMemoryWorker(MemoryBaseWorker): return i = 0 - nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all") + nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all") for node in nodes: if node.content == query: i += 1 @@ -88,7 +88,7 @@ class UpdateMemoryWorker(MemoryBaseWorker): return i = 0 - nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all") + nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all") for node in nodes: if node.memory_id == memory_id: i += 1 @@ -109,4 +109,4 @@ class UpdateMemoryWorker(MemoryBaseWorker): if not hasattr(self, method): self.logger.info(f"method={method} is missing!") return - self.memory_handler.update_memories(nodes=getattr(self, method)()) + self.memory_manager.update_memories(nodes=getattr(self, method)()) diff --git a/memoryscope/memory/worker/frontend/fuse_rerank_worker.py b/memoryscope/memory/worker/frontend/fuse_rerank_worker.py index 373d82b0..3d394790 100644 --- a/memoryscope/memory/worker/frontend/fuse_rerank_worker.py +++ b/memoryscope/memory/worker/frontend/fuse_rerank_worker.py @@ -62,7 +62,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.memory_handler.get_memories(RANKED_MEMORY_NODES) + memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RANKED_MEMORY_NODES) # Check if memory nodes are available; warn and return if not if not memory_node_list: diff --git a/memoryscope/memory/worker/frontend/print_memory_worker.py b/memoryscope/memory/worker/frontend/print_memory_worker.py index 6617c49c..e50adad9 100644 --- a/memoryscope/memory/worker/frontend/print_memory_worker.py +++ b/memoryscope/memory/worker/frontend/print_memory_worker.py @@ -22,7 +22,7 @@ class PrintMemoryWorker(MemoryBaseWorker): 3. Set the formatted string back into the worker's context """ # get long-term memory - memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RETRIEVE_MEMORY_NODES) + memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RETRIEVE_MEMORY_NODES) memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True) observation_memory_list: List[str] = [] diff --git a/memoryscope/memory/worker/frontend/retrieve_memory_worker.py b/memoryscope/memory/worker/frontend/retrieve_memory_worker.py index fd928f7d..e51bbd2d 100644 --- a/memoryscope/memory/worker/frontend/retrieve_memory_worker.py +++ b/memoryscope/memory/worker/frontend/retrieve_memory_worker.py @@ -139,4 +139,4 @@ class RetrieveMemoryWorker(MemoryBaseWorker): self.logger.info(f"recall_stage: content={node.content} score={node.score_rerank} 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) + self.memory_manager.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list) diff --git a/memoryscope/memory/worker/frontend/semantic_rank_worker.py b/memoryscope/memory/worker/frontend/semantic_rank_worker.py index be5d67dc..9a78b96c 100644 --- a/memoryscope/memory/worker/frontend/semantic_rank_worker.py +++ b/memoryscope/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.memory_handler.get_memories(RETRIEVE_MEMORY_NODES) + memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RETRIEVE_MEMORY_NODES) if not memory_node_list: self.logger.warning("Retrieve memory nodes is empty!") return @@ -58,4 +58,4 @@ class SemanticRankWorker(MemoryBaseWorker): self.logger.info(f"Rank stage: Content={node.content}, Score={node.score_rank}") # save ranked nodes back to memory - self.memory_handler.set_memories(RANKED_MEMORY_NODES, memory_node_list, log_repeat=False) + self.memory_manager.set_memories(RANKED_MEMORY_NODES, memory_node_list, log_repeat=False) diff --git a/memoryscope/memory/worker/memory_base_worker.py b/memoryscope/memory/worker/memory_base_worker.py index 00fb72be..e800f8b3 100644 --- a/memoryscope/memory/worker/memory_base_worker.py +++ b/memoryscope/memory/worker/memory_base_worker.py @@ -1,14 +1,16 @@ from abc import ABCMeta from typing import List, Dict, Any -from memoryscope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, MEMORY_HANDLER +from memoryscope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, MEMORYSCOPE_CONTEXT, \ + WORKFLOW_NAME, MEMORY_MANAGER +from memoryscope.enumeration.language_enum import LanguageEnum from memoryscope.memory.worker.base_worker import BaseWorker +from memoryscope.memory.worker.memory_manager import MemoryManager +from memoryscope.memoryscope_context import MemoryscopeContext from memoryscope.models.base_model import BaseModel from memoryscope.scheme.message import Message from memoryscope.storage.base_memory_store import BaseMemoryStore from memoryscope.storage.base_monitor import BaseMonitor -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.memory_handler import MemoryHandler from memoryscope.utils.prompt_handler import PromptHandler @@ -77,6 +79,18 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): """ return self.get_context(CHAT_KWARGS) + @property + def workflow_name(self) -> str: + return self.get_context(WORKFLOW_NAME) + + @property + def memoryscope_context(self) -> MemoryscopeContext: + return self.get_context(MEMORYSCOPE_CONTEXT) + + @property + def language(self) -> LanguageEnum: + return self.memoryscope_context.language + @property def embedding_model(self) -> BaseModel: """ @@ -87,8 +101,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseModel: The embedding model used for converting text into vector representations. """ if isinstance(self._embedding_model, str): - self._embedding_model = G_CONTEXT.model_conf_dict[self._embedding_model] - # ⭐ Retrieve the actual model instance when the attribute is a string reference + self._embedding_model = self.memoryscope_context.model_dict[self._embedding_model] return self._embedding_model @property @@ -101,8 +114,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseModel: The model used for text generation. """ if isinstance(self._generation_model, str): - self._generation_model = G_CONTEXT.model_conf_dict[self._generation_model] - # ⭐ Retrieve the model instance if currently a string reference + self._generation_model = self.memoryscope_context.model_dict[self._generation_model] return self._generation_model @property @@ -115,7 +127,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseModel: The rank model instance used for ranking tasks. """ if isinstance(self._rank_model, str): - self._rank_model = G_CONTEXT.model_conf_dict[self._rank_model] # Fetch model instance if string reference + self._rank_model = self.memoryscope_context.model_dict[self._rank_model] return self._rank_model @property @@ -128,7 +140,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseMemoryStore: The memory store instance used for inserting, updating, retrieving and deleting operations. """ if self._memory_store is None: - self._memory_store = G_CONTEXT.memory_store_conf + self._memory_store = self.memoryscope_context.memory_store return self._memory_store @property @@ -141,7 +153,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseMonitor: The monitoring component instance. """ if self._monitor is None: - self._monitor = G_CONTEXT.monitor_conf + self._monitor = self.memoryscope_context.monitor return self._monitor @property @@ -154,7 +166,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): str: The name of the assistant. """ if self._user_name is None: - self._user_name = G_CONTEXT.meta_data["assistant_name"] + self._user_name = self.memoryscope_context.meta_data["assistant_name"] return self._user_name @property @@ -166,7 +178,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): str: The readable name of the human. """ if self._target_name is None: - self._target_name = G_CONTEXT.meta_data["human_name"] + self._target_name = self.memoryscope_context.meta_data["human_name"] return self._target_name @property @@ -182,19 +194,18 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): return self._prompt_handler @property - def memory_handler(self) -> MemoryHandler: + def memory_manager(self) -> MemoryManager: """ Lazily initializes and returns the MemoryHandler instance. Returns: MemoryHandler: An instance of 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) + if not self.has_content(MEMORY_MANAGER): + self.set_context(MEMORY_MANAGER, MemoryManager(self.memoryscope_context)) + return self.get_context(MEMORY_MANAGER) - @staticmethod - def get_language_value(languages: dict | List[dict]) -> Any | List[Any]: + def get_language_value(self, languages: dict | List[dict]) -> Any | List[Any]: """ Retrieves the value(s) corresponding to the current language context. @@ -205,5 +216,5 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): Any | list[Any]: The value or list of values matching the current language setting. """ if isinstance(languages, list): - return [x[G_CONTEXT.language] for x in languages] - return languages[G_CONTEXT.language] + return [x[self.language] for x in languages] + return languages[self.language] diff --git a/memoryscope/utils/memory_handler.py b/memoryscope/memory/worker/memory_manager.py similarity index 96% rename from memoryscope/utils/memory_handler.py rename to memoryscope/memory/worker/memory_manager.py index a3cbf889..f705d832 100644 --- a/memoryscope/utils/memory_handler.py +++ b/memoryscope/memory/worker/memory_manager.py @@ -2,21 +2,20 @@ from typing import Dict, List from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.store_status_enum import StoreStatusEnum +from memoryscope.memoryscope_context import MemoryscopeContext from memoryscope.scheme.memory_node import MemoryNode from memoryscope.storage.base_memory_store import BaseMemoryStore -from memoryscope.utils.global_context import G_CONTEXT from memoryscope.utils.logger import Logger -class MemoryHandler(object): +class MemoryManager(object): """ The `MemoryHandler` class manages memory nodes with memory store. """ - def __init__(self): - """ - Initializes the MemoryHandler. - """ + def __init__(self, memoryscope_context: MemoryscopeContext): + self.memoryscope_context: MemoryscopeContext = memoryscope_context + self._memory_store: BaseMemoryStore | None = None # dict: memory_id -> MemoryNode @@ -36,7 +35,7 @@ class MemoryHandler(object): BaseMemoryStore: The memory store instance associated with this worker. """ if self._memory_store is None: - self._memory_store = G_CONTEXT.memory_store_conf + self._memory_store = self.memoryscope_context.memory_store return self._memory_store def clear(self): diff --git a/memoryscope/models/dummy_generation_model.py b/memoryscope/models/dummy_generation_model.py index 5949bf8e..ee6eead6 100644 --- a/memoryscope/models/dummy_generation_model.py +++ b/memoryscope/models/dummy_generation_model.py @@ -20,9 +20,6 @@ class DummyGenerationModel(BaseModel): m_type: ModelEnum = ModelEnum.GENERATION_MODEL class DummyModel: - """ - An inner class representing the dummy model placeholder. - """ pass MODEL_REGISTRY.register("dummy_generation", DummyModel) @@ -79,12 +76,12 @@ class DummyGenerationModel(BaseModel): for delta in call_result: model_response.message.content += delta model_response.delta = delta - time.sleep(0.1) # ⭐ Introduce a delay to simulate streaming + time.sleep(0.1) yield model_response return gen() else: - model_response.message.content = "".join(call_result) # ⭐ Concatenate results for non-streaming + model_response.message.content = "".join(call_result) return model_response def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: diff --git a/memoryscope/utils/datetime_handler.py b/memoryscope/utils/datetime_handler.py index 6f2d6d2f..44a0b3ad 100644 --- a/memoryscope/utils/datetime_handler.py +++ b/memoryscope/utils/datetime_handler.py @@ -1,8 +1,9 @@ import datetime import re +from typing import List from memoryscope.constants.language_constants import WEEKDAYS, DATATIME_WORD_LIST, MONTH_DICT -from memoryscope.utils.global_context import G_CONTEXT +from memoryscope.enumeration.language_enum import LanguageEnum from memoryscope.utils.logger import Logger @@ -40,7 +41,7 @@ class DatetimeHandler(object): self._dt_info_dict: dict | None = None - def _parse_dt_info(self): + def _parse_dt_info(self, language: LanguageEnum): """ Parses the datetime object (_dt) into a dictionary containing detailed date and time components, including language-specific weekday representation. @@ -52,17 +53,16 @@ class DatetimeHandler(object): """ return { "year": self._dt.year, - "month": MONTH_DICT[G_CONTEXT.language][self._dt.month - 1], + "month": MONTH_DICT[language][self._dt.month - 1], "day": self._dt.day, "hour": self._dt.hour, "minute": self._dt.minute, "second": self._dt.second, "week": self._dt.isocalendar().week, - "weekday": WEEKDAYS[G_CONTEXT.language][self._dt.isocalendar().weekday - 1], + "weekday": WEEKDAYS[language][self._dt.isocalendar().weekday - 1], } - @property - def dt_info_dict(self): + def get_dt_info_dict(self, language: LanguageEnum): """ Property method to get the dictionary containing parsed datetime information. If None, initialize using `_parse_dt_info`. @@ -71,7 +71,7 @@ class DatetimeHandler(object): dict: A dictionary with parsed datetime information. """ if self._dt_info_dict is None: - self._dt_info_dict = self._parse_dt_info() + self._dt_info_dict = self._parse_dt_info(language=language) return self._dt_info_dict @classmethod @@ -207,7 +207,7 @@ class DatetimeHandler(object): return date_info @classmethod - def extract_date_parts(cls, input_string: str) -> dict: + def extract_date_parts(cls, input_string: str, language: LanguageEnum) -> dict: """ Extracts various date components from the input string based on the current language context. @@ -217,48 +217,51 @@ class DatetimeHandler(object): Args: input_string (str): The string containing date information to be parsed. + language (str): current language. Returns: dict: A dictionary containing extracted date components, or an empty dictionary if parsing fails. """ - func_name = f"extract_date_parts_{G_CONTEXT.language.value}" + func_name = f"extract_date_parts_{language}" if not hasattr(cls, func_name): - cls.logger.warning(f"language={G_CONTEXT.language.value} needs to complete extract_date_parts func!") + cls.logger.warning(f"language={language} needs to complete extract_date_parts func!") return {} return getattr(cls, func_name)(input_string=input_string) @classmethod - def has_time_word_cn(cls, query: str) -> bool: + def has_time_word_cn(cls, query: str, datetime_word_list: List[str]) -> bool: """ Check if the input query contains any datetime-related words based on the cn language context. Args: query (str): The input string to check for datetime-related words. + datetime_word_list (list[str]): datetime keywords Returns: bool: True if the query contains at least one datetime-related word, False otherwise. """ contain_datetime = False # TODO use re - for datetime_word in DATATIME_WORD_LIST[G_CONTEXT.language]: + for datetime_word in datetime_word_list: if datetime_word in query: contain_datetime = True break return contain_datetime @classmethod - def has_time_word_en(cls, query: str) -> bool: + def has_time_word_en(cls, query: str, datetime_word_list: List[str]) -> bool: """ Check if the input query contains any datetime-related words based on the en language context. Args: query (str): The input string to check for datetime-related words. + datetime_word_list (list[str]): datetime keywords Returns: bool: True if the query contains at least one datetime-related word, False otherwise. """ contain_datetime = False - for datetime_word in DATATIME_WORD_LIST[G_CONTEXT.language]: + for datetime_word in datetime_word_list: datetime_word = datetime_word.lower() # TODO fix strip if datetime_word in [x.strip().lower().strip(",").strip(".").strip("?").strip(":") @@ -268,12 +271,18 @@ class DatetimeHandler(object): return contain_datetime @classmethod - def has_time_word(cls, query: str) -> bool: - func_name = f"has_time_word_{G_CONTEXT.language.value}" + def has_time_word(cls, query: str, language: LanguageEnum) -> bool: + func_name = f"has_time_word_{language}" if not hasattr(cls, func_name): - cls.logger.warning(f"language={G_CONTEXT.language.value} needs to complete has_time_word func!") + cls.logger.warning(f"language={language} needs to complete has_time_word function!") return False - return getattr(cls, func_name)(query=query) + + if language not in DATATIME_WORD_LIST: + cls.logger.warning(f"language={language} is missing in DATATIME_WORD_LIST!") + return False + + datetime_word_list = DATATIME_WORD_LIST[language] + return getattr(cls, func_name)(query=query, datetime_word_list=datetime_word_list) def datetime_format(self, dt_format: str = "%Y%m%d") -> str: """ @@ -287,7 +296,7 @@ class DatetimeHandler(object): """ return self._dt.strftime(dt_format) - def string_format(self, string_format: str) -> str: + def string_format(self, string_format: str, language: LanguageEnum) -> str: """ Format the datetime information stored in the instance using a custom string format. @@ -297,7 +306,7 @@ class DatetimeHandler(object): Returns: str: A formatted datetime string. """ - return string_format.format(**self.dt_info_dict) + return string_format.format(**self.get_dt_info_dict(language=language)) @property def timestamp(self) -> int: diff --git a/tests/worker/test_workers_cn.py b/tests/worker/test_workers_cn.py index a8114b35..c841d73d 100644 --- a/tests/worker/test_workers_cn.py +++ b/tests/worker/test_workers_cn.py @@ -140,7 +140,7 @@ class TestWorkersCn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -172,7 +172,7 @@ class TestWorkersCn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -202,7 +202,7 @@ class TestWorkersCn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_WITH_TIME_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_WITH_TIME_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -224,30 +224,30 @@ class TestWorkersCn(unittest.TestCase): MemoryNode(user_name="AI", target_name="用户", content="用户在阿里巴巴工作"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result1 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作"), MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result2 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作或有工作经验。"), MemoryNode(user_name="AI", target_name="用户", content="用户跳槽至openai工作。"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result3 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="我喜欢吃西瓜"), @@ -257,10 +257,10 @@ class TestWorkersCn(unittest.TestCase): MemoryNode(user_name="AI", target_name="用户", content="我爱吃苹果和香蕉"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result4 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) worker.logger.info(f"result1={result1}") worker.logger.info(f"result2={result2}") @@ -293,11 +293,11 @@ class TestWorkersCn(unittest.TestCase): MemoryNode(content="用户想知道维持广泛社交关系的方法。"), ] - worker.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes) - worker.memory_handler.set_memories(INSIGHT_NODES, []) + worker.memory_manager.set_memories(NOT_REFLECTED_NODES, nodes) + worker.memory_manager.set_memories(INSIGHT_NODES, []) worker.run() - result = [node.key for node in worker.memory_handler.get_memories(INSIGHT_NODES)] + result = [node.key for node in worker.memory_manager.get_memories(INSIGHT_NODES)] result = "\n".join(result) worker.logger.info(f"result.get_reflection={result}") return worker @@ -320,10 +320,10 @@ class TestWorkersCn(unittest.TestCase): nodes = [ MemoryNode(content="用户喜欢打王者荣耀"), ] - worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(INSIGHT_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(INSIGHT_NODES)] result = "\n".join(result) worker.logger.info(f"result.update_insight={result}") @@ -345,10 +345,10 @@ class TestWorkersCn(unittest.TestCase): MemoryNode(content="用户在北京工作,感到压力大,寻求放松方式。"), MemoryNode(content="用户在上海工作。"), ] - worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) worker.unit_test_flag = True worker.run() - result = [node.content for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result.long_contra_repeat={result}") diff --git a/tests/worker/test_workers_en.py b/tests/worker/test_workers_en.py index 6b1c012a..356c3f21 100644 --- a/tests/worker/test_workers_en.py +++ b/tests/worker/test_workers_en.py @@ -150,7 +150,7 @@ class TestWorkersEn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -193,7 +193,7 @@ class TestWorkersEn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -226,7 +226,7 @@ class TestWorkersEn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_WITH_TIME_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_WITH_TIME_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -248,30 +248,30 @@ class TestWorkersEn(unittest.TestCase): MemoryNode(user_name="AI", target_name="用户", content="User works at Alibaba"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result1 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="User works at JD.com"), MemoryNode(user_name="AI", target_name="用户", content="Users working in Meituan"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result2 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="User works at JD.com"), MemoryNode(user_name="AI", target_name="用户", content="User is working in Meituan"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result3 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="I like to eat watermelon"), @@ -279,10 +279,10 @@ class TestWorkersEn(unittest.TestCase): MemoryNode(user_name="AI", target_name="用户", content="I don't like watermelon"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result4 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) worker.logger.info(f"result1={result1}") worker.logger.info(f"result2={result2}") @@ -316,11 +316,11 @@ class TestWorkersEn(unittest.TestCase): MemoryNode(content="Users want to know how to maintain extensive social relationships."), ] - worker.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes) - worker.memory_handler.set_memories(INSIGHT_NODES, []) + worker.memory_manager.set_memories(NOT_REFLECTED_NODES, nodes) + worker.memory_manager.set_memories(INSIGHT_NODES, []) worker.run() - result = [node.key for node in worker.memory_handler.get_memories(INSIGHT_NODES)] + result = [node.key for node in worker.memory_manager.get_memories(INSIGHT_NODES)] result = "\n".join(result) worker.logger.info(f"result.get_reflection={result}") return worker @@ -343,10 +343,10 @@ class TestWorkersEn(unittest.TestCase): nodes = [ MemoryNode(content="Users like to play King of Glory"), ] - worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(INSIGHT_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(INSIGHT_NODES)] result = "\n".join(result) worker.logger.info(f"result.update_insight={result}") @@ -368,10 +368,10 @@ class TestWorkersEn(unittest.TestCase): MemoryNode(content="The user works in Beijing, feels stressed, and is looking for ways to relax."), MemoryNode(content="User works in Shanghai."), ] - worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) worker.unit_test_flag = True worker.run() - result = [node.content for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result.long_contra_repeat={result}")