mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
[dev] change dummy workflow to actual workflow
This commit is contained in:
parent
835a9d5752
commit
4d6be46a1e
27 changed files with 188 additions and 187 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -1,9 +0,0 @@
|
|||
from enum import Enum
|
||||
|
||||
|
||||
class MemoryRecallType(str, Enum):
|
||||
SIMILAR = "similar"
|
||||
|
||||
KEYWORD = "keyword"
|
||||
|
||||
PROFILE = "profile"
|
||||
|
|
@ -2,6 +2,12 @@ from enum import Enum
|
|||
|
||||
|
||||
class MemoryNodeStatus(str, Enum):
|
||||
NEW = "new"
|
||||
|
||||
MODIFIED = "modified"
|
||||
|
||||
CONTENT_MODIFIED = "content_modified"
|
||||
|
||||
ACTIVE = "active"
|
||||
|
||||
EXPIRED = "expired"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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] = []
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
33
memory_scope/storage/base_memory_store.py
Normal file
33
memory_scope/storage/base_memory_store.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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
|
||||
24
memory_scope/storage/dummy_memory_store.py
Normal file
24
memory_scope/storage/dummy_memory_store.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
|
|
@ -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(?P<month>January|February|March|April|May|June|July|August|September|October|November|December) (?P<day>\d{1,2}),? (?P<year>\d{4})\b",
|
||||
"day_month": r"\b(?P<month>January|February|March|April|May|June|July|August|September|October|November|December) (?P<day>\d{1,2})\b",
|
||||
"day_month_year": r"\b(?P<month>January|February|March|April|May|June|July|August|September|October"
|
||||
r"|November|December) (?P<day>\d{1,2}),? (?P<year>\d{4})\b",
|
||||
"day_month": r"\b(?P<month>January|February|March|April|May|June|July|August|September|October|November"
|
||||
r"|December) (?P<day>\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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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={
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue