From 523b752a3ddbbb656d122f500c0ec9ebee4aca7a Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Mon, 8 Jul 2024 23:26:10 +0800 Subject: [PATCH] [dev] fix demo config --- config/demo_config.yaml | 5 +- memory_scope/chat/cli_memory_chat.py | 25 +++++--- memory_scope/chat/cli_memory_chat.yaml | 7 ++- .../operation/base_backend_operation.py | 57 +++++++++++++++++++ .../memory/operation/base_operation.py | 6 -- memory_scope/memory/operation/read_memory.py | 1 + .../memory/operation/summary_memory.py | 50 +++------------- memory_scope/memory/operation/write_memory.py | 49 ++++------------ .../memory/service/base_memory_service.py | 8 ++- .../memory/worker/memory_base_worker.py | 7 ++- .../memory/worker/read/extract_time_worker.py | 1 + .../memory/worker/read/print_memory_worker.py | 14 ++--- .../memory/worker/read/set_query_worker.py | 10 ++-- .../summary/get_reflection_subject_worker.py | 1 + .../summary/long_contra_repeat_worker.py | 1 + .../worker/summary/update_insight_worker.py | 1 + .../worker/write/contra_repeat_worker.py | 2 + .../write/get_observation_with_time_worker.py | 1 + .../worker/write/get_observation_worker.py | 2 + .../memory/worker/write/info_filter_worker.py | 4 ++ .../memory/worker/write/load_memory_worker.py | 2 +- .../models/llama_index_generation_model.py | 7 ++- memory_scope/scheme/memory_node.py | 3 +- memory_scope/utils/prompt_handler.py | 4 +- 24 files changed, 145 insertions(+), 123 deletions(-) create mode 100644 memory_scope/memory/operation/base_backend_operation.py diff --git a/config/demo_config.yaml b/config/demo_config.yaml index ab1e8c46..3b2bc4c0 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -13,7 +13,6 @@ memory_service: class: memory.service.chat_memory_service history_msg_count: 32 contextual_msg_count: 6 - read_memory_key: read_memory memory_operations: read_message: class: memory.operation.read_message @@ -30,12 +29,12 @@ memory_service: class: memory.operation.write_memory workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory description: "write observation memories of the user" - interval_time: 60 + interval_time: 5 summary_memory: class: memory.operation.summary_memory workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory description: "summary observation memories of the user" - interval_time: 300 + interval_time: 60 worker: dummy: diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index a9501cf2..585ff831 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -1,5 +1,6 @@ import os import time +from typing import List import questionary @@ -75,21 +76,29 @@ class CliMemoryChat(BaseMemoryChat): self._generation_model = G_CONTEXT.model_dict[self._generation_model] return self._generation_model - @property - def system_prompt_with_memory(self) -> Message: - system_prompt = self.prompt_handler.system_prompt + def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen: + new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query) + self.memory_service.add_messages(new_message) + messages: List[Message] = [] + + # add memory to system prompt + system_prompt = self.prompt_handler.system_prompt memories: str = self.memory_service.read_memory() if memories: memory_prompt = self.prompt_handler.memory_prompt system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]]) + messages.append(Message(role=MessageRoleEnum.SYSTEM, content=system_prompt)) - return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt) + # add history messages + history_messages = self.memory_service.read_message() + if history_messages: + messages.extend(history_messages) - def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen: - new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query) - self.memory_service.add_messages(new_message) - return self.generation_model.call(messages=[self.system_prompt_with_memory, new_message], stream=self.stream) + # add new_message + messages.append(new_message) + self.logger.info(f"messages={messages}") + return self.generation_model.call(messages=messages, stream=self.stream) @staticmethod def parse_query_command(query: str): diff --git a/memory_scope/chat/cli_memory_chat.yaml b/memory_scope/chat/cli_memory_chat.yaml index 37b7b02a..5bf8da82 100644 --- a/memory_scope/chat/cli_memory_chat.yaml +++ b/memory_scope/chat/cli_memory_chat.yaml @@ -1,8 +1,11 @@ system_prompt: cn: | - 你是一个可靠的小助手,你的名字叫MemoryScope - + 你是一个可靠的小助手,你的名字叫MemoryScope。 + en: | + You are a helpful assistant, your name is MemoryScope. memory_prompt: cn: | 请记住以下信息,他们可以帮助更好地理解用户的问题。 + en: | + Please remember the following information, as they can help better understand the user's question. \ No newline at end of file diff --git a/memory_scope/memory/operation/base_backend_operation.py b/memory_scope/memory/operation/base_backend_operation.py new file mode 100644 index 00000000..0cfc43a2 --- /dev/null +++ b/memory_scope/memory/operation/base_backend_operation.py @@ -0,0 +1,57 @@ +import time +from abc import abstractmethod + +from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE +from memory_scope.utils.global_context import G_CONTEXT +from memory_scope.utils.logger import Logger + + +class BaseBackendOperation(BaseOperation): + operation_type: OPERATION_TYPE = "backend" + + def __init__(self, interval_time: int, **kwargs): + super(BaseBackendOperation, self).__init__(**kwargs) + + self.interval_time: int = interval_time + + self._operation_status_run: bool = False + self._loop_switch: bool = False + self._run_thread = None + + self.logger = Logger.get_logger() + + @abstractmethod + def _run_operation(self, **kwargs): + raise NotImplementedError + + def run_operation(self, **kwargs): + if self._operation_status_run: + return + + self._operation_status_run = True + result = None + try: + result = self._run_operation(**kwargs) + except Exception as e: + self.logger.exception(f"{self.name} encounter exception. args={e.args}") + + self._operation_status_run = False + return result + + def _loop_operation(self): + while self._loop_switch: + for _ in range(self.interval_time): + if self._loop_switch: + time.sleep(1) + else: + break + if self._loop_switch: + self.run_operation() + + def run_operation_backend(self): + if not self._loop_switch: + self._loop_switch = True + 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/base_operation.py b/memory_scope/memory/operation/base_operation.py index a58569f7..24f3f536 100644 --- a/memory_scope/memory/operation/base_operation.py +++ b/memory_scope/memory/operation/base_operation.py @@ -18,9 +18,3 @@ class BaseOperation(metaclass=ABCMeta): @abstractmethod def run_operation(self, **kwargs): raise NotImplementedError - - def run_operation_backend(self): - pass - - def stop_operation_backend(self): - pass diff --git a/memory_scope/memory/operation/read_memory.py b/memory_scope/memory/operation/read_memory.py index b51bbd6c..0b8cd944 100644 --- a/memory_scope/memory/operation/read_memory.py +++ b/memory_scope/memory/operation/read_memory.py @@ -17,6 +17,7 @@ class ReadMemory(BaseWorkflow, BaseOperation): **kwargs): super().__init__(name=name, **kwargs) BaseOperation.__init__(self, name=name, description=description) + self.chat_messages: List[Message] = chat_messages self.his_msg_count: int = his_msg_count diff --git a/memory_scope/memory/operation/summary_memory.py b/memory_scope/memory/operation/summary_memory.py index ffb1c297..b713ebea 100644 --- a/memory_scope/memory/operation/summary_memory.py +++ b/memory_scope/memory/operation/summary_memory.py @@ -1,58 +1,22 @@ -import time - from memory_scope.constants.common_constants import RESULT, CHAT_KWARGS -from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE +from memory_scope.memory.operation.base_backend_operation import BaseBackendOperation +from memory_scope.memory.operation.base_operation import OPERATION_TYPE from memory_scope.memory.operation.base_workflow import BaseWorkflow -from memory_scope.utils.global_context import G_CONTEXT -class SummaryMemory(BaseWorkflow, BaseOperation): +class SummaryMemory(BaseWorkflow, BaseBackendOperation): operation_type: OPERATION_TYPE = "backend" - def __init__(self, - name: str, - description: str, - interval_time: int = 300, - **kwargs): - super().__init__(name=name, **kwargs) - BaseOperation.__init__(self, name=name, description=description) - - self.interval_time: int = interval_time - - self._operation_status_run: bool = False - self._loop_switch: bool = False - self._run_thread = None + def __init__(self, **kwargs): + super().__init__(**kwargs) + BaseBackendOperation.__init__(self, **kwargs) def init_workflow(self): self.init_workers() - def run_operation(self, **kwargs): - if self._operation_status_run: - return - - self._operation_status_run = True + def _run_operation(self, **kwargs): self.context[CHAT_KWARGS] = kwargs self.run_workflow() result = self.context.get(RESULT) self.context.clear() - self._operation_status_run = False - return result - - def _loop_operation(self): - while self._loop_switch: - for _ in range(self.interval_time): - if self._loop_switch: - time.sleep(1) - else: - break - if self._loop_switch: - self.run_operation() - - def run_operation_backend(self): - if not self._loop_switch: - self._loop_switch = True - 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 9a90db0b..e32eb7a3 100644 --- a/memory_scope/memory/operation/write_memory.py +++ b/memory_scope/memory/operation/write_memory.py @@ -1,38 +1,30 @@ -import time from typing import List from memory_scope.constants.common_constants import CHAT_MESSAGES, RESULT, CHAT_KWARGS -from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE +from memory_scope.memory.operation.base_backend_operation import BaseBackendOperation +from memory_scope.memory.operation.base_operation import OPERATION_TYPE from memory_scope.memory.operation.base_workflow import BaseWorkflow from memory_scope.scheme.message import Message -from memory_scope.utils.global_context import G_CONTEXT -class WriteMemory(BaseWorkflow, BaseOperation): +class WriteMemory(BaseWorkflow, BaseBackendOperation): operation_type: OPERATION_TYPE = "backend" def __init__(self, - name: str, - description: str, chat_messages: List[Message], his_msg_count: int = 0, message_lock=None, - interval_time: int = 60, contextual_msg_count: int = 6, **kwargs): - super().__init__(name=name, **kwargs) - BaseOperation.__init__(self, name=name, description=description) + super().__init__(**kwargs) + BaseBackendOperation.__init__(self, **kwargs) self.chat_messages: List[Message] = chat_messages self.his_msg_count: int = his_msg_count self.message_lock = message_lock - self.interval_time: int = interval_time self.contextual_msg_count: int = contextual_msg_count - self._operation_status_run: bool = False - self._loop_switch: bool = False - @property def not_memorized_size(self): return sum([not x.memorized for x in self.chat_messages]) @@ -46,40 +38,19 @@ class WriteMemory(BaseWorkflow, BaseOperation): def init_workflow(self): self.init_workers() - def run_operation(self, **kwargs): - if self._operation_status_run: - return - - self._operation_status_run = True + def _run_operation(self, **kwargs): self.context[CHAT_KWARGS] = kwargs not_memorized_size = self.not_memorized_size if not_memorized_size < self.contextual_msg_count: + self.logger.info(f"not_memorized_size={not_memorized_size} < " + f"contextual_msg_count={self.contextual_msg_count}, skip.") return + self._operation_status_run = True max_count = not_memorized_size + self.his_msg_count self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]] self.run_workflow() result = self.context.get(RESULT) self.context.clear() self.set_memorized() - self._operation_status_run = False - - return result - - def _loop_operation(self): - while self._loop_switch: - for _ in range(self.interval_time): - if self._loop_switch: - time.sleep(1) - else: - break - if self._loop_switch: - self.run_operation() - - def run_operation_backend(self): - 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 + return result \ No newline at end of file diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index bfac4780..17e18530 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -11,14 +11,16 @@ class BaseMemoryService(metaclass=ABCMeta): def __init__(self, memory_operations: Dict[str, dict], read_memory_key: str = "read_memory", + read_message_key: str = "read_message", **kwargs): self.memory_operations: Dict[str, dict] = memory_operations self.read_memory_key: str = read_memory_key + self.read_message_key: str = read_message_key self._operation_dict: Dict[str, BaseOperation] = {} self._op_description_dict: Dict[str, str] = {} self.chat_messages: List[Message] = [] - self.message_lock = threading.Lock + self.message_lock = threading.Lock() self.logger = Logger.get_logger() self.kwargs = kwargs @@ -48,5 +50,9 @@ class BaseMemoryService(metaclass=ABCMeta): assert self.read_memory_key in self._operation_dict, f"op={self.read_memory_key} is not inited!" return self.do_operation(self.read_memory_key) + def read_message(self): + assert self.read_message_key in self._operation_dict, f"op={self.read_message_key} is not inited!" + return self.do_operation(self.read_message_key) + def stop_service(self): pass diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py index e98af783..876e67a7 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -13,6 +13,7 @@ from memory_scope.utils.prompt_handler import PromptHandler class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): + FILE_PATH: str = __file__ def __init__(self, embedding_model: str = "", @@ -119,19 +120,19 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): @property def user_name(self) -> str: if self._user_name is None: - self._user_name = G_CONTEXT.meta_data["human_name"] + self._user_name = G_CONTEXT.meta_data["assistant_name"] return self._user_name @property def target_name(self) -> str: if self._target_name is None: - self._target_name = G_CONTEXT.meta_data["assistant_name"] + self._target_name = G_CONTEXT.meta_data["human_name"] return self._target_name @property def prompt_handler(self) -> PromptHandler: if self._prompt_handler is None: - self._prompt_handler = PromptHandler(__file__, **self.kwargs) + self._prompt_handler = PromptHandler(self.FILE_PATH, **self.kwargs) return self._prompt_handler def __getattr__(self, key: str): diff --git a/memory_scope/memory/worker/read/extract_time_worker.py b/memory_scope/memory/worker/read/extract_time_worker.py index 2369be77..4a0d68d9 100644 --- a/memory_scope/memory/worker/read/extract_time_worker.py +++ b/memory_scope/memory/worker/read/extract_time_worker.py @@ -10,6 +10,7 @@ from memory_scope.utils.tool_functions import prompt_to_msg class ExtractTimeWorker(MemoryBaseWorker): EXTRACT_TIME_PATTERN = r'-\s*(\S+):(\d+)' + FILE_PATH: str = __file__ def _run(self): query, query_timestamp = self.get_context(QUERY_WITH_TS) diff --git a/memory_scope/memory/worker/read/print_memory_worker.py b/memory_scope/memory/worker/read/print_memory_worker.py index 0e532144..437eb4d7 100644 --- a/memory_scope/memory/worker/read/print_memory_worker.py +++ b/memory_scope/memory/worker/read/print_memory_worker.py @@ -14,9 +14,9 @@ class PrintMemoryWorker(MemoryBaseWorker): memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES) memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True) - expired_content_list: List[str] = [] - obs_content_list: List[str] = [] - insight_content_list: List[str] = [] + expired_content_list: List[str] = ["----- expired -----"] + obs_content_list: List[str] = ["----- observation -----"] + insight_content_list: List[str] = ["----- insight -----"] i = 0 j = 0 k = 0 @@ -44,16 +44,12 @@ class PrintMemoryWorker(MemoryBaseWorker): result: str = f""" The memories of {self.user_name} about {self.target_name}. ------ observation ----- {obs_content} ------ observation ----- ------ insight ----- + {insight_content} ------ insight ----- ------ expired ----- + {expired_content} ------ expired ----- """.strip() self.set_context(RESULT, result) diff --git a/memory_scope/memory/worker/read/set_query_worker.py b/memory_scope/memory/worker/read/set_query_worker.py index 28f2da4b..4dec4bbb 100644 --- a/memory_scope/memory/worker/read/set_query_worker.py +++ b/memory_scope/memory/worker/read/set_query_worker.py @@ -7,12 +7,14 @@ from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker class SetQueryWorker(MemoryBaseWorker): def _run(self): + query = "_" + query_timestamp = int(datetime.datetime.now().timestamp()) + if "query" in self.chat_kwargs: - """ cli test query - """ + # cli test query query = self.chat_kwargs["query"] - query_timestamp = int(datetime.datetime.now().timestamp()) - else: + + elif self.chat_messages: query = self.chat_messages[-1].content query_timestamp = self.chat_messages[-1].time_created diff --git a/memory_scope/memory/worker/summary/get_reflection_subject_worker.py b/memory_scope/memory/worker/summary/get_reflection_subject_worker.py index 95e5f977..365318c2 100644 --- a/memory_scope/memory/worker/summary/get_reflection_subject_worker.py +++ b/memory_scope/memory/worker/summary/get_reflection_subject_worker.py @@ -12,6 +12,7 @@ from memory_scope.utils.tool_functions import prompt_to_msg class GetReflectionSubjectWorker(MemoryBaseWorker): + FILE_PATH: str = __file__ def new_insight_node(self, insight_key: str) -> MemoryNode: dt_handler = DatetimeHandler() diff --git a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py index 61ced731..e4b46ad8 100644 --- a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py +++ b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py @@ -11,6 +11,7 @@ from memory_scope.utils.tool_functions import prompt_to_msg class LongContraRepeatWorker(MemoryBaseWorker): + FILE_PATH: str = __file__ async def retrieve_similar_content(self, node: MemoryNode) -> (MemoryNode, List[MemoryNode]): filter_dict = { diff --git a/memory_scope/memory/worker/summary/update_insight_worker.py b/memory_scope/memory/worker/summary/update_insight_worker.py index cc1bff92..6d4a5052 100644 --- a/memory_scope/memory/worker/summary/update_insight_worker.py +++ b/memory_scope/memory/worker/summary/update_insight_worker.py @@ -11,6 +11,7 @@ from memory_scope.utils.tool_functions import prompt_to_msg class UpdateInsightWorker(MemoryBaseWorker): + FILE_PATH: str = __file__ def filter_obs_nodes(self, insight_node: MemoryNode, diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py index 58d0b626..8b96fc90 100644 --- a/memory_scope/memory/worker/write/contra_repeat_worker.py +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -9,6 +9,8 @@ from memory_scope.utils.response_text_parser import ResponseTextParser class ContraRepeatWorker(MemoryBaseWorker): + FILE_PATH: str = __file__ + def _run(self): all_obs_nodes: List[MemoryNode] = self.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES]) if not all_obs_nodes: diff --git a/memory_scope/memory/worker/write/get_observation_with_time_worker.py b/memory_scope/memory/worker/write/get_observation_with_time_worker.py index 8ab3b71b..27e3af94 100644 --- a/memory_scope/memory/worker/write/get_observation_with_time_worker.py +++ b/memory_scope/memory/worker/write/get_observation_with_time_worker.py @@ -10,6 +10,7 @@ from memory_scope.utils.tool_functions import prompt_to_msg class GetObservationWithTimeWorker(GetObservationWorker): + FILE_PATH: str = __file__ def build_prompt(self) -> List[Message]: # build prompt diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index fb9690bd..31f61618 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -13,6 +13,8 @@ from memory_scope.utils.tool_functions import prompt_to_msg class GetObservationWorker(MemoryBaseWorker): + FILE_PATH: str = __file__ + def add_observation(self, message: Message, time_infer: str, obs_content: str, keywords: str): dt_handler = DatetimeHandler(dt=message.time_created) diff --git a/memory_scope/memory/worker/write/info_filter_worker.py b/memory_scope/memory/worker/write/info_filter_worker.py index a085cd48..657b8ada 100644 --- a/memory_scope/memory/worker/write/info_filter_worker.py +++ b/memory_scope/memory/worker/write/info_filter_worker.py @@ -9,6 +9,7 @@ from memory_scope.utils.tool_functions import prompt_to_msg class InfoFilterWorker(MemoryBaseWorker): + FILE_PATH: str = __file__ def _run(self): # filter user msg @@ -22,6 +23,8 @@ class InfoFilterWorker(MemoryBaseWorker): msg.content = msg.content[: half_size] + msg.content[-half_size:] info_messages.append(msg) + self.logger.warning(info_messages) + if not info_messages: self.logger.warning("info_messages is empty!") self.continue_run = False @@ -31,6 +34,7 @@ class InfoFilterWorker(MemoryBaseWorker): user_query_list = [] for i, msg in enumerate(info_messages): user_query_list.append(f"{i + 1} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}") + self.logger.warning(self.prompt_handler.prompt_dict) system_prompt = self.prompt_handler.info_filter_system.format(batch_size=len(info_messages), user_name=self.target_name) few_shot = self.prompt_handler.info_filter_few_shot.format(user_name=self.target_name) diff --git a/memory_scope/memory/worker/write/load_memory_worker.py b/memory_scope/memory/worker/write/load_memory_worker.py index c441d51d..6717f7a5 100644 --- a/memory_scope/memory/worker/write/load_memory_worker.py +++ b/memory_scope/memory/worker/write/load_memory_worker.py @@ -86,7 +86,7 @@ class LoadMemoryWorker(MemoryBaseWorker): self.set_memories(TODAY_NODES, nodes) - async def _run(self): + def _run(self): mock_query = "-" self.submit_async_task(self.retrieve_not_reflected_memory, query=mock_query) self.submit_async_task(self.retrieve_not_updated_memory, query=mock_query) diff --git a/memory_scope/models/llama_index_generation_model.py b/memory_scope/models/llama_index_generation_model.py index e0accfa2..462f88c7 100644 --- a/memory_scope/models/llama_index_generation_model.py +++ b/memory_scope/models/llama_index_generation_model.py @@ -17,12 +17,15 @@ class LlamaIndexGenerationModel(BaseModel): def before_call(self, **kwargs): prompt: str = kwargs.pop("prompt", "") - messages: List[Message] = kwargs.pop("messages", []) + messages: List[Message] | List[dict] = kwargs.pop("messages", []) if prompt: self.data = {"prompt": prompt} elif messages: - self.data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]} + if isinstance(messages[0], dict): + self.data = {"messages": [ChatMessage(role=msg["role"], content=msg["content"]) for msg in messages]} + else: + self.data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]} else: raise RuntimeError("prompt and messages is both empty!") diff --git a/memory_scope/scheme/memory_node.py b/memory_scope/scheme/memory_node.py index d1366397..e1dcb71f 100644 --- a/memory_scope/scheme/memory_node.py +++ b/memory_scope/scheme/memory_node.py @@ -28,7 +28,8 @@ class MemoryNode(BaseModel): memory_type: str = Field("", description="conversation/observation/insight...") - status: str = Field("active", description="active or expired") + status: str = Field("active", description="db status: active / expired; modification_status: " + "new / content_modified / modified / active / expired") vector: List[float] = Field([], description="content embedding result, return empty") diff --git a/memory_scope/utils/prompt_handler.py b/memory_scope/utils/prompt_handler.py index 521b6ed1..b4f1dc90 100644 --- a/memory_scope/utils/prompt_handler.py +++ b/memory_scope/utils/prompt_handler.py @@ -15,6 +15,7 @@ class PromptHandler(object): self.kwargs = kwargs file_path = self._class_path.strip(".py") + self.add_prompt_file(file_path) if prompt_file: @@ -25,6 +26,7 @@ class PromptHandler(object): @staticmethod def file_path_completion(file_path: str) -> str: + if file_path.endswith(".yaml") or file_path.endswith(".json"): return file_path @@ -56,7 +58,7 @@ class PromptHandler(object): prompts = language_dict.get(G_CONTEXT.language) if not prompts: raise RuntimeError(f"{key}.prompt.{G_CONTEXT.language} is empty!") - self._prompt_dict[key] = prompts + self._prompt_dict[key] = prompts.strip() @property def prompt_dict(self):