diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py index 29cb09b6..7309af49 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -28,7 +28,8 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): self._vector_store: BaseVectorStore | None = None self._monitor: BaseMonitor | None = None - self._user_id: str | None = None + self._user_name: str | None = None + self._target_name: str | None = None self._prompt_handler: PromptHandler | None = None @property @@ -74,11 +75,20 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): return self._monitor @property - def user_id(self) -> str: - if self._user_id is None: + def user_name(self) -> str: + # FIXME: complex situations require adjustment + if self._user_name is None: + message = [x for x in self.messages if x.role == MessageRoleEnum.ASSISTANT.value][-1] + self._user_name = message.role_name + return self._user_name + + @property + def target_name(self) -> str: + # FIXME: complex situations require adjustment + if self._target_name is None: message = [x for x in self.messages if x.role == MessageRoleEnum.USER.value][-1] - self._user_id = message.role_name - return self._user_id + self._target_name = message.role_name + return self._target_name @property def prompt_handler(self) -> PromptHandler: diff --git a/memory_scope/memory/worker/read/extract_time_worker.py b/memory_scope/memory/worker/read/extract_time_worker.py index 25fbe6b0..f2368f11 100644 --- a/memory_scope/memory/worker/read/extract_time_worker.py +++ b/memory_scope/memory/worker/read/extract_time_worker.py @@ -30,11 +30,7 @@ class ExtractTimeWorker(MemoryBaseWorker): self.logger.info(f"extract_time_prompt={extract_time_prompt}") # call sft model - response = self.generation_model.call(prompt=extract_time_prompt, - model_name=self.extra_time_model, - max_token=self.extra_time_max_token, - temperature=self.extra_time_temperature, - top_k=self.extra_time_top_k) + response = self.generation_model.call(prompt=extract_time_prompt, top_k=self.extra_time_top_k) # if empty, return if not response.status or not response.message.content: @@ -47,6 +43,5 @@ class ExtractTimeWorker(MemoryBaseWorker): for key, value in matches: if key in DATATIME_KEY_MAP.keys(): extract_time_dict[DATATIME_KEY_MAP[key]] = value - self.set_context(EXTRACT_TIME_DICT, extract_time_dict) - self.logger.info(f"response_text={response_text} filters={extract_time_dict}") + self.set_context(EXTRACT_TIME_DICT, extract_time_dict) diff --git a/memory_scope/memory/worker/read/fuse_rerank_worker.py b/memory_scope/memory/worker/read/fuse_rerank_worker.py index b993e593..639afeee 100644 --- a/memory_scope/memory/worker/read/fuse_rerank_worker.py +++ b/memory_scope/memory/worker/read/fuse_rerank_worker.py @@ -1,8 +1,10 @@ from typing import Dict, List from memory_scope.constants.common_constants import EXTRACT_TIME_DICT, RANKED_MEMORY_NODES, RESULT +from memory_scope.constants.language_constants import COLON_WORD from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.utils.datetime_handler import DatetimeHandler class FuseRerankWorker(MemoryBaseWorker): @@ -107,10 +109,10 @@ class FuseRerankWorker(MemoryBaseWorker): content = node.content if f_event or f_msg: - time_infer = self.format_time_infer(time_infer="", - extract_time_dict=extract_time_dict, - meta_data=node.memory_node.metaData) - content = f"{time_infer}: {content}" + time_infer = DatetimeHandler.format_time_by_extract_time(extract_time_dict=extract_time_dict, + meta_data=node.memory_node.metaData) + if time_infer: + content = f"{time_infer}{self.get_language_value(COLON_WORD)}{content}" memories.append(content) self.set_context(RESULT, "\n".join(memories)) diff --git a/memory_scope/memory/worker/read/retrieve_store_worker.py b/memory_scope/memory/worker/read/retrieve_store_worker.py index 6462a1e2..20e21dd5 100644 --- a/memory_scope/memory/worker/read/retrieve_store_worker.py +++ b/memory_scope/memory/worker/read/retrieve_store_worker.py @@ -11,7 +11,8 @@ class RetrieveStoreWorker(MemoryBaseWorker): async def retrieve_from_observation(self, query: str) -> List[MemoryNode]: filter_dict = { - "user_id": self.user_id, + "user_name": self.user_name, + "target_name": self.target_name, "status": MemoryNodeStatus.ACTIVE.value, "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], } @@ -21,7 +22,8 @@ class RetrieveStoreWorker(MemoryBaseWorker): async def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]: filter_dict = { - "user_id": self.user_id, + "user_name": self.user_name, + "target_name": self.target_name, "status": MemoryNodeStatus.ACTIVE.value, "memory_type": [MemoryTypeEnum.INSIGHT.value, MemoryTypeEnum.PROFILE.value], } diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index cc027a08..b9897a4c 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -1,4 +1,3 @@ -from datetime import datetime from typing import List from memory_scope.constants.common_constants import NEW_OBS_NODES, TIME_INFER @@ -15,21 +14,21 @@ from memory_scope.utils.tool_functions import prompt_to_msg class GetObservationWorker(MemoryBaseWorker): def add_observation(self, message: Message, time_infer: str, obs_content: str, keywords: str): - created_dt: datetime = datetime.fromtimestamp(float(message.time_created)) - dt_handler = DatetimeHandler(dt=created_dt) + dt_handler = DatetimeHandler(dt=message.time_created) - # 组合meta_data + # buidl meta data meta_data = { MemoryTypeEnum.CONVERSATION.value: message.content, TIME_INFER: time_infer, - **{f"msg_{k}": str(v) for k, v in dt_handler.dt_info_dict.items()}, + **dt_handler.dt_info_dict.items(), } if time_infer: - dt_infer_handler = DatetimeHandler(dt=time_infer) - meta_data.update({f"event_{k}": str(v) for k, v in dt_infer_handler.dt_info_dict.items()}) + dt_info_dict = DatetimeHandler.extract_date_parts(input_string=time_infer) + meta_data.update({f"event_{k}": str(v) for k, v in dt_info_dict.items()}) - node = MemoryNode(user_id=self.user_id, + node = MemoryNode(user_name=self.user_name, + target_name=self.target_name, meta_data=meta_data, content=obs_content, memoryType=MemoryTypeEnum.OBSERVATION.value, @@ -38,7 +37,7 @@ class GetObservationWorker(MemoryBaseWorker): obs_dt=dt_handler.datetime_format(), obs_reflected=False, obs_profile_updated=False, - keywords=keywords) + obs_keyword=keywords) node.gen_memory_id() return node diff --git a/memory_scope/scheme/memory_node.py b/memory_scope/scheme/memory_node.py index 76a88370..f193a124 100644 --- a/memory_scope/scheme/memory_node.py +++ b/memory_scope/scheme/memory_node.py @@ -7,11 +7,13 @@ from memory_scope.utils.tool_functions import md5_hash class MemoryNode(BaseModel): - memory_id: str = Field("", description="unique id for memory item") + memory_id: str = Field("", description="unique id for memory") - user_id: str = Field("", description="unique memory id for user") + user_name: str = Field("", description="the user who owns the memory") - meta_data: Dict[str, str] = Field({}, description="other data infos") + target_name: str = Field("", description="target name described by the memory") + + meta_data: Dict[str, str] = Field({}, description="meta data infos") content: str = Field("", description="memory content") @@ -35,8 +37,7 @@ class MemoryNode(BaseModel): obs_profile_updated: bool = Field(False, description="if the observation has updated user profile") - keyword: str = Field("", description="keywords of the content") - + obs_keyword: str = Field("", description="keywords of the content") @property def node_keys(self): @@ -46,4 +47,4 @@ class MemoryNode(BaseModel): return self.model_dump().get(key) def gen_memory_id(self): - self.memory_id = f"{self.user_id}_{self.timestamp}_{md5_hash(self.content)[:8]}" + self.memory_id = f"{self.user_name}_{self.target_name}_{self.timestamp}_{md5_hash(self.content)[:8]}" diff --git a/memory_scope/storage/llama_index_elastic_search_store.py b/memory_scope/storage/llama_index_elastic_search_store.py index 0c8f481f..f428eb79 100644 --- a/memory_scope/storage/llama_index_elastic_search_store.py +++ b/memory_scope/storage/llama_index_elastic_search_store.py @@ -92,9 +92,7 @@ class LlamaIndexElasticSearchStore(BaseVectorStore): filter_dict = {} es_filter = _to_elasticsearch_filter(filter_dict) - retriever = self.index.as_retriever( - vector_store_kwargs={"es_filter": es_filter}, - similarity_top_k=top_k) + retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter}, similarity_top_k=top_k) text_nodes = retriever.retrieve(query) return [self._text_node_2_memory_node(n) for n in text_nodes] diff --git a/memory_scope/utils/datetime_handler.py b/memory_scope/utils/datetime_handler.py index a36861ab..58cebd15 100644 --- a/memory_scope/utils/datetime_handler.py +++ b/memory_scope/utils/datetime_handler.py @@ -1,5 +1,6 @@ import datetime import re +from typing import Dict from memory_scope.constants.language_constants import WEEKDAYS from memory_scope.utils.global_context import G_CONTEXT @@ -7,6 +8,7 @@ from memory_scope.utils.logger import Logger class DatetimeHandler(object): + logger = Logger.get_logger() def __init__(self, dt: datetime.datetime | str | int | float = None): if isinstance(dt, str | int | float): @@ -19,7 +21,6 @@ class DatetimeHandler(object): self._dt: datetime.datetime = datetime.datetime.now() self._dt_info_dict: dict | None = None - self.logger = Logger.get_logger() def _parse_dt_info(self): return { @@ -39,8 +40,8 @@ class DatetimeHandler(object): self._dt_info_dict = self._parse_dt_info() return self._dt_info_dict - @staticmethod - def extract_date_parts_cn(input_string: str): + @classmethod + def extract_date_parts_cn(cls, input_string: str) -> dict: # Extending our pattern to handle every/每 as a possible value. patterns = { "year": r"(\d+|每)年", @@ -64,12 +65,36 @@ class DatetimeHandler(object): extracted_data[key] = int(match.group(1)) return extracted_data - def extract_date_parts(self): + @classmethod + def extract_date_parts(cls, input_string: str) -> dict: func_name = f"extract_date_parts_{G_CONTEXT.language}" - if not hasattr(self, func_name): - self.logger.warning(f"language={G_CONTEXT.language} needs to complete extract_date_parts function!") + if not hasattr(cls, func_name): + cls.logger.warning(f"language={G_CONTEXT.language} needs to complete extract_date_parts func!") return {} - return getattr(self, func_name)() + return getattr(cls, func_name)(input_string=input_string) + + @classmethod + def format_time_by_extract_time_cn(cls, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]) -> str: + cn_key_dict = {"year": "年", "month": "月", "day": "日", "weekday": ""} + format_time_str = "" + for key, value_cn in cn_key_dict.items(): + if key in extract_time_dict and key in meta_data: + value = meta_data[key] + if value_cn: + if value == "-1": + value = "每" + format_time_str += f"{value}{value_cn}" + else: + format_time_str += value + return format_time_str + + @classmethod + def format_time_by_extract_time(cls, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]) -> str: + func_name = f"format_time_by_extract_time_{G_CONTEXT.language}" + if not hasattr(cls, func_name): + cls.logger.warning(f"language={G_CONTEXT.language} needs to complete format_time_by_extract_time func!") + return "" + return getattr(cls, func_name)(extract_time_dict, meta_data) def datetime_format(self, dt_format: str = "%Y%m%d"): return self._dt.strftime(dt_format)