diff --git a/config/demo_config.yaml b/config/demo_config.yaml index d22f0723..600cf834 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -1,5 +1,5 @@ global_config: - language: en + language: cn max_workers: 5 dash_scope_apikey: open_ai_apikey: @@ -8,8 +8,8 @@ memory_chat: class: chat.cli_memory_chat memory_service: memory_chat_service generation_model: dashscope_generation - human_name: human - assistant_name: assistant + human_name: 用户 + assistant_name: AI memory_service: memory_chat_service: class: memory.service.chat_memory_service diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index a62cc1f7..bb8c6317 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -60,5 +60,8 @@ class BaseWorker(metaclass=ABCMeta): else: self.context[key] = value + def has_content(self, key: str): + return key in self.context + def __getattr__(self, key: str): return self.kwargs[key] diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py index aec0d71e..abcf8945 100644 --- a/memory_scope/memory/worker/write/contra_repeat_worker.py +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -1,6 +1,6 @@ from typing import List -from memory_scope.constants.common_constants import NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES +from memory_scope.constants.common_constants import NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, MERGE_OBS_NODES from memory_scope.constants.language_constants import NONE_WORD, CONTRADICTORY_WORD, INCLUDED_WORD from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum @@ -116,4 +116,4 @@ class ContraRepeatWorker(MemoryBaseWorker): self.logger.info(f"contra_repeat stage: {node.content} {node.status}") # save context - self.vector_store.update_batch(merge_obs_nodes) + self.set_context(MERGE_OBS_NODES, merge_obs_nodes) diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index fc6f9e44..9f3eae19 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -31,7 +31,7 @@ class GetObservationWorker(MemoryBaseWorker): target_name=self.target_name, meta_data=meta_data, content=obs_content, - memoryType=MemoryTypeEnum.OBSERVATION.value, + memory_type=MemoryTypeEnum.OBSERVATION.value, status=MemoryNodeStatus.ACTIVE.value, timestamp=message.time_created, obs_dt=dt_handler.datetime_format(), diff --git a/memory_scope/memory/worker/write/store_memory_worker.py b/memory_scope/memory/worker/write/store_memory_worker.py new file mode 100644 index 00000000..a9742ad7 --- /dev/null +++ b/memory_scope/memory/worker/write/store_memory_worker.py @@ -0,0 +1,36 @@ +from typing import List + +from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus +from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum +from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.utils.datetime_handler import DatetimeHandler + + +class StoreMemoryWorker(MemoryBaseWorker): + + def _run(self): + store_key: str = self.store_key + + if self.has_content(store_key): + memory_nodes: List[MemoryNode] = self.get_context(store_key) + self.vector_store.update_batch(memory_nodes) + + elif store_key in self.chat_kwargs: + query = self.chat_kwargs[store_key] + query = query.strip() + if not query: + return + + dt_handler = DatetimeHandler() + node = MemoryNode( + user_name=self.user_name, + target_name=self.target_name, + content=query, + memory_type=MemoryTypeEnum.OBSERVATION.value, + status=MemoryNodeStatus.ACTIVE.value, + timestamp=dt_handler.timestamp, + obs_dt=dt_handler.datetime_format(), + obs_reflected=False, + obs_profile_updated=False) + self.vector_store.update(node) diff --git a/memory_scope/utils/datetime_handler.py b/memory_scope/utils/datetime_handler.py index 58cebd15..b0384644 100644 --- a/memory_scope/utils/datetime_handler.py +++ b/memory_scope/utils/datetime_handler.py @@ -101,3 +101,7 @@ class DatetimeHandler(object): def string_format(self, string_format: str): return string_format.format(**self.dt_info_dict) + + @property + def timestamp(self) -> int: + return int(self._dt.timestamp())