diff --git a/config/demo_config.yaml b/config/demo_config.yaml index 33ed7954..c2aa5890 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -65,4 +65,8 @@ worker: generation_model: dashscope_generation embedding_model: dashscope_embedding rank_model: dashscope_rank + retrieve_store_worker: + class: memory.worker.read.retrieve_store_worker + retrieve_obs_top_k: 100 + retrieve_ins_pf_top_k: 100 diff --git a/memory_scope/constants/common_constants.py b/memory_scope/constants/common_constants.py index c38ed926..7863f55c 100644 --- a/memory_scope/constants/common_constants.py +++ b/memory_scope/constants/common_constants.py @@ -8,6 +8,7 @@ CHAT_KWARGS = "chat_kwargs" QUERY_WITH_TS = "query_with_ts" +RETRIEVE_MEMORY_NODES = "RETRIEVE_MEMORY_NODES" diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index 85426cc3..a62cc1f7 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -1,3 +1,4 @@ +import asyncio from abc import ABCMeta, abstractmethod from typing import Any, Dict @@ -25,6 +26,13 @@ class BaseWorker(metaclass=ABCMeta): self.continue_run: bool = True self.logger: Logger = Logger.get_logger() + @staticmethod + def _async_run(fn_list, *args, **kwargs): + async def async_gather(): + return await asyncio.gather(*[fn(*args, **kwargs) for fn in fn_list]) + + return asyncio.run(async_gather()) + @abstractmethod def _run(self): raise NotImplementedError diff --git a/memory_scope/memory/worker/read/extract_time_worker.py b/memory_scope/memory/worker/read/extract_time_worker.py index 34107113..1cbc0381 100644 --- a/memory_scope/memory/worker/read/extract_time_worker.py +++ b/memory_scope/memory/worker/read/extract_time_worker.py @@ -7,7 +7,7 @@ from memory_scope.utils.tool_functions import time_to_formatted_str class ExtractTimeWorker(MemoryBaseWorker): - Extract_Time_PATTERN = r'-\s*(\S+):(\d+)' + EXTRACT_TIME_PATTERN = r'-\s*(\S+):(\d+)' def _run(self): query, query_timestamp = self.get_context(QUERY_WITH_TS) @@ -44,7 +44,7 @@ class ExtractTimeWorker(MemoryBaseWorker): # re-match time info to dict extract_time_dict = {} - matches = re.findall(self.Extract_Time_PATTERN, response_text) + matches = re.findall(self.EXTRACT_TIME_PATTERN, response_text) for key, value in matches: if key in DATATIME_KEY_MAP.keys(): extract_time_dict[DATATIME_KEY_MAP[key]] = value diff --git a/memory_scope/memory/worker/read/retrieve_store_worker.py b/memory_scope/memory/worker/read/retrieve_store_worker.py new file mode 100644 index 00000000..4da85f70 --- /dev/null +++ b/memory_scope/memory/worker/read/retrieve_store_worker.py @@ -0,0 +1,44 @@ +from typing import List + +from memory_scope.constants.common_constants import QUERY_WITH_TS, RETRIEVE_MEMORY_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 + + +class RetrieveStoreWorker(MemoryBaseWorker): + + async def retrieve_from_observation(self, query: str) -> List[MemoryNode]: + filter_dict = { + "user_id": self.user_id, + "status": MemoryNodeStatus.ACTIVE.value, + "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], + } + return await self.vector_store.async_retrieve(query=query, + top_k=self.retrieve_obs_top_k, + filter_dict=filter_dict) + + async def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]: + filter_dict = { + "user_id": self.user_id, + "status": MemoryNodeStatus.ACTIVE.value, + "memory_type": [MemoryTypeEnum.INSIGHT.value, MemoryTypeEnum.PROFILE.value], + } + return await self.vector_store.async_retrieve(query=query, + top_k=self.retrieve_ins_pf_top_k, + filter_dict=filter_dict) + + def _run(self): + query, _ = self.get_context(QUERY_WITH_TS) + memory_node_list: List[MemoryNode] = [] + fn_list = [self.retrieve_from_observation, self.retrieve_from_insight_and_profile] + for result in self._async_run(fn_list=fn_list, query=query): + if result: + memory_node_list.extend(result) + memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True) + + self.logger.info(f"memory_node_list.size={len(memory_node_list)}") + for i, node in enumerate(memory_node_list): + self.logger.info(f"{i}: node={node.content} score={node.score_similar} type={node.memory_type}") + self.set_context(RETRIEVE_MEMORY_NODES, memory_node_list) diff --git a/memory_scope/storage/base_vector_store.py b/memory_scope/storage/base_vector_store.py index 9c5589b8..d0ba4f6f 100644 --- a/memory_scope/storage/base_vector_store.py +++ b/memory_scope/storage/base_vector_store.py @@ -7,11 +7,11 @@ from memory_scope.scheme.memory_node import MemoryNode class BaseVectorStore(metaclass=ABCMeta): @abstractmethod - def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]): + def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]: pass @abstractmethod - async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]): + async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]: pass @abstractmethod