diff --git a/config/demo_config.yaml b/config/demo_config.yaml index f3edf11d..8a5e2589 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -112,7 +112,7 @@ worker: retrieve_not_reflected_top_k: 0 retrieve_not_updated_top_k: 0 retrieve_insight_top_k: 0 - today_obs_top_k: 100 + retrieve_today_top_k: 100 get_observation: class: memory.worker.write.get_observation_worker generation_model: dashscope_generation @@ -135,7 +135,7 @@ worker: retrieve_not_reflected_top_k: 100 retrieve_not_updated_top_k: 100 retrieve_insight_top_k: 100 - today_obs_top_k: 0 + retrieve_today_top_k: 0 get_reflection_subject: class: memory.worker.summary.get_reflection_subject_worker retrieve_top_k: 100 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 304c8ab5..9532a05e 100644 --- a/memory_scope/memory/worker/summary/get_reflection_subject_worker.py +++ b/memory_scope/memory/worker/summary/get_reflection_subject_worker.py @@ -2,7 +2,7 @@ from typing import List from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_NODES from memory_scope.constants.language_constants import COMMA_WORD -from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus +from memory_scope.enumeration.action_status_enum import ActionStatusEnum from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker from memory_scope.scheme.memory_node import MemoryNode @@ -30,14 +30,15 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): MemoryNode: A new MemoryNode instance representing the insight, marked as new and of type INSIGHT. """ dt_handler = DatetimeHandler() - meta_data = {k: str(v) for k, v in dt_handler.dt_info_dict.items()} # ⭐ Prepare metadata with current datetime info + # Prepare metadata with current datetime info + meta_data = {k: str(v) for k, v in dt_handler.dt_info_dict.items()} - return MemoryNode(user_name=self.user_name, # ⭐ Populate MemoryNode attributes + return MemoryNode(user_name=self.user_name, target_name=self.target_name, meta_data=meta_data, key=insight_key, memory_type=MemoryTypeEnum.INSIGHT.value, - status=MemoryNodeStatus.NEW.value) + action_status=ActionStatusEnum.NEW.value) def _run(self): """ @@ -58,7 +59,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): # Count unaudited nodes not_reflected_count = len(not_reflected_nodes) if not_reflected_count <= self.reflect_obs_cnt_threshold: - self.logger.info(f"not_reflected_count={not_reflected_count} is not enough, stopping process.") + self.logger.info(f"not_reflected_count={not_reflected_count} is not enough, skip.") self.continue_run = False return 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 59cfefb9..91979d76 100644 --- a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py +++ b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py @@ -2,8 +2,9 @@ from typing import List, Dict from memory_scope.constants.common_constants import NOT_UPDATED_NODES, MERGE_OBS_NODES from memory_scope.constants.language_constants import NONE_WORD, INCLUDED_WORD, CONTRADICTORY_WORD -from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus +from memory_scope.enumeration.action_status_enum import ActionStatusEnum from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum +from memory_scope.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.response_text_parser import ResponseTextParser @@ -34,7 +35,7 @@ class LongContraRepeatWorker(MemoryBaseWorker): filter_dict = { "user_name": self.user_name, "target_name": self.target_name, - "status": MemoryNodeStatus.ACTIVE.value, + "store_status": StoreStatusEnum.VALID.value, "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value] } # Retrieve memories similar to the node's content, limited by top_k and filtered by filter_dict @@ -129,13 +130,13 @@ class LongContraRepeatWorker(MemoryBaseWorker): node: MemoryNode = all_obs_nodes[idx] if status == self.get_language_value(CONTRADICTORY_WORD): if not content: - node.status = MemoryNodeStatus.EXPIRED.value + node.store_status = StoreStatusEnum.EXPIRED.value else: node.content = content - node.status = MemoryNodeStatus.CONTENT_MODIFIED.value + node.action_status = ActionStatusEnum.CONTENT_MODIFIED.value elif status == self.get_language_value(INCLUDED_WORD): - node.status = MemoryNodeStatus.EXPIRED.value + node.store_status = StoreStatusEnum.EXPIRED.value merge_obs_nodes.append(node) self.logger.info(f"after_long_contra_repeat: {node.content} {node.status}") diff --git a/memory_scope/memory/worker/summary/update_insight_worker.py b/memory_scope/memory/worker/summary/update_insight_worker.py index 69f7e7cf..919a28dc 100644 --- a/memory_scope/memory/worker/summary/update_insight_worker.py +++ b/memory_scope/memory/worker/summary/update_insight_worker.py @@ -2,7 +2,7 @@ from typing import List from memory_scope.constants.common_constants import INSIGHT_NODES, NOT_UPDATED_NODES, NOT_REFLECTED_NODES from memory_scope.constants.language_constants import COLON_WORD, NONE_WORD, REPEATED_WORD -from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus +from memory_scope.enumeration.action_status_enum import ActionStatusEnum 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 @@ -76,8 +76,8 @@ class UpdateInsightWorker(MemoryBaseWorker): insight_node.meta_data.update({k: str(v) for k, v in dt_handler.dt_info_dict.items()}) insight_node.timestamp = dt_handler.timestamp insight_node.dt = dt_handler.datetime_format() - if insight_node.status == MemoryNodeStatus.ACTIVE.value: - insight_node.status = MemoryNodeStatus.CONTENT_MODIFIED.value + if insight_node.action_status == ActionStatusEnum.NONE.value: + insight_node.action_status = ActionStatusEnum.CONTENT_MODIFIED.value self.logger.info(f"after_update_{insight_node.key} value={insight_value}") return insight_node @@ -162,7 +162,7 @@ class UpdateInsightWorker(MemoryBaseWorker): # Process active insight nodes with corresponding not updated nodes for node in insight_nodes: - if node.status == MemoryNodeStatus.ACTIVE.value: + if node.action_status == ActionStatusEnum.NONE.value: self.submit_thread_task(fn=self.filter_obs_nodes, insight_node=node, not_updated_nodes=not_updated_nodes) diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py index e6354f4e..d441d927 100644 --- a/memory_scope/memory/worker/write/contra_repeat_worker.py +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -1,7 +1,8 @@ from typing import List + from memory_scope.constants.common_constants import NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, MERGE_OBS_NODES, TODAY_NODES from memory_scope.constants.language_constants import NONE_WORD, CONTRADICTORY_WORD, INCLUDED_WORD -from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus +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.response_text_parser import ResponseTextParser @@ -106,8 +107,8 @@ class ContraRepeatWorker(MemoryBaseWorker): node: MemoryNode = all_obs_nodes[idx] if keep_flag != self.get_language_value(NONE_WORD): - node.status = MemoryNodeStatus.EXPIRED.value - self.logger.info(f"contra_repeat stage: {node.content} {node.status}") + node.store_status = StoreStatusEnum.EXPIRED.value + self.logger.info(f"contra_repeat stage: {node.content} {node.store_status} {node.action_status}") merge_obs_nodes.append(node) # save context diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index 40334d95..2a4ed53f 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -2,7 +2,7 @@ from typing import List from memory_scope.constants.common_constants import NEW_OBS_NODES, TIME_INFER from memory_scope.constants.language_constants import REPEATED_WORD, NONE_WORD, COLON_WORD, TIME_INFER_WORD -from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus +from memory_scope.enumeration.action_status_enum import ActionStatusEnum from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker from memory_scope.scheme.memory_node import MemoryNode @@ -38,10 +38,8 @@ class GetObservationWorker(MemoryBaseWorker): meta_data=meta_data, content=obs_content, memory_type=MemoryTypeEnum.OBSERVATION.value, - status=MemoryNodeStatus.NEW.value, - timestamp=message.time_created, - obs_reflected=False, - obs_updated=False) + action_status=ActionStatusEnum.NEW.value, + timestamp=message.time_created) def filter_messages(self) -> List[Message]: filter_messages = [] @@ -134,6 +132,7 @@ class GetObservationWorker(MemoryBaseWorker): idx, time_infer, obs_content, keywords = obs_content_list # Skips processing if content indicates no meaningful observation + obs_content = obs_content.lower() if obs_content in self.get_language_value([NONE_WORD, REPEATED_WORD]): continue @@ -142,7 +141,7 @@ class GetObservationWorker(MemoryBaseWorker): self.logger.warning(f"idx={idx} is invalid!") continue - if time_infer == self.get_language_value(NONE_WORD): + if time_infer.lower() == self.get_language_value(NONE_WORD): time_infer = "" # Adjusts index to zero-based and checks validity against filtered messages diff --git a/memory_scope/memory/worker/write/info_filter_worker.py b/memory_scope/memory/worker/write/info_filter_worker.py index 97c04c50..81b90a85 100644 --- a/memory_scope/memory/worker/write/info_filter_worker.py +++ b/memory_scope/memory/worker/write/info_filter_worker.py @@ -33,16 +33,18 @@ class InfoFilterWorker(MemoryBaseWorker): # filter user msg info_messages: List[Message] = [] for msg in self.chat_messages: + if msg.memorized: + continue + # TODO: add memory for assistant if msg.role != MessageRoleEnum.USER.value: continue + if len(msg.content) >= self.info_filter_msg_max_size: half_size = int(self.info_filter_msg_max_size * 0.5 + 0.5) msg.content = msg.content[: half_size] + msg.content[-half_size:] info_messages.append(msg) - self.logger.warning(info_messages) - if not info_messages: self.logger.warning("info_messages is empty!") self.continue_run = False @@ -79,12 +81,25 @@ class InfoFilterWorker(MemoryBaseWorker): # filter messages filtered_messages: List[Message] = [] - for msg, info_score in zip(info_messages, info_score_list): + for info_score in info_score_list: if not info_score: continue - score = info_score[0] + if len(info_score) != 2: + self.logger.warning(f"info_score={info_score} is invalid!") + continue + + idx, score = info_score + + idx = int(idx) - 1 + if idx >= len(info_messages): + self.logger.warning(f"idx={idx} is invalid! info_messages.size={len(info_messages)}") + continue + message = info_messages[idx] + if score in self.preserved_scores: - msg.meta_data["info_score"] = score - filtered_messages.append(msg) + message.meta_data["info_score"] = score + filtered_messages.append(message) + self.logger.info(f"info filter stage: keep {message.content}") + self.chat_messages = filtered_messages diff --git a/memory_scope/memory/worker/write/load_memory_worker.py b/memory_scope/memory/worker/write/load_memory_worker.py index 690a6deb..51425fd3 100644 --- a/memory_scope/memory/worker/write/load_memory_worker.py +++ b/memory_scope/memory/worker/write/load_memory_worker.py @@ -81,7 +81,7 @@ class LoadMemoryWorker(MemoryBaseWorker): "dt": dt_handler.datetime_format(), } nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=message.content, - top_k=self.today_obs_top_k, + top_k=self.retrieve_today_top_k, filter_dict=filter_dict) self.memory_handler.set_memories(TODAY_NODES, nodes) diff --git a/memory_scope/utils/memory_handler.py b/memory_scope/utils/memory_handler.py index a3d187df..0a4fca3f 100644 --- a/memory_scope/utils/memory_handler.py +++ b/memory_scope/utils/memory_handler.py @@ -1,6 +1,7 @@ from typing import Dict, List, Set from memory_scope.enumeration.action_status_enum import ActionStatusEnum +from memory_scope.enumeration.store_status_enum import StoreStatusEnum from memory_scope.scheme.memory_node import MemoryNode from memory_scope.storage.base_memory_store import BaseMemoryStore from memory_scope.utils.global_context import G_CONTEXT @@ -123,6 +124,11 @@ class MemoryHandler(object): if not nodes: return + for node in nodes: + # Non-deleted expired memory nodes need to be changed to a modified state. + if node.store_status == StoreStatusEnum.EXPIRED.value and node.action_status != ActionStatusEnum.DELETE: + node.action_status = ActionStatusEnum.MODIFIED + # emb & insert new memories new_memories = [n for n in nodes if n.action_status == ActionStatusEnum.NEW.value] if new_memories: diff --git a/memory_scope/utils/response_text_parser.py b/memory_scope/utils/response_text_parser.py index fbcb0889..fb2ca139 100644 --- a/memory_scope/utils/response_text_parser.py +++ b/memory_scope/utils/response_text_parser.py @@ -1,6 +1,8 @@ import re from memory_scope.utils.logger import Logger +from memory_scope.constants.language_constants import NONE_WORD +from memory_scope.utils.global_context import G_CONTEXT class ResponseTextParser(object): @@ -30,19 +32,15 @@ class ResponseTextParser(object): matches = [match.group(1) for match in self.pattern_v1.finditer(line)] if matches: result.append(matches) - self.logger.info( - f"{prefix} response_text={self.response_text} result={result}", stacklevel=2 - ) + self.logger.info(f"{prefix} response_text={self.response_text} result={result}", stacklevel=2) return result def parse_v2(self, prefix: str = ""): result = [] for line in self.response_text.split("\n"): line = line.strip() - if not line or line == "无": + if not line or line.lower() == NONE_WORD.get(G_CONTEXT.language): continue result.append(line) - self.logger.info( - f"{prefix} response_text={self.response_text} result={result}", stacklevel=2 - ) + self.logger.info(f"{prefix} response_text={self.response_text} result={result}", stacklevel=2) return result