[dev] add memory handler and get memories

This commit is contained in:
jinli.yl 2024-07-08 00:36:12 +08:00
parent e758cf41e6
commit bc5278ffa0
5 changed files with 55 additions and 8 deletions

View file

@ -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:

View file

@ -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)

View file

@ -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)

View file

@ -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)

View file

@ -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)