From aeeb6012da870519f1159612becaf2b039bb8e8f Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 27 Jun 2024 11:14:16 +0800 Subject: [PATCH] [dev] update cli memory chat --- memory_scope/chat_v2/cli_memory_chat.py | 8 +++++--- memory_scope/memory/operation/base_operation.py | 3 +++ memory_scope/memory/operation/summary_memory.py | 6 +++++- memory_scope/memory/operation/write_memory.py | 3 +++ memory_scope/memory/service/base_memory_service.py | 6 +++++- memory_scope/memory/worker/base_worker.py | 5 ++--- memory_scope/memory/worker/dummy_worker.py | 2 +- 7 files changed, 24 insertions(+), 9 deletions(-) diff --git a/memory_scope/chat_v2/cli_memory_chat.py b/memory_scope/chat_v2/cli_memory_chat.py index f9e60d8a..1cf6db9e 100644 --- a/memory_scope/chat_v2/cli_memory_chat.py +++ b/memory_scope/chat_v2/cli_memory_chat.py @@ -21,9 +21,9 @@ class CliMemoryChat(BaseMemoryChat): } def __init__(self, memory_service: str, generation_model: str, **kwargs): + super().__init__(**kwargs) self._memory_service: BaseMemoryService | str = memory_service self._generation_model: BaseModel | str = generation_model - self.kwargs: dict = kwargs @property def memory_service(self) -> BaseMemoryService: @@ -43,7 +43,9 @@ class CliMemoryChat(BaseMemoryChat): system_prompt = SYSTEM_PROMPT[G_CONTEXT.language] if related_memories: memory_prompt = MEMORY_PROMPT[G_CONTEXT.language] - system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt] + related_memories]) + all_prompt_list = [system_prompt, memory_prompt] + all_prompt_list.extend(related_memories) + system_prompt = "\n".join([x.strip() for x in all_prompt_list]) return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created) def chat_with_memory(self, query: str): @@ -53,7 +55,7 @@ class CliMemoryChat(BaseMemoryChat): time_created = int(datetime.datetime.now().timestamp()) new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created) - related_memories: List[str] = self.memory_service.do_operation("read_memory") + related_memories: List[str] = self.memory_service.read_memory() system_message: Message = self.get_system_prompt(related_memories, time_created) return self.generation_model.call(messages=[system_message, new_message], stream=True) diff --git a/memory_scope/memory/operation/base_operation.py b/memory_scope/memory/operation/base_operation.py index 3f531b4e..2287f9e3 100644 --- a/memory_scope/memory/operation/base_operation.py +++ b/memory_scope/memory/operation/base_operation.py @@ -20,3 +20,6 @@ class BaseOperation(metaclass=ABCMeta): def run_operation_backend(self): pass + + def stop_operation_backend(self): + pass diff --git a/memory_scope/memory/operation/summary_memory.py b/memory_scope/memory/operation/summary_memory.py index 30dc96c8..9eac30aa 100644 --- a/memory_scope/memory/operation/summary_memory.py +++ b/memory_scope/memory/operation/summary_memory.py @@ -21,6 +21,7 @@ class SummaryMemory(BaseWorkflow, BaseOperation): self._operation_status_run: bool = False self._loop_switch: bool = False + self._run_thread = None def init_workflow(self): self.init_workers() @@ -44,4 +45,7 @@ class SummaryMemory(BaseWorkflow, BaseOperation): def run_operation_backend(self): if not self._loop_switch: self._loop_switch = True - return G_CONTEXT.thread_pool.submit(self._loop_operation) + self._run_thread = G_CONTEXT.thread_pool.submit(self._loop_operation) + + def stop_operation_backend(self): + self._loop_switch = False diff --git a/memory_scope/memory/operation/write_memory.py b/memory_scope/memory/operation/write_memory.py index 1e9070b7..92bc8b47 100644 --- a/memory_scope/memory/operation/write_memory.py +++ b/memory_scope/memory/operation/write_memory.py @@ -73,3 +73,6 @@ class WriteMemory(BaseOperation, BaseWorkflow): if not self._loop_switch: self._loop_switch = True return G_CONTEXT.thread_pool.submit(self._loop_operation) + + def stop_operation_backend(self): + self._loop_switch = False diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index e2d2a452..05ee4b54 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -6,7 +6,8 @@ from memory_scope.utils.logger import Logger class BaseMemoryService(metaclass=ABCMeta): - def __init__(self, **kwargs): + def __init__(self, read_memory_key: str = "read_memory", **kwargs): + self.read_memory_key: str = read_memory_key self.logger = Logger.get_logger() self.kwargs = kwargs @@ -23,3 +24,6 @@ class BaseMemoryService(metaclass=ABCMeta): @abstractmethod def get_op_description_dict(self) -> Dict[str, str]: pass + + def read_memory(self): + return self.do_operation(self.read_memory_key) diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index 78bfae5e..9f9ed979 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -43,15 +43,14 @@ class BaseWorker(metaclass=ABCMeta): self.logger.info(f"----- worker_{self.name}_end cost={t.cost_str}-----") def get_context(self, key: str, default=None): - return self.context_dict.get(key, default) + return self.context.get(key, default) def set_context(self, key: str, value: Any): if self.is_multi_thread: with self.context_lock: self.context_dict[key] = value else: - self.context_dict[key] = value + self.context[key] = value def __getattr__(self, key): - # raise exception if not exists return self.kwargs[key] diff --git a/memory_scope/memory/worker/dummy_worker.py b/memory_scope/memory/worker/dummy_worker.py index 9513bc79..d0cb7d93 100644 --- a/memory_scope/memory/worker/dummy_worker.py +++ b/memory_scope/memory/worker/dummy_worker.py @@ -5,4 +5,4 @@ from memory_scope.memory.worker.base_worker import BaseWorker class DummyWorker(BaseWorker): def _run(self): self.set_context(RESULT, ["test 123"]) - self.logger.info("enter dummy worker!") \ No newline at end of file + self.logger.info("enter dummy worker!")