diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index 036da933..4debfed7 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -4,6 +4,7 @@ import time import questionary from memory_scope.chat.base_memory_chat import BaseMemoryChat +from memory_scope.constants.language_constants import DEFAULT_HUMAN_NAME from memory_scope.enumeration.message_role_enum import MessageRoleEnum from memory_scope.memory.service.base_memory_service import BaseMemoryService from memory_scope.models.base_model import BaseModel @@ -27,7 +28,7 @@ class CliMemoryChat(BaseMemoryChat): memory_service: str, generation_model: str, stream: bool = True, - human_name: str = "用户", + human_name: str = DEFAULT_HUMAN_NAME[G_CONTEXT.language], assistant_name: str = "AI", **kwargs): @@ -192,9 +193,7 @@ class CliMemoryChat(BaseMemoryChat): except Exception as e: import traceback traceback.print_exc() - line = f"An exception occurred when running cli memory chat. args={e.args}." - questionary.print(line) - self.logger.exception(line) + self.logger.exception(f"An exception occurred when running cli memory chat. args={e.args}.") continue questionary.print(f"A memory writing thread is still running, please be patient and wait!") diff --git a/memory_scope/constants/common_constants.py b/memory_scope/constants/common_constants.py index 5f645e02..ba6be131 100644 --- a/memory_scope/constants/common_constants.py +++ b/memory_scope/constants/common_constants.py @@ -37,6 +37,8 @@ NEW_OBS_WITH_TIME_NODES = "new_obs_with_time_nodes" INSIGHT_NODES = "insight_nodes" +TODAY_NODES = "today_nodes" + MERGE_OBS_NODES = "merge_obs_nodes" NEW_INSIGHT_NODES = "new_insight_nodes" diff --git a/memory_scope/constants/language_constants.py b/memory_scope/constants/language_constants.py index caa5aa9a..82fad839 100644 --- a/memory_scope/constants/language_constants.py +++ b/memory_scope/constants/language_constants.py @@ -74,3 +74,8 @@ COMMA_WORD = { LanguageEnum.CN: ",", LanguageEnum.EN: "," } + +DEFAULT_HUMAN_NAME = { + LanguageEnum.CN: "用户", + LanguageEnum.EN: "user" +} diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py index 644e77fa..3cef85c2 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -1,5 +1,5 @@ from abc import ABCMeta -from typing import List, Dict +from typing import List, Dict, Set from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS from memory_scope.memory.worker.base_worker import BaseWorker @@ -70,18 +70,21 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): self._memory_store = G_CONTEXT.memory_store return self._memory_store - def get_memories(self, key: str) -> List[MemoryNode]: + def get_memories(self, keys: str | List[str]) -> List[MemoryNode]: memories: List[MemoryNode] = [] - memory_ids: List[str] = self.get_context(key) - if memory_ids: - memories.extend([self._contex_memory_dict[x] for x in memory_ids]) + if isinstance(keys, str): + keys = [keys] + + for key in keys: + memory_ids: List[str] = self.get_context(key) + if memory_ids: + memories.extend([self._contex_memory_dict[x] for x in memory_ids]) return memories - def set_memories(self, key: str, nodes: List[MemoryNode] | MemoryNode): - if not nodes: - return - - if isinstance(nodes, MemoryNode): + def set_memories(self, key: str, nodes: MemoryNode | List[MemoryNode]): + if nodes is None: + nodes = [] + elif isinstance(nodes, MemoryNode): nodes = [nodes] for node in nodes: @@ -90,6 +93,23 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): self._contex_memory_dict[node.memory_id] = node self.set_context(key, [n.memory_id for n in nodes]) + def save_memories(self, keys: str | List[str] = None): + if keys is None: + self.memory_store.update_memories(list(self._contex_memory_dict.values())) + self._contex_memory_dict.clear() + return + + if isinstance(keys, str): + keys = [keys] + + ids: Set[str] = Set[str]() + for key in keys: + t_ids: List[str] = self.get_context(key) + if t_ids: + ids.update(t_ids) + nodes = [self._contex_memory_dict.pop(_) for _ in ids] + self.memory_store.update_memories(nodes) + @property def monitor(self) -> BaseMonitor: if self._monitor is None: diff --git a/memory_scope/memory/worker/read/print_memory_worker.py b/memory_scope/memory/worker/read/print_memory_worker.py index 3ddce845..0e532144 100644 --- a/memory_scope/memory/worker/read/print_memory_worker.py +++ b/memory_scope/memory/worker/read/print_memory_worker.py @@ -6,51 +6,38 @@ 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 -from memory_scope.utils.timer import timer class PrintMemoryWorker(MemoryBaseWorker): - @timer - def retrieve_expired_memory(self, query: str): - filter_dict = { - "user_name": self.user_name, - "target_name": self.target_name, - "status": MemoryNodeStatus.EXPIRED.value, - "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], - } - return self.memory_store.retrieve_memories(query=query, - top_k=self.retrieve_expired_top_k, - filter_dict=filter_dict) - def _run(self): - expired_memories: List[MemoryNode] = self.retrieve_expired_memory(query="_") memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES) memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True) + expired_content_list: List[str] = [] obs_content_list: List[str] = [] insight_content_list: List[str] = [] - expired_content_list: List[str] = [] i = 0 j = 0 + k = 0 for node in memory_node_list: + if MemoryNodeStatus(node.status) is MemoryNodeStatus.EXPIRED: + i += 1 + line = f" {i} {node.content}" + expired_content_list.append(line) if MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.OBSERVATION, MemoryTypeEnum.OBS_CUSTOMIZED]: - i += 1 + j += 1 dt_handler = DatetimeHandler(node.timestamp) dt = dt_handler.datetime_format("%Y%m%d %H:%M:%S") - line = f" {i} {dt} {node.content}" + line = f" {j} {dt} {node.content}" obs_content_list.append(line) - elif MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.INSIGHT, ]: - j += 1 - line = f" {j} {node.content}" + elif MemoryTypeEnum(node.memory_type) is MemoryTypeEnum.INSIGHT: + k += 1 + line = f" {k} {node.content}" insight_content_list.append(line) - for i, node in enumerate(expired_memories): - line = f" {j} {node.content}" - expired_content_list.append(line) - obs_content = "\n".join(obs_content_list) insight_content = "\n".join(insight_content_list) expired_content = "\n".join(expired_content_list) diff --git a/memory_scope/memory/worker/read/retrieve_store_worker.py b/memory_scope/memory/worker/read/retrieve_memory_worker.py similarity index 67% rename from memory_scope/memory/worker/read/retrieve_store_worker.py rename to memory_scope/memory/worker/read/retrieve_memory_worker.py index 87b4d50b..42f7fe03 100644 --- a/memory_scope/memory/worker/read/retrieve_store_worker.py +++ b/memory_scope/memory/worker/read/retrieve_memory_worker.py @@ -5,11 +5,16 @@ from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus 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.timer import timer -class RetrieveStoreWorker(MemoryBaseWorker): +class RetrieveMemoryWorker(MemoryBaseWorker): + @timer async def retrieve_from_observation(self, query: str) -> List[MemoryNode]: + if not self.retrieve_obs_top_k: + return [] + filter_dict = { "user_name": self.user_name, "target_name": self.target_name, @@ -20,7 +25,11 @@ class RetrieveStoreWorker(MemoryBaseWorker): top_k=self.retrieve_obs_top_k, filter_dict=filter_dict) + @timer async def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]: + if not self.retrieve_ins_pf_top_k: + return [] + filter_dict = { "user_name": self.user_name, "target_name": self.target_name, @@ -31,10 +40,26 @@ class RetrieveStoreWorker(MemoryBaseWorker): top_k=self.retrieve_ins_pf_top_k, filter_dict=filter_dict) + @timer + async def retrieve_expired_memory(self, query: str) -> List[MemoryNode]: + if not self.retrieve_expired_top_k: + return [] + + filter_dict = { + "user_name": self.user_name, + "target_name": self.target_name, + "status": MemoryNodeStatus.EXPIRED.value, + "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], + } + return await self.memory_store.a_retrieve_memories(query=query, + top_k=self.retrieve_expired_top_k, + filter_dict=filter_dict) + def _run(self): query, _ = self.get_context(QUERY_WITH_TS) self.submit_async_task(self.retrieve_from_observation, query=query) self.submit_async_task(self.retrieve_from_insight_and_profile, query=query) + self.submit_async_task(self.retrieve_expired_memory, query=query) memory_node_list: List[MemoryNode] = [] for result in self.gather_async_result(): @@ -44,5 +69,6 @@ class RetrieveStoreWorker(MemoryBaseWorker): memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True) for node in memory_node_list: - self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} type={node.memory_type}") + self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} type={node.memory_type}" + f"status={node.status}") self.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list) 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 601fd0da..95e5f977 100644 --- a/memory_scope/memory/worker/summary/get_reflection_subject_worker.py +++ b/memory_scope/memory/worker/summary/get_reflection_subject_worker.py @@ -25,8 +25,8 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): status=MemoryNodeStatus.NEW.value) def _run(self): - not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES) - insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES) + not_reflected_nodes: List[MemoryNode] = self.get_memories(NOT_REFLECTED_NODES) + insight_nodes: List[MemoryNode] = self.get_memories(INSIGHT_NODES) # count not_reflected_count = len(not_reflected_nodes) 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 1193646c..61ced731 100644 --- a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py +++ b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py @@ -1,8 +1,6 @@ from typing import List, Dict -from memory_scope.constants.common_constants import ( - NOT_UPDATED_NODES, MERGE_OBS_NODES, -) +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.memory_type_enum import MemoryTypeEnum @@ -27,7 +25,7 @@ class LongContraRepeatWorker(MemoryBaseWorker): return node, [n for n in retrieve_nodes if n.score_similar >= self.long_contra_repeat_threshold] def _run(self): - not_updated_nodes: List[MemoryNode] = self.get_context(NOT_UPDATED_NODES) + not_updated_nodes: List[MemoryNode] = self.get_memories(NOT_UPDATED_NODES) for node in not_updated_nodes: self.submit_async_task(fn=self.retrieve_similar_content, node=node) @@ -103,10 +101,13 @@ class LongContraRepeatWorker(MemoryBaseWorker): node.status = MemoryNodeStatus.EXPIRED.value else: node.content = content + node.status = MemoryNodeStatus.CONTENT_MODIFIED.value + elif status == self.get_language_value(INCLUDED_WORD): node.status = MemoryNodeStatus.EXPIRED.value + merge_obs_nodes.append(node) self.logger.info(f"after_long_contra_repeat: {node.content} {node.status}") # save context - self.set_context(MERGE_OBS_NODES, merge_obs_nodes) + self.set_memories(MERGE_OBS_NODES, merge_obs_nodes) diff --git a/memory_scope/memory/worker/summary/summary_collect_worker.py b/memory_scope/memory/worker/summary/summary_collect_worker.py deleted file mode 100644 index 2b757d59..00000000 --- a/memory_scope/memory/worker/summary/summary_collect_worker.py +++ /dev/null @@ -1,40 +0,0 @@ -from typing import List, Dict - -from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_NODES, MERGE_OBS_NODES, \ - NOT_UPDATED_NODES -from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker -from memory_scope.scheme.memory_node import MemoryNode - - -class SummaryCollectWorker(MemoryBaseWorker): - - def _run(self): - update_memories: Dict[str, MemoryNode] = {} - insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES) - if insight_nodes: - update_memories.update({n.memory_id: n for n in insight_nodes}) - - not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES) - if not_reflected_nodes: - update_memories.update({n.memory_id: n for n in not_reflected_nodes}) - - not_updated_nodes: List[MemoryNode] = self.get_context(NOT_UPDATED_NODES) - if not_updated_nodes: - for node in not_updated_nodes: - if node.memory_id in update_memories: - - - - - keys = [ - INSIGHT_NODES, - MERGE_OBS_NODES, - NOT_UPDATED_NODES, - NOT_REFLECTED_NODES, - ] - - memory_nodes: List[MemoryNode] = [] - for key in keys: - memory_nodes.extend(self.get_context(key)) - - self.memory_store.update_memories(update_memories) diff --git a/memory_scope/memory/worker/summary/update_insight_worker.py b/memory_scope/memory/worker/summary/update_insight_worker.py index b235cd25..cc1bff92 100644 --- a/memory_scope/memory/worker/summary/update_insight_worker.py +++ b/memory_scope/memory/worker/summary/update_insight_worker.py @@ -100,9 +100,9 @@ class UpdateInsightWorker(MemoryBaseWorker): return insight_node def _run(self): - insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES) - not_updated_nodes: List[MemoryNode] = self.get_context(NOT_UPDATED_NODES) - not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES) + insight_nodes: List[MemoryNode] = self.get_memories(INSIGHT_NODES) + not_updated_nodes: List[MemoryNode] = self.get_memories(NOT_UPDATED_NODES) + not_reflected_nodes: List[MemoryNode] = self.get_memories(NOT_REFLECTED_NODES) if not insight_nodes: self.logger.warning("insight_nodes is empty, stop.") @@ -136,3 +136,6 @@ class UpdateInsightWorker(MemoryBaseWorker): for node in not_updated_nodes: node.obs_updated = True + + for node in not_reflected_nodes: + node.obs_updated = True diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py index b436d3da..58d0b626 100644 --- a/memory_scope/memory/worker/write/contra_repeat_worker.py +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -1,51 +1,22 @@ from typing import List -from memory_scope.constants.common_constants import NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, MERGE_OBS_NODES +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.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.scheme.message import Message -from memory_scope.utils.datetime_handler import DatetimeHandler from memory_scope.utils.response_text_parser import ResponseTextParser -from memory_scope.utils.timer import timer class ContraRepeatWorker(MemoryBaseWorker): - - @timer - def retrieve_today_memory(self) -> List[MemoryNode]: - if not self.chat_messages: - self.logger.warning("chat_messages is empty!") - return [] - - message: Message = self.chat_messages[-1] - dt_handler = DatetimeHandler(message.time_created) - filter_dict = { - "user_name": self.user_name, - "target_name": self.target_name, - "status": MemoryNodeStatus.ACTIVE.value, - "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], - "dt": dt_handler.datetime_format(), - } - return self.memory_store.retrieve_memories(query=message.content, top_k=self.today_obs_top_k, filter_dict=filter_dict) - def _run(self): - all_obs_nodes: List[MemoryNode] = [] - - new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) - if new_obs_nodes: - all_obs_nodes.extend(new_obs_nodes) - new_obs_with_time_nodes: List[MemoryNode] = self.get_context(NEW_OBS_WITH_TIME_NODES) - if new_obs_with_time_nodes: - all_obs_nodes.extend(new_obs_with_time_nodes) + all_obs_nodes: List[MemoryNode] = self.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.retrieve_today_memory() + today_obs_nodes: List[MemoryNode] = self.get_memories(TODAY_NODES) if today_obs_nodes: all_obs_nodes.extend(today_obs_nodes) all_obs_nodes = sorted(all_obs_nodes, key=lambda x: x.timestamp, reverse=True)[:self.contra_repeat_max_count] @@ -109,11 +80,11 @@ class ContraRepeatWorker(MemoryBaseWorker): node: MemoryNode = all_obs_nodes[idx] if keep_flag != self.get_language_value(NONE_WORD): - node.memory_node.status = MemoryNodeStatus.EXPIRED.value + node.status = MemoryNodeStatus.EXPIRED.value merge_obs_nodes.append(node) # forbid keyword self.logger.info(f"contra_repeat stage: {node.content} {node.status}") # save context - self.set_context(MERGE_OBS_NODES, merge_obs_nodes) + self.set_memories(MERGE_OBS_NODES, merge_obs_nodes) diff --git a/memory_scope/memory/worker/write/get_observation_with_time_worker.py b/memory_scope/memory/worker/write/get_observation_with_time_worker.py index 6afc3eeb..8ab3b71b 100644 --- a/memory_scope/memory/worker/write/get_observation_with_time_worker.py +++ b/memory_scope/memory/worker/write/get_observation_with_time_worker.py @@ -38,4 +38,4 @@ class GetObservationWithTimeWorker(GetObservationWorker): return obtain_obs_message def save(self, new_obs_nodes: List[MemoryNode]): - self.set_context(NEW_OBS_WITH_TIME_NODES, new_obs_nodes) + self.set_memories(NEW_OBS_WITH_TIME_NODES, new_obs_nodes) diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index afae2724..fb9690bd 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -62,7 +62,7 @@ class GetObservationWorker(MemoryBaseWorker): return obtain_obs_message def save(self, new_obs_nodes: List[MemoryNode]): - self.set_context(NEW_OBS_NODES, new_obs_nodes) + self.set_memories(NEW_OBS_NODES, new_obs_nodes) def _run(self): obtain_obs_message = self.build_prompt() diff --git a/memory_scope/memory/worker/summary/load_memory_worker.py b/memory_scope/memory/worker/write/load_memory_worker.py similarity index 63% rename from memory_scope/memory/worker/summary/load_memory_worker.py rename to memory_scope/memory/worker/write/load_memory_worker.py index e99224aa..c441d51d 100644 --- a/memory_scope/memory/worker/summary/load_memory_worker.py +++ b/memory_scope/memory/worker/write/load_memory_worker.py @@ -1,10 +1,12 @@ from typing import List -from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, NOT_UPDATED_NODES, INSIGHT_NODES +from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, NOT_UPDATED_NODES, INSIGHT_NODES, TODAY_NODES from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.scheme.message import Message +from memory_scope.utils.datetime_handler import DatetimeHandler from memory_scope.utils.timer import timer @@ -12,6 +14,9 @@ class LoadMemoryWorker(MemoryBaseWorker): @timer async def retrieve_not_reflected_memory(self, query: str): + if not self.retrieve_not_reflected_top_k: + return + filter_dict = { "user_name": self.user_name, "target_name": self.target_name, @@ -22,10 +27,13 @@ class LoadMemoryWorker(MemoryBaseWorker): nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query, top_k=self.retrieve_not_reflected_top_k, filter_dict=filter_dict) - self.set_context(NOT_REFLECTED_NODES, nodes) + self.set_memories(NOT_REFLECTED_NODES, nodes) @timer async def retrieve_not_updated_memory(self, query: str): + if not self.retrieve_not_updated_top_k: + return + filter_dict = { "user_name": self.user_name, "target_name": self.target_name, @@ -36,10 +44,13 @@ class LoadMemoryWorker(MemoryBaseWorker): nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query, top_k=self.retrieve_not_updated_top_k, filter_dict=filter_dict) - self.set_context(NOT_UPDATED_NODES, nodes) + self.set_memories(NOT_UPDATED_NODES, nodes) @timer async def retrieve_insight_memory(self, query: str): + if not self.retrieve_insight_top_k: + return + filter_dict = { "user_name": self.user_name, "target_name": self.target_name, @@ -49,12 +60,36 @@ class LoadMemoryWorker(MemoryBaseWorker): nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query, top_k=self.retrieve_insight_top_k, filter_dict=filter_dict) - self.set_context(INSIGHT_NODES, nodes) + self.set_memories(INSIGHT_NODES, nodes) + + @timer + async def retrieve_today_memory(self): + if not self.today_obs_top_k: + return + + if not self.chat_messages: + self.logger.warning("chat_messages is empty!") + return + + message: Message = self.chat_messages[-1] + dt_handler = DatetimeHandler(message.time_created) + filter_dict = { + "user_name": self.user_name, + "target_name": self.target_name, + "status": MemoryNodeStatus.ACTIVE.value, + "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], + "dt": dt_handler.datetime_format(), + } + nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=message.content, + top_k=self.today_obs_top_k, + filter_dict=filter_dict) + + self.set_memories(TODAY_NODES, nodes) async def _run(self): mock_query = "-" self.submit_async_task(self.retrieve_not_reflected_memory, query=mock_query) self.submit_async_task(self.retrieve_not_updated_memory, query=mock_query) self.submit_async_task(self.retrieve_insight_memory, query=mock_query) - + self.submit_async_task(self.retrieve_today_memory) self.gather_async_result() diff --git a/memory_scope/memory/worker/write/store_memory_worker.py b/memory_scope/memory/worker/write/store_memory_worker.py index 13bf8783..d97a9151 100644 --- a/memory_scope/memory/worker/write/store_memory_worker.py +++ b/memory_scope/memory/worker/write/store_memory_worker.py @@ -1,5 +1,3 @@ -from typing import List - from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker @@ -12,9 +10,11 @@ class StoreMemoryWorker(MemoryBaseWorker): def _run(self): store_key: str = self.store_key - if self.has_content(store_key): - memory_nodes: List[MemoryNode] = self.get_context(store_key) - self.memory_store.update_memories(memory_nodes) + if store_key == "all": + self.save_memories() + + elif self.has_content(store_key): + self.save_memories(store_key) elif store_key in self.chat_kwargs: query = self.chat_kwargs[store_key]