[dev] change dummy workflow to actual workflow

This commit is contained in:
jinli.yl 2024-07-07 23:18:31 +08:00
parent 835a9d5752
commit 4d6be46a1e
27 changed files with 188 additions and 187 deletions

View file

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

View file

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

View file

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

View file

@ -1,9 +0,0 @@
from enum import Enum
class MemoryRecallType(str, Enum):
SIMILAR = "similar"
KEYWORD = "keyword"
PROFILE = "profile"

View file

@ -2,6 +2,12 @@ from enum import Enum
class MemoryNodeStatus(str, Enum):
NEW = "new"
MODIFIED = "modified"
CONTENT_MODIFIED = "content_modified"
ACTIVE = "active"
EXPIRED = "expired"

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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] = []

View file

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

View file

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

View file

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

View file

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

View 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

View file

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

View 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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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={

View file

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