diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py index 3898b167..644e77fa 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -4,6 +4,7 @@ from typing import List, Dict from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS from memory_scope.memory.worker.base_worker import BaseWorker from memory_scope.models.base_model import BaseModel +from memory_scope.scheme.memory_node import MemoryNode from memory_scope.scheme.message import Message from memory_scope.storage.base_memory_store import BaseMemoryStore from memory_scope.storage.base_monitor import BaseMonitor @@ -31,6 +32,8 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): self._target_name: str | None = None self._prompt_handler: PromptHandler | None = None + self._contex_memory_dict: Dict[str, MemoryNode] = {} + @property def chat_messages(self) -> List[Message]: return self.get_context(CHAT_MESSAGES) @@ -67,6 +70,26 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): self._memory_store = G_CONTEXT.memory_store return self._memory_store + def get_memories(self, key: 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]) + return memories + + def set_memories(self, key: str, nodes: List[MemoryNode] | MemoryNode): + if not nodes: + return + + if isinstance(nodes, MemoryNode): + nodes = [nodes] + + for node in nodes: + if node.memory_id in self._contex_memory_dict: + continue + self._contex_memory_dict[node.memory_id] = node + self.set_context(key, [n.memory_id for n in 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 d6ca7672..0485f1e5 100644 --- a/memory_scope/memory/worker/read/print_memory_worker.py +++ b/memory_scope/memory/worker/read/print_memory_worker.py @@ -10,7 +10,7 @@ from memory_scope.utils.datetime_handler import DatetimeHandler class PrintMemoryWorker(MemoryBaseWorker): def _run(self): - memory_node_list: List[MemoryNode] = self.get_context(RETRIEVE_MEMORY_NODES) + 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) obs_content_list: List[str] = [] @@ -29,13 +29,20 @@ class PrintMemoryWorker(MemoryBaseWorker): obs_content = "\n".join(obs_content_list) insight_content = "\n".join(insight_content_list) + expired_content = "\n".join() result: str = f""" The memories of {self.user_name} about {self.target_name}. -observation: +----- observation ----- {obs_content} +----- observation ----- -insight: +----- insight ----- {insight_content} +----- insight ----- + +----- expired ----- +{expired_content} +----- expired ----- """.strip() self.set_context(RESULT, result) diff --git a/memory_scope/memory/worker/read/retrieve_store_worker.py b/memory_scope/memory/worker/read/retrieve_store_worker.py index 28f4110f..87b4d50b 100644 --- a/memory_scope/memory/worker/read/retrieve_store_worker.py +++ b/memory_scope/memory/worker/read/retrieve_store_worker.py @@ -45,4 +45,4 @@ 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.set_context(RETRIEVE_MEMORY_NODES, memory_node_list) + self.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list) diff --git a/memory_scope/memory/worker/read/semantic_rank_worker.py b/memory_scope/memory/worker/read/semantic_rank_worker.py index 5efa1571..2574b32a 100644 --- a/memory_scope/memory/worker/read/semantic_rank_worker.py +++ b/memory_scope/memory/worker/read/semantic_rank_worker.py @@ -10,7 +10,7 @@ class SemanticRankWorker(MemoryBaseWorker): def _run(self): # query query, _ = self.get_context(QUERY_WITH_TS) - memory_node_list: List[MemoryNode] = self.get_context(RETRIEVE_MEMORY_NODES) + memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES) if not memory_node_list: self.logger.warning(f"retrieve memory nodes is empty!") return @@ -33,4 +33,4 @@ class SemanticRankWorker(MemoryBaseWorker): memory_node_list = sorted(memory_node_list, key=lambda n: n.score_rank, reverse=True) for node in memory_node_list: self.logger.info(f"rank_stage: content={node.content} score={node.score_rank}") - self.get_context(RANKED_MEMORY_NODES, memory_node_list) + self.set_memories(RANKED_MEMORY_NODES, memory_node_list) diff --git a/memory_scope/memory/worker/summary/summary_collect_worker.py b/memory_scope/memory/worker/summary/summary_collect_worker.py index ae07a237..2b757d59 100644 --- a/memory_scope/memory/worker/summary/summary_collect_worker.py +++ b/memory_scope/memory/worker/summary/summary_collect_worker.py @@ -1,4 +1,4 @@ -from typing import List +from typing import List, Dict from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_NODES, MERGE_OBS_NODES, \ NOT_UPDATED_NODES @@ -9,6 +9,23 @@ 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, @@ -20,4 +37,4 @@ class SummaryCollectWorker(MemoryBaseWorker): for key in keys: memory_nodes.extend(self.get_context(key)) - self.memory_store.update_memories(memory_nodes) + self.memory_store.update_memories(update_memories)