From 4d6be46a1e82c3b5125862b44bc72d4b3361c55b Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Sun, 7 Jul 2024 23:18:31 +0800 Subject: [PATCH] [dev] change dummy workflow to actual workflow --- config/demo_config.yaml | 30 +++++++--- memory_scope/cli.py | 8 +-- .../enumeration/memory_method_enum.py | 13 ---- .../enumeration/memory_recall_enum.py | 9 --- .../enumeration/memory_status_enum.py | 6 ++ .../memory/operation/base_operation.py | 1 + .../memory/worker/memory_base_worker.py | 12 ++-- .../worker/read/retrieve_store_worker.py | 12 ++-- .../worker/summary/load_memory_worker.py | 18 +++--- .../summary/long_contra_repeat_worker.py | 6 +- .../worker/write/contra_repeat_worker.py | 2 +- .../write/get_observation_with_time_worker.py | 9 +-- .../worker/write/get_observation_worker.py | 11 +--- .../worker/write/store_memory_worker.py | 4 +- memory_scope/scheme/memory_node.py | 7 +-- memory_scope/storage/base_memory_store.py | 33 +++++++++++ memory_scope/storage/base_vector_store.py | 37 ------------ memory_scope/storage/dummy_memory_store.py | 24 ++++++++ memory_scope/storage/dummy_vector_store.py | 21 ------- ...tore.py => llama_index_es_memory_store.py} | 59 +++++++++++++++---- memory_scope/utils/datetime_handler.py | 39 +++--------- memory_scope/utils/global_context.py | 4 +- old/worker/es/es_insight_worker.py | 2 +- old/worker/es/es_new_obs_worker.py | 2 +- old/worker/es/es_not_reflected_worker.py | 2 +- old/worker/es/es_similar_worker.py | 2 +- old/worker/es/es_today_obs_worker.py | 2 +- 27 files changed, 188 insertions(+), 187 deletions(-) delete mode 100644 memory_scope/enumeration/memory_method_enum.py delete mode 100644 memory_scope/enumeration/memory_recall_enum.py create mode 100644 memory_scope/storage/base_memory_store.py delete mode 100644 memory_scope/storage/base_vector_store.py create mode 100644 memory_scope/storage/dummy_memory_store.py delete mode 100644 memory_scope/storage/dummy_vector_store.py rename memory_scope/storage/{llama_index_elastic_search_store.py => llama_index_es_memory_store.py} (69%) diff --git a/config/demo_config.yaml b/config/demo_config.yaml index a0d7f376..83d51494 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -28,12 +28,12 @@ memory_service: description: "read all memories of the user" write_memory: class: memory.operation.write_memory - workflow: dummy_worker + workflow: info_filter_worker,[get_observation_worker|get_observation_with_time_worker],contra_repeat_worker,store_memory_worker description: "write observation memories of the user" interval_time: 60 summary_memory: class: memory.operation.summary_memory - workflow: dummy_worker + workflow: load_memory_worker,get_reflection_subject_worker,update_insight_worker,long_contra_repeat_worker,summary_collect_worker description: "summary observation memories of the user" interval_time: 300 worker: @@ -44,13 +44,13 @@ worker: rank_model: dashscope_rank set_query_worker: class: memory.worker.read.set_query_worker - extract_time_worker: - class: memory.worker.read.extract_time_worker - generation_model_top_k: 1 retrieve_store_worker: class: memory.worker.read.retrieve_store_worker retrieve_obs_top_k: 100 retrieve_ins_pf_top_k: 100 + extract_time_worker: + class: memory.worker.read.extract_time_worker + generation_model_top_k: 1 semantic_rank_worker: class: memory.worker.read.semantic_rank_worker fuse_rerank_worker: @@ -63,6 +63,8 @@ worker: insight: 2.0 fuse_time_ratio: 2.0 fuse_rerank_top_k: 10 + print_memory_worker: + class: memory.worker.read.print_memory_worker info_filter_worker: class: memory.worker.write.info_filter_worker generation_model: dashscope_generation @@ -82,10 +84,20 @@ worker: generation_model_top_k: 1 retrieve_top_k: 30 contra_repeat_max_count: 50 - get_reflection_worker: - class: memory.worker.summary.get_reflection_worker + store_memory_worker: + class: memory.worker.write.store_memory_worker + load_memory_worker: + class: memory.worker.summary.load_memory_worker + get_reflection_subject_worker: + class: memory.worker.summary.get_reflection_subject_worker retrieve_top_k: 100 reflect_obs_cnt_threshold: 32 + update_insight_worker: + class: memory.worker.summary.update_insight_worker + long_contra_repeat_worker: + class: memory.worker.summary.long_contra_repeat_worker + summary_collect_worker: + class: memory.worker.summary.summary_collect_worker models: dashscope_generation: class: models.llama_index_generation_model @@ -99,8 +111,8 @@ models: class: models.llama_index_rank_model module_name: dashscope_rank model_name: gte-rerank -vector_store: - class: storage.llama_index_elastic_search_store +memory_store: + class: storage.llama_index_es_memory_store embedding_model: dashscope_embedding index_name: memory_index es_url: http://localhost:9200 diff --git a/memory_scope/cli.py b/memory_scope/cli.py index 95e8d372..23519e6f 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -56,9 +56,9 @@ class CliJob(object): G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name) # init vector_store - vector_store_config = self.config["vector_store"] - embedding_model = G_CONTEXT.model_dict[vector_store_config[ModelEnum.EMBEDDING_MODEL.value]] - G_CONTEXT.vector_store = init_instance_by_config(vector_store_config, embedding_model=embedding_model) + memory_store_config = self.config["memory_store"] + embedding_model = G_CONTEXT.model_dict[memory_store_config[ModelEnum.EMBEDDING_MODEL.value]] + G_CONTEXT.memory_store = init_instance_by_config(memory_store_config, embedding_model=embedding_model) # init monitor G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) @@ -73,7 +73,7 @@ class CliJob(object): memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0] memory_chat.run() - G_CONTEXT.vector_store.close() + G_CONTEXT.memory_store.close() G_CONTEXT.monitor.close() diff --git a/memory_scope/enumeration/memory_method_enum.py b/memory_scope/enumeration/memory_method_enum.py deleted file mode 100644 index b30adf31..00000000 --- a/memory_scope/enumeration/memory_method_enum.py +++ /dev/null @@ -1,13 +0,0 @@ -from enum import Enum - - -class MemoryMethodEnum(str, Enum): - SUMMARY = "summary" - - RETRIEVE = "retrieve" - - RETRIEVE_ALL = "retrieve_all" - - SUMMARY_SHORT = "summary_short" - - SUMMARY_LONG = "summary_long" diff --git a/memory_scope/enumeration/memory_recall_enum.py b/memory_scope/enumeration/memory_recall_enum.py deleted file mode 100644 index 956ee4f8..00000000 --- a/memory_scope/enumeration/memory_recall_enum.py +++ /dev/null @@ -1,9 +0,0 @@ -from enum import Enum - - -class MemoryRecallType(str, Enum): - SIMILAR = "similar" - - KEYWORD = "keyword" - - PROFILE = "profile" diff --git a/memory_scope/enumeration/memory_status_enum.py b/memory_scope/enumeration/memory_status_enum.py index f8c706b0..60703d95 100644 --- a/memory_scope/enumeration/memory_status_enum.py +++ b/memory_scope/enumeration/memory_status_enum.py @@ -2,6 +2,12 @@ from enum import Enum class MemoryNodeStatus(str, Enum): + NEW = "new" + + MODIFIED = "modified" + + CONTENT_MODIFIED = "content_modified" + ACTIVE = "active" EXPIRED = "expired" diff --git a/memory_scope/memory/operation/base_operation.py b/memory_scope/memory/operation/base_operation.py index 72d360ff..a58569f7 100644 --- a/memory_scope/memory/operation/base_operation.py +++ b/memory_scope/memory/operation/base_operation.py @@ -10,6 +10,7 @@ class BaseOperation(metaclass=ABCMeta): def __init__(self, name: str, description: str = "", **kwargs): self.name: str = name self.description: str = description + self.kwargs: dict = kwargs def init_workflow(self): pass diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py index 15656ed8..05074f91 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -6,7 +6,7 @@ from memory_scope.memory.worker.base_worker import BaseWorker from memory_scope.models.base_model import BaseModel from memory_scope.scheme.message import Message from memory_scope.storage.base_monitor import BaseMonitor -from memory_scope.storage.base_vector_store import BaseVectorStore +from memory_scope.storage.base_memory_store import BaseMemoryStore from memory_scope.utils.global_context import G_CONTEXT from memory_scope.utils.prompt_handler import PromptHandler @@ -24,7 +24,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): self._generation_model: BaseModel | str = generation_model self._rank_model: BaseModel | str = rank_model - self._vector_store: BaseVectorStore | None = None + self._memory_store: BaseMemoryStore | None = None self._monitor: BaseMonitor | None = None self._user_name: str | None = None @@ -62,10 +62,10 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): return self._rank_model @property - def vector_store(self) -> BaseVectorStore: - if self._vector_store is None: - self._vector_store = G_CONTEXT.vector_store - return self._vector_store + def memory_store(self) -> BaseMemoryStore: + if self._memory_store is None: + self._memory_store = G_CONTEXT.memory_store + return self._memory_store @property def monitor(self) -> BaseMonitor: diff --git a/memory_scope/memory/worker/read/retrieve_store_worker.py b/memory_scope/memory/worker/read/retrieve_store_worker.py index ab2c1a38..28f4110f 100644 --- a/memory_scope/memory/worker/read/retrieve_store_worker.py +++ b/memory_scope/memory/worker/read/retrieve_store_worker.py @@ -16,9 +16,9 @@ class RetrieveStoreWorker(MemoryBaseWorker): "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) + return await self.memory_store.a_retrieve_memories(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 = { @@ -27,9 +27,9 @@ class RetrieveStoreWorker(MemoryBaseWorker): "status": MemoryNodeStatus.ACTIVE.value, "memory_type": MemoryTypeEnum.INSIGHT.value, } - return await self.vector_store.async_retrieve(query=query, - top_k=self.retrieve_ins_pf_top_k, - filter_dict=filter_dict) + return await self.memory_store.a_retrieve_memories(query=query, + top_k=self.retrieve_ins_pf_top_k, + filter_dict=filter_dict) def _run(self): query, _ = self.get_context(QUERY_WITH_TS) diff --git a/memory_scope/memory/worker/summary/load_memory_worker.py b/memory_scope/memory/worker/summary/load_memory_worker.py index 00258f15..e99224aa 100644 --- a/memory_scope/memory/worker/summary/load_memory_worker.py +++ b/memory_scope/memory/worker/summary/load_memory_worker.py @@ -19,9 +19,9 @@ class LoadMemoryWorker(MemoryBaseWorker): "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], "obs_reflected": False, } - nodes: List[MemoryNode] = await self.vector_store.async_retrieve(query=query, - top_k=self.retrieve_not_reflected_top_k, - filter_dict=filter_dict) + 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) @timer @@ -33,9 +33,9 @@ class LoadMemoryWorker(MemoryBaseWorker): "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], "obs_updated": False, } - nodes: List[MemoryNode] = await self.vector_store.async_retrieve(query=query, - top_k=self.retrieve_not_updated_top_k, - filter_dict=filter_dict) + 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) @timer @@ -46,9 +46,9 @@ class LoadMemoryWorker(MemoryBaseWorker): "status": MemoryNodeStatus.ACTIVE.value, "memory_type": MemoryTypeEnum.INSIGHT.value, } - nodes: List[MemoryNode] = await self.vector_store.async_retrieve(query=query, - top_k=self.retrieve_insight_top_k, - filter_dict=filter_dict) + 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) async def _run(self): 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 a11d700d..1193646c 100644 --- a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py +++ b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py @@ -21,9 +21,9 @@ class LongContraRepeatWorker(MemoryBaseWorker): "status": MemoryNodeStatus.ACTIVE.value, "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value] } - retrieve_nodes = await self.vector_store.async_retrieve(query=node.content, - top_k=self.long_contra_repeat_top_k, - filter_dict=filter_dict) + retrieve_nodes = await self.memory_store.a_retrieve_memories(query=node.content, + top_k=self.long_contra_repeat_top_k, + filter_dict=filter_dict) return node, [n for n in retrieve_nodes if n.score_similar >= self.long_contra_repeat_threshold] def _run(self): diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py index c3c9f21d..b436d3da 100644 --- a/memory_scope/memory/worker/write/contra_repeat_worker.py +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -29,7 +29,7 @@ class ContraRepeatWorker(MemoryBaseWorker): "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], "dt": dt_handler.datetime_format(), } - return self.vector_store.retrieve(query=message.content, top_k=self.today_obs_top_k, filter_dict=filter_dict) + 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] = [] 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 eb9defea..6afc3eeb 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 @@ -1,7 +1,7 @@ from typing import List from memory_scope.constants.common_constants import NEW_OBS_WITH_TIME_NODES -from memory_scope.constants.language_constants import DATATIME_WORD_LIST, COLON_WORD +from memory_scope.constants.language_constants import COLON_WORD from memory_scope.memory.worker.write.get_observation_worker import GetObservationWorker from memory_scope.scheme.memory_node import MemoryNode from memory_scope.scheme.message import Message @@ -16,12 +16,7 @@ class GetObservationWithTimeWorker(GetObservationWorker): user_query_list = [] i = 1 for msg in self.chat_messages: - match = False - for time_keyword in self.get_language_value(DATATIME_WORD_LIST): - if time_keyword in msg.content: - match = True - break - if match: + if DatetimeHandler.has_time_word(query=msg.content): dt_handler = DatetimeHandler(dt=msg.time_created) dt = dt_handler.string_format(self.prompt_handler.time_string_format) user_query_list.append(f"{i} {dt} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}") diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index e7c94991..afae2724 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -1,7 +1,7 @@ from typing import List from memory_scope.constants.common_constants import NEW_OBS_NODES, TIME_INFER -from memory_scope.constants.language_constants import DATATIME_WORD_LIST, REPEATED_WORD, NONE_WORD, COLON_WORD +from memory_scope.constants.language_constants import REPEATED_WORD, NONE_WORD, COLON_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 @@ -33,7 +33,7 @@ class GetObservationWorker(MemoryBaseWorker): meta_data=meta_data, content=obs_content, memory_type=MemoryTypeEnum.OBSERVATION.value, - status=MemoryNodeStatus.ACTIVE.value, + status=MemoryNodeStatus.NEW.value, timestamp=message.time_created, obs_reflected=False, obs_updated=False) @@ -43,12 +43,7 @@ class GetObservationWorker(MemoryBaseWorker): user_query_list = [] i = 1 for msg in self.chat_messages: - match = False - for time_keyword in self.get_language_value(DATATIME_WORD_LIST): - if time_keyword in msg.content: - match = True - break - if not match: + if not DatetimeHandler.has_time_word(query=msg.content): user_query_list.append(f"{i} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}") i += 1 diff --git a/memory_scope/memory/worker/write/store_memory_worker.py b/memory_scope/memory/worker/write/store_memory_worker.py index 92e4ecbf..95f8ac05 100644 --- a/memory_scope/memory/worker/write/store_memory_worker.py +++ b/memory_scope/memory/worker/write/store_memory_worker.py @@ -14,7 +14,7 @@ class StoreMemoryWorker(MemoryBaseWorker): if self.has_content(store_key): memory_nodes: List[MemoryNode] = self.get_context(store_key) - self.vector_store.update_batch(memory_nodes) + self.memory_store.update_batch(memory_nodes) elif store_key in self.chat_kwargs: query = self.chat_kwargs[store_key] @@ -31,4 +31,4 @@ class StoreMemoryWorker(MemoryBaseWorker): timestamp=dt_handler.timestamp, obs_reflected=False, obs_updated=False) - self.vector_store.update(node) + self.memory_store.update(node) diff --git a/memory_scope/scheme/memory_node.py b/memory_scope/scheme/memory_node.py index c5f13fa0..d1366397 100644 --- a/memory_scope/scheme/memory_node.py +++ b/memory_scope/scheme/memory_node.py @@ -1,13 +1,12 @@ import datetime from typing import Dict, List +from uuid import uuid4 from pydantic import Field, BaseModel -from memory_scope.utils.tool_functions import md5_hash - class MemoryNode(BaseModel): - memory_id: str = Field("", description="unique id for memory") + memory_id: str = Field(uuid4(), description="unique id for memory") user_name: str = Field("", description="the user who owns the memory") @@ -43,7 +42,6 @@ class MemoryNode(BaseModel): def __init__(self, **kwargs): super().__init__(**kwargs) - self.memory_id = f"{self.user_name}_{self.target_name}_{self.timestamp}_{md5_hash(self.content)[:8]}" self.dt = datetime.datetime.fromtimestamp(self.timestamp).strftime("%Y%m%d") @property @@ -52,4 +50,3 @@ class MemoryNode(BaseModel): def __getitem__(self, key: str): return self.model_dump().get(key) - diff --git a/memory_scope/storage/base_memory_store.py b/memory_scope/storage/base_memory_store.py new file mode 100644 index 00000000..bc8aff37 --- /dev/null +++ b/memory_scope/storage/base_memory_store.py @@ -0,0 +1,33 @@ +from abc import ABCMeta, abstractmethod +from typing import Dict, List + +from memory_scope.scheme.memory_node import MemoryNode + + +class BaseMemoryStore(metaclass=ABCMeta): + + @abstractmethod + def retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]: + pass + + @abstractmethod + async def a_retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]: + pass + + @abstractmethod + def update_memories(self, nodes: MemoryNode | List[MemoryNode]): + """ + status: + 1. new: emb & insert + 2. modified: update + 3. content_modified: emb & update + 4. active: do nothing + 5. expired: update + """ + + def flush(self): + pass + + @abstractmethod + def close(self): + pass diff --git a/memory_scope/storage/base_vector_store.py b/memory_scope/storage/base_vector_store.py deleted file mode 100644 index 34be90ca..00000000 --- a/memory_scope/storage/base_vector_store.py +++ /dev/null @@ -1,37 +0,0 @@ -from abc import ABCMeta, abstractmethod -from typing import Dict, List - -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]]) -> List[MemoryNode]: - pass - - @abstractmethod - async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]: - pass - - @abstractmethod - def insert(self, node: MemoryNode): - pass - - def insert_batch(self, nodes: List[MemoryNode]): - pass - - def delete(self, node: MemoryNode): - pass - - def update(self, node: MemoryNode): - pass - - def update_batch(self, nodes: List[MemoryNode]): - pass - - def flush(self): - pass - - def close(self): - pass diff --git a/memory_scope/storage/dummy_memory_store.py b/memory_scope/storage/dummy_memory_store.py new file mode 100644 index 00000000..40802d77 --- /dev/null +++ b/memory_scope/storage/dummy_memory_store.py @@ -0,0 +1,24 @@ +from typing import Dict, List + +from memory_scope.models.base_model import BaseModel +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.storage.base_memory_store import BaseMemoryStore + + +class DummyMemoryStore(BaseMemoryStore): + + def __init__(self, embedding_model: BaseModel, **kwargs): + self.embedding_model: BaseModel = embedding_model + self.kwargs = kwargs + + def retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]: + pass + + async def a_retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]: + pass + + def update_memories(self, nodes: MemoryNode | List[MemoryNode]): + pass + + def close(self): + pass diff --git a/memory_scope/storage/dummy_vector_store.py b/memory_scope/storage/dummy_vector_store.py deleted file mode 100644 index a766dccb..00000000 --- a/memory_scope/storage/dummy_vector_store.py +++ /dev/null @@ -1,21 +0,0 @@ -from typing import Dict, List - -from memory_scope.models.base_model import BaseModel -from memory_scope.scheme.memory_node import MemoryNode -from memory_scope.storage.base_vector_store import BaseVectorStore - - -class DummyVectorStore(BaseVectorStore): - - def __init__(self, embedding_model: BaseModel, **kwargs): - self.embedding_model: BaseModel = embedding_model - self.kwargs = kwargs - - def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]: - pass - - async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]: - pass - - def insert(self, node: MemoryNode): - pass diff --git a/memory_scope/storage/llama_index_elastic_search_store.py b/memory_scope/storage/llama_index_es_memory_store.py similarity index 69% rename from memory_scope/storage/llama_index_elastic_search_store.py rename to memory_scope/storage/llama_index_es_memory_store.py index d36ac5f8..68885615 100644 --- a/memory_scope/storage/llama_index_elastic_search_store.py +++ b/memory_scope/storage/llama_index_es_memory_store.py @@ -4,9 +4,10 @@ from llama_index.core import VectorStoreIndex from llama_index.core.schema import TextNode, NodeWithScore from llama_index.vector_stores.elasticsearch import ElasticsearchStore, AsyncDenseVectorStrategy +from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus from memory_scope.models.base_model import BaseModel from memory_scope.scheme.memory_node import MemoryNode -from memory_scope.storage.base_vector_store import BaseVectorStore +from memory_scope.storage.base_memory_store import BaseMemoryStore from memory_scope.utils.logger import Logger @@ -69,7 +70,7 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str return result -class LlamaIndexElasticSearchStore(BaseVectorStore): +class LlamaIndexEsMemoryStore(BaseMemoryStore): def __init__(self, embedding_model: BaseModel, index_name: str, @@ -87,10 +88,10 @@ class LlamaIndexElasticSearchStore(BaseVectorStore): self.index.build_index_from_nodes([TextNode(text="text")]) self.logger = Logger.get_logger() - def retrieve(self, - query: str, - top_k: int, - filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]: + def retrieve_memories(self, + query: str, + top_k: int, + filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]: if filter_dict is None: filter_dict = {} @@ -99,10 +100,10 @@ class LlamaIndexElasticSearchStore(BaseVectorStore): text_nodes = retriever.retrieve(query) return [self._text_node_2_memory_node(n) for n in text_nodes] - async def async_retrieve(self, - query: str, - top_k: int, - filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]: + async def a_retrieve_memories(self, + query: str, + top_k: int, + filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]: self.logger.info(f"query={query} top_k={top_k} filter_dict={filter_dict}") if filter_dict is None: @@ -132,6 +133,44 @@ class LlamaIndexElasticSearchStore(BaseVectorStore): def close(self): self.es_store.close() + def update_memories(self, nodes: MemoryNode | List[MemoryNode]): + if isinstance(nodes, MemoryNode): + nodes = [nodes] + + # emb & insert new memories + # TODO batch insert + new_memories = [n for n in nodes if n.status == MemoryNodeStatus.NEW] + if new_memories: + for n in new_memories: + n.status = MemoryNodeStatus.ACTIVE.value + self.insert(n) + + # emb & update new memories + # TODO insert overwrite + c_modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.CONTENT_MODIFIED] + if c_modified_memories: + for n in c_modified_memories: + n.status = MemoryNodeStatus.ACTIVE.value + self.delete(n) + self.insert(n) + + # update new memories + # TODO no emb + modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.MODIFIED] + if modified_memories: + for n in modified_memories: + n.status = MemoryNodeStatus.ACTIVE.value + self.delete(n) + self.insert(n) + + # set memories expired + expired_memories = [n for n in nodes if n.status == MemoryNodeStatus.EXPIRED] + if expired_memories: + for n in expired_memories: + n.status = MemoryNodeStatus.ACTIVE.value + self.delete(n) + self.insert(n) + @staticmethod def _memory_node_2_text_node(memory_node: MemoryNode) -> TextNode: return TextNode(id_=memory_node.memory_id, diff --git a/memory_scope/utils/datetime_handler.py b/memory_scope/utils/datetime_handler.py index 73c938ba..cdc93fc1 100644 --- a/memory_scope/utils/datetime_handler.py +++ b/memory_scope/utils/datetime_handler.py @@ -3,7 +3,6 @@ import re from typing import Dict from memory_scope.constants.language_constants import WEEKDAYS, DATATIME_WORD_LIST -from memory_scope.enumeration.language_enum import LanguageEnum from memory_scope.utils.global_context import G_CONTEXT from memory_scope.utils.logger import Logger @@ -78,26 +77,25 @@ class DatetimeHandler(object): "weekday": -1 } - # Patterns to extract the parts of the date/time # Patterns to extract the parts of the date/time patterns = { "year": r"\b(\d{4})\b", "month": r"\b(January|February|March|April|May|June|July|August|September|October|November|December)\b", - "day_month_year": r"\b(?PJanuary|February|March|April|May|June|July|August|September|October|November|December) (?P\d{1,2}),? (?P\d{4})\b", - "day_month": r"\b(?PJanuary|February|March|April|May|June|July|August|September|October|November|December) (?P\d{1,2})\b", + "day_month_year": r"\b(?PJanuary|February|March|April|May|June|July|August|September|October" + r"|November|December) (?P\d{1,2}),? (?P\d{4})\b", + "day_month": r"\b(?PJanuary|February|March|April|May|June|July|August|September|October|November" + r"|December) (?P\d{1,2})\b", "hour_12": r"\b(\d{1,2})\s*(AM|PM|am|pm)\b", "hour_24": r"\b(\d{1,2}):(\d{2}):(\d{2})\b" } month_mapping = { - "January": 1, "February": 2, "March": 3, "April": 4, - "May": 5, "June": 6, "July": 7, "August": 8, + "January": 1, "February": 2, "March": 3, "April": 4, "May": 5, "June": 6, "July": 7, "August": 8, "September": 9, "October": 10, "November": 11, "December": 12 } weekday_mapping = { - "Monday": 1, "Tuesday": 2, "Wednesday": 3, "Thursday": 4, - "Friday": 5, "Saturday": 6, "Sunday": 7 + "Monday": 1, "Tuesday": 2, "Wednesday": 3, "Thursday": 4, "Friday": 5, "Saturday": 6, "Sunday": 7 } day_month_year_match = re.search(patterns["day_month_year"], input_string) @@ -135,13 +133,6 @@ class DatetimeHandler(object): hour = 0 date_info["hour"] = hour - # # Extract 24-hour format time - # hour_24_match = re.search(patterns["hour_24"], input_string) - # if hour_24_match: - # date_info["hour"] = int(hour_24_match.group(1)) - # date_info["minute"] = int(hour_24_match.group(2)) - # date_info["second"] = int(hour_24_match.group(3)) - # Extract weekday for week_day, value in weekday_mapping.items(): if week_day in input_string: @@ -182,27 +173,15 @@ class DatetimeHandler(object): return getattr(cls, func_name)(extract_time_dict, meta_data) @classmethod - def has_time_word_cn(cls, query: str) -> bool: - # find datetime keyword + def has_time_word(cls, query: str) -> bool: contain_datetime = False - for datetime_word in DATATIME_WORD_LIST[LanguageEnum.CN]: + # TODO use re + for datetime_word in DATATIME_WORD_LIST[G_CONTEXT.language]: if datetime_word in query: contain_datetime = True break return contain_datetime - @classmethod - def has_time_word_en(cls, query: str) -> bool: - pass - - @classmethod - def has_time_word(cls, query: str) -> bool: - func_name = f"has_time_word_{G_CONTEXT.language}" - if not hasattr(cls, func_name): - cls.logger.warning(f"language={G_CONTEXT.language} needs to complete has_time_word func!") - return False - return getattr(cls, func_name)(query=query) - def datetime_format(self, dt_format: str = "%Y%m%d"): return self._dt.strftime(dt_format) diff --git a/memory_scope/utils/global_context.py b/memory_scope/utils/global_context.py index 5bfc53ca..807ae4f5 100644 --- a/memory_scope/utils/global_context.py +++ b/memory_scope/utils/global_context.py @@ -6,7 +6,7 @@ from memory_scope.enumeration.language_enum import LanguageEnum from memory_scope.memory.service.base_memory_service import BaseMemoryService from memory_scope.models.base_model import BaseModel from memory_scope.storage.base_monitor import BaseMonitor -from memory_scope.storage.base_vector_store import BaseVectorStore +from memory_scope.storage.base_memory_store import BaseMemoryStore class GlobalContext(object): @@ -18,7 +18,7 @@ class GlobalContext(object): self.model_dict: Dict[str, BaseModel] = {} self.memory_chat_dict: Dict[str, BaseMemoryChat] = {} - self.vector_store: BaseVectorStore | None = None + self.memory_store: BaseMemoryStore | None = None self.monitor: BaseMonitor | None = None self.thread_pool: ThreadPoolExecutor | None = None self.language: LanguageEnum = LanguageEnum.EN diff --git a/old/worker/es/es_insight_worker.py b/old/worker/es/es_insight_worker.py index 985aec13..cd658119 100644 --- a/old/worker/es/es_insight_worker.py +++ b/old/worker/es/es_insight_worker.py @@ -10,7 +10,7 @@ from cli import GLOBAL_CONTEXT class EsInsightWorker(MemoryBaseWorker): def _run(self): - insight_nodes = self.vector_store.retrieve( + insight_nodes = self.vector_store.retrieve_memories( size=self.kwargs.es_insight_top_k, filter_dict={ "memory_id": self.memory_id, diff --git a/old/worker/es/es_new_obs_worker.py b/old/worker/es/es_new_obs_worker.py index 35823e96..021e963f 100644 --- a/old/worker/es/es_new_obs_worker.py +++ b/old/worker/es/es_new_obs_worker.py @@ -9,7 +9,7 @@ from worker.memory_base_worker import MemoryBaseWorker class EsNewObsWorker(MemoryBaseWorker): def _run(self): - new_obs_nodes = self.vector_store.retrieve( + new_obs_nodes = self.vector_store.retrieve_memories( size=self.kwargs.es_new_obs_top_k, filter_dict={ "memory_id": self.memory_id, diff --git a/old/worker/es/es_not_reflected_worker.py b/old/worker/es/es_not_reflected_worker.py index 1a9ab0cd..10c96d7f 100644 --- a/old/worker/es/es_not_reflected_worker.py +++ b/old/worker/es/es_not_reflected_worker.py @@ -11,7 +11,7 @@ class EsNotReflectedWorker(MemoryBaseWorker): def _run(self): - not_reflected_obs_nodes = self.vector_store.retrieve( + not_reflected_obs_nodes = self.vector_store.retrieve_memories( size=self.kwargs.es_new_obs_top_k, filter_dict={ "memory_id": self.memory_id, diff --git a/old/worker/es/es_similar_worker.py b/old/worker/es/es_similar_worker.py index ddca69b9..2b419398 100644 --- a/old/worker/es/es_similar_worker.py +++ b/old/worker/es/es_similar_worker.py @@ -15,7 +15,7 @@ class EsSimilarWorker(MemoryBaseWorker): def _run(self): query = self.messages[-1].content - similar_obs_nodes = self.vector_store.retrieve( + similar_obs_nodes = self.vector_store.retrieve_memories( text=query, size=self.es_similar_top_k, filter_dict={ diff --git a/old/worker/es/es_today_obs_worker.py b/old/worker/es/es_today_obs_worker.py index 8cd1ce54..34f65f02 100644 --- a/old/worker/es/es_today_obs_worker.py +++ b/old/worker/es/es_today_obs_worker.py @@ -18,7 +18,7 @@ class EsTodayObsWorker(MemoryBaseWorker): self.logger.warning("messages is empty!") return msg_time_created = self.messages[-1].time_created - today_obs_nodes = self.vector_store.retrieve( + today_obs_nodes = self.vector_store.retrieve_memories( size=self.es_today_obs_top_k, filter_dict={ "memory_id": self.memory_id,