diff --git a/config/cli_chat_dash_cn.yaml b/config/cli_chat_dash_cn.yaml deleted file mode 100644 index dd04463e..00000000 --- a/config/cli_chat_dash_cn.yaml +++ /dev/null @@ -1,192 +0,0 @@ -global_config: - language: cn - max_workers: 5 - -logger_config: - logger_name: memoryscope - logger_suffix: time - -memory_chat: - cli_memory_chat: - class: chat.cli_memory_chat - memory_service: memory_scope_service - generation_model: dashscope_generation - -memory_service: - memory_scope_service: - class: memory.service.memory_scope_service - memory_operations: - read_message: - class: memory.operation.frontend_operation - workflow: read_message - description: "read short memory" - - retrieve_memory: - class: memory.operation.frontend_operation - workflow: set_query,[extract_time|retrieve_obs_ins,semantic_rank],fuse_rerank - description: "retrieve long-term memory" - - list_memory: - class: memory.operation.frontend_operation - workflow: set_query,retrieve_top_memory,print_memory - description: "read all long-term memory of the user" - - delete_memory: - class: memory.operation.frontend_operation - workflow: set_query,retrieve_all_memory,delete_memory - description: "delete a single long-term memory" - - delete_all: - class: memory.operation.frontend_operation - workflow: set_query,retrieve_all_memory,delete_all - description: "delete all long-term memory" - - add_memory: - class: memory.operation.frontend_operation - workflow: add_memory - description: "add a single observation" - - summary_observation_memory: - class: memory.operation.summary_observation_op - workflow: info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat,store_memory - description: "summary user's observation memory" - interval_time: 1 - - summary_insight_memory: - class: memory.operation.backend_operation - workflow: load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat,store_memory - description: "summary user's insight memory" - interval_time: 15 - -worker: - dummy: - class: memory.worker.dummy_worker - generation_model: dashscope_generation - embedding_model: dashscope_embedding - rank_model: dashscope_rank - read_message: - class: memory.worker.frontend.read_message_worker - set_query: - class: memory.worker.frontend.set_query_worker - retrieve_obs_ins: - class: memory.worker.frontend.retrieve_memory_worker - retrieve_obs_top_k: 100 - retrieve_ins_top_k: 100 - extract_time: - class: memory.worker.frontend.extract_time_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 - semantic_rank: - class: memory.worker.frontend.semantic_rank_worker - rank_model: dashscope_rank - fuse_rerank: - class: memory.worker.frontend.fuse_rerank_worker - fuse_score_threshold: 0.01 - fuse_ratio_dict: - conversation: 0.5 - observation: 1 - obs_customized: 1.2 - insight: 2.0 - fuse_time_ratio: 2.0 - fuse_rerank_top_k: 10 - retrieve_top_memory: - class: memory.worker.frontend.retrieve_memory_worker - retrieve_obs_top_k: 100 - retrieve_ins_top_k: 100 - retrieve_expired_top_k: 100 - print_memory: - class: memory.worker.frontend.print_memory_worker - retrieve_all_memory: - class: memory.worker.frontend.retrieve_memory_worker - retrieve_obs_top_k: 100 - retrieve_ins_top_k: 100 - retrieve_expired_top_k: 100 - delete_memory: - class: memory.worker.backend.update_memory_worker - method: delete_memory - delete_all: - class: memory.worker.backend.update_memory_worker - method: delete_all - add_memory: - class: memory.worker.backend.update_memory_worker - method: from_query - info_filter: - class: memory.worker.backend.info_filter_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 - load_today_memory: - class: memory.worker.backend.load_memory_worker - retrieve_today_top_k: 100 - get_observation: - class: memory.worker.backend.get_observation_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 - get_observation_with_time: - class: memory.worker.backend.get_observation_with_time_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 - contra_repeat: - class: memory.worker.backend.contra_repeat_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 - store_memory: - class: memory.worker.backend.update_memory_worker - method: from_memory_key - memory_key: all - load_obs_and_insight: - class: memory.worker.backend.load_memory_worker - retrieve_not_reflected_top_k: 100 - retrieve_not_updated_top_k: 100 - retrieve_insight_top_k: 100 - get_reflection_subject: - class: memory.worker.backend.get_reflection_subject_worker - generation_model: dashscope_generation - reflect_obs_cnt_threshold: 10 - generation_model_kwargs: - top_k: 1 - update_insight: - class: memory.worker.backend.update_insight_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 - rank_model: dashscope_rank - long_contra_repeat: - class: memory.worker.backend.long_contra_repeat_worker - generation_model: dashscope_generation - generation_model_kwargs: - top_k: 1 - -models: - dashscope_generation: - class: models.llama_index_generation_model - module_name: dashscope_generation - model_name: qwen-max - max_tokens: 2000 - dashscope_embedding: - class: models.llama_index_embedding_model - module_name: dashscope_embedding - model_name: text-embedding-v2 - dashscope_rank: - class: models.llama_index_rank_model - module_name: dashscope_rank - model_name: gte-rerank - top_n: 10 - dummy_generation: - class: models.dummy_generation_model - module_name: dummy_generation - model_name: dummy_generation_model - -memory_store: - class: storage.llama_index_es_memory_store - embedding_model: dashscope_embedding - index_name: memory_index - es_url: http://localhost:9200 - use_hybrid: true - -monitor: - class: storage.dummy_monitor \ No newline at end of file diff --git a/config/demo_config_no_stream.yaml b/config/demo_config_no_stream.yaml deleted file mode 100644 index b9154131..00000000 --- a/config/demo_config_no_stream.yaml +++ /dev/null @@ -1,148 +0,0 @@ -global_config: - language: cn - max_workers: 5 - -memory_chat: - cli_memory_chat: - class: chat.cli_memory_chat - stream: false - memory_service: memory_chat_service - generation_model: dashscope_generation - -memory_service: - memory_chat_service: - class: memory.service.chat_memory_service - history_msg_count: 32 - contextual_msg_count: 6 - memory_operations: - read_message: - class: memory.operation.read_message - description: "read session messages of the user" - read_memory: - class: memory.operation.read_memory - workflow: set_query,retrieve_memory1,[extract_time|semantic_rank],fuse_rerank - description: "read related memories of the user" - list_memory: - class: memory.operation.read_memory - workflow: set_query,retrieve_memory2,print_memory - description: "read all memories of the user" - write_memory: - 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: 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: 60 - -worker: - dummy: - class: memory.worker.dummy_worker - generation_model: dashscope_generation - embedding_model: dashscope_embedding - rank_model: dashscope_rank - set_query: - class: memory.worker.read.set_query_worker - retrieve_memory1: - class: memory.worker.read.retrieve_memory_worker - retrieve_obs_top_k: 100 - retrieve_ins_pf_top_k: 100 - retrieve_expired_top_k: 0 - extract_time: - class: memory.worker.read.extract_time_worker - generation_model: dashscope_generation - generation_model_top_k: 1 - semantic_rank: - class: memory.worker.read.semantic_rank_worker - fuse_rerank: - class: memory.worker.read.fuse_rerank_worker - fuse_score_threshold: 0.1 - fuse_ratio_dict: - conversation: 0.5 - observation: 1 - obs_customized: 1.2 - insight: 2.0 - fuse_time_ratio: 2.0 - fuse_rerank_top_k: 10 - retrieve_memory2: - class: memory.worker.read.retrieve_memory_worker - retrieve_obs_top_k: 100 - retrieve_ins_pf_top_k: 100 - retrieve_expired_top_k: 100 - print_memory: - class: memory.worker.read.print_memory_worker - info_filter: - class: memory.worker.write.info_filter_worker - generation_model: dashscope_generation - info_filter_msg_max_size: 200 - generation_model_top_k: 1 - load_memory1: - class: memory.worker.write.load_memory_worker - retrieve_not_reflected_top_k: 0 - retrieve_not_updated_top_k: 0 - retrieve_insight_top_k: 0 - today_obs_top_k: 100 - get_observation: - class: memory.worker.write.get_observation_worker - generation_model: dashscope_generation - generation_model_top_k: 1 - get_observation_with_time: - class: memory.worker.write.get_observation_with_time_worker - generation_model: dashscope_generation - generation_model_top_k: 1 - contra_repeat: - class: memory.worker.write.contra_repeat_worker - generation_model: dashscope_generation - generation_model_top_k: 1 - retrieve_top_k: 30 - contra_repeat_max_count: 50 - store_memory: - class: memory.worker.write.store_memory_worker - store_key: all - load_memory2: - class: memory.worker.write.load_memory_worker - retrieve_not_reflected_top_k: 100 - retrieve_not_updated_top_k: 100 - retrieve_insight_top_k: 100 - today_obs_top_k: 0 - get_reflection_subject: - class: memory.worker.summary.get_reflection_subject_worker - retrieve_top_k: 100 - reflect_obs_cnt_threshold: 32 - generation_model_top_k: 1 - update_insight: - class: memory.worker.summary.update_insight_worker - update_insight_threshold: 0.1 - generation_model_top_k: 1 - update_insight_max_thread: 10 - long_contra_repeat: - class: memory.worker.summary.long_contra_repeat_worker - long_contra_repeat_top_k: 2 - long_contra_repeat_threshold: 0.1 - generation_model_top_k: 1 - -models: - dashscope_generation: - class: models.llama_index_generation_model - module_name: dashscope_generation - model_name: qwen-max - dashscope_embedding: - class: models.llama_index_embedding_model - module_name: dashscope_embedding - model_name: text-embedding-v2 - dashscope_rank: - class: models.llama_index_rank_model - module_name: dashscope_rank - model_name: gte-rerank - -memory_store: - class: storage.llama_index_es_memory_store - embedding_model: dashscope_embedding - index_name: memory_index - es_url: http://localhost:9200 - use_hybrid: false - -monitor: - class: storage.dummy_monitor \ No newline at end of file diff --git a/config/docker_config.yaml b/config/docker_config.yaml deleted file mode 100644 index 2668a5e7..00000000 --- a/config/docker_config.yaml +++ /dev/null @@ -1,83 +0,0 @@ -global_config: - language: cn - max_workers: 5 - dash_scope_apikey: - open_ai_apikey: - -memory_chat: - cli_memory_chat: - class: chat.cli_memory_chat # select class - memory_service: memory_chat_service - generation_model: dashscope_generation - human_name: human - assistant_name: assistant - -memory_service: - memory_chat_service: - class: memory.service.chat_memory_service # select class - history_msg_count: 32 - contextual_msg_count: 6 - read_memory_key: read_memory - memory_operations: - read_message: # define operation - class: memory.operation.read_memory - workflow: dummy_workflow # select workflow - description: "read session messages of the user" - read_memory: - class: memory.operation.read_memory - workflow: dummy_workflow - description: "read related memories of the user" - list_memory: - class: memory.operation.read_memory - workflow: dummy_workflow - description: "read all memories of the user" - write_memory: - class: memory.operation.write_memory - workflow: dummy_workflow - description: "write observation memories of the user" - interval_time: 60 - summary_memory: - class: memory.operation.summary_memory - workflow: dummy_workflow - description: "summary observation memories of the user" - interval_time: 300 - -models: - dashscope_generation: - class: models.llama_index_generation_model # select class - module_name: dashscope_generation - model_name: qwen-max - dashscope_embedding: - class: models.llama_index_embedding_model # select class - module_name: dashscope_embedding - model_name: text-embedding-v2 - dashscope_rank: - class: models.llama_index_rank_model # select class - module_name: dashscope_rank - model_name: gte-rerank - -vector_store: - class: storage.dummy_vector_store # select class - embedding_model: dashscope_embedding - -monitor: - class: storage.dummy_monitor # select class - -worker: - dummy_workflow: - class: memory.worker.dummy_worker - generation_model: dashscope_generation - embedding_model: dashscope_embedding - rank_model: dashscope_rank - retrieve_store_worker: - class: memory.worker.read.retrieve_store_worker - retrieve_obs_top_k: 100 - retrieve_ins_pf_top_k: 100 - fuse_rerank_worker: - class: memory.worker.read.fuse_rerank_worker - fuse_score_threshold: 0.1 - fuse_ratio_dict: - observation: 1 - fuse_time_ratio: 2.0 - fuse_rerank_top_k: 10 - diff --git a/examples/api/chat_example.py b/examples/api/chat_example.py new file mode 100644 index 00000000..393b004e --- /dev/null +++ b/examples/api/chat_example.py @@ -0,0 +1,66 @@ +import sys + +sys.path.append(".") + +from memoryscope import MemoryScope, Arguments + +arguments = Arguments( + language="cn", + logger_to_screen=False, + memory_chat_class="api_memory_chat", + generation_backend="dashscope_generation", + generation_model="qwen-max", + embedding_backend="dashscope_embedding", + embedding_model="text-embedding-v2", + use_dummy_ranker=False, + rank_backend="dashscope_rank", + rank_model="gte-rerank", +) + + +def chat_example1(): + with MemoryScope(config_path="memoryscope/core/config/demo_config.yaml") as ms: + memory_chat = ms.default_memory_chat + response = memory_chat.chat_with_memory(query="我的爱好是弹琴。") + print("回答:\n" + response.message.content) + print("记忆:\n" + response.meta_data["memories"]) + + +def chat_example2(): + with MemoryScope(arguments=arguments) as ms: + memory_chat = ms.default_memory_chat + + response = memory_chat.chat_with_memory(query="我的爱好是弹琴。") + print("回答:\n" + response.message.content) + print("记忆:\n" + response.meta_data["memories"]) + + +def chat_example3(): + with MemoryScope(arguments=arguments) as ms: + memory_chat = ms.default_memory_chat + + response = memory_chat.chat_with_memory( + query="你知道我消灭了几个星系吗?", + extra_memories="用户是一只哥斯拉,这是真是的,不是虚构的。\n前天用户消灭了星系0x4be?\n大前天用户消灭了星系0x5be?") + print("回答:\n" + response.message.content) + + +def chat_example4(): + with MemoryScope(arguments=arguments) as ms: + memory_chat = ms.default_memory_chat + + response = memory_chat.chat_with_memory(query="我的爱好是弹琴。") + print("回答1:\n" + response.message.content) + memory_chat.memory_service.consolidate_memory() + + response = memory_chat.chat_with_memory(query="你知道我的乐器爱好是什么?", + history_message_strategy=None) + print("回答2:\n" + response.message.content) + print("记忆2:\n" + response.meta_data["memories"]) + + +if __name__ == "__main__": + chat_example1() + # chat_example2() + # chat_example3() + # chat_example4() diff --git a/examples/cli/dash_cli_cn1.sh b/examples/cli/dash_cli_cn1.sh new file mode 100644 index 00000000..2998bd29 --- /dev/null +++ b/examples/cli/dash_cli_cn1.sh @@ -0,0 +1 @@ +python memoryscope/cli.py -config_path=memoryscope/core/config/demo_config.yaml \ No newline at end of file diff --git a/examples/cli/dash_cli_cn2.sh b/examples/cli/dash_cli_cn2.sh new file mode 100644 index 00000000..14c92cee --- /dev/null +++ b/examples/cli/dash_cli_cn2.sh @@ -0,0 +1,10 @@ +python memoryscope/cli.py \ + -language="cn" \ + -memory_chat_class="cli_memory_chat" \ + -generation_backend="dashscope_generation" \ + -generation_model="qwen-max" \ + -embedding_backend="dashscope_embedding" \ + -embedding_model="text-embedding-v2" \ + -use_dummy_ranker=False \ + -rank_backend="dashscope_rank" \ + -rank_model="gte-rerank" \ No newline at end of file diff --git a/examples/__init__.py b/examples/docker/docker_config.yaml similarity index 100% rename from examples/__init__.py rename to examples/docker/docker_config.yaml diff --git a/memoryscope/__init__.py b/memoryscope/__init__.py index d8b7815a..0dc03c13 100644 --- a/memoryscope/__init__.py +++ b/memoryscope/__init__.py @@ -1,3 +1,5 @@ -""" Version of MemoryScope.""" +from memoryscope.core.config.arguments import Arguments +from memoryscope.core.memoryscope import MemoryScope -__version__ = "0.1.0-alpha.1" +""" Version of MemoryScope.""" +__version__ = "0.1.0" diff --git a/memoryscope/chat/base_memory_chat.py b/memoryscope/chat/base_memory_chat.py deleted file mode 100644 index 2be6ad6b..00000000 --- a/memoryscope/chat/base_memory_chat.py +++ /dev/null @@ -1,43 +0,0 @@ -from abc import ABCMeta, abstractmethod - -from memoryscope.memory.service.base_memory_service import BaseMemoryService - - -class BaseMemoryChat(metaclass=ABCMeta): - """ - An abstract base class representing a chat system integrated with memory services. - It outlines the method to initiate a chat session leveraging memory data, which concrete subclasses must implement. - """ - - @abstractmethod - def chat_with_memory(self, query: str): - """ - Initiates a chat interaction using the memory service, with the provided query as input. - - Args: - query (str): The user's query or message to start the chat. - - Returns: - This method should return the chat response generated after processing the query - with the associated memory context. The actual return type and content are defined by the implementing - subclass. - """ - - @property - def memory_service(self) -> BaseMemoryService: - """ - Abstract property to access the memory service. - - Raises: - NotImplementedError: This method should be implemented in a subclass. - """ - raise NotImplementedError - - def run(self): - """ - Abstract method to run the chat system. - - This method should contain the logic to initiate and manage the chat process, - utilizing the memory service as needed. It must be implemented by subclasses. - """ - pass diff --git a/memoryscope/chat/cli_memory_chat.py b/memoryscope/chat/cli_memory_chat.py deleted file mode 100644 index e06fa0fc..00000000 --- a/memoryscope/chat/cli_memory_chat.py +++ /dev/null @@ -1,346 +0,0 @@ -import os -import time -from typing import List - -import questionary - -from memoryscope.chat.base_memory_chat import BaseMemoryChat -from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME -from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.service.base_memory_service import BaseMemoryService -from memoryscope.models.base_model import BaseModel -from memoryscope.scheme.message import Message -from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger -from memoryscope.utils.prompt_handler import PromptHandler -from memoryscope.utils.tool_functions import char_logo - - -class CliMemoryChat(BaseMemoryChat): - """ - Command-line interface for chatting with an AI that integrates memory functionality. - Allows users to interact, manage chat history, adjust streaming settings, and view commands' help. - """ - USER_COMMANDS = { - "exit": "Exit the CLI.", - "clear": "Clear the command history.", - "help": "Display available CLI commands and their descriptions.", - "stream": "Toggle between getting streamed responses from the model." - } - - def __init__(self, - memory_service: str, - generation_model: str, - stream: bool = True, - human_name: str = DEFAULT_HUMAN_NAME[G_CONTEXT.language], - assistant_name: str = "AI", - **kwargs): - """ - Initializes the CLI chat instance with specified services, models, and personalized settings. - - Args: - memory_service (str | BaseMemoryService): The memory service to be used for storing conversation history. - generation_model (str | BaseModel): The model responsible for generating AI responses. - stream (bool, optional): Flag indicating whether responses should be streamed. Defaults to True. - human_name (str, optional): The name assigned to the human user. Defaults to a language-specific user. - assistant_name (str, optional): The name of the AI assistant. Defaults to "AI". - **kwargs: Additional keyword arguments for flexibility or future extensions. - - Side Effects: - - Updates global context with human and AI names. - - Initializes logging for the instance. - """ - self._memory_service: BaseMemoryService | str = memory_service - self._generation_model: BaseModel | str = generation_model - self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {}) - - self.stream: bool = stream - self.human_name: str = human_name - self.assistant_name: str = assistant_name - self.kwargs: dict = kwargs - - self._logo = char_logo("MemoryScope") - self._prompt_handler: PromptHandler | None = None - G_CONTEXT.meta_data.update({ - "human_name": human_name, - "assistant_name": assistant_name, - }) - - self.logger = Logger.get_logger() - - @property - def prompt_handler(self) -> PromptHandler: - """ - Lazy initialization property for the prompt handler. - - This property ensures that the `_prompt_handler` attribute is only instantiated when it is first accessed. - It uses the current file's path and additional keyword arguments for configuration. - - Returns: - PromptHandler: An instance of the PromptHandler configured for this CLI session. - """ - if self._prompt_handler is None: - self._prompt_handler = PromptHandler(__file__, **self.kwargs) - return self._prompt_handler - - def print_logo(self): - """ - Prints the logo of the CLI application to the console. - - The logo is composed of multiple lines, which are iterated through - and printed one by one to provide a visual identity for the chat interface. - """ - for line in self._logo: - print(line) - - @property - def memory_service(self) -> BaseMemoryService: - """ - Property to access the memory service. If the service is initially set as a string, - it will be looked up in the memory service dictionary of global context, initialized, - and then returned as an instance of `BaseMemoryService`. Ensures the memory service - is properly started before use. - - Returns: - BaseMemoryService: An active memory service instance. - - Raises: - ValueError: If the declaration of memory service is not found in the memory service dictionary of global context. - """ - if isinstance(self._memory_service, str): - if self._memory_service not in G_CONTEXT.memory_service_dict: - raise ValueError("Missing declaration of memory_service in yaml configuration: " + self._memory_service) - self._memory_service = G_CONTEXT.memory_service_dict[self._memory_service] - self._memory_service.init_service() - self._memory_service.start_backend_service() - return self._memory_service - - @property - def generation_model(self) -> BaseModel: - """ - Property to get the generation model. If the model is set as a string, it will be resolved from the global - context's model dictionary. - - Raises: - ValueError: If the declaration of generation model is not found in the model dictionary of global context . - - Returns: - BaseModel: An actual generation model instance. - """ - if isinstance(self._generation_model, str): - if self._generation_model not in G_CONTEXT.model_dict: - raise ValueError(f"Missing declaration of generation model in yaml config: {self._generation_model}") - self._generation_model = G_CONTEXT.model_dict[self._generation_model] - return self._generation_model - - def chat_with_memory(self, query: str, remember_response: bool = False) -> ModelResponse | ModelResponseGen: - """ - Engages in a conversation with the AI model, utilizing conversation memory. - The function sends the user's query, incorporates conversation history and memory, - and optionally remembers the AI's response based on the user's preference. - - Args: - query (str): The user's input or query for the AI. - remember_response (bool, optional): Flag indicating whether to save the AI's response to memory. - Defaults to False. - - Returns: - - ModelResponse: In non-streaming mode, returns a complete AI response. - - ModelResponseGen: In streaming mode, returns a generator yielding AI response parts. - - Side Effects: - - Updates the conversation memory with the query of user and (optionally) the response of AI. - - Retrieves and includes historical messages and memory content in the context of conversation. - """ - 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] = [] - - # Incorporate memory into the system prompt if available - system_prompt = self.prompt_handler.system_prompt - memories: str = self.memory_service.retrieve_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)) - - # Include past conversation history in the message list - history_messages = self.memory_service.read_message() - if history_messages: - messages.extend(history_messages) - - # Append the current user's message to the conversation context - messages.append(new_message) - self.logger.info(f"messages={messages}") - - # Invoke the Language Model with the constructed message context, respecting streaming setting - generated = self.generation_model.call(messages=messages, stream=self.stream, **self.generation_model_kwargs) - - # In non-streaming interactions, explicitly save the AI's reply to memory if instructed - if remember_response: - assert not self.stream # Ensure we're not in streaming mode when remembering responses - generated.message.role_name = self.assistant_name - self.memory_service.add_messages(generated.message) - - # Return the AI's response directly or as a generator based on the streaming mode - return generated - - @staticmethod - def parse_query_command(query: str): - """ - Parses the user's input query command, separating it into the command and its associated keyword arguments. - - Args: - query (str): The raw input string from the user which includes the command and its arguments. - - Returns: - tuple: A tuple containing the command (str) as the first element and a dictionary (kwargs) of keyword - arguments as the second element. - """ - query_split = query.lstrip("/").lower().split(" ") # Split and preprocess the input command - command = query_split[0] # Extract the command - args = query_split[1:] # Extract the arguments following the command - kwargs = {} # Initialize dictionary to hold keyword arguments - - for arg in args: - # Skip if no arguments exist (unnecessary check due to prior assignment, but retained as per original) - if not args: - continue - arg_split = arg.split("=") # Split argument into key-value pair - if len(arg_split) >= 2: # Ensure there's both a key and value - k = arg_split[0] # Extract key - v = arg_split[1] # Extract value - if k and v: # Only add to kwargs if both key and value are non-empty - kwargs[k] = v - - return command, kwargs # Return the parsed command and keyword arguments - - def process_commands(self, query: str) -> bool: - """ - Parses and executes commands from user input in the CLI chat interface. - Supports operations like exiting, clearing screen, showing help, toggling stream mode, - executing predefined memory operations, and handling unknown commands. - - Args: - query (str): The user's input command string. - - Returns: - bool: Indicates whether to continue running the CLI after processing the command. - """ - continue_run = True - command, kwargs = self.parse_query_command(query) - - # Print prompt for AI's response - questionary.print("> ", end="", style="fg:yellow") - questionary.print(f"{self.assistant_name}: ", end="", style="bold") - - if command == "exit": - self.memory_service.stop_backend_service() - continue_run = False - - elif command == "clear": - os.system("clear") - - elif command == "help": - questionary.print("CLI commands", "bold") - for cmd, desc in self.USER_COMMANDS.items(): - questionary.print(text=f" /{cmd}:", style="bold") - questionary.print(text=f" {desc}") - - elif command == "stream": - self.stream = not self.stream - questionary.print(f"set stream: {self.stream}") - - elif command in self.memory_service.op_description_dict: - refresh_time = kwargs.pop("refresh_time", "") - if refresh_time and refresh_time.isdigit(): - refresh_time = int(refresh_time) - self.memory_service.stop_backend_service() - while True: - result = self.memory_service.do_operation(op_name=command, **kwargs) - os.system("clear") - self.print_logo() - if result: - if isinstance(result, list): - result = "\n".join([str(x) for x in result]) - questionary.print(result) - else: - questionary.print(f"command={command} result is empty! kwargs={kwargs}") - time.sleep(refresh_time) - - else: - result = self.memory_service.do_operation(op_name=command, **kwargs) - if result: - if isinstance(result, list): - result = "\n".join([str(x) for x in result]) - questionary.print(result) - else: - questionary.print(f"command={command} result is empty! kwargs={kwargs}") - - else: - questionary.print(f"Unknown command={command} received.") - - return continue_run - - def run(self): - """ - Runs the CLI chat loop, which handles user input, processes commands, - communicates with the AI model, manages conversation memory, and controls - the chat session including streaming responses, command execution, and error handling. - - The loop continues until the user explicitly chooses to exit. - """ - self.print_logo() - self.USER_COMMANDS.update(self.memory_service.op_description_dict) - - while True: - try: - query = questionary.text(message=f"{self.human_name}:", multiline=False, qmark=">").unsafe_ask() - if not query: - continue - - query: str = query.strip() - - # Handle special commands prefixed with '/' - if query.startswith("/"): - if self.process_commands(query=query): - continue - else: - break - - # Print prompt for AI's response - questionary.print("> ", end="", style="fg:yellow") - questionary.print(f"{self.assistant_name}: ", end="", style="bold") - - # Fetch and display AI's response, with support for streaming - self.memory_service.start_backend_service() - if self.stream: - model_response = None - for model_response in self.chat_with_memory(query=query): - questionary.print(model_response.delta, end="") - questionary.print("") - - else: - model_response = self.chat_with_memory(query=query) - questionary.print(model_response.message.content) - - # Append AI's response to the conversation memory - model_response.message.role_name = self.assistant_name - self.memory_service.add_messages(model_response.message) - - except KeyboardInterrupt: - # Handle user interruption and confirm exit - questionary.print("User interrupt occurred.") - is_exit = questionary.confirm("Continue exit?").unsafe_ask() - if is_exit: - self.memory_service.stop_backend_service() - break - - except Exception as e: - # Log and handle any unanticipated exceptions - import traceback - traceback.print_exc() - self.logger.exception(f"An exception occurred when running cli memory chat. args={e.args}.") - continue diff --git a/memoryscope/cli.py b/memoryscope/cli.py index 14b82099..a429fda0 100644 --- a/memoryscope/cli.py +++ b/memoryscope/cli.py @@ -1,105 +1,17 @@ -import datetime import sys -import questionary - sys.path.append(".") # noqa: E402 -import json -from concurrent.futures import ThreadPoolExecutor -from typing import Dict, Any - import fire -import yaml -import atexit -from memoryscope.enumeration.language_enum import LanguageEnum -from memoryscope.enumeration.model_enum import ModelEnum -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger -from memoryscope.utils.timer import timer -from memoryscope.utils.tool_functions import init_instance_by_config, camelcase_to_underscore +from memoryscope.core.memoryscope import MemoryScope -class MemoryScope(object): - - def __init__(self): - self.config: Dict[str, Any] = {} - datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S') - class_name = camelcase_to_underscore(self.__class__.__name__) - self.logger: Logger = Logger.get_logger(f"{class_name}_{datetime_suffix}", to_stream=False) - - def load_config(self, path: str): - with open(path) as f: - if path.endswith("yaml"): - self.config = yaml.load(f, yaml.FullLoader) - elif path.endswith("json"): - self.config = json.load(f) - else: - raise RuntimeError("not supported config file type!") - self.init_global_content_by_config() - atexit.register(self.shutdown) # register clean up function - return self - - @staticmethod - def shutdown(): - questionary.print('Gracefully executing the shutdown function...') - G_CONTEXT.memory_store.close() - G_CONTEXT.monitor.close() - G_CONTEXT.thread_pool.shutdown() - - def set_global_config(self): - G_CONTEXT.global_config = global_config = self.config["global_config"] - G_CONTEXT.language = LanguageEnum(global_config["language"]) - G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(global_config["max_workers"])) - - @timer - def init_global_content_by_config(self): - # set global config - self.set_global_config() - - # init memory_chat - for name, conf in self.config["memory_chat"].items(): - G_CONTEXT.memory_chat_dict[name] = init_instance_by_config(conf, name=name) - - # set memory_service - for name, conf in self.config["memory_service"].items(): - G_CONTEXT.memory_service_dict[name] = init_instance_by_config(conf, name=name) - - # init models - for name, conf in self.config["models"].items(): - G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name) - - # init vector_store - if "memory_store" not in self.config: - raise RuntimeError("memory_store config is required!") - 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"]) - - # set worker config - G_CONTEXT.worker_config = self.config["worker"] - - @property - def default_chat_handle(self): - return list(G_CONTEXT.memory_chat_dict.values())[0] - - @property - def default_service(self): - return self.default_chat_handle.memory_service - - -class CliJob(MemoryScope): - - def run(self, config: str): - self.load_config(config) - self.init_global_content_by_config() - self.default_chat_handle.run() +def cli_job(**kwargs): + with MemoryScope(**kwargs) as ms: + memory_chat = ms.default_memory_chat + memory_chat.run() if __name__ == "__main__": - cli_job = CliJob() - fire.Fire(cli_job.run) + fire.Fire(cli_job) diff --git a/memoryscope/constants/common_constants.py b/memoryscope/constants/common_constants.py index adf9e548..df217846 100644 --- a/memoryscope/constants/common_constants.py +++ b/memoryscope/constants/common_constants.py @@ -5,11 +5,15 @@ WORKFLOW_NAME = "workflow_name" +MEMORYSCOPE_CONTEXT = "memoryscope_context" + RESULT = "result" +MEMORIES = "memories" + CHAT_MESSAGES = "chat_messages" -MEMORY_HANDLER = "memory_handler" +MEMORY_MANAGER = "memory_manager" CHAT_KWARGS = "chat_kwargs" diff --git a/memoryscope/chat/__init__.py b/memoryscope/core/__init__.py similarity index 100% rename from memoryscope/chat/__init__.py rename to memoryscope/core/__init__.py diff --git a/memoryscope/memory/__init__.py b/memoryscope/core/chat/__init__.py similarity index 100% rename from memoryscope/memory/__init__.py rename to memoryscope/core/chat/__init__.py diff --git a/memoryscope/core/chat/api_memory_chat.py b/memoryscope/core/chat/api_memory_chat.py new file mode 100644 index 00000000..6a8078f7 --- /dev/null +++ b/memoryscope/core/chat/api_memory_chat.py @@ -0,0 +1,217 @@ +from typing import List, Optional, Literal + +from memoryscope.constants.common_constants import MEMORIES +from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME +from memoryscope.core.chat.base_memory_chat import BaseMemoryChat +from memoryscope.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.models.base_model import BaseModel +from memoryscope.core.service.base_memory_service import BaseMemoryService +from memoryscope.core.utils.prompt_handler import PromptHandler +from memoryscope.enumeration.message_role_enum import MessageRoleEnum +from memoryscope.scheme.message import Message +from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen + + +class ApiMemoryChat(BaseMemoryChat): + + def __init__(self, + memory_service: str, + generation_model: str, + context: MemoryscopeContext, + stream: bool = False, + human_name: str = None, + assistant_name: str = None, + **kwargs): + + super().__init__(**kwargs) + + self._memory_service: BaseMemoryService | str = memory_service + self._generation_model: BaseModel | str = generation_model + self.context: MemoryscopeContext = context + self.stream: bool = stream + self.generation_model_kwargs: dict = kwargs.pop("generation_model_kwargs", {}) + + self.human_name: str = human_name + if not self.human_name: + self.human_name = DEFAULT_HUMAN_NAME[self.context.language] + self.context.meta_data["human_name"] = self.human_name + + self.assistant_name: str = assistant_name + if not self.assistant_name: + self.assistant_name = "AI" + self.context.meta_data["assistant_name"] = self.assistant_name + + self._prompt_handler: PromptHandler | None = None + + @property + def prompt_handler(self) -> PromptHandler: + """ + Lazy initialization property for the prompt handler. + + This property ensures that the `_prompt_handler` attribute is only instantiated when it is first accessed. + It uses the current file's path and additional keyword arguments for configuration. + + Returns: + PromptHandler: An instance of the PromptHandler configured for this CLI session. + """ + if self._prompt_handler is None: + self._prompt_handler = PromptHandler(__file__, + language=self.context.language, + prompt_file="memory_chat_prompt", + **self.kwargs) + return self._prompt_handler + + @property + def memory_service(self) -> BaseMemoryService: + """ + Property to access the memory service. If the service is initially set as a string, + it will be looked up in the memory service dictionary of context, initialized, + and then returned as an instance of `BaseMemoryService`. Ensures the memory service + is properly started before use. + + Returns: + BaseMemoryService: An active memory service instance. + + Raises: + ValueError: If the declaration of memory service is not found in the memory service dictionary of context. + """ + if isinstance(self._memory_service, str): + if self._memory_service not in self.context.memory_service_dict: + raise ValueError(f"Missing declaration of memory_service in context: {self._memory_service}") + + self._memory_service: BaseMemoryService = self.context.memory_service_dict[self._memory_service] + # init service & update kwargs + self._memory_service.init_service() + return self._memory_service + + @property + def generation_model(self) -> BaseModel: + """ + Property to get the generation model. If the model is set as a string, it will be resolved from the global + context's model dictionary. + + Raises: + ValueError: If the declaration of generation model is not found in the model dictionary of context . + + Returns: + BaseModel: An actual generation model instance. + """ + if isinstance(self._generation_model, str): + if self._generation_model not in self.context.model_dict: + raise ValueError(f"Missing declaration of generation model in yaml config: {self._generation_model}") + self._generation_model = self.context.model_dict[self._generation_model] + return self._generation_model + + def iter_response(self, + remember_response: bool, + resp: ModelResponseGen, + memories: str, + query_message: Message) -> ModelResponseGen: + + model_response: ModelResponse | None = None + for model_response in resp: + yield model_response + + if remember_response: + if model_response and model_response.message: + model_response.message.role_name = self.assistant_name + model_response.meta_data[MEMORIES] = memories + self.memory_service.add_messages([query_message, model_response.message]) + else: + self.logger.warning("model_response or model_response.message is empty!") + + def chat_with_memory(self, + query: str, + role_name: Optional[str] = None, + system_prompt: Optional[str] = None, + memory_prompt: Optional[str] = None, + extra_memories: Optional[str] = None, + history_message_strategy: Literal["auto", None] | int = "auto", + remember_response: bool = True, + **kwargs): + """ + The core function that carries out conversation with memory accepts user queries through query and returns the + conversation results through model_response. The retrieved memories are stored in the memories within meta_data. + Args: + query (str, optional): User's query, includes the user's question. + role_name (str, optional): User's role name. + system_prompt (str, optional): System prompt. Defaults to the system_prompt in "memory_chat_prompt.yaml". + memory_prompt (str, optional): Memory prompt. Defaults to the memory_prompt in "memory_chat_prompt.yaml". + extra_memories (str, optional): Manually added user memory in this function. + history_message_strategy ("auto", None, int): + - If it is set to "auto", the history messages in the conversation will retain those that have not + yet been summarized. Default to "auto". + - If it is set to None, no conversation history will be saved. + - If it is set to an integer value "n", the most recent "n" messages will be retained. + remember_response (bool, optional): Flag indicating whether to save the AI's response to memory. + Defaults to False. + Returns: + - ModelResponse: In non-streaming mode, returns a complete AI response. + - ModelResponseGen: In streaming mode, returns a generator yielding AI response parts. + - Memories: To obtain the memory by invoking the method of model_response.meta_data[MEMORIES] + """ + chat_messages: List[Message] = [] + + # prepare query message + if not role_name: + role_name = self.human_name + query_message = Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query) + + # To retrieve memory, prepare the query timestamp and role name by adding query_message. + memories: str = self.memory_service.retrieve_memory(query=query_message.content, + role_name=query_message.role_name, + timestamp=query_message.time_created) + + # format system_message with memories + system_prompt_list = [] + if system_prompt: + system_prompt_list.append(system_prompt) + else: + system_prompt_list.append(self.prompt_handler.system_prompt) + + if memories: + # add memory prompt + if memory_prompt: + system_prompt_list.append(memory_prompt) + else: + system_prompt_list.append(self.prompt_handler.memory_prompt) + system_prompt_list.append(memories) + + if extra_memories: + system_prompt_list.extend(extra_memories) + + system_prompt_join = "\n".join([x.strip() for x in system_prompt_list]) + system_message = Message(role=MessageRoleEnum.SYSTEM, content=system_prompt_join) + chat_messages.append(system_message) + + # Include past conversation history in the message list + if history_message_strategy: + history_messages = [] + + if history_message_strategy == "auto": + history_messages = self.memory_service.read_message() + + elif isinstance(history_message_strategy, int): + history_messages = self.memory_service.chat_messages[-history_message_strategy:] + + if history_messages: + chat_messages.extend(history_messages) + + # Append the current user's message to the conversation context + chat_messages.append(query_message) + self.logger.info(f"chat_messages={chat_messages}") + + resp = self.generation_model.call(messages=chat_messages, stream=self.stream, **self.generation_model_kwargs) + if self.stream: + return self.iter_response(remember_response, resp, memories, query_message) + + else: + model_response: ModelResponse = resp + if remember_response: + if model_response and model_response.message: + model_response.message.role_name = self.assistant_name + model_response.meta_data[MEMORIES] = memories + self.memory_service.add_messages([query_message, model_response.message]) + else: + self.logger.warning("model_response or model_response.message is empty!") + return model_response diff --git a/memoryscope/core/chat/base_memory_chat.py b/memoryscope/core/chat/base_memory_chat.py new file mode 100644 index 00000000..e311f072 --- /dev/null +++ b/memoryscope/core/chat/base_memory_chat.py @@ -0,0 +1,68 @@ +from abc import ABCMeta, abstractmethod +from typing import Optional, Literal + +from memoryscope.core.service.base_memory_service import BaseMemoryService +from memoryscope.core.utils.logger import Logger + + +class BaseMemoryChat(metaclass=ABCMeta): + """ + An abstract base class representing a chat system integrated with memory services. + It outlines the method to initiate a chat session leveraging memory data, which concrete subclasses must implement. + """ + + def __init__(self, **kwargs): + self.kwargs: dict = kwargs + self.logger = Logger.get_logger() + + @property + def memory_service(self) -> BaseMemoryService: + """ + Abstract property to access the memory service. + + Raises: + NotImplementedError: This method should be implemented in a subclass. + """ + raise NotImplementedError + + @abstractmethod + def chat_with_memory(self, + query: str, + role_name: Optional[str] = None, + system_prompt: Optional[str] = None, + memory_prompt: Optional[str] = None, + extra_memories: Optional[str] = None, + history_message_strategy: Literal["auto", None] | int = "auto", + remember_response: bool = True, + **kwargs): + """ + The core function that carries out conversation with memory accepts user queries through query and returns the + conversation results through model_response. The retrieved memories are stored in the memories within meta_data. + Args: + query (str, optional): User's query, includes the user's question. + role_name (str, optional): User's role name. + system_prompt (str, optional): System prompt. Defaults to the system_prompt in "memory_chat_prompt.yaml". + memory_prompt (str, optional): Memory prompt. Defaults to the memory_prompt in "memory_chat_prompt.yaml". + extra_memories (str, optional): Manually added user memory in this function. + history_message_strategy ("auto", None, int): + - If it is set to "auto", the history messages in the conversation will retain those that have not + yet been summarized. Default to "auto". + - If it is set to None, no conversation history will be saved. + - If it is set to an integer value "n", the most recent "n" messages will be retained. + remember_response (bool, optional): Flag indicating whether to save the AI's response to memory. + Defaults to False. + Returns: + - ModelResponse: In non-streaming mode, returns a complete AI response. + - ModelResponseGen: In streaming mode, returns a generator yielding AI response parts. + - Memories: To obtain the memory by invoking the method of model_response.meta_data[MEMORIES] + """ + raise NotImplementedError + + def run(self): + """ + Abstract method to run the chat system. + + This method should contain the logic to initiate and manage the chat process, + utilizing the memory service as needed. It must be implemented by subclasses. + """ + pass diff --git a/memoryscope/core/chat/cli_memory_chat.py b/memoryscope/core/chat/cli_memory_chat.py new file mode 100644 index 00000000..150c79b3 --- /dev/null +++ b/memoryscope/core/chat/cli_memory_chat.py @@ -0,0 +1,206 @@ +import os +import time +from typing import Optional, Literal + +import questionary + +from memoryscope.core.chat.api_memory_chat import ApiMemoryChat +from memoryscope.core.utils.tool_functions import char_logo + + +class CliMemoryChat(ApiMemoryChat): + """ + Command-line interface for chatting with an AI that integrates memory functionality. + Allows users to interact, manage chat history, adjust streaming settings, and view commands' help. + """ + USER_COMMANDS = { + "exit": "Exit the CLI.", + "clear": "Clear the command history.", + "help": "Display available CLI commands and their descriptions.", + "stream": "Toggle between getting streamed responses from the model." + } + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self._logo = char_logo("MemoryScope") + + def print_logo(self): + """ + Prints the logo of the CLI application to the console. + + The logo is composed of multiple lines, which are iterated through + and printed one by one to provide a visual identity for the chat interface. + """ + for line in self._logo: + print(line) + + def chat_with_memory(self, + query: str, + role_name: Optional[str] = None, + system_prompt: Optional[str] = None, + memory_prompt: Optional[str] = None, + extra_memories: Optional[str] = None, + history_message_strategy: Literal["auto", None] | int = "auto", + remember_response: bool = True, + **kwargs): + resp = super().chat_with_memory(query=query, + role_name=role_name, + system_prompt=system_prompt, + memory_prompt=memory_prompt, + extra_memories=extra_memories, + history_message_strategy=history_message_strategy, + remember_response=remember_response, + **kwargs) + + if self.stream: + for _resp in resp: + questionary.print(_resp.delta, end="") + questionary.print("") + else: + questionary.print(resp.message.content) + + @staticmethod + def parse_query_command(query: str): + """ + Parses the user's input query command, separating it into the command and its associated keyword arguments. + + Args: + query (str): The raw input string from the user which includes the command and its arguments. + + Returns: + tuple: A tuple containing the command (str) as the first element and a dictionary (kwargs) of keyword + arguments as the second element. + """ + query_split = query.lstrip("/").lower().split(" ") # Split and preprocess the input command + command = query_split[0] # Extract the command + args = query_split[1:] # Extract the arguments following the command + kwargs = {} # Initialize dictionary to hold keyword arguments + + for arg in args: + # Skip if no arguments exist (unnecessary check due to prior assignment, but retained as per original) + if not args: + continue + arg_split = arg.split("=") # Split argument into key-value pair + if len(arg_split) >= 2: # Ensure there's both a key and value + k = arg_split[0] # Extract key + v = arg_split[1] # Extract value + if k and v: # Only add to kwargs if both key and value are non-empty + kwargs[k] = v + + return command, kwargs # Return the parsed command and keyword arguments + + def process_commands(self, query: str) -> bool: + """ + Parses and executes commands from user input in the CLI chat interface. + Supports operations like exiting, clearing screen, showing help, toggling stream mode, + executing predefined memory operations, and handling unknown commands. + + Args: + query (str): The user's input command string. + + Returns: + bool: Indicates whether to continue running the CLI after processing the command. + """ + continue_run = True + command, kwargs = self.parse_query_command(query) + + # Print prompt for AI's response + questionary.print("> ", end="", style="fg:yellow") + questionary.print(f"{self.assistant_name}: ", end="", style="bold") + + if command == "exit": + self.memory_service.stop_backend_service() + continue_run = False + + elif command == "clear": + os.system("clear") + + elif command == "help": + questionary.print("CLI commands", "bold") + for cmd, desc in self.USER_COMMANDS.items(): + questionary.print(text=f" /{cmd}:", style="bold") + questionary.print(text=f" {desc}") + + elif command == "stream": + self.stream = not self.stream + questionary.print(f"set stream: {self.stream}") + + elif command in self.memory_service.op_description_dict: + refresh_time = kwargs.pop("refresh_time", "") + if refresh_time and refresh_time.isdigit(): + refresh_time = int(refresh_time) + self.memory_service.stop_backend_service() + while True: + result = self.memory_service.do_operation(name=command, **kwargs) + os.system("clear") + self.print_logo() + if result: + if isinstance(result, list): + result = "\n".join([str(x) for x in result]) + questionary.print(result) + else: + questionary.print(f"command={command} result is empty! kwargs={kwargs}") + time.sleep(refresh_time) + + else: + result = self.memory_service.do_operation(name=command, **kwargs) + if result: + if isinstance(result, list): + result = "\n".join([str(x) for x in result]) + questionary.print(result) + else: + questionary.print(f"command={command} result is empty! kwargs={kwargs}") + + else: + questionary.print(f"Unknown command={command} received.") + + return continue_run + + def run(self): + """ + Runs the CLI chat loop, which handles user input, processes commands, + communicates with the AI model, manages conversation memory, and controls + the chat session including streaming responses, command execution, and error handling. + + The loop continues until the user explicitly chooses to exit. + """ + self.print_logo() + self.USER_COMMANDS.update(self.memory_service.op_description_dict) + + while True: + try: + query = questionary.text(message=f"{self.human_name}:", multiline=False, qmark=">").unsafe_ask() + if not query: + continue + + query: str = query.strip() + + # Handle special commands prefixed with '/' + if query.startswith("/"): + if self.process_commands(query=query): + continue + else: + break + + # Print prompt for AI's response + questionary.print("> ", end="", style="fg:yellow") + questionary.print(f"{self.assistant_name}: ", end="", style="bold") + + # Fetch and display AI's response + self.memory_service.start_backend_service() + self.chat_with_memory(query=query) + + except KeyboardInterrupt: + # Handle user interruption and confirm exit + questionary.print("User interrupt occurred.") + is_exit = questionary.confirm("Continue exit?").unsafe_ask() + if is_exit: + self.memory_service.stop_backend_service() + break + + except Exception as e: + # Log and handle any unanticipated exceptions + import traceback + traceback.print_exc() + self.logger.exception(f"An exception occurred when running cli memory chat. args={e.args}.") + continue diff --git a/memoryscope/chat/cli_memory_chat.yaml b/memoryscope/core/chat/memory_chat_prompt.yaml similarity index 100% rename from memoryscope/chat/cli_memory_chat.yaml rename to memoryscope/core/chat/memory_chat_prompt.yaml diff --git a/memoryscope/memory/operation/__init__.py b/memoryscope/core/config/__init__.py similarity index 100% rename from memoryscope/memory/operation/__init__.py rename to memoryscope/core/config/__init__.py diff --git a/memoryscope/core/config/arguments.py b/memoryscope/core/config/arguments.py new file mode 100644 index 00000000..44ed574f --- /dev/null +++ b/memoryscope/core/config/arguments.py @@ -0,0 +1,61 @@ +from dataclasses import dataclass, field +from typing import Literal, Dict + + +@dataclass +class Arguments(object): + language: Literal["cn", "en"] = field(default="en", metadata={"help": "support en & cn now"}) + + thread_pool_max_workers: int = field(default=5, metadata={"help": "thread pool max workers"}) + + logger_name: str = field(default="memoryscope") + + logger_name_time_suffix: str = field(default="%Y%m%d_%H%M%S") + + logger_to_screen: bool = field(default=False, metadata={"help": "If false, it does not print to the screen."}) + + memory_chat_class: str = field(default="cli_memory_chat", metadata={ + "help": "cli_memory_chat(Command-line interaction), api_memory_chat(API interface interaction), etc."}) + + consolidate_memory_interval_time: int = field(default=1, metadata={ + "help": "If you feel that the token consumption is relatively high, please increase the time interval."}) + + reflect_and_reconsolidate_interval_time: int = field(default=15, metadata={ + "help": "If you feel that the token consumption is relatively high, please increase the time interval."}) + + worker_params: Dict[str, dict] = field(default_factory=lambda: {}, metadata={ + "help": "dict format: worker_name -> param_key -> param_value"}) + + generation_backend: str = field(default="openai_generation", metadata={ + "help": "global generation backend: openai_generation, dashscope_generation, etc."}) + + generation_model: str = field(default="gpt-4o", metadata={ + "help": "global generation model: gpt-4o, gpt-4, qwen-max, etc."}) + + generation_params: dict = field(default_factory=lambda: {}, metadata={ + "help": "global generation params: max_tokens, top_p, temperature, etc."}) + + embedding_backend: str = field(default="openai_embedding", metadata={ + "help": "global embedding backend: openai_embedding, dashscope_embedding, etc."}) + + embedding_model: str = field(default="text-embedding-ada-002", metadata={ + "help": "global embedding model: text-embedding-ada-002, text-embedding-v2, etc."}) + + embedding_params: dict = field(default_factory=lambda: {}) + + use_dummy_ranker: bool = field(default=True, metadata={ + "help": "If a semantic ranking model is not available, MemoryScope will use cosine similarity scoring as a " + "substitute. However, the ranking effectiveness will be somewhat compromised."}) + + rank_backend: str = field(default="dashscope_rank", metadata={"help": "global rank backend: dashscope_rank, etc."}) + + rank_model: str = field(default="gte-rerank", metadata={"help": "global rank model: gte-rerank, etc."}) + + rank_params: dict = field(default_factory=lambda: {}) + + es_index_name: str = field(default="memory_index") + + es_url: str = field(default="http://localhost:9200") + + retrieve_mode: str = field(default="dense", metadata={ + "help": "retrieve_mode: dense, sparse(not implemented), hybrid(not implemented)"}) diff --git a/memoryscope/core/config/config_manager.py b/memoryscope/core/config/config_manager.py new file mode 100644 index 00000000..e30d51e3 --- /dev/null +++ b/memoryscope/core/config/config_manager.py @@ -0,0 +1,189 @@ +import json +from dataclasses import fields +from datetime import datetime +from pathlib import Path +from typing import Optional, Literal + +import yaml + +from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME +from memoryscope.core.config.arguments import Arguments +from memoryscope.core.utils.logger import Logger +from memoryscope.enumeration.language_enum import LanguageEnum + + +class ConfigManager(object): + + def __init__(self, + config: dict = None, + config_path: Optional[str] = None, + arguments: Optional[Arguments] = None, + demo_config_name: str = "demo_config.yaml", + **kwargs): + self.config: dict = {} + self.kwargs = kwargs + + if config: + self.config = config + self.logger = self._init_logger() + self.logger.info("init by config mode:") + + elif config_path: + self.read_config(config_path) + self.logger = self._init_logger() + self.logger.info("init by config_path mode:") + + else: + self.read_demo_config(demo_config_name) + if arguments: + self.update_config_by_arguments(arguments) + self.logger = self._init_logger() + self.logger.info(f"init by arguments mode: {arguments.__dict__}") + + elif kwargs: + kwargs = {k: v for k, v in kwargs.items() if k in [x.name for x in fields(Arguments)]} + arguments = Arguments(**kwargs) + self.update_config_by_arguments(arguments) + self.logger = self._init_logger() + self.logger.info(f"init by kwargs mode: {kwargs}") + + else: + raise RuntimeError("can not init config manager without kwargs!") + self.logger.info(self.dump_config()) + + def _init_logger(self) -> Logger: + global_config = self.config["global"] + logger_name = global_config["logger_name"] + logger_name_time_suffix = global_config["logger_name_time_suffix"] + if logger_name_time_suffix: + suffix = datetime.now().strftime(logger_name_time_suffix) + logger_name = f"{logger_name}_{suffix}" + return Logger.get_logger(logger_name, to_stream=global_config["logger_to_screen"]) + + def read_config(self, config_path: str): + if config_path.endswith(".yaml"): + with open(config_path) as f: + self.config = yaml.load(f, yaml.FullLoader) + + elif config_path.endswith(".json"): + with open(config_path) as f: + self.config = json.load(f) + + def read_demo_config(self, demo_config_name: str): + file_path = Path(__file__) + demo_config_path = (file_path.parent / demo_config_name).__str__() + with open(demo_config_path) as f: + self.config = yaml.load(f, yaml.FullLoader) + + @staticmethod + def update_global_by_arguments(config: dict, arguments: Arguments): + config.update({ + "language": arguments.language, + "thread_pool_max_workers": arguments.thread_pool_max_workers, + "logger_name": arguments.logger_name, + "logger_name_time_suffix": arguments.logger_name_time_suffix, + "logger_to_screen": arguments.logger_to_screen, + "use_dummy_ranker": arguments.use_dummy_ranker, + }) + + @staticmethod + def update_memory_chat_by_arguments(config: dict, arguments: Arguments): + memory_chat_class_split = config["class"].split(".") + stream = arguments.memory_chat_class in ["cli_memory_chat", ] + config.update({ + "class": ".".join(memory_chat_class_split[:-1] + [arguments.memory_chat_class]), + "human_name": DEFAULT_HUMAN_NAME[LanguageEnum(arguments.language)], + "assistant_name": "AI", + "stream": stream, + }) + + @staticmethod + def update_memory_service_by_arguments(config: dict, arguments: Arguments): + config.update({ + "human_name": DEFAULT_HUMAN_NAME[LanguageEnum(arguments.language)], + "assistant_name": "AI", + }) + config["memory_operations"]["consolidate_memory"]["interval_time"] = \ + arguments.consolidate_memory_interval_time + config["memory_operations"]["reflect_and_reconsolidate"]["interval_time"] = \ + arguments.reflect_and_reconsolidate_interval_time + + @staticmethod + def update_worker_by_arguments(config: dict, arguments: Arguments): + for worker_name, kv_dict in arguments.worker_params.items(): + if worker_name not in config: + continue + config[worker_name].update(kv_dict) + + @staticmethod + def update_model_by_arguments(config: dict, arguments: Arguments): + config["generation_model"].update({ + "module_name": arguments.generation_backend, + "model_name": arguments.generation_model, + **arguments.generation_params, + }) + + config["embedding_model"].update({ + "module_name": arguments.embedding_backend, + "model_name": arguments.embedding_model, + **arguments.embedding_params, + }) + + config["rank_model"].update({ + "module_name": arguments.rank_backend, + "model_name": arguments.rank_model, + **arguments.rank_params, + }) + + @staticmethod + def update_memory_store_by_arguments(config: dict, arguments: Arguments): + config.update({ + "index_name": arguments.es_index_name, + "es_url": arguments.es_url, + "retrieve_mode": arguments.retrieve_mode}) + + def update_config_by_arguments(self, arguments: Arguments): + # prepare global + self.update_global_by_arguments(self.config["global"], arguments) + + # prepare memory chat + memory_chat_conf_dict = self.config["memory_chat"] + memory_chat_config = list(memory_chat_conf_dict.values())[0] + self.update_memory_chat_by_arguments(memory_chat_config, arguments) + + # prepare memory service + memory_service_conf_dict = self.config["memory_service"] + memory_service_config = list(memory_service_conf_dict.values())[0] + self.update_memory_service_by_arguments(memory_service_config, arguments) + + # prepare worker + self.update_worker_by_arguments(self.config["worker"], arguments) + + # prepare model + self.update_model_by_arguments(self.config["model"], arguments) + + # prepare memory store + self.update_memory_store_by_arguments(self.config["memory_store"], arguments) + + def add_node_object(self, node: str, name: str, config: dict): + self.config[node][name] = config + + def pop_node_object(self, node: str, name: str): + return self.config[node].pop(name, None) + + def clear_node_all(self, node: str): + self.config[node].clear() + + def dump_config(self, file_type: Literal["json", "yaml"] = "yaml", file_path: Optional[str] = None) -> str: + if file_type == "json": + content = json.dumps(self.config, indent=2, ensure_ascii=False) + elif file_type == "yaml": + content = yaml.dump(self.config, indent=2, allow_unicode=True) + else: + raise NotImplementedError + + if file_path: + with open(file_path, "w") as f: + f.write(content) + + return content diff --git a/memoryscope/core/config/demo_config.yaml b/memoryscope/core/config/demo_config.yaml new file mode 100644 index 00000000..60db4b34 --- /dev/null +++ b/memoryscope/core/config/demo_config.yaml @@ -0,0 +1,176 @@ +global: + language: en + thread_pool_max_workers: 5 + logger_name: memoryscope + logger_name_time_suffix: "%Y%m%d_%H%M%S" + logger_to_screen: false + use_dummy_ranker: false + +memory_chat: + cli_memory_chat: + class: core.chat.cli_memory_chat + memory_service: memoryscope_service + generation_model: generation_model + +memory_service: + memoryscope_service: + class: core.service.memory_scope_service + memory_operations: + read_message: + class: core.operation.frontend_operation + workflow: read_message + description: "read short memory" + + retrieve_memory: + class: core.operation.frontend_operation + workflow: set_query,[extract_time|retrieve_obs_ins,semantic_rank],fuse_rerank + description: "retrieve long-term memory" + + list_memory: + class: core.operation.frontend_operation + workflow: set_query,retrieve_top_memory,print_memory + description: "read all long-term memory of the user" + + delete_memory: + class: core.operation.frontend_operation + workflow: set_query,retrieve_all_memory,delete_memory + description: "delete a single long-term memory" + + delete_all: + class: core.operation.frontend_operation + workflow: set_query,retrieve_all_memory,delete_all + description: "delete all long-term memory" + + add_memory: + class: core.operation.frontend_operation + workflow: add_memory + description: "add a single observation" + + consolidate_memory: + class: core.operation.consolidate_memory_op + workflow: info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat,store_memory + description: "summary user's observation memory" + interval_time: 1 + + reflect_and_reconsolidate: + class: core.operation.backend_operation + workflow: load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat,store_memory + description: "summary user's insight memory" + interval_time: 15 + +worker: + dummy: + class: core.worker.dummy_worker + generation_model: generation_model + embedding_model: embedding_model + rank_model: rank_model + read_message: + class: core.worker.frontend.read_message_worker + set_query: + class: core.worker.frontend.set_query_worker + retrieve_obs_ins: + class: core.worker.frontend.retrieve_memory_worker + retrieve_obs_top_k: 100 + retrieve_ins_top_k: 100 + extract_time: + class: core.worker.frontend.extract_time_worker + generation_model: generation_model + semantic_rank: + class: core.worker.frontend.semantic_rank_worker + rank_model: rank_model + fuse_rerank: + class: core.worker.frontend.fuse_rerank_worker + fuse_score_threshold: 0.01 + fuse_ratio_dict: + conversation: 0.5 + observation: 1 + obs_customized: 1.2 + insight: 2.0 + fuse_time_ratio: 2.0 + fuse_rerank_top_k: 20 + retrieve_top_memory: + class: core.worker.frontend.retrieve_memory_worker + retrieve_obs_top_k: 100 + retrieve_ins_top_k: 100 + retrieve_expired_top_k: 100 + print_memory: + class: core.worker.frontend.print_memory_worker + retrieve_all_memory: + class: core.worker.frontend.retrieve_memory_worker + retrieve_obs_top_k: 1000 + retrieve_ins_top_k: 1000 + retrieve_expired_top_k: 1000 + delete_memory: + class: core.worker.backend.update_memory_worker + method: delete_memory + delete_all: + class: core.worker.backend.update_memory_worker + method: delete_all + add_memory: + class: core.worker.backend.update_memory_worker + method: from_query + info_filter: + class: core.worker.backend.info_filter_worker + generation_model: generation_model + load_today_memory: + class: core.worker.backend.load_memory_worker + retrieve_today_top_k: 100 + get_observation: + class: core.worker.backend.get_observation_worker + generation_model: generation_model + get_observation_with_time: + class: core.worker.backend.get_observation_with_time_worker + generation_model: generation_model + contra_repeat: + class: core.worker.backend.contra_repeat_worker + generation_model: generation_model + store_memory: + class: core.worker.backend.update_memory_worker + method: from_memory_key + memory_key: all + load_obs_and_insight: + class: core.worker.backend.load_memory_worker + retrieve_not_reflected_top_k: 100 + retrieve_not_updated_top_k: 100 + retrieve_insight_top_k: 100 + get_reflection_subject: + class: core.worker.backend.get_reflection_subject_worker + generation_model: generation_model + reflect_obs_cnt_threshold: 10 + update_insight: + class: core.worker.backend.update_insight_worker + generation_model: generation_model + rank_model: rank_model + long_contra_repeat: + class: core.worker.backend.long_contra_repeat_worker + generation_model: generation_model + +model: + generation_model: + class: core.models.llama_index_generation_model + module_name: dashscope_generation + model_name: qwen-max + max_tokens: 2000 + embedding_model: + class: core.models.llama_index_embedding_model + module_name: dashscope_embedding + model_name: text-embedding-v2 + rank_model: + class: core.models.llama_index_rank_model + module_name: dashscope_rank + model_name: gte-rerank + top_n: 500 + dummy_generation: + class: core.models.dummy_generation_model + module_name: dummy_generation + model_name: dummy_generation_model + +memory_store: + class: core.storage.llama_index_es_memory_store + embedding_model: embedding_model + index_name: memory_index + es_url: http://localhost:9200 + retrieve_mode: dense + +monitor: + class: core.storage.dummy_monitor \ No newline at end of file diff --git a/memoryscope/core/memoryscope.py b/memoryscope/core/memoryscope.py new file mode 100644 index 00000000..1f1ece7b --- /dev/null +++ b/memoryscope/core/memoryscope.py @@ -0,0 +1,94 @@ +from concurrent.futures import ThreadPoolExecutor + +from memoryscope.core.chat.base_memory_chat import BaseMemoryChat +from memoryscope.core.config.config_manager import ConfigManager +from memoryscope.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.service.base_memory_service import BaseMemoryService +from memoryscope.core.utils.tool_functions import init_instance_by_config +from memoryscope.enumeration.language_enum import LanguageEnum +from memoryscope.enumeration.model_enum import ModelEnum + + +class MemoryScope(ConfigManager): + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.context: MemoryscopeContext = MemoryscopeContext() + self.init_context_by_config() + + def init_context_by_config(self): + # set global config + global_conf = self.config["global"] + self.context.language = LanguageEnum(global_conf["language"]) + self.context.thread_pool = ThreadPoolExecutor(max_workers=global_conf["thread_pool_max_workers"]) + self.context.meta_data["use_dummy_ranker"] = global_conf["use_dummy_ranker"] + + # init memory_chat + memory_chat_conf_dict = self.config["memory_chat"] + if memory_chat_conf_dict: + for name, conf in memory_chat_conf_dict.items(): + self.context.memory_chat_dict[name] = init_instance_by_config(conf, name=name, context=self.context) + + # set memory_service + memory_service_conf_dict = self.config["memory_service"] + assert memory_service_conf_dict + for name, conf in memory_service_conf_dict.items(): + self.context.memory_service_dict[name] = init_instance_by_config(conf, name=name, context=self.context) + + # init model + model_conf_dict = self.config["model"] + assert model_conf_dict + for name, conf in model_conf_dict.items(): + self.context.model_dict[name] = init_instance_by_config(conf, name=name) + + # init memory_store + memory_store_conf = self.config["memory_store"] + assert memory_store_conf + emb_model_name: str = memory_store_conf[ModelEnum.EMBEDDING_MODEL.value] + embedding_model = self.context.model_dict[emb_model_name] + self.context.memory_store = init_instance_by_config(memory_store_conf, embedding_model=embedding_model) + + # init monitor + monitor_conf = self.config["monitor"] + if monitor_conf: + self.context.monitor = init_instance_by_config(monitor_conf) + + # set worker config + self.context.worker_conf_dict = self.config["worker"] + + def close(self): + # wait service to stop + for _, service in self.context.memory_service_dict.items(): + service.stop_backend_service(wait_service_end=True) + + self.context.thread_pool.shutdown() + + self.context.memory_store.close() + + if self.context.monitor: + self.context.monitor.close() + + self.logger.close() + + def __enter__(self): + self.init_context_by_config() + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + self.close() + + @property + def memory_chat_dict(self): + return self.context.memory_chat_dict + + @property + def memory_service_dict(self): + return self.context.memory_service_dict + + @property + def default_memory_chat(self) -> BaseMemoryChat: + return list(self.memory_chat_dict.values())[0] + + @property + def default_service(self) -> BaseMemoryService: + return list(self.memory_service_dict.values())[0] diff --git a/memoryscope/core/memoryscope_context.py b/memoryscope/core/memoryscope_context.py new file mode 100644 index 00000000..ac8d89fb --- /dev/null +++ b/memoryscope/core/memoryscope_context.py @@ -0,0 +1,29 @@ +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field + +from memoryscope.enumeration.language_enum import LanguageEnum + + +@dataclass +class MemoryscopeContext(object): + """ + The context class archives all configs utilized by store, monitor, services and workers. + """ + + language: LanguageEnum = LanguageEnum.EN + + thread_pool: ThreadPoolExecutor | None = None + + memory_store = None + + monitor = None + + memory_chat_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> memory_chat"}) + + memory_service_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> memory_service"}) + + model_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> model"}) + + worker_conf_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> worker_conf"}) + + meta_data: dict = field(default_factory=lambda: {}) diff --git a/memoryscope/memory/service/__init__.py b/memoryscope/core/models/__init__.py similarity index 100% rename from memoryscope/memory/service/__init__.py rename to memoryscope/core/models/__init__.py diff --git a/memoryscope/models/base_model.py b/memoryscope/core/models/base_model.py similarity index 96% rename from memoryscope/models/base_model.py rename to memoryscope/core/models/base_model.py index 83b5aa50..07cbae69 100644 --- a/memoryscope/models/base_model.py +++ b/memoryscope/core/models/base_model.py @@ -3,11 +3,11 @@ import time from abc import abstractmethod, ABCMeta from typing import Any +from memoryscope.core.utils.logger import Logger +from memoryscope.core.utils.registry import Registry +from memoryscope.core.utils.timer import Timer from memoryscope.enumeration.model_enum import ModelEnum from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen -from memoryscope.utils.logger import Logger -from memoryscope.utils.registry import Registry -from memoryscope.utils.timer import Timer MODEL_REGISTRY = Registry("models") diff --git a/memoryscope/models/dummy_generation_model.py b/memoryscope/core/models/dummy_generation_model.py similarity index 51% rename from memoryscope/models/dummy_generation_model.py rename to memoryscope/core/models/dummy_generation_model.py index 5949bf8e..e1d81438 100644 --- a/memoryscope/models/dummy_generation_model.py +++ b/memoryscope/core/models/dummy_generation_model.py @@ -3,9 +3,9 @@ from typing import List from llama_index.core.base.llms.types import ChatMessage +from memoryscope.core.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.enumeration.message_role_enum import MessageRoleEnum from memoryscope.enumeration.model_enum import ModelEnum -from memoryscope.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.scheme.message import Message from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen @@ -19,38 +19,37 @@ class DummyGenerationModel(BaseModel): """ m_type: ModelEnum = ModelEnum.GENERATION_MODEL - class DummyModel: - """ - An inner class representing the dummy model placeholder. - """ - pass + MODEL_REGISTRY.register("dummy_generation", object) - MODEL_REGISTRY.register("dummy_generation", DummyModel) - - def before_call(self, **kwargs): + def before_call(self, model_response: ModelResponse, **kwargs): """ - Prepares the input data before making a call to the model's generate function. - Accepts either a 'prompt' or a list of 'messages'. If both are provided or missing, - a RuntimeError is raised. Transforms the input into a standardized format for processing. + Prepares the input data before making a call to the language model. + It accepts either a 'prompt' directly or a list of 'messages'. + If 'prompt' is provided, it sets the data accordingly. + If 'messages' are provided, it constructs a list of ChatMessage objects from the list. + Raises an error if neither 'prompt' nor 'messages' are supplied. Args: - **kwargs: Arbitrary keyword arguments including 'prompt' or 'messages'. - + model_response: model_response + **kwargs: Arbitrary keyword arguments including 'prompt' and 'messages'. + Raises: - RuntimeError: If neither 'prompt' nor 'messages' is provided, or both are provided. + RuntimeError: When both 'prompt' and 'messages' inputs are not provided. """ prompt: str = kwargs.pop("prompt", "") messages: List[Message] | List[dict] = kwargs.pop("messages", []) if prompt: - self.data = {"prompt": prompt} + data = {"prompt": prompt} elif messages: if isinstance(messages[0], dict): - self.data = {"messages": [ChatMessage(role=msg["role"], content=msg["content"]) for msg in messages]} + 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]} + data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]} else: - raise RuntimeError("Both 'prompt' and 'messages' are empty!") + raise RuntimeError("prompt and messages are both empty!") + data.update(**kwargs) + model_response.meta_data["data"] = data def after_call(self, model_response: ModelResponse, @@ -79,43 +78,16 @@ class DummyGenerationModel(BaseModel): for delta in call_result: model_response.message.content += delta model_response.delta = delta - time.sleep(0.1) # ⭐ Introduce a delay to simulate streaming + time.sleep(0.1) yield model_response return gen() else: - model_response.message.content = "".join(call_result) # ⭐ Concatenate results for non-streaming + model_response.message.content = "".join(call_result) return model_response - def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen: - """ - Generates a dummy response based on the input data, supporting both immediate - and streamed response types. + def _call(self, model_response: ModelResponse, stream: bool = False, **kwargs): + return model_response - Args: - stream (bool, optional): If True, indicates the response should be generated - in a streaming manner. Defaults to False. - **kwargs: Additional keyword arguments not used in this dummy implementation. - - Returns: - Union[ModelResponse, ModelResponseGen]: A dummy response object or a generator - object capable of streaming responses. - """ - assert "prompt" in self.data or "messages" in self.data - results = ModelResponse(m_type=self.m_type) - return results - - async def _async_call(self, **kwargs) -> ModelResponse: - """ - Asynchronous version of `_call`, providing the same functionality but designed - to be used in asynchronous contexts. - - Args: - **kwargs: Additional keyword arguments not used in this dummy implementation. - - Returns: - ModelResponse: A dummy response object suitable for asynchronous use. - """ - assert "prompt" in self.data or "messages" in self.data - results = ModelResponse(m_type=self.m_type) - return results + async def _async_call(self, model_response: ModelResponse, **kwargs): + return model_response diff --git a/memoryscope/models/llama_index_embedding_model.py b/memoryscope/core/models/llama_index_embedding_model.py similarity index 97% rename from memoryscope/models/llama_index_embedding_model.py rename to memoryscope/core/models/llama_index_embedding_model.py index 822f619d..efde959b 100644 --- a/memoryscope/models/llama_index_embedding_model.py +++ b/memoryscope/core/models/llama_index_embedding_model.py @@ -2,8 +2,8 @@ from typing import List from llama_index.embeddings.dashscope import DashScopeEmbedding +from memoryscope.core.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.enumeration.model_enum import ModelEnum -from memoryscope.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.scheme.model_response import ModelResponse diff --git a/memoryscope/models/llama_index_generation_model.py b/memoryscope/core/models/llama_index_generation_model.py similarity index 98% rename from memoryscope/models/llama_index_generation_model.py rename to memoryscope/core/models/llama_index_generation_model.py index 63a249d9..230bde67 100644 --- a/memoryscope/models/llama_index_generation_model.py +++ b/memoryscope/core/models/llama_index_generation_model.py @@ -3,9 +3,9 @@ from typing import List from llama_index.core.base.llms.types import ChatMessage, ChatResponse, CompletionResponse from llama_index.llms.dashscope import DashScope +from memoryscope.core.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.enumeration.message_role_enum import MessageRoleEnum from memoryscope.enumeration.model_enum import ModelEnum -from memoryscope.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.scheme.message import Message from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen diff --git a/memoryscope/models/llama_index_rank_model.py b/memoryscope/core/models/llama_index_rank_model.py similarity index 98% rename from memoryscope/models/llama_index_rank_model.py rename to memoryscope/core/models/llama_index_rank_model.py index 38debfd7..9e54c545 100644 --- a/memoryscope/models/llama_index_rank_model.py +++ b/memoryscope/core/models/llama_index_rank_model.py @@ -4,8 +4,8 @@ from llama_index.core.data_structs import Node from llama_index.core.schema import NodeWithScore from llama_index.postprocessor.dashscope_rerank import DashScopeRerank +from memoryscope.core.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.enumeration.model_enum import ModelEnum -from memoryscope.models.base_model import BaseModel, MODEL_REGISTRY from memoryscope.scheme.model_response import ModelResponse diff --git a/memoryscope/memory/worker/__init__.py b/memoryscope/core/operation/__init__.py similarity index 100% rename from memoryscope/memory/worker/__init__.py rename to memoryscope/core/operation/__init__.py diff --git a/memoryscope/memory/operation/backend_operation.py b/memoryscope/core/operation/backend_operation.py similarity index 79% rename from memoryscope/memory/operation/backend_operation.py rename to memoryscope/core/operation/backend_operation.py index 2fcbe398..e4b4a41c 100644 --- a/memoryscope/memory/operation/backend_operation.py +++ b/memoryscope/core/operation/backend_operation.py @@ -2,11 +2,10 @@ import time from typing import List from memoryscope.constants.common_constants import CHAT_KWARGS, RESULT, CHAT_MESSAGES -from memoryscope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE -from memoryscope.memory.operation.base_workflow import BaseWorkflow +from memoryscope.core.operation.base_operation import BaseOperation, OPERATION_TYPE +from memoryscope.core.operation.base_workflow import BaseWorkflow +from memoryscope.core.utils.logger import Logger from memoryscope.scheme.message import Message -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger class BackendOperation(BaseWorkflow, BaseOperation): @@ -30,7 +29,7 @@ class BackendOperation(BaseWorkflow, BaseOperation): self._operation_status_run: bool = False self._loop_switch: bool = False - self._run_thread = None + self._backend_task = None self.logger = Logger.get_logger() @@ -54,17 +53,14 @@ class BackendOperation(BaseWorkflow, BaseOperation): Returns: Any: The result obtained after executing the workflow. """ - self.context.clear() - - # Add additional arguments to the context - kwargs.update(**self.kwargs) - self.context[CHAT_KWARGS] = kwargs - - # Include the most recent messages in the operation context - self.context[CHAT_MESSAGES] = self.chat_messages + # prepare kwargs + workflow_kwargs = { + CHAT_MESSAGES: self.chat_messages, + CHAT_KWARGS: {**kwargs, **self.kwargs}, + } # Execute the workflow with the prepared context - self.run_workflow() + self.run_workflow(**workflow_kwargs) # Retrieve the result from the context after workflow execution return self.context.get(RESULT) @@ -107,17 +103,24 @@ class BackendOperation(BaseWorkflow, BaseOperation): if self._loop_switch: self.run_operation() - def run_operation_backend(self): + def start_operation_backend(self): """ Initiates the background operation loop if it's not already running. Sets the _loop_switch to True and submits the _loop_operation to a thread from the global thread pool. """ if not self._loop_switch: self._loop_switch = True - self._run_thread = G_CONTEXT.thread_pool.submit(self._loop_operation) + self._backend_task = self.thread_pool.submit(self._loop_operation) + self.logger.info(f"start operation={self.name}...") - def stop_operation_backend(self): + def stop_operation_backend(self, wait_task_end: bool = False): """ Stops the background operation loop by setting the _loop_switch to False. """ self._loop_switch = False + if self._backend_task: + if wait_task_end: + self._backend_task.result() + self.logger.info(f"stop operation={self.name}...") + else: + self.logger.info(f"send stop signal to operation={self.name}...") diff --git a/memoryscope/memory/operation/base_operation.py b/memoryscope/core/operation/base_operation.py similarity index 92% rename from memoryscope/memory/operation/base_operation.py rename to memoryscope/core/operation/base_operation.py index 8904530c..b5276a7c 100644 --- a/memoryscope/memory/operation/base_operation.py +++ b/memoryscope/core/operation/base_operation.py @@ -12,7 +12,6 @@ class BaseOperation(metaclass=ABCMeta): operation_type (OPERATION_TYPE): Specifies the type of operation, defaulting to "frontend". name (str): The name of the operation. description (str): A description of the operation. - kwargs (dict): Additional keyword arguments for operation configuration. """ operation_type: OPERATION_TYPE = "frontend" @@ -51,14 +50,14 @@ class BaseOperation(metaclass=ABCMeta): """ raise NotImplementedError - def run_operation_backend(self): + def start_operation_backend(self): """ Placeholder method for running an operation specific to the backend. Intended to be overridden by subclasses if backend operations are required. """ pass - def stop_operation_backend(self): + def stop_operation_backend(self, wait_task_end: bool = False): """ Placeholder method to stop any ongoing backend operations. Should be implemented in subclasses where backend operations are managed. diff --git a/memoryscope/memory/operation/base_workflow.py b/memoryscope/core/operation/base_workflow.py similarity index 86% rename from memoryscope/memory/operation/base_workflow.py rename to memoryscope/core/operation/base_workflow.py index 926b59a1..4c6ec8fa 100644 --- a/memoryscope/memory/operation/base_workflow.py +++ b/memoryscope/core/operation/base_workflow.py @@ -4,25 +4,26 @@ from concurrent.futures import ThreadPoolExecutor, as_completed from itertools import zip_longest from typing import Dict, Any, List -from memoryscope.constants.common_constants import WORKFLOW_NAME -from memoryscope.memory.worker.base_worker import BaseWorker -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger -from memoryscope.utils.timer import Timer -from memoryscope.utils.tool_functions import init_instance_by_config +from memoryscope.constants.common_constants import WORKFLOW_NAME, MEMORYSCOPE_CONTEXT +from memoryscope.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.utils.logger import Logger +from memoryscope.core.utils.timer import Timer +from memoryscope.core.utils.tool_functions import init_instance_by_config +from memoryscope.core.worker.base_worker import BaseWorker class BaseWorkflow(object): def __init__(self, name: str, + memoryscope_context: MemoryscopeContext, workflow: str = "", - thread_pool: ThreadPoolExecutor = G_CONTEXT.thread_pool, **kwargs): self.name: str = name + self.memoryscope_context: MemoryscopeContext = memoryscope_context + self.thread_pool: ThreadPoolExecutor = self.memoryscope_context.thread_pool self.workflow: str = workflow - self.thread_pool: ThreadPoolExecutor = thread_pool self.kwargs = kwargs self.workflow_worker_list: List[List[List[str]]] = [] @@ -128,17 +129,16 @@ class BaseWorkflow(object): This method modifies `self.worker_dict` in-place, replacing the keys with actual worker instances. """ for name in list(self.worker_dict.keys()): - if name not in G_CONTEXT.worker_config: - raise RuntimeError(f"worker={name} is not exists in worker_config!") + if name not in self.memoryscope_context.worker_conf_dict: + raise RuntimeError(f"worker={name} is not exists in worker config!") self.worker_dict[name] = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.memoryscope_context.worker_conf_dict[name], name=name, is_multi_thread=is_backend or self.worker_dict[name], context=self.context, context_lock=self.context_lock, - thread_pool=G_CONTEXT.thread_pool, + thread_pool=self.thread_pool, **kwargs) def _run_sub_workflow(self, worker_list: List[str]) -> bool: @@ -150,7 +150,7 @@ class BaseWorkflow(object): return False return True - def run_workflow(self): + def run_workflow(self, **kwargs): """ Executes the workflow by orchestrating the steps defined in `self.workflow_worker_list`. This method supports both sequential and parallel execution of sub-workflows based on the structure @@ -159,9 +159,18 @@ class BaseWorkflow(object): If a workflow part consists of a single item, it is executed sequentially. For parts with multiple items, they are submitted for parallel execution using a thread pool. The workflow will stop if any sub-workflow returns False. + + Args: + **kwargs: Additional keyword arguments to be passed to context. """ with Timer(f"workflow.{self.name}", time_log_type="wrap"): - self.context[WORKFLOW_NAME] = self.name + self.context.clear() + + self.context.update({ + WORKFLOW_NAME: self.name, + MEMORYSCOPE_CONTEXT: self.memoryscope_context, + **kwargs, + }) # Iterate over each part of the workflow for workflow_part in self.workflow_worker_list: @@ -174,7 +183,7 @@ class BaseWorkflow(object): t_list = [] # Submit tasks to the thread pool for sub_workflow in workflow_part: - t_list.append(G_CONTEXT.thread_pool.submit(self._run_sub_workflow, sub_workflow)) + t_list.append(self.thread_pool.submit(self._run_sub_workflow, sub_workflow)) # Check results; if any task returns False, stop the workflow flag = True diff --git a/memoryscope/memory/operation/summary_observation_op.py b/memoryscope/core/operation/consolidate_memory_op.py similarity index 79% rename from memoryscope/memory/operation/summary_observation_op.py rename to memoryscope/core/operation/consolidate_memory_op.py index 1da2e65c..fec60e12 100644 --- a/memoryscope/memory/operation/summary_observation_op.py +++ b/memoryscope/core/operation/consolidate_memory_op.py @@ -1,12 +1,12 @@ from memoryscope.constants.common_constants import CHAT_KWARGS, CHAT_MESSAGES, RESULT +from memoryscope.core.operation.backend_operation import BackendOperation from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.operation.backend_operation import BackendOperation -class SummaryObservationOp(BackendOperation): +class ConsolidateMemoryOp(BackendOperation): def __init__(self, **kwargs): - super(SummaryObservationOp, self).__init__(**kwargs) + super(ConsolidateMemoryOp, self).__init__(**kwargs) self.message_lock = kwargs.get("message_lock", None) self.contextual_msg_min_count: int = kwargs.get("contextual_msg_min_count", 0) @@ -43,17 +43,14 @@ class SummaryObservationOp(BackendOperation): f"contextual_msg_min_count({self.contextual_msg_min_count}), skip.") return - self.context.clear() - - # Add additional arguments to the context - kwargs.update(**self.kwargs) - self.context[CHAT_KWARGS] = kwargs - - # Include the most recent messages in the operation context - self.context[CHAT_MESSAGES] = chat_messages + # prepare kwargs + workflow_kwargs = { + CHAT_MESSAGES: chat_messages, + CHAT_KWARGS: {**kwargs, **self.kwargs}, + } # Execute the workflow with the prepared context - self.run_workflow() + self.run_workflow(**workflow_kwargs) # Retrieve the result from the context after workflow execution result = self.context.get(RESULT) diff --git a/memoryscope/memory/operation/frontend_operation.py b/memoryscope/core/operation/frontend_operation.py similarity index 75% rename from memoryscope/memory/operation/frontend_operation.py rename to memoryscope/core/operation/frontend_operation.py index 8cde7977..bed28875 100644 --- a/memoryscope/memory/operation/frontend_operation.py +++ b/memoryscope/core/operation/frontend_operation.py @@ -1,8 +1,8 @@ from typing import List from memoryscope.constants.common_constants import RESULT, CHAT_MESSAGES, CHAT_KWARGS -from memoryscope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE -from memoryscope.memory.operation.base_workflow import BaseWorkflow +from memoryscope.core.operation.base_operation import BaseOperation, OPERATION_TYPE +from memoryscope.core.operation.base_workflow import BaseWorkflow from memoryscope.scheme.message import Message @@ -39,17 +39,15 @@ class FrontendOperation(BaseWorkflow, BaseOperation): Returns: Any: The result obtained from executing the workflow. """ - self.context.clear() - # Include the most recent messages in the operation context - self.context[CHAT_MESSAGES] = self.chat_messages - - # Add additional arguments to the context - kwargs.update(**self.kwargs) - self.context[CHAT_KWARGS] = kwargs + # prepare kwargs + workflow_kwargs = { + CHAT_MESSAGES: self.chat_messages, + CHAT_KWARGS: {**kwargs, **self.kwargs}, + } # Execute the workflow with the prepared context - self.run_workflow() + self.run_workflow(**workflow_kwargs) # Retrieve the result from the context after workflow execution return self.context.get(RESULT) diff --git a/memoryscope/memory/worker/backend/__init__.py b/memoryscope/core/service/__init__.py similarity index 100% rename from memoryscope/memory/worker/backend/__init__.py rename to memoryscope/core/service/__init__.py diff --git a/memoryscope/core/service/base_memory_service.py b/memoryscope/core/service/base_memory_service.py new file mode 100644 index 00000000..68acc566 --- /dev/null +++ b/memoryscope/core/service/base_memory_service.py @@ -0,0 +1,82 @@ +from abc import ABCMeta, abstractmethod +from typing import List, Dict + +from memoryscope.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.operation.base_operation import BaseOperation +from memoryscope.core.utils.logger import Logger +from memoryscope.scheme.message import Message + + +class BaseMemoryService(metaclass=ABCMeta): + """ + An abstract base class for managing memory operations within a multithreaded context. + It sets up the infrastructure for operation handling, message storage, and synchronization, + along with logging capabilities and customizable configurations. + """ + + def __init__(self, memory_operations: Dict[str, dict], context: MemoryscopeContext, **kwargs): + """ + Initializes the BaseMemoryService with operation definitions, keys for memory access, + and additional keyword arguments for flexibility. + + Args: + memory_operations (Dict[str, dict]): A dictionary defining available memory operations. + **kwargs: Additional parameters to customize service behavior. + """ + self.memory_operations_conf: Dict[str, dict] = memory_operations + self.context: MemoryscopeContext = context + self.kwargs = kwargs + + self._operation_dict: Dict[str, BaseOperation] = {} + self.chat_messages: List[Message] = [] + self.logger = Logger.get_logger() + + @property + def op_description_dict(self) -> Dict[str, str]: + """ + Property to retrieve a dictionary mapping operation keys to their descriptions. + Returns: + Dict[str, str]: A dictionary where keys are operation identifiers and values are their descriptions. + """ + return {k: v.description for k, v in self._operation_dict.items()} + + @abstractmethod + def add_messages(self, messages: List[Message] | Message): + raise NotImplementedError + + @abstractmethod + def register_operation(self, name: str, operation_config: dict, **kwargs): + raise NotImplementedError + + @abstractmethod + def init_service(self, **kwargs): + raise NotImplementedError + + def start_backend_service(self, name: str = None): + pass + + def stop_backend_service(self, wait_service_end: bool = False): + pass + + def do_operation(self, name: str, **kwargs): + """ + Executes a specific operation by its name with provided keyword arguments. + + Args: + name (str): The name of the operation to execute. + **kwargs: Keyword arguments for the operation's execution. + + Returns: + The result of the operation execution, if any. Otherwise, None. + + Raises: + Warning: If the operation name is not initialized in `_operation_dict`. + """ + if name not in self._operation_dict: + self.logger.warning(f"operation={name} is not registered!") + return + return self._operation_dict[name].run_operation(**kwargs) + + def __getattr__(self, name: str): + assert name in self._operation_dict, f"operation={name} is not registered!" + return lambda **kwargs: self.do_operation(name=name, **kwargs) diff --git a/memoryscope/memory/service/memory_scope_service.py b/memoryscope/core/service/memory_scope_service.py similarity index 53% rename from memoryscope/memory/service/memory_scope_service.py rename to memoryscope/core/service/memory_scope_service.py index da4f4b76..e60dfd2f 100644 --- a/memoryscope/memory/service/memory_scope_service.py +++ b/memoryscope/core/service/memory_scope_service.py @@ -1,9 +1,10 @@ +import threading from typing import List -from memoryscope.memory.operation.base_operation import BaseOperation -from memoryscope.memory.service.base_memory_service import BaseMemoryService +from memoryscope.core.operation.base_operation import BaseOperation +from memoryscope.core.service.base_memory_service import BaseMemoryService +from memoryscope.core.utils.tool_functions import init_instance_by_config from memoryscope.scheme.message import Message -from memoryscope.utils.tool_functions import init_instance_by_config class MemoryScopeService(BaseMemoryService): @@ -11,6 +12,8 @@ class MemoryScopeService(BaseMemoryService): history_msg_count: int = 100, contextual_msg_max_count: int = 20, contextual_msg_min_count: int = 0, + human_name: str = None, + assistant_name: str = None, **kwargs): """ init function. @@ -20,13 +23,21 @@ class MemoryScopeService(BaseMemoryService): it will not be included in the context to prevent token overflow. contextual_msg_min_count (int): The minimum context length in a conversation. If it is shorter than this length, no conversation summary will be made and no long-term memory will be generated. - kwargs (dict): other kwargs + human_name (str): human name. + assistant_name (str): assistant name. + kwargs (dict): other kwargs. """ super().__init__(**kwargs) self.history_msg_count: int = history_msg_count self.contextual_msg_max_count: int = contextual_msg_max_count self.contextual_msg_min_count: int = contextual_msg_min_count assert history_msg_count >= contextual_msg_max_count >= contextual_msg_min_count + if human_name: + self.context.meta_data["human_name"] = human_name + if assistant_name: + self.context.meta_data["assistant_name"] = assistant_name + + self.message_lock = threading.Lock() def add_messages(self, messages: List[Message] | Message): """ @@ -54,60 +65,45 @@ class MemoryScopeService(BaseMemoryService): for _ in range(gap_size): self.chat_messages.pop(0) - def do_operation(self, op_name: str, **kwargs): - """ - Executes a specific operation by its name with provided keyword arguments. - - Args: - op_name (str): The name of the operation to execute. - **kwargs: Keyword arguments for the operation's execution. - - Returns: - The result of the operation execution, if any. Otherwise, None. - - Raises: - Warning: If the operation name is not initialized in `_operation_dict`. - """ - if op_name not in self._operation_dict: - self.logger.warning(f"op_name={op_name} is not inited!") # Warn if operation not initialized + def register_operation(self, name: str, operation_config: dict, **kwargs): + if name in self._operation_dict: + self.logger.warning(f"op_name={name} is registered before!") return - return self._operation_dict[op_name].run_operation(**kwargs) # Execute the operation + + operation: BaseOperation = init_instance_by_config( + config=operation_config, + name=name, + chat_messages=self.chat_messages, + message_lock=self.message_lock, + memoryscope_context=self.context, + contextual_msg_max_count=self.contextual_msg_max_count, + contextual_msg_min_count=self.contextual_msg_min_count) + + # Initialize workflow for each operation + operation.init_workflow(**kwargs) + self._operation_dict[name] = operation + self.logger.info(f"service={self.__class__.__name__} init operation={name}") def init_service(self, **kwargs): - for name, operation_config in self.memory_operations.items(): - if name in self._operation_dict: - self.logger.warning(f"memory operation={name} is repeated!") - continue + for name, operation_config in self.memory_operations_conf.items(): + self.register_operation(name, operation_config, **kwargs) - # ⭐ Initialize operation instance by its config - operation: BaseOperation = init_instance_by_config( - config=operation_config, - name=name, - chat_messages=self.chat_messages, - message_lock=self.message_lock, - contextual_msg_max_count=self.contextual_msg_max_count, - contextual_msg_min_count=self.contextual_msg_min_count) - operation.init_workflow(**kwargs) # Initialize workflow for each operation - - self._operation_dict[name] = operation - self.logger.info(f"service={self.__class__.__name__} init operation={name}") - - def start_backend_service(self): + def start_backend_service(self, name: str = None): """ Start all backend operations. """ - for _, operation in self._operation_dict.items(): - if operation.operation_type == "backend": - # Run backend operations - operation.run_operation_backend() - self.logger.info(f"start operation={operation.name}...") + for op_name, operation in self._operation_dict.items(): + if name: + if op_name == name: + operation.start_operation_backend() + else: + if operation.operation_type == "backend": + operation.start_operation_backend() - def stop_backend_service(self): + def stop_backend_service(self, wait_service_end: bool = False): """ Stops all backend operations that are currently running. """ for _, operation in self._operation_dict.items(): if operation.operation_type == "backend": - # Stop backend operations - operation.stop_operation_backend() - self.logger.info(f"stop operation={operation.name}...") + operation.stop_operation_backend(wait_task_end=wait_service_end) diff --git a/memoryscope/memory/worker/frontend/__init__.py b/memoryscope/core/storage/__init__.py similarity index 100% rename from memoryscope/memory/worker/frontend/__init__.py rename to memoryscope/core/storage/__init__.py diff --git a/memoryscope/storage/base_memory_store.py b/memoryscope/core/storage/base_memory_store.py similarity index 100% rename from memoryscope/storage/base_memory_store.py rename to memoryscope/core/storage/base_memory_store.py diff --git a/memoryscope/storage/base_monitor.py b/memoryscope/core/storage/base_monitor.py similarity index 100% rename from memoryscope/storage/base_monitor.py rename to memoryscope/core/storage/base_monitor.py diff --git a/memoryscope/storage/dummy_memory_store.py b/memoryscope/core/storage/dummy_memory_store.py similarity index 92% rename from memoryscope/storage/dummy_memory_store.py rename to memoryscope/core/storage/dummy_memory_store.py index 304818bb..2d5eeede 100644 --- a/memoryscope/storage/dummy_memory_store.py +++ b/memoryscope/core/storage/dummy_memory_store.py @@ -1,8 +1,8 @@ from typing import Dict, List -from memoryscope.models.base_model import BaseModel +from memoryscope.core.models.base_model import BaseModel +from memoryscope.core.storage.base_memory_store import BaseMemoryStore from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.storage.base_memory_store import BaseMemoryStore class DummyMemoryStore(BaseMemoryStore): @@ -12,6 +12,17 @@ class DummyMemoryStore(BaseMemoryStore): semantic retrieval. Actual storage operations are not implemented. """ + def __init__(self, embedding_model: BaseModel, **kwargs): + """ + Initializes the DummyMemoryStore with an embedding model and additional keyword arguments. + + Args: + embedding_model (BaseModel): The model used to embed data for potential similarity-based retrieval. + **kwargs: Additional keyword arguments for configuration or future expansion. + """ + self.embedding_model: BaseModel = embedding_model + self.kwargs = kwargs + def retrieve_memories(self, query: str = "", top_k: int = 3, @@ -24,17 +35,6 @@ class DummyMemoryStore(BaseMemoryStore): filter_dict: Dict[str, List[str]] = None) -> List[MemoryNode]: pass - def __init__(self, embedding_model: BaseModel, **kwargs): - """ - Initializes the DummyMemoryStore with an embedding model and additional keyword arguments. - - Args: - embedding_model (BaseModel): The model used to embed data for potential similarity-based retrieval. - **kwargs: Additional keyword arguments for configuration or future expansion. - """ - self.embedding_model: BaseModel = embedding_model - self.kwargs = kwargs - def batch_insert(self, nodes: List[MemoryNode]): pass diff --git a/memoryscope/storage/dummy_monitor.py b/memoryscope/core/storage/dummy_monitor.py similarity index 92% rename from memoryscope/storage/dummy_monitor.py rename to memoryscope/core/storage/dummy_monitor.py index 4818e0fb..816202e5 100644 --- a/memoryscope/storage/dummy_monitor.py +++ b/memoryscope/core/storage/dummy_monitor.py @@ -1,4 +1,4 @@ -from memoryscope.storage.base_monitor import BaseMonitor +from memoryscope.core.storage.base_monitor import BaseMonitor class DummyMonitor(BaseMonitor): diff --git a/memoryscope/storage/llama_index_es_memory_store.py b/memoryscope/core/storage/llama_index_es_memory_store.py similarity index 81% rename from memoryscope/storage/llama_index_es_memory_store.py rename to memoryscope/core/storage/llama_index_es_memory_store.py index 96b1542f..dabdbc60 100644 --- a/memoryscope/storage/llama_index_es_memory_store.py +++ b/memoryscope/core/storage/llama_index_es_memory_store.py @@ -1,15 +1,17 @@ import random -from typing import Dict, List, Optional +from typing import Dict, List from llama_index.core import VectorStoreIndex from llama_index.core.schema import TextNode, NodeWithScore, QueryBundle -from memoryscope.models.base_model import BaseModel +from memoryscope.core.models.base_model import BaseModel +from memoryscope.core.storage.base_memory_store import BaseMemoryStore +from memoryscope.core.storage.llama_index_sync_elasticsearch import (SyncElasticsearchStore, + ESCombinedRetrieveStrategy, + _to_elasticsearch_filter, + SPECIAL_QUERY) +from memoryscope.core.utils.logger import Logger from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.storage.base_memory_store import BaseMemoryStore -from memoryscope.storage.llama_index_sync_elasticsearch import SyncElasticsearchStore, ESCombinedRetrieveStrategy, \ - _to_elasticsearch_filter -from memoryscope.utils.logger import Logger class LlamaIndexEsMemoryStore(BaseMemoryStore): @@ -19,15 +21,15 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): index_name: str, es_url: str, retrieve_mode: str = "dense", - hybrid_alpha: float = None, + hybrid_alpha: float = None, **kwargs): self.emb_dims = None self.index_name = index_name self.embedding_model: BaseModel = embedding_model + retrieval_strategy = ESCombinedRetrieveStrategy(retrieve_mode=retrieve_mode, hybrid_alpha=hybrid_alpha) self.es_store = SyncElasticsearchStore(index_name=index_name, es_url=es_url, - retrieval_strategy=ESCombinedRetrieveStrategy(retrieve_mode=retrieve_mode, - hybrid_alpha=hybrid_alpha), + retrieval_strategy=retrieval_strategy, **kwargs) # TODO The llamaIndex utilizes some deprecated functions, hence langchain logs warning messages. By @@ -38,7 +40,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): self.logger = Logger.get_logger() def retrieve_memories(self, - query: str = "**--**", + query: str = "", top_k: int = 3, filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]: # if index is not created, return [] @@ -53,8 +55,12 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter, "fields": ['embedding']}, similarity_top_k=top_k, sparse_top_k=top_k) + + if not query: + query = SPECIAL_QUERY + if not query and self.emb_dims: - query = QueryBundle(query_str='**--**', embedding=self.dummy_query_vector()) + query = QueryBundle(query_str=SPECIAL_QUERY, embedding=self.dummy_query_vector()) text_nodes = retriever.retrieve(query) if text_nodes and text_nodes[0].embedding: @@ -80,7 +86,10 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): sparse_top_k=top_k) if not query: - query = QueryBundle(query_str='**--**', embedding=self.dummy_query_vector()) + query = SPECIAL_QUERY + + if not query: + query = QueryBundle(query_str=SPECIAL_QUERY, embedding=self.dummy_query_vector()) text_nodes: List[NodeWithScore] = await retriever.aretrieve(query) @@ -144,7 +153,11 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): return TextNode(id_=memory_node.memory_id, text=memory_node.content, embedding=embedding, - metadata=memory_node.model_dump(exclude={"content", "vector", "score_recall", "score_rank", "score_rerank"})) + metadata=memory_node.model_dump(exclude={"content", + "vector", + "score_recall", + "score_rank", + "score_rerank"})) @staticmethod def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode: diff --git a/memoryscope/storage/llama_index_sync_elasticsearch.py b/memoryscope/core/storage/llama_index_sync_elasticsearch.py similarity index 98% rename from memoryscope/storage/llama_index_sync_elasticsearch.py rename to memoryscope/core/storage/llama_index_sync_elasticsearch.py index 0794b284..a94c9168 100644 --- a/memoryscope/storage/llama_index_sync_elasticsearch.py +++ b/memoryscope/core/storage/llama_index_sync_elasticsearch.py @@ -38,6 +38,8 @@ DISTANCE_STRATEGIES = Literal[ "EUCLIDEAN_DISTANCE", ] +SPECIAL_QUERY: str = "**--**" + def get_elasticsearch_client( url: Optional[str] = None, @@ -133,17 +135,16 @@ def _mode_must_match_retrieval_strategy( class ESCombinedRetrieveStrategy(AsyncDenseVectorStrategy): - def __init__( - self, - *, - distance: DistanceMetric = DistanceMetric.COSINE, - model_id: Optional[str] = None, - retrieve_mode: str = "dense", - rrf: Union[bool, Dict[str, Any]] = True, - text_field: Optional[str] = "text_field", - hybrid_alpha: Optional[float] = None, - ): + self, + *, + distance: DistanceMetric = DistanceMetric.COSINE, + model_id: Optional[str] = None, + retrieve_mode: str = "dense", + rrf: Union[bool, Dict[str, Any]] = True, + text_field: Optional[str] = "text_field", + hybrid_alpha: Optional[float] = None, + ): if retrieve_mode == "dense": self.alpha = 1.0 elif retrieve_mode == "sparse": @@ -152,7 +153,7 @@ class ESCombinedRetrieveStrategy(AsyncDenseVectorStrategy): elif retrieve_mode == "hybrid": # self.alpha = hybrid_alpha raise NotImplementedError - + super().__init__(distance=distance, model_id=model_id, hybrid=True, rrf=rrf, text_field=text_field) def _hybrid(self, query: str, knn: Dict[str, Any], filter: List[Dict[str, Any]], top_k: int) -> Dict[str, Any]: @@ -160,7 +161,7 @@ class ESCombinedRetrieveStrategy(AsyncDenseVectorStrategy): # RRF is used to even the score from the knn query and text query # RRF has two optional parameters: {'rank_constant':int, 'window_size':int} # https://www.elastic.co/guide/en/elasticsearch/reference/current/rrf.html - if query == "**--**": + if query == SPECIAL_QUERY: query_body = { "query": { "bool": { @@ -268,7 +269,7 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str } } ) - result['bool'].update({"should": operands}) # ⭐ Add 'should' clause for OR logic + result['bool'].update({"should": operands}) # Add 'should' clause for OR logic result['bool'].update({"minimum_should_match": 1}) # Ensure at least one 'should' match else: key_str = f"metadata.{key}.keyword" if isinstance(value, str) else f"metadata.{key}" @@ -613,7 +614,6 @@ class SyncElasticsearchStore(BasePydanticVectorStore): ] = None, es_filter: Optional[List[Dict]] = None, fields: List[str] = [], - **kwargs: Any, ) -> VectorStoreQueryResult: """ Asynchronously queries the Elasticsearch index for the top k most similar nodes @@ -626,6 +626,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore): A custom function to modify the Elasticsearch query body. Defaults to None. es_filter (List[Dict], optional): Additional filters to apply during the query. If filters are present in the query, these filters will not be used. Defaults to None. + fields (List[str], optional): . Returns: VectorStoreQueryResult: The result of the query, including nodes, their IDs, diff --git a/memoryscope/models/__init__.py b/memoryscope/core/utils/__init__.py similarity index 100% rename from memoryscope/models/__init__.py rename to memoryscope/core/utils/__init__.py diff --git a/memoryscope/utils/datetime_handler.py b/memoryscope/core/utils/datetime_handler.py similarity index 86% rename from memoryscope/utils/datetime_handler.py rename to memoryscope/core/utils/datetime_handler.py index 6f2d6d2f..21ab3f31 100644 --- a/memoryscope/utils/datetime_handler.py +++ b/memoryscope/core/utils/datetime_handler.py @@ -1,9 +1,10 @@ import datetime import re +from typing import List from memoryscope.constants.language_constants import WEEKDAYS, DATATIME_WORD_LIST, MONTH_DICT -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger +from memoryscope.core.utils.logger import Logger +from memoryscope.enumeration.language_enum import LanguageEnum class DatetimeHandler(object): @@ -40,7 +41,7 @@ class DatetimeHandler(object): self._dt_info_dict: dict | None = None - def _parse_dt_info(self): + def _parse_dt_info(self, language: LanguageEnum): """ Parses the datetime object (_dt) into a dictionary containing detailed date and time components, including language-specific weekday representation. @@ -52,17 +53,16 @@ class DatetimeHandler(object): """ return { "year": self._dt.year, - "month": MONTH_DICT[G_CONTEXT.language][self._dt.month - 1], + "month": MONTH_DICT[language][self._dt.month - 1], "day": self._dt.day, "hour": self._dt.hour, "minute": self._dt.minute, "second": self._dt.second, "week": self._dt.isocalendar().week, - "weekday": WEEKDAYS[G_CONTEXT.language][self._dt.isocalendar().weekday - 1], + "weekday": WEEKDAYS[language][self._dt.isocalendar().weekday - 1], } - @property - def dt_info_dict(self): + def get_dt_info_dict(self, language: LanguageEnum): """ Property method to get the dictionary containing parsed datetime information. If None, initialize using `_parse_dt_info`. @@ -71,7 +71,7 @@ class DatetimeHandler(object): dict: A dictionary with parsed datetime information. """ if self._dt_info_dict is None: - self._dt_info_dict = self._parse_dt_info() + self._dt_info_dict = self._parse_dt_info(language=language) return self._dt_info_dict @classmethod @@ -207,7 +207,7 @@ class DatetimeHandler(object): return date_info @classmethod - def extract_date_parts(cls, input_string: str) -> dict: + def extract_date_parts(cls, input_string: str, language: LanguageEnum) -> dict: """ Extracts various date components from the input string based on the current language context. @@ -217,48 +217,51 @@ class DatetimeHandler(object): Args: input_string (str): The string containing date information to be parsed. + language (str): current language. Returns: dict: A dictionary containing extracted date components, or an empty dictionary if parsing fails. """ - func_name = f"extract_date_parts_{G_CONTEXT.language.value}" + func_name = f"extract_date_parts_{language.value}" if not hasattr(cls, func_name): - cls.logger.warning(f"language={G_CONTEXT.language.value} needs to complete extract_date_parts func!") + cls.logger.warning(f"language={language.value} needs to complete extract_date_parts func!") return {} return getattr(cls, func_name)(input_string=input_string) @classmethod - def has_time_word_cn(cls, query: str) -> bool: + def has_time_word_cn(cls, query: str, datetime_word_list: List[str]) -> bool: """ Check if the input query contains any datetime-related words based on the cn language context. Args: query (str): The input string to check for datetime-related words. + datetime_word_list (list[str]): datetime keywords Returns: bool: True if the query contains at least one datetime-related word, False otherwise. """ contain_datetime = False # TODO use re - for datetime_word in DATATIME_WORD_LIST[G_CONTEXT.language]: + for datetime_word in datetime_word_list: if datetime_word in query: contain_datetime = True break return contain_datetime @classmethod - def has_time_word_en(cls, query: str) -> bool: + def has_time_word_en(cls, query: str, datetime_word_list: List[str]) -> bool: """ Check if the input query contains any datetime-related words based on the en language context. Args: query (str): The input string to check for datetime-related words. + datetime_word_list (list[str]): datetime keywords Returns: bool: True if the query contains at least one datetime-related word, False otherwise. """ contain_datetime = False - for datetime_word in DATATIME_WORD_LIST[G_CONTEXT.language]: + for datetime_word in datetime_word_list: datetime_word = datetime_word.lower() # TODO fix strip if datetime_word in [x.strip().lower().strip(",").strip(".").strip("?").strip(":") @@ -268,12 +271,18 @@ class DatetimeHandler(object): return contain_datetime @classmethod - def has_time_word(cls, query: str) -> bool: - func_name = f"has_time_word_{G_CONTEXT.language.value}" + def has_time_word(cls, query: str, language: LanguageEnum) -> bool: + func_name = f"has_time_word_{language.value}" if not hasattr(cls, func_name): - cls.logger.warning(f"language={G_CONTEXT.language.value} needs to complete has_time_word func!") + cls.logger.warning(f"language={language.value} needs to complete has_time_word function!") return False - return getattr(cls, func_name)(query=query) + + if language not in DATATIME_WORD_LIST: + cls.logger.warning(f"language={language.value} is missing in DATATIME_WORD_LIST!") + return False + + datetime_word_list = DATATIME_WORD_LIST[language] + return getattr(cls, func_name)(query=query, datetime_word_list=datetime_word_list) def datetime_format(self, dt_format: str = "%Y%m%d") -> str: """ @@ -287,17 +296,18 @@ class DatetimeHandler(object): """ return self._dt.strftime(dt_format) - def string_format(self, string_format: str) -> str: + def string_format(self, string_format: str, language: LanguageEnum) -> str: """ Format the datetime information stored in the instance using a custom string format. Args: string_format (str): A format string where placeholders are keys from `dt_info_dict`. + language (str): current language. Returns: str: A formatted datetime string. """ - return string_format.format(**self.dt_info_dict) + return string_format.format(**self.get_dt_info_dict(language=language)) @property def timestamp(self) -> int: diff --git a/memoryscope/utils/logger.py b/memoryscope/core/utils/logger.py similarity index 98% rename from memoryscope/utils/logger.py rename to memoryscope/core/utils/logger.py index 9558b3b0..52d286f8 100644 --- a/memoryscope/utils/logger.py +++ b/memoryscope/core/utils/logger.py @@ -26,7 +26,7 @@ class Logger(logging.Logger): max_bytes: int = 1024 * 1024 * 1024, backup_count: int = 10): """ - Initializes the Logger instance, setting up handlers for console and/or file logging based on provided parameters. + Initializes the Logger instance, setting up handlers for console and file logging based on provided parameters. Args: name (str): Identifier for the logger. @@ -105,7 +105,8 @@ class Logger(logging.Logger): by the handlers are freed properly. """ for handler in self.handlers: - handler.close() # ⭐ Close each handler to release resources + # Close each handler to release resources + handler.close() def clear(self): """ diff --git a/memoryscope/utils/prompt_handler.py b/memoryscope/core/utils/prompt_handler.py similarity index 74% rename from memoryscope/utils/prompt_handler.py rename to memoryscope/core/utils/prompt_handler.py index f2e9b380..a60497e4 100644 --- a/memoryscope/utils/prompt_handler.py +++ b/memoryscope/core/utils/prompt_handler.py @@ -1,19 +1,25 @@ import json import os.path +from pathlib import Path from typing import Dict import yaml -from memoryscope.utils.global_context import G_CONTEXT +from memoryscope.enumeration.language_enum import LanguageEnum class PromptHandler(object): """ The `PromptHandler` class manages prompt messages by loading them from YAML or JSON files and dictionaries, - supporting language selection based on a global context, and providing dictionary-like access to the prompt messages. + supporting language selection based on a context, and providing dictionary-like access to the prompt messages. """ - def __init__(self, class_path: str, prompt_file: str = "", prompt_dict: dict = None, **kwargs): + def __init__(self, + class_path: str, + language: LanguageEnum | str, + prompt_file: str = "", + prompt_dict: dict = None, + **kwargs): """ Initializes the PromptHandler with paths to prompt sources and additional keyword arguments. @@ -21,30 +27,32 @@ class PromptHandler(object): class_path (str): The path to the class where prompts are utilized. prompt_file (str, optional): The path to an external file containing prompts. Defaults to "". prompt_dict (dict, optional): A dictionary directly containing prompt definitions. Defaults to None. + language (LanguageEnum, str): context language. **kwargs: Additional keyword arguments that might be used in prompt handling. """ - self._class_path: str = class_path - self._prompt_dict: Dict[str, str] = {} + class_path: Path = Path(class_path) + self._class_dir: Path = class_path.parent + self._class_name: str = class_path.stem + self._language_enum: LanguageEnum = LanguageEnum(language) self.kwargs = kwargs - file_path = self._class_path.strip(".py") - - self.add_prompt_file(file_path) + self._prompt_dict: Dict[str, str] = {} + self.add_prompt_file((self._class_dir / self._class_name).__str__(), raise_exception=False) if prompt_file: - self.add_prompt_file(prompt_file) - + self.add_prompt_file((self._class_dir / prompt_file).__str__()) if prompt_dict: self.add_prompt_dict(prompt_dict) @staticmethod - def file_path_completion(file_path: str) -> str: + def file_path_completion(file_path: str, raise_exception: bool = True) -> str: """ Attempts to complete the given file path by appending either a `.yaml` or `.json` extension based on the existence of the respective file. If neither exists, an exception is raised. Args: file_path (str): The base path of the file to be completed. + raise_exception (bool): If the file cannot be found, report an error. Returns: str: The completed file path with the appropriate extension. @@ -61,9 +69,10 @@ class PromptHandler(object): if os.path.exists(f"{file_path}.json"): return f"{file_path}.json" - raise RuntimeError(f"{file_path}/yaml/json is not exists!") + if raise_exception: + raise RuntimeError(f"{file_path}/yaml/json is not exists!") - def add_prompt_file(self, file_path: str): + def add_prompt_file(self, file_path: str, raise_exception: bool = True): """ Adds prompt messages from a YAML or JSON file to the internal dictionary. @@ -72,8 +81,11 @@ class PromptHandler(object): Args: file_path (str): The path to the YAML or JSON file containing the prompts. + raise_exception (bool): If the file cannot be found, report an error. """ - file_path = self.file_path_completion(file_path) + file_path = self.file_path_completion(file_path, raise_exception=raise_exception) + if not file_path: + return prompt_dict = {} @@ -102,9 +114,9 @@ class PromptHandler(object): RuntimeError: If a prompt message for the current language is not found. """ for key, language_dict in prompt_dict.items(): - prompts = language_dict.get(G_CONTEXT.language) + prompts = language_dict.get(self._language_enum.value) if not prompts: - raise RuntimeError(f"{key}.prompt.{G_CONTEXT.language} is empty!") + raise RuntimeError(f"{key}.prompt.{self._language_enum.value} is empty!") self._prompt_dict[key] = prompts.strip() @property diff --git a/memoryscope/utils/registry.py b/memoryscope/core/utils/registry.py similarity index 91% rename from memoryscope/utils/registry.py rename to memoryscope/core/utils/registry.py index df396306..935a6c04 100644 --- a/memoryscope/utils/registry.py +++ b/memoryscope/core/utils/registry.py @@ -12,7 +12,8 @@ class Registry(object): Attributes: name (str): The name of the registry. - module_dict (Dict[str, Any]): A dictionary holding registered modules where keys are module names and values are the modules themselves. + module_dict (Dict[str, Any]): A dictionary holding registered modules where keys are module names and values are + the modules themselves. """ def __init__(self, name: str): @@ -31,7 +32,7 @@ class Registry(object): Args: module_name (str): The name of module to be registered. - modules (List[Any] | Dict[str, Any]): The module to be registered. + module (List[Any] | Dict[str, Any]): The module to be registered. Raises: NotImplementedError: If the input is already registered. @@ -46,7 +47,8 @@ class Registry(object): def batch_register(self, modules: List[Any] | Dict[str, Any]): """ - Registers multiple modules in the registry in a single call. Accepts either a list of modules or a dictionary mapping names to modules. + Registers multiple modules in the registry in a single call. Accepts either a list of modules or a dictionary + mapping names to modules. Args: modules (List[Any] | Dict[str, Any]): A list of modules or a dictionary mapping module names to the modules. diff --git a/memoryscope/utils/response_text_parser.py b/memoryscope/core/utils/response_text_parser.py similarity index 67% rename from memoryscope/utils/response_text_parser.py rename to memoryscope/core/utils/response_text_parser.py index b7de8982..d74b3141 100644 --- a/memoryscope/utils/response_text_parser.py +++ b/memoryscope/core/utils/response_text_parser.py @@ -2,28 +2,22 @@ import re from typing import List from memoryscope.constants.language_constants import NONE_WORD -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger +from memoryscope.core.utils.logger import Logger +from memoryscope.enumeration.language_enum import LanguageEnum class ResponseTextParser(object): """ - The `ResponseTextParser` class is designed to parse and process response texts. It provides methods to extract specific + The `ResponseTextParser` class is designed to parse and process response texts. It provides methods to extract patterns from the text and filter out unnecessary information, while also logging the processing steps and outcomes. """ - pattern_v1 = re.compile(r"<(.*?)>") # Regular expression pattern to match content within angle brackets - - def __init__(self, response_text: str, logger_prefix: str = ""): - """ - Initializes the `ResponseTextParser` instance with the provided response text and sets up a logger. - - Args: - response_text (str): The raw response text that needs to be parsed and processed. - """ + PATTERN_V1 = re.compile(r"<(.*?)>") # Regular expression pattern to match content within angle brackets + def __init__(self, response_text: str, language: LanguageEnum, logger_prefix: str = ""): # Strips leading and trailing whitespace from the response text self.response_text: str = response_text.strip() + self.language: LanguageEnum = language # The prefix of log. Defaults to "". self.logger_prefix: str = logger_prefix @@ -43,7 +37,7 @@ class ResponseTextParser(object): line = line.strip() if not line: continue - matches = [match.group(1) for match in self.pattern_v1.finditer(line)] + matches = [match.group(1) for match in self.PATTERN_V1.finditer(line)] if matches: result.append(matches) self.logger.info(f"{self.logger_prefix} response_text={self.response_text} result={result}", stacklevel=2) @@ -51,18 +45,15 @@ class ResponseTextParser(object): def parse_v2(self) -> List[str]: """ - Extract lines which contain NONE_WORD in Chinese or English. + Extract lines which contain NONE_WORD. - Args: - prefix (str): The prefix of log. Defaults to "". - Returns: Contents match the specific patterns. """ result = [] for line in self.response_text.split("\n"): line = line.strip() - if not line or line.lower() == NONE_WORD.get(G_CONTEXT.language): + if not line or line.lower() == NONE_WORD.get(self.language): continue result.append(line) self.logger.info(f"{self.logger_prefix} response_text={self.response_text} result={result}", stacklevel=2) diff --git a/memoryscope/utils/timer.py b/memoryscope/core/utils/timer.py similarity index 95% rename from memoryscope/utils/timer.py rename to memoryscope/core/utils/timer.py index f63ff880..4ba40903 100644 --- a/memoryscope/utils/timer.py +++ b/memoryscope/core/utils/timer.py @@ -1,7 +1,7 @@ import time from typing import Literal -from memoryscope.utils.logger import Logger +from memoryscope.core.utils.logger import Logger TIME_LOG_TYPE = Literal["end", "wrap", "none"] @@ -26,7 +26,7 @@ class Timer(object): Args: name (str): The log name. time_log_type (str): The log type. Defaults to 'End'. - use_ms (bool): Use 'ms' as the time scale or not. Defaults to True. + use_ms (bool): Use 'ms' as the timescale or not. Defaults to True. stack_level (int): The stack level of log. Defaults to 2. float_precision (int): The precision of cost time. Defaults to 4. @@ -75,7 +75,7 @@ class Timer(object): self.logger.info(f"----- {self.name}.begin -----") return self - def __exit__(self, *args, **kwargs): + def __exit__(self, exc_type, exc_value, exc_tb): """ End timing and print the formatted log. """ diff --git a/memoryscope/utils/tool_functions.py b/memoryscope/core/utils/tool_functions.py similarity index 79% rename from memoryscope/utils/tool_functions.py rename to memoryscope/core/utils/tool_functions.py index 2fd432c1..6d5a6834 100644 --- a/memoryscope/utils/tool_functions.py +++ b/memoryscope/core/utils/tool_functions.py @@ -6,6 +6,7 @@ from copy import deepcopy from importlib import import_module from typing import List +import numpy as np import pyfiglet from termcolor import colored @@ -18,7 +19,7 @@ ALL_COLORS = ["red", "green", "yellow", "blue", "magenta", "cyan", "light_grey", def underscore_to_camelcase(name: str, is_first_title: bool = True) -> str: """ - Converts a underscore_notation string to CamelCase. + Converts an underscore_notation string to CamelCase. Args: name (str): The underscore_notation string to be converted. @@ -47,10 +48,7 @@ def camelcase_to_underscore(name: str) -> str: return re.sub(r'(? bool: escaped_keywords = map(re.escape, keywords) pattern = re.compile('|'.join(escaped_keywords), re.IGNORECASE) return pattern.search(text) is not None + + +def cosine_similarity(query: List[float], documents: List[List[float]]): + query = np.array(query) + documents = np.array(documents) + + query_norm = np.linalg.norm(query) + if query_norm == 0: + raise ValueError("Query vector norm is zero, which will result in a division by zero") + + documents_norm = np.linalg.norm(documents, axis=1) + if np.any(documents_norm == 0): + raise ValueError("One of the document vectors has zero norm, which will result in a division by zero") + + dot_product = np.dot(documents, query) + + cosine_similarities = dot_product / (query_norm * documents_norm) + return cosine_similarities.tolist() diff --git a/memoryscope/storage/__init__.py b/memoryscope/core/worker/__init__.py similarity index 100% rename from memoryscope/storage/__init__.py rename to memoryscope/core/worker/__init__.py diff --git a/memoryscope/utils/__init__.py b/memoryscope/core/worker/backend/__init__.py similarity index 100% rename from memoryscope/utils/__init__.py rename to memoryscope/core/worker/backend/__init__.py diff --git a/memoryscope/memory/worker/backend/contra_repeat_worker.py b/memoryscope/core/worker/backend/contra_repeat_worker.py similarity index 92% rename from memoryscope/memory/worker/backend/contra_repeat_worker.py rename to memoryscope/core/worker/backend/contra_repeat_worker.py index 5146887f..177523eb 100644 --- a/memoryscope/memory/worker/backend/contra_repeat_worker.py +++ b/memoryscope/core/worker/backend/contra_repeat_worker.py @@ -2,11 +2,11 @@ from typing import List from memoryscope.constants.common_constants import NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, MERGE_OBS_NODES, TODAY_NODES from memoryscope.constants.language_constants import NONE_WORD, CONTRADICTORY_WORD, CONTAINED_WORD +from memoryscope.core.utils.response_text_parser import ResponseTextParser +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.store_status_enum import StoreStatusEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.response_text_parser import ResponseTextParser -from memoryscope.utils.tool_functions import prompt_to_msg class ContraRepeatWorker(MemoryBaseWorker): @@ -43,13 +43,13 @@ class ContraRepeatWorker(MemoryBaseWorker): 6. Updates the status of nodes accordingly. 7. Persists the changes back to memory storage. """ - all_obs_nodes: List[MemoryNode] = self.memory_handler.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES]) + all_obs_nodes: List[MemoryNode] = self.memory_manager.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES]) if not all_obs_nodes: self.logger.info("all_obs_nodes is empty!") # self.continue_run = False return - today_obs_nodes: List[MemoryNode] = self.memory_handler.get_memories(TODAY_NODES) + today_obs_nodes: List[MemoryNode] = self.memory_manager.get_memories(TODAY_NODES) if today_obs_nodes: all_obs_nodes.extend(today_obs_nodes) @@ -80,7 +80,7 @@ class ContraRepeatWorker(MemoryBaseWorker): response_text = response.message.content # parse text - idx_merge_obs_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1() + idx_merge_obs_list = ResponseTextParser(response_text, self.language, self.__class__.__name__).parse_v1() if len(idx_merge_obs_list) <= 0: self.logger.warning("idx_merge_obs_list is empty!") return @@ -121,4 +121,4 @@ class ContraRepeatWorker(MemoryBaseWorker): merge_obs_nodes.append(node) # save context - self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes, log_repeat=False) + self.memory_manager.set_memories(MERGE_OBS_NODES, merge_obs_nodes, log_repeat=False) diff --git a/memoryscope/memory/worker/backend/contra_repeat_worker.yaml b/memoryscope/core/worker/backend/contra_repeat_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/contra_repeat_worker.yaml rename to memoryscope/core/worker/backend/contra_repeat_worker.yaml diff --git a/memoryscope/memory/worker/backend/get_observation_with_time_worker.py b/memoryscope/core/worker/backend/get_observation_with_time_worker.py similarity index 88% rename from memoryscope/memory/worker/backend/get_observation_with_time_worker.py rename to memoryscope/core/worker/backend/get_observation_with_time_worker.py index 9655fce0..b1fa4c79 100644 --- a/memoryscope/memory/worker/backend/get_observation_with_time_worker.py +++ b/memoryscope/core/worker/backend/get_observation_with_time_worker.py @@ -2,10 +2,10 @@ from typing import List from memoryscope.constants.common_constants import NEW_OBS_WITH_TIME_NODES from memoryscope.constants.language_constants import COLON_WORD -from memoryscope.memory.worker.backend.get_observation_worker import GetObservationWorker +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.backend.get_observation_worker import GetObservationWorker from memoryscope.scheme.message import Message -from memoryscope.utils.datetime_handler import DatetimeHandler -from memoryscope.utils.tool_functions import prompt_to_msg class GetObservationWithTimeWorker(GetObservationWorker): @@ -26,7 +26,7 @@ class GetObservationWithTimeWorker(GetObservationWorker): filter_messages = [] for msg in self.chat_messages: # Checks if the message content has any time reference words - if DatetimeHandler.has_time_word(query=msg.content): + if DatetimeHandler.has_time_word(query=msg.content, language=self.language): filter_messages.append(msg) return filter_messages @@ -49,7 +49,7 @@ class GetObservationWithTimeWorker(GetObservationWorker): for i, msg in enumerate(filter_messages): # Create a DatetimeHandler instance for each message's timestamp and format it dt_handler = DatetimeHandler(dt=msg.time_created) - dt = dt_handler.string_format(self.prompt_handler.time_string_format) + dt = dt_handler.string_format(string_format=self.prompt_handler.time_string_format, language=self.language) # Append formatted timestamp-query pairs to the user_query_list user_query_list.append(f"{i + 1} {dt} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}") diff --git a/memoryscope/memory/worker/backend/get_observation_with_time_worker.yaml b/memoryscope/core/worker/backend/get_observation_with_time_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/get_observation_with_time_worker.yaml rename to memoryscope/core/worker/backend/get_observation_with_time_worker.yaml diff --git a/memoryscope/memory/worker/backend/get_observation_worker.py b/memoryscope/core/worker/backend/get_observation_worker.py similarity index 93% rename from memoryscope/memory/worker/backend/get_observation_worker.py rename to memoryscope/core/worker/backend/get_observation_worker.py index b36d2b6f..78c7eeda 100644 --- a/memoryscope/memory/worker/backend/get_observation_worker.py +++ b/memoryscope/core/worker/backend/get_observation_worker.py @@ -2,14 +2,14 @@ from typing import List from memoryscope.constants.common_constants import NEW_OBS_NODES, TIME_INFER from memoryscope.constants.language_constants import REPEATED_WORD, NONE_WORD, COLON_WORD, TIME_INFER_WORD +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.utils.response_text_parser import ResponseTextParser +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode from memoryscope.scheme.message import Message -from memoryscope.utils.datetime_handler import DatetimeHandler -from memoryscope.utils.response_text_parser import ResponseTextParser -from memoryscope.utils.tool_functions import prompt_to_msg class GetObservationWorker(MemoryBaseWorker): @@ -42,11 +42,11 @@ class GetObservationWorker(MemoryBaseWorker): MemoryTypeEnum.CONVERSATION.value: message.content, TIME_INFER: time_infer, "keywords": keywords, - **{k: str(v) for k, v in dt_handler.dt_info_dict.items()}, + **{k: str(v) for k, v in dt_handler.get_dt_info_dict(self.language).items()}, } if time_infer: - dt_info_dict = DatetimeHandler.extract_date_parts(input_string=time_infer) + dt_info_dict = DatetimeHandler.extract_date_parts(input_string=time_infer, language=self.language) meta_data.update({f"event_{k}": str(v) for k, v in dt_info_dict.items()}) obs_content = (f"{obs_content} ({self.get_language_value(TIME_INFER_WORD)}" f"{self.get_language_value(COLON_WORD)} {time_infer})") @@ -68,7 +68,7 @@ class GetObservationWorker(MemoryBaseWorker): """ filter_messages = [] for msg in self.chat_messages: - if not DatetimeHandler.has_time_word(query=msg.content): + if not DatetimeHandler.has_time_word(query=msg.content, language=self.language): filter_messages.append(msg) self.logger.info(f"after filter_messages.size from {len(self.chat_messages)} to {len(filter_messages)}") @@ -139,7 +139,7 @@ class GetObservationWorker(MemoryBaseWorker): response_text = response.message.content # Parses the generated text to extract observation indices, times, contents, and keywords - idx_obs_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1() + idx_obs_list = ResponseTextParser(response_text, self.language, self.__class__.__name__).parse_v1() if len(idx_obs_list) <= 0: self.logger.warning("idx_obs_list is empty!") return @@ -184,4 +184,4 @@ class GetObservationWorker(MemoryBaseWorker): keywords=keywords)) # Stores the extracted and structured observations in the conversation memory - self.memory_handler.set_memories(self.OBS_STORE_KEY, new_obs_nodes) + self.memory_manager.set_memories(self.OBS_STORE_KEY, new_obs_nodes) diff --git a/memoryscope/memory/worker/backend/get_observation_worker.yaml b/memoryscope/core/worker/backend/get_observation_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/get_observation_worker.yaml rename to memoryscope/core/worker/backend/get_observation_worker.yaml diff --git a/memoryscope/memory/worker/backend/get_reflection_subject_worker.py b/memoryscope/core/worker/backend/get_reflection_subject_worker.py similarity index 87% rename from memoryscope/memory/worker/backend/get_reflection_subject_worker.py rename to memoryscope/core/worker/backend/get_reflection_subject_worker.py index 4c9d7b27..f9323372 100644 --- a/memoryscope/memory/worker/backend/get_reflection_subject_worker.py +++ b/memoryscope/core/worker/backend/get_reflection_subject_worker.py @@ -2,13 +2,13 @@ from typing import List from memoryscope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_NODES from memoryscope.constants.language_constants import COMMA_WORD +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.utils.response_text_parser import ResponseTextParser +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.datetime_handler import DatetimeHandler -from memoryscope.utils.response_text_parser import ResponseTextParser -from memoryscope.utils.tool_functions import prompt_to_msg class GetReflectionSubjectWorker(MemoryBaseWorker): @@ -36,7 +36,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): """ dt_handler = DatetimeHandler() # Prepare metadata with current datetime info - meta_data = {k: str(v) for k, v in dt_handler.dt_info_dict.items()} + meta_data = {k: str(v) for k, v in dt_handler.get_dt_info_dict(self.language).items()} return MemoryNode(user_name=self.user_name, target_name=self.target_name, @@ -58,8 +58,8 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): - Parsing the model's responses for new insight keys. - Creating new insight nodes and updating the memory status accordingly. """ - not_reflected_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_REFLECTED_NODES) - insight_nodes: List[MemoryNode] = self.memory_handler.get_memories(INSIGHT_NODES) + not_reflected_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_REFLECTED_NODES) + insight_nodes: List[MemoryNode] = self.memory_manager.get_memories(INSIGHT_NODES) # Count unaudited nodes not_reflected_count = len(not_reflected_nodes) @@ -101,10 +101,11 @@ class GetReflectionSubjectWorker(MemoryBaseWorker): return # Parse LLM response for new insight keys and update memory - new_insight_keys = ResponseTextParser(response.message.content, self.__class__.__name__).parse_v2() + new_insight_keys = ResponseTextParser(response.message.content, self.language, + self.__class__.__name__).parse_v2() if new_insight_keys: for insight_key in new_insight_keys: - self.memory_handler.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key)) + self.memory_manager.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key)) # Mark unaudited nodes as reflected for node in not_reflected_nodes: diff --git a/memoryscope/memory/worker/backend/get_reflection_subject_worker.yaml b/memoryscope/core/worker/backend/get_reflection_subject_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/get_reflection_subject_worker.yaml rename to memoryscope/core/worker/backend/get_reflection_subject_worker.yaml diff --git a/memoryscope/memory/worker/backend/info_filter_worker.py b/memoryscope/core/worker/backend/info_filter_worker.py similarity index 94% rename from memoryscope/memory/worker/backend/info_filter_worker.py rename to memoryscope/core/worker/backend/info_filter_worker.py index ae28c13d..78307c68 100644 --- a/memoryscope/memory/worker/backend/info_filter_worker.py +++ b/memoryscope/core/worker/backend/info_filter_worker.py @@ -1,11 +1,11 @@ from typing import List from memoryscope.constants.language_constants import COLON_WORD +from memoryscope.core.utils.response_text_parser import ResponseTextParser +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.message import Message -from memoryscope.utils.response_text_parser import ResponseTextParser -from memoryscope.utils.tool_functions import prompt_to_msg class InfoFilterWorker(MemoryBaseWorker): @@ -76,7 +76,7 @@ class InfoFilterWorker(MemoryBaseWorker): response_text = response.message.content # parse text - info_score_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1() + info_score_list = ResponseTextParser(response_text, self.language, self.__class__.__name__).parse_v1() if len(info_score_list) != len(info_messages): self.logger.warning(f"score_size != messages_size, {len(info_score_list)} vs {len(info_messages)}") diff --git a/memoryscope/memory/worker/backend/info_filter_worker.yaml b/memoryscope/core/worker/backend/info_filter_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/info_filter_worker.yaml rename to memoryscope/core/worker/backend/info_filter_worker.yaml diff --git a/memoryscope/memory/worker/backend/load_memory_worker.py b/memoryscope/core/worker/backend/load_memory_worker.py similarity index 91% rename from memoryscope/memory/worker/backend/load_memory_worker.py rename to memoryscope/core/worker/backend/load_memory_worker.py index 1f5e4233..00c24ab3 100644 --- a/memoryscope/memory/worker/backend/load_memory_worker.py +++ b/memoryscope/core/worker/backend/load_memory_worker.py @@ -1,12 +1,12 @@ from typing import List from memoryscope.constants.common_constants import NOT_REFLECTED_NODES, NOT_UPDATED_NODES, INSIGHT_NODES, TODAY_NODES +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.utils.timer import timer +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum from memoryscope.enumeration.store_status_enum import StoreStatusEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.datetime_handler import DatetimeHandler -from memoryscope.utils.timer import timer class LoadMemoryWorker(MemoryBaseWorker): @@ -33,7 +33,7 @@ class LoadMemoryWorker(MemoryBaseWorker): } nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_not_reflected_top_k, filter_dict=filter_dict) - self.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes) + self.memory_manager.set_memories(NOT_REFLECTED_NODES, nodes) @timer def retrieve_not_updated_memory(self): @@ -52,7 +52,7 @@ class LoadMemoryWorker(MemoryBaseWorker): } nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_not_updated_top_k, filter_dict=filter_dict) - self.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + self.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) @timer def retrieve_insight_memory(self): @@ -70,7 +70,7 @@ class LoadMemoryWorker(MemoryBaseWorker): } nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_insight_top_k, filter_dict=filter_dict) - self.memory_handler.set_memories(INSIGHT_NODES, nodes) + self.memory_manager.set_memories(INSIGHT_NODES, nodes) @timer def retrieve_today_memory(self, dt: str): @@ -93,7 +93,7 @@ class LoadMemoryWorker(MemoryBaseWorker): nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_today_top_k, filter_dict=filter_dict) - self.memory_handler.set_memories(TODAY_NODES, nodes) + self.memory_manager.set_memories(TODAY_NODES, nodes) def _run(self): """ diff --git a/memoryscope/memory/worker/backend/long_contra_repeat_worker.py b/memoryscope/core/worker/backend/long_contra_repeat_worker.py similarity index 94% rename from memoryscope/memory/worker/backend/long_contra_repeat_worker.py rename to memoryscope/core/worker/backend/long_contra_repeat_worker.py index 94b306b4..527e1201 100644 --- a/memoryscope/memory/worker/backend/long_contra_repeat_worker.py +++ b/memoryscope/core/worker/backend/long_contra_repeat_worker.py @@ -2,13 +2,13 @@ from typing import List, Dict from memoryscope.constants.common_constants import NOT_UPDATED_NODES, MERGE_OBS_NODES from memoryscope.constants.language_constants import NONE_WORD, CONTAINED_WORD, CONTRADICTORY_WORD +from memoryscope.core.utils.response_text_parser import ResponseTextParser +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum from memoryscope.enumeration.store_status_enum import StoreStatusEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.response_text_parser import ResponseTextParser -from memoryscope.utils.tool_functions import prompt_to_msg class LongContraRepeatWorker(MemoryBaseWorker): @@ -63,7 +63,7 @@ class LongContraRepeatWorker(MemoryBaseWorker): The process helps in maintaining conversation coherence by resolving contradictions and redundancies. """ - not_updated_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_UPDATED_NODES) + not_updated_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_UPDATED_NODES) for node in not_updated_nodes: self.submit_thread_task(fn=self.retrieve_similar_content, node=node) @@ -111,7 +111,8 @@ class LongContraRepeatWorker(MemoryBaseWorker): return # Parses the model's response text to identify updates for memory nodes - idx_obs_info_list = ResponseTextParser(response.message.content, self.__class__.__name__).parse_v1() + idx_obs_info_list = ResponseTextParser(response.message.content, self.language, + self.__class__.__name__).parse_v1() if len(idx_obs_info_list) <= 0: self.logger.warning("idx_obs_info_list is empty!") return @@ -157,4 +158,4 @@ class LongContraRepeatWorker(MemoryBaseWorker): f"action_status={node.action_status}") # save context - self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes) + self.memory_manager.set_memories(MERGE_OBS_NODES, merge_obs_nodes) diff --git a/memoryscope/memory/worker/backend/long_contra_repeat_worker.yaml b/memoryscope/core/worker/backend/long_contra_repeat_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/long_contra_repeat_worker.yaml rename to memoryscope/core/worker/backend/long_contra_repeat_worker.yaml diff --git a/memoryscope/memory/worker/backend/update_insight_worker.py b/memoryscope/core/worker/backend/update_insight_worker.py similarity index 73% rename from memoryscope/memory/worker/backend/update_insight_worker.py rename to memoryscope/core/worker/backend/update_insight_worker.py index b56ed6f7..51f8d7e9 100644 --- a/memoryscope/memory/worker/backend/update_insight_worker.py +++ b/memoryscope/core/worker/backend/update_insight_worker.py @@ -3,12 +3,12 @@ from typing import List from memoryscope.constants.common_constants import INSIGHT_NODES, NOT_UPDATED_NODES, NOT_REFLECTED_NODES from memoryscope.constants.language_constants import COLON_WORD, NONE_WORD, REPEATED_WORD +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.utils.response_text_parser import ResponseTextParser +from memoryscope.core.utils.tool_functions import prompt_to_msg, cosine_similarity +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.datetime_handler import DatetimeHandler -from memoryscope.utils.response_text_parser import ResponseTextParser -from memoryscope.utils.tool_functions import prompt_to_msg class UpdateInsightWorker(MemoryBaseWorker): @@ -27,13 +27,15 @@ class UpdateInsightWorker(MemoryBaseWorker): def filter_obs_nodes(self, insight_node: MemoryNode, - obs_nodes: List[MemoryNode]) -> (MemoryNode, List[MemoryNode], float): + obs_nodes: List[MemoryNode], + use_dummy_ranker: bool) -> (MemoryNode, List[MemoryNode], float): """ Filters observed nodes based on their relevance to a given insight node using a ranking model. Args: insight_node (MemoryNode): The insight node used as the basis for filtering. obs_nodes (List[MemoryNode]): A list of observed nodes to be filtered. + use_dummy_ranker (bool): Global parameters, whether to use rank model or not. Returns: tuple: A tuple containing: @@ -53,24 +55,48 @@ class UpdateInsightWorker(MemoryBaseWorker): self.logger.warning("obs_nodes is empty!") return insight_node, filtered_nodes, max_score - # Call the ranking model to get scores for each observed node's content against the insight key - documents = [x.content for x in obs_nodes] - self.logger.debug(f"update.insight.rank key={insight_node.key} \n docs={'|'.join(documents)}") - response = self.rank_model.call(query=insight_node.key, documents=documents) - if not response.status: - return insight_node, filtered_nodes, max_score + if use_dummy_ranker: + if not insight_node.key_vector: + key_vector: List[float] = self.embedding_model.call(text=insight_node.key).embedding_results + if not key_vector: + self.logger.warning(f"embedding call {insight_node.key} failed!") + return insight_node, filtered_nodes, max_score - # Iterate over the ranked scores to filter nodes - for index, score in response.rank_scores.items(): - node = obs_nodes[index] - # Determine if the node should be kept based on the threshold - keep_flag = score >= self.update_insight_threshold - if keep_flag: - filtered_nodes.append(node) - max_score = max(max_score, score) - # Log information about each node's processing - self.logger.info(f"insight_key={insight_node.key} content={node.content} " - f"score={score} keep_flag={keep_flag}") + insight_node.key_vector = key_vector + + score_recall_list = cosine_similarity(insight_node.key_vector, [x.vector for x in obs_nodes]) + assert len(score_recall_list) == len(obs_nodes), \ + f"size is not as excepted. {len(score_recall_list)} v.s. {len(obs_nodes)}" + + for score, node in zip(score_recall_list, obs_nodes): + keep_flag = score >= self.update_insight_threshold + if keep_flag: + filtered_nodes.append(node) + max_score = max(max_score, score) + + # Log information about each node's processing + self.logger.info(f"insight_key={insight_node.key} content={node.content} " + f"score={score} keep_flag={keep_flag}") + + else: + # Call the ranking model to get scores for each observed node's content against the insight key + documents = [x.content for x in obs_nodes] + self.logger.debug(f"update.insight.rank key={insight_node.key} \n docs={'|'.join(documents)}") + response = self.rank_model.call(query=insight_node.key, documents=documents) + if not response.status: + return insight_node, filtered_nodes, max_score + + # Iterate over the ranked scores to filter nodes + for index, score in response.rank_scores.items(): + node = obs_nodes[index] + # Determine if the node should be kept based on the threshold + keep_flag = score >= self.update_insight_threshold + if keep_flag: + filtered_nodes.append(node) + max_score = max(max_score, score) + # Log information about each node's processing + self.logger.info(f"insight_key={insight_node.key} content={node.content} " + f"score={score} keep_flag={keep_flag}") # Warn if no nodes were filtered if not filtered_nodes: @@ -95,7 +121,7 @@ class UpdateInsightWorker(MemoryBaseWorker): content = f"{key}{self.get_language_value(COLON_WORD)} {insight_value}" insight_node.content = content insight_node.value = insight_value - insight_node.meta_data.update({k: str(v) for k, v in dt_handler.dt_info_dict.items()}) + insight_node.meta_data.update({k: str(v) for k, v in dt_handler.get_dt_info_dict(self.language).items()}) insight_node.timestamp = dt_handler.timestamp insight_node.dt = dt_handler.datetime_format() if insight_node.action_status == ActionStatusEnum.NONE.value: @@ -136,7 +162,7 @@ class UpdateInsightWorker(MemoryBaseWorker): if not response.status or not response.message.content: return insight_node - insight_value_list = ResponseTextParser(response.message.content, + insight_value_list = ResponseTextParser(response.message.content, self.language, f"update_{insight_node.key}").parse_v1() if not insight_value_list: self.logger.warning(f"update_{insight_node.key} insight_value_list is empty!") @@ -175,26 +201,30 @@ class UpdateInsightWorker(MemoryBaseWorker): 6. Gather the results of all update tasks. 7. Mark processed nodes as updated in memory. """ - insight_nodes: List[MemoryNode] = self.memory_handler.get_memories(INSIGHT_NODES) - not_updated_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_UPDATED_NODES) - not_reflected_nodes: List[MemoryNode] = self.memory_handler.get_memories(keys=[NOT_REFLECTED_NODES, + insight_nodes: List[MemoryNode] = self.memory_manager.get_memories(INSIGHT_NODES) + not_updated_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_UPDATED_NODES) + not_reflected_nodes: List[MemoryNode] = self.memory_manager.get_memories(keys=[NOT_REFLECTED_NODES, NOT_UPDATED_NODES]) if not insight_nodes: self.logger.warning("insight_nodes is empty, stopping processing.") return + use_dummy_ranker: bool = self.memoryscope_context.meta_data["use_dummy_ranker"] + # Process active insight nodes with corresponding not updated nodes for node in insight_nodes: time.sleep(1) if node.action_status == ActionStatusEnum.NEW.value: self.submit_thread_task(fn=self.filter_obs_nodes, insight_node=node, - obs_nodes=not_reflected_nodes) + obs_nodes=not_reflected_nodes, + use_dummy_ranker=use_dummy_ranker) else: self.submit_thread_task(fn=self.filter_obs_nodes, insight_node=node, - obs_nodes=not_updated_nodes) + obs_nodes=not_updated_nodes, + use_dummy_ranker=use_dummy_ranker) # select top n result_list = [] @@ -216,12 +246,8 @@ class UpdateInsightWorker(MemoryBaseWorker): # delete empty nodes empty_nodes = [n for n in insight_nodes if not n.content.strip()] - self.memory_handler.delete_memories(empty_nodes) + self.memory_manager.delete_memories(empty_nodes) for node in not_updated_nodes: node.obs_updated = 1 node.action_status = ActionStatusEnum.MODIFIED - - # for node in not_reflected_nodes: - # node.obs_updated = 1 - # node.action_status = ActionStatusEnum.MODIFIED diff --git a/memoryscope/memory/worker/backend/update_insight_worker.yaml b/memoryscope/core/worker/backend/update_insight_worker.yaml similarity index 100% rename from memoryscope/memory/worker/backend/update_insight_worker.yaml rename to memoryscope/core/worker/backend/update_insight_worker.yaml diff --git a/memoryscope/memory/worker/backend/update_memory_worker.py b/memoryscope/core/worker/backend/update_memory_worker.py similarity index 88% rename from memoryscope/memory/worker/backend/update_memory_worker.py rename to memoryscope/core/worker/backend/update_memory_worker.py index 87b182a7..b6803776 100644 --- a/memoryscope/memory/worker/backend/update_memory_worker.py +++ b/memoryscope/core/worker/backend/update_memory_worker.py @@ -1,10 +1,10 @@ from typing import List +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.datetime_handler import DatetimeHandler class UpdateMemoryWorker(MemoryBaseWorker): @@ -46,7 +46,7 @@ class UpdateMemoryWorker(MemoryBaseWorker): if not self.memory_key: return - return self.memory_handler.get_memories(keys=self.memory_key) + return self.memory_manager.get_memories(keys=self.memory_key) def delete_all(self): """ @@ -55,7 +55,7 @@ class UpdateMemoryWorker(MemoryBaseWorker): Returns: List[MemoryNode]: A list of all MemoryNode objects marked for deletion. """ - nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all") + nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all") for node in nodes: node.action_status = ActionStatusEnum.DELETE.value self.logger.info(f"delete_all.size={len(nodes)}") @@ -74,7 +74,7 @@ class UpdateMemoryWorker(MemoryBaseWorker): return i = 0 - nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all") + nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all") for node in nodes: if node.content == query: i += 1 @@ -88,7 +88,7 @@ class UpdateMemoryWorker(MemoryBaseWorker): return i = 0 - nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all") + nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all") for node in nodes: if node.memory_id == memory_id: i += 1 @@ -109,4 +109,4 @@ class UpdateMemoryWorker(MemoryBaseWorker): if not hasattr(self, method): self.logger.info(f"method={method} is missing!") return - self.memory_handler.update_memories(nodes=getattr(self, method)()) + self.memory_manager.update_memories(nodes=getattr(self, method)()) diff --git a/memoryscope/memory/worker/base_worker.py b/memoryscope/core/worker/base_worker.py similarity index 98% rename from memoryscope/memory/worker/base_worker.py rename to memoryscope/core/worker/base_worker.py index e5bf691c..fc92ce71 100644 --- a/memoryscope/memory/worker/base_worker.py +++ b/memoryscope/core/worker/base_worker.py @@ -3,8 +3,8 @@ from abc import ABCMeta, abstractmethod from concurrent.futures import ThreadPoolExecutor, as_completed from typing import Any, Dict -from memoryscope.utils.logger import Logger -from memoryscope.utils.timer import Timer +from memoryscope.core.utils.logger import Logger +from memoryscope.core.utils.timer import Timer class BaseWorker(metaclass=ABCMeta): diff --git a/memoryscope/memory/worker/dummy_worker.py b/memoryscope/core/worker/dummy_worker.py similarity index 92% rename from memoryscope/memory/worker/dummy_worker.py rename to memoryscope/core/worker/dummy_worker.py index b8eecb7b..bd2c1760 100644 --- a/memoryscope/memory/worker/dummy_worker.py +++ b/memoryscope/core/worker/dummy_worker.py @@ -1,7 +1,7 @@ import datetime from memoryscope.constants.common_constants import RESULT, WORKFLOW_NAME, CHAT_KWARGS -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker class DummyWorker(MemoryBaseWorker): diff --git a/memoryscope/core/worker/frontend/__init__.py b/memoryscope/core/worker/frontend/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/memoryscope/memory/worker/frontend/extract_time_worker.py b/memoryscope/core/worker/frontend/extract_time_worker.py similarity index 87% rename from memoryscope/memory/worker/frontend/extract_time_worker.py rename to memoryscope/core/worker/frontend/extract_time_worker.py index 087431c5..b2d864e2 100644 --- a/memoryscope/memory/worker/frontend/extract_time_worker.py +++ b/memoryscope/core/worker/frontend/extract_time_worker.py @@ -3,9 +3,9 @@ from typing import Dict from memoryscope.constants.common_constants import QUERY_WITH_TS, EXTRACT_TIME_DICT from memoryscope.constants.language_constants import DATATIME_KEY_MAP -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker -from memoryscope.utils.datetime_handler import DatetimeHandler -from memoryscope.utils.tool_functions import prompt_to_msg +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.utils.tool_functions import prompt_to_msg +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker class ExtractTimeWorker(MemoryBaseWorker): @@ -33,13 +33,14 @@ class ExtractTimeWorker(MemoryBaseWorker): query, query_timestamp = self.get_context(QUERY_WITH_TS) # Identify if the query contains datetime keywords - contain_datetime = DatetimeHandler.has_time_word(query) + contain_datetime = DatetimeHandler.has_time_word(query, self.language) if not contain_datetime: self.logger.info(f"contain_datetime={contain_datetime}") return # Prepare the prompt with necessary contextual details - query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_string_format) + query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_string_format, + self.language) system_prompt = self.prompt_handler.extract_time_system few_shot = self.prompt_handler.extract_time_few_shot user_query = self.prompt_handler.extract_time_user_query.format(query=query, query_time_str=query_time_str) diff --git a/memoryscope/memory/worker/frontend/extract_time_worker.yaml b/memoryscope/core/worker/frontend/extract_time_worker.yaml similarity index 100% rename from memoryscope/memory/worker/frontend/extract_time_worker.yaml rename to memoryscope/core/worker/frontend/extract_time_worker.yaml diff --git a/memoryscope/memory/worker/frontend/fuse_rerank_worker.py b/memoryscope/core/worker/frontend/fuse_rerank_worker.py similarity index 95% rename from memoryscope/memory/worker/frontend/fuse_rerank_worker.py rename to memoryscope/core/worker/frontend/fuse_rerank_worker.py index 373d82b0..cd749211 100644 --- a/memoryscope/memory/worker/frontend/fuse_rerank_worker.py +++ b/memoryscope/core/worker/frontend/fuse_rerank_worker.py @@ -1,9 +1,9 @@ from typing import Dict, List from memoryscope.constants.common_constants import EXTRACT_TIME_DICT, RANKED_MEMORY_NODES, RESULT -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.datetime_handler import DatetimeHandler class FuseRerankWorker(MemoryBaseWorker): @@ -62,7 +62,7 @@ class FuseRerankWorker(MemoryBaseWorker): """ # Parse input parameters from the worker's context extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT) - memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RANKED_MEMORY_NODES) + memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RANKED_MEMORY_NODES) # Check if memory nodes are available; warn and return if not if not memory_node_list: diff --git a/memoryscope/memory/worker/frontend/print_memory_worker.py b/memoryscope/core/worker/frontend/print_memory_worker.py similarity index 92% rename from memoryscope/memory/worker/frontend/print_memory_worker.py rename to memoryscope/core/worker/frontend/print_memory_worker.py index 6617c49c..5154371f 100644 --- a/memoryscope/memory/worker/frontend/print_memory_worker.py +++ b/memoryscope/core/worker/frontend/print_memory_worker.py @@ -1,11 +1,11 @@ from typing import List from memoryscope.constants.common_constants import RETRIEVE_MEMORY_NODES, RESULT +from memoryscope.core.utils.datetime_handler import DatetimeHandler +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum from memoryscope.enumeration.store_status_enum import StoreStatusEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.datetime_handler import DatetimeHandler class PrintMemoryWorker(MemoryBaseWorker): @@ -22,7 +22,7 @@ class PrintMemoryWorker(MemoryBaseWorker): 3. Set the formatted string back into the worker's context """ # get long-term memory - memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RETRIEVE_MEMORY_NODES) + memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RETRIEVE_MEMORY_NODES) memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True) observation_memory_list: List[str] = [] diff --git a/memoryscope/memory/worker/frontend/print_memory_worker.yaml b/memoryscope/core/worker/frontend/print_memory_worker.yaml similarity index 100% rename from memoryscope/memory/worker/frontend/print_memory_worker.yaml rename to memoryscope/core/worker/frontend/print_memory_worker.yaml diff --git a/memoryscope/memory/worker/frontend/read_message_worker.py b/memoryscope/core/worker/frontend/read_message_worker.py similarity index 91% rename from memoryscope/memory/worker/frontend/read_message_worker.py rename to memoryscope/core/worker/frontend/read_message_worker.py index dfbe4a12..32f3380b 100644 --- a/memoryscope/memory/worker/frontend/read_message_worker.py +++ b/memoryscope/core/worker/frontend/read_message_worker.py @@ -1,6 +1,6 @@ from memoryscope.constants.common_constants import RESULT +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker class ReadMessageWorker(MemoryBaseWorker): diff --git a/memoryscope/memory/worker/frontend/retrieve_memory_worker.py b/memoryscope/core/worker/frontend/retrieve_memory_worker.py similarity index 95% rename from memoryscope/memory/worker/frontend/retrieve_memory_worker.py rename to memoryscope/core/worker/frontend/retrieve_memory_worker.py index fd928f7d..6b5ddea4 100644 --- a/memoryscope/memory/worker/frontend/retrieve_memory_worker.py +++ b/memoryscope/core/worker/frontend/retrieve_memory_worker.py @@ -1,12 +1,12 @@ from typing import List from memoryscope.constants.common_constants import QUERY_WITH_TS, RETRIEVE_MEMORY_NODES +from memoryscope.core.utils.timer import timer +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum from memoryscope.enumeration.store_status_enum import StoreStatusEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.utils.timer import timer class RetrieveMemoryWorker(MemoryBaseWorker): @@ -120,6 +120,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker): 7. Stores the processed memory nodes for further use. """ query, _ = self.get_context(QUERY_WITH_TS) + self.logger.info(f"retrieve memory with query={query}.") self.submit_thread_task(self.retrieve_from_observation, query=query) self.submit_thread_task(self.retrieve_from_insight, query=query) self.submit_thread_task(self.retrieve_expired_memory, query=query) @@ -136,7 +137,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker): memory_node_list = sorted(memory_node_list, key=lambda x: x.score_recall, reverse=True) for node in memory_node_list: node.action_status = ActionStatusEnum.NONE.value - self.logger.info(f"recall_stage: content={node.content} score={node.score_rerank} type={node.memory_type} " + self.logger.info(f"recall_stage: content={node.content} score={node.score_recall} type={node.memory_type} " f"store_status={node.store_status} action_status={node.action_status}") - self.memory_handler.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list) + self.memory_manager.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list) diff --git a/memoryscope/memory/worker/frontend/semantic_rank_worker.py b/memoryscope/core/worker/frontend/semantic_rank_worker.py similarity index 57% rename from memoryscope/memory/worker/frontend/semantic_rank_worker.py rename to memoryscope/core/worker/frontend/semantic_rank_worker.py index be5d67dc..e1ca0a04 100644 --- a/memoryscope/memory/worker/frontend/semantic_rank_worker.py +++ b/memoryscope/core/worker/frontend/semantic_rank_worker.py @@ -1,7 +1,7 @@ from typing import List, Dict from memoryscope.constants.common_constants import RETRIEVE_MEMORY_NODES, QUERY_WITH_TS, RANKED_MEMORY_NODES -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode @@ -29,25 +29,33 @@ class SemanticRankWorker(MemoryBaseWorker): """ # query query, _ = self.get_context(QUERY_WITH_TS) - memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RETRIEVE_MEMORY_NODES) + memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RETRIEVE_MEMORY_NODES) if not memory_node_list: self.logger.warning("Retrieve memory nodes is empty!") return - # drop repeated - memory_node_dict: Dict[str, MemoryNode] = {n.content.strip(): n for n in memory_node_list if n.content.strip()} - memory_node_list = list(memory_node_dict.values()) + use_dummy_ranker: bool = self.memoryscope_context.meta_data["use_dummy_ranker"] + if use_dummy_ranker: + for node in memory_node_list: + node.score_rank = node.score_recall + self.logger.warning("use score_recall instead of score_rank!") - response = self.rank_model.call(query=query, documents=[n.content for n in memory_node_list]) - if not response.status or not response.rank_scores: - return + else: + # drop repeated + memory_node_dict: Dict[str, MemoryNode] = {n.content.strip(): n for n in memory_node_list if + n.content.strip()} + memory_node_list = list(memory_node_dict.values()) - # set score - for idx, score in response.rank_scores.items(): - if idx >= len(memory_node_list): - self.logger.warning(f"Idx={idx} exceeds the maximum length of rank_scores!") - continue - memory_node_list[idx].score_rank = score + response = self.rank_model.call(query=query, documents=[n.content for n in memory_node_list]) + if not response.status or not response.rank_scores: + return + + # set score + for idx, score in response.rank_scores.items(): + if idx >= len(memory_node_list): + self.logger.warning(f"Idx={idx} exceeds the maximum length of rank_scores!") + continue + memory_node_list[idx].score_rank = score # sort by score memory_node_list = sorted(memory_node_list, key=lambda n: n.score_rank, reverse=True) @@ -58,4 +66,4 @@ class SemanticRankWorker(MemoryBaseWorker): self.logger.info(f"Rank stage: Content={node.content}, Score={node.score_rank}") # save ranked nodes back to memory - self.memory_handler.set_memories(RANKED_MEMORY_NODES, memory_node_list, log_repeat=False) + self.memory_manager.set_memories(RANKED_MEMORY_NODES, memory_node_list, log_repeat=False) diff --git a/memoryscope/memory/worker/frontend/set_query_worker.py b/memoryscope/core/worker/frontend/set_query_worker.py similarity index 58% rename from memoryscope/memory/worker/frontend/set_query_worker.py rename to memoryscope/core/worker/frontend/set_query_worker.py index 552baf5e..85961bc9 100644 --- a/memoryscope/memory/worker/frontend/set_query_worker.py +++ b/memoryscope/core/worker/frontend/set_query_worker.py @@ -1,8 +1,8 @@ import datetime from memoryscope.constants.common_constants import QUERY_WITH_TS +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker class SetQueryWorker(MemoryBaseWorker): @@ -22,22 +22,39 @@ class SetQueryWorker(MemoryBaseWorker): along with its creation timestamp. """ query = "" # Default query value - query_timestamp = int(datetime.datetime.now().timestamp()) # Current timestamp as default + timestamp = int(datetime.datetime.now().timestamp()) # Current timestamp as default if "query" in self.chat_kwargs: - # Check if a specific 'query' has been provided via chat kwargs + # set query if exists query = self.chat_kwargs["query"] if not query: query = "" query = query.strip() + # set ts if exists + _timestamp = self.chat_kwargs.get("timestamp") + if _timestamp and isinstance(_timestamp, int): + timestamp = _timestamp + + # check role_name + role_name = self.chat_kwargs.get("role_name") + if role_name: + assert role_name == self.target_name, \ + f"role_name={role_name} is not supported in human/assistant memory workflow!" + elif self.chat_messages: # If no explicit query is given, use the content of the latest chat message chat_messages = [msg for msg in self.chat_messages if msg.role == MessageRoleEnum.USER.value] if chat_messages: message = chat_messages[-1] query = message.content - query_timestamp = message.time_created + timestamp = message.time_created + + # check role_name + role_name = message.role_name + if role_name: + assert role_name == self.target_name, \ + f"role_name={role_name} is not supported in human/assistant memory workflow!" # Store the determined query and its timestamp in the context - self.set_context(QUERY_WITH_TS, (query, query_timestamp)) + self.set_context(QUERY_WITH_TS, (query, timestamp)) diff --git a/memoryscope/memory/worker/memory_base_worker.py b/memoryscope/core/worker/memory_base_worker.py similarity index 76% rename from memoryscope/memory/worker/memory_base_worker.py rename to memoryscope/core/worker/memory_base_worker.py index 5d0ee99a..85b1ad41 100644 --- a/memoryscope/memory/worker/memory_base_worker.py +++ b/memoryscope/core/worker/memory_base_worker.py @@ -1,15 +1,17 @@ from abc import ABCMeta from typing import List, Dict, Any -from memoryscope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, MEMORY_HANDLER -from memoryscope.memory.worker.base_worker import BaseWorker -from memoryscope.models.base_model import BaseModel +from memoryscope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, MEMORYSCOPE_CONTEXT, \ + WORKFLOW_NAME, MEMORY_MANAGER +from memoryscope.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.models.base_model import BaseModel +from memoryscope.core.storage.base_memory_store import BaseMemoryStore +from memoryscope.core.storage.base_monitor import BaseMonitor +from memoryscope.core.utils.prompt_handler import PromptHandler +from memoryscope.core.worker.base_worker import BaseWorker +from memoryscope.core.worker.memory_manager import MemoryManager +from memoryscope.enumeration.language_enum import LanguageEnum from memoryscope.scheme.message import Message -from memoryscope.storage.base_memory_store import BaseMemoryStore -from memoryscope.storage.base_monitor import BaseMonitor -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.memory_handler import MemoryHandler -from memoryscope.utils.prompt_handler import PromptHandler class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): @@ -77,6 +79,18 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): """ return self.get_context(CHAT_KWARGS) + @property + def workflow_name(self) -> str: + return self.get_context(WORKFLOW_NAME) + + @property + def memoryscope_context(self) -> MemoryscopeContext: + return self.get_context(MEMORYSCOPE_CONTEXT) + + @property + def language(self) -> LanguageEnum: + return self.memoryscope_context.language + @property def embedding_model(self) -> BaseModel: """ @@ -87,8 +101,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseModel: The embedding model used for converting text into vector representations. """ if isinstance(self._embedding_model, str): - self._embedding_model = G_CONTEXT.model_dict[self._embedding_model] - # ⭐ Retrieve the actual model instance when the attribute is a string reference + self._embedding_model = self.memoryscope_context.model_dict[self._embedding_model] return self._embedding_model @property @@ -101,8 +114,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseModel: The model used for text generation. """ if isinstance(self._generation_model, str): - self._generation_model = G_CONTEXT.model_dict[self._generation_model] - # ⭐ Retrieve the model instance if currently a string reference + self._generation_model = self.memoryscope_context.model_dict[self._generation_model] return self._generation_model @property @@ -115,7 +127,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseModel: The rank model instance used for ranking tasks. """ if isinstance(self._rank_model, str): - self._rank_model = G_CONTEXT.model_dict[self._rank_model] # Fetch model instance if string reference + self._rank_model = self.memoryscope_context.model_dict[self._rank_model] return self._rank_model @property @@ -128,7 +140,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseMemoryStore: The memory store instance used for inserting, updating, retrieving and deleting operations. """ if self._memory_store is None: - self._memory_store = G_CONTEXT.memory_store + self._memory_store = self.memoryscope_context.memory_store return self._memory_store @property @@ -141,7 +153,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): BaseMonitor: The monitoring component instance. """ if self._monitor is None: - self._monitor = G_CONTEXT.monitor + self._monitor = self.memoryscope_context.monitor return self._monitor @property @@ -154,7 +166,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): str: The name of the assistant. """ if self._user_name is None: - self._user_name = G_CONTEXT.meta_data["assistant_name"] + self._user_name = self.memoryscope_context.meta_data["assistant_name"] return self._user_name @property @@ -166,7 +178,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): str: The readable name of the human. """ if self._target_name is None: - self._target_name = G_CONTEXT.meta_data["human_name"] + self._target_name = self.memoryscope_context.meta_data["human_name"] return self._target_name @property @@ -178,23 +190,22 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): PromptHandler: An instance of PromptHandler initialized with specific file path and keyword arguments. """ if self._prompt_handler is None: - self._prompt_handler = PromptHandler(self.FILE_PATH, **self.kwargs) + self._prompt_handler = PromptHandler(self.FILE_PATH, language=self.language, **self.kwargs) return self._prompt_handler @property - def memory_handler(self) -> MemoryHandler: + def memory_manager(self) -> MemoryManager: """ Lazily initializes and returns the MemoryHandler instance. Returns: MemoryHandler: An instance of MemoryHandler. """ - if not self.has_content(MEMORY_HANDLER): - self.set_context(MEMORY_HANDLER, MemoryHandler()) # Initialize the memory handler if not present - return self.get_context(MEMORY_HANDLER) + if not self.has_content(MEMORY_MANAGER): + self.set_context(MEMORY_MANAGER, MemoryManager(self.memoryscope_context)) + return self.get_context(MEMORY_MANAGER) - @staticmethod - def get_language_value(languages: dict | List[dict]) -> Any | List[Any]: + def get_language_value(self, languages: dict | List[dict]) -> Any | List[Any]: """ Retrieves the value(s) corresponding to the current language context. @@ -205,5 +216,5 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): Any | list[Any]: The value or list of values matching the current language setting. """ if isinstance(languages, list): - return [x[G_CONTEXT.language] for x in languages] - return languages[G_CONTEXT.language] + return [x[self.language] for x in languages] + return languages[self.language] diff --git a/memoryscope/utils/memory_handler.py b/memoryscope/core/worker/memory_manager.py similarity index 94% rename from memoryscope/utils/memory_handler.py rename to memoryscope/core/worker/memory_manager.py index 36ec9b9d..af7b3016 100644 --- a/memoryscope/utils/memory_handler.py +++ b/memoryscope/core/worker/memory_manager.py @@ -1,22 +1,21 @@ from typing import Dict, List +from memoryscope.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.storage.base_memory_store import BaseMemoryStore +from memoryscope.core.utils.logger import Logger from memoryscope.enumeration.action_status_enum import ActionStatusEnum from memoryscope.enumeration.store_status_enum import StoreStatusEnum from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.storage.base_memory_store import BaseMemoryStore -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger -class MemoryHandler(object): +class MemoryManager(object): """ The `MemoryHandler` class manages memory nodes with memory store. """ - def __init__(self): - """ - Initializes the MemoryHandler. - """ + def __init__(self, memoryscope_context: MemoryscopeContext): + self.memoryscope_context: MemoryscopeContext = memoryscope_context + self._memory_store: BaseMemoryStore | None = None # dict: memory_id -> MemoryNode @@ -36,7 +35,7 @@ class MemoryHandler(object): BaseMemoryStore: The memory store instance associated with this worker. """ if self._memory_store is None: - self._memory_store = G_CONTEXT.memory_store + self._memory_store = self.memoryscope_context.memory_store return self._memory_store def clear(self): diff --git a/memoryscope/enumeration/model_enum.py b/memoryscope/enumeration/model_enum.py index 4f7cfb44..8cc76d2f 100644 --- a/memoryscope/enumeration/model_enum.py +++ b/memoryscope/enumeration/model_enum.py @@ -7,8 +7,9 @@ class ModelEnum(str, Enum): Members: GENERATION_MODEL: Represents a model responsible for generating content. - EMBEDDING_MODEL: Represents a model tasked with creating embeddings, typically used for transforming data into a numerical form suitable for machine learning tasks. - RANK_MODEL: Denotes a model that specializes in ranking, often used to order items based on relevance or importance. + EMBEDDING_MODEL: Represents a model tasked with creating embeddings, typically used for transforming data into a + numerical form suitable for machine learning tasks. + RANK_MODEL: Denotes a model that specializes in ranking, often used to order items based on relevance. """ GENERATION_MODEL = "generation_model" diff --git a/memoryscope/memory/service/base_memory_service.py b/memoryscope/memory/service/base_memory_service.py deleted file mode 100644 index 31b2df5e..00000000 --- a/memoryscope/memory/service/base_memory_service.py +++ /dev/null @@ -1,106 +0,0 @@ -import threading -from abc import ABCMeta, abstractmethod -from typing import List, Dict - -from memoryscope.memory.operation.base_operation import BaseOperation -from memoryscope.scheme.message import Message -from memoryscope.utils.logger import Logger - - -class BaseMemoryService(metaclass=ABCMeta): - """ - An abstract base class for managing memory operations within a multi-threaded context. - It sets up the infrastructure for operation handling, message storage, and synchronization, - along with logging capabilities and customizable configurations. - """ - - def __init__(self, - memory_operations: Dict[str, dict], - retrieve_memory_key: str = "retrieve_memory", - read_message_key: str = "read_message", - **kwargs): - """ - Initializes the BaseMemoryService with operation definitions, keys for memory access, - and additional keyword arguments for flexibility. - - Args: - memory_operations (Dict[str, dict]): A dictionary defining available memory operations. - retrieve_memory_key (str): The key indicating a retrieve memory operation. Defaults to "retrieve_memory". - read_message_key (str): The key for reading messages. Defaults to "read_message". - **kwargs: Additional parameters to customize service behavior. - """ - self.memory_operations: Dict[str, dict] = memory_operations - self.retrieve_memory_key: str = retrieve_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.logger = Logger.get_logger() - self.kwargs = kwargs - - @abstractmethod - def add_messages(self, messages: List[Message] | Message): - raise NotImplementedError - - @abstractmethod - def do_operation(self, op_name: str, **kwargs): - """ - Abstract method defining the interface for executing a specific operation by its name. - This method must be implemented by subclasses to provide the actual operation logic. - - Args: - op_name (str): The name identifying the operation to be performed. - **kwargs: Additional keyword arguments required for the operation execution. - - Raises: - NotImplementedError: This exception is raised when the method is not overridden in a subclass. - """ - raise NotImplementedError - - @property - def op_description_dict(self) -> Dict[str, str]: - """ - Property to retrieve a dictionary mapping operation keys to their descriptions. - Lazily initializes the dictionary on first access. - - Returns: - Dict[str, str]: A dictionary where keys are operation identifiers and values are their descriptions. - """ - if not self._op_description_dict: - self._op_description_dict = {k: v.description for k, v in self._operation_dict.items()} - return self._op_description_dict - - def retrieve_memory(self): - """ - Executes the operation associated with retrieved memory. - Asserts that the operation for retrieved memory has been initialized. - - Returns: - Any: The result of the retrieved memory operation. - """ - assert self.retrieve_memory_key in self._operation_dict, f"op={self.retrieve_memory_key} is not inited!" - return self.do_operation(self.retrieve_memory_key) - - def read_message(self): - """ - Executes the operation associated with reading messages. - Asserts that the operation for reading messages has been initialized. - - Returns: - Any: The result of the read message operation. - """ - 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) - - @abstractmethod - def init_service(self, **kwargs): - raise NotImplementedError - - def start_backend_service(self): - pass - - def stop_backend_service(self): - pass diff --git a/memoryscope/scheme/memory_node.py b/memoryscope/scheme/memory_node.py index de884ae7..32205a54 100644 --- a/memoryscope/scheme/memory_node.py +++ b/memoryscope/scheme/memory_node.py @@ -23,6 +23,8 @@ class MemoryNode(BaseModel): key: str = Field("", description="memory key") + key_vector: List[float] = Field([], description="memory key embedding result") + value: str = Field("", description="memory value") score_recall: float = Field(0, description="embedding similarity score used in recall stage") @@ -37,7 +39,7 @@ class MemoryNode(BaseModel): store_status: str = Field("valid", description="store_status: valid / expired") - vector: List[float] = Field([], description="content embedding result, return empty") + vector: List[float] = Field([], description="content embedding result") timestamp: int = Field(default_factory=lambda: int(datetime.datetime.now().timestamp()), description="timestamp of the memory node") diff --git a/memoryscope/utils/global_context.py b/memoryscope/utils/global_context.py deleted file mode 100644 index 81ac9e07..00000000 --- a/memoryscope/utils/global_context.py +++ /dev/null @@ -1,33 +0,0 @@ -from concurrent.futures import ThreadPoolExecutor -from typing import Dict, Any - -from memoryscope.chat.base_memory_chat import BaseMemoryChat -from memoryscope.enumeration.language_enum import LanguageEnum -from memoryscope.memory.service.base_memory_service import BaseMemoryService -from memoryscope.models.base_model import BaseModel -from memoryscope.storage.base_memory_store import BaseMemoryStore -from memoryscope.storage.base_monitor import BaseMonitor - - -class GlobalContext(object): - """ - The GlobalContext class archives all configs utilized by store, monitor, services and workers. - """ - - def __init__(self): - self.global_config: Dict[str, Any] = {} - self.worker_config: Dict[str, Dict[str, Any]] = {} - - self.memory_service_dict: Dict[str, BaseMemoryService] = {} - self.model_dict: Dict[str, BaseModel] = {} - self.memory_chat_dict: Dict[str, BaseMemoryChat] = {} - - self.memory_store: BaseMemoryStore | None = None - self.monitor: BaseMonitor | None = None - self.thread_pool: ThreadPoolExecutor | None = None - self.language: LanguageEnum = LanguageEnum.EN - - self.meta_data: Dict[str, Any] = {} - - -G_CONTEXT = GlobalContext() diff --git a/tests/models/test_models_lli_embedding.py b/tests/models/test_models_lli_embedding.py index f21d0bc7..67d3bb13 100644 --- a/tests/models/test_models_lli_embedding.py +++ b/tests/models/test_models_lli_embedding.py @@ -5,8 +5,8 @@ sys.path.append(".") # noqa: E402 import asyncio import unittest -from memoryscope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel -from memoryscope.utils.logger import Logger +from memoryscope.core.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.utils.logger import Logger class TestLLIEmbedding(unittest.TestCase): diff --git a/tests/models/test_models_lli_generation.py b/tests/models/test_models_lli_generation.py index 474851b4..1fbe41b5 100644 --- a/tests/models/test_models_lli_generation.py +++ b/tests/models/test_models_lli_generation.py @@ -6,8 +6,8 @@ import unittest import time import asyncio from memoryscope.scheme.message import Message -from memoryscope.models.llama_index_generation_model import LlamaIndexGenerationModel -from memoryscope.utils.logger import Logger +from memoryscope.core.models.llama_index_generation_model import LlamaIndexGenerationModel +from memoryscope.core.utils.logger import Logger class TestLLILLM(unittest.TestCase): diff --git a/tests/models/test_models_lli_rank.py b/tests/models/test_models_lli_rank.py index 25c59fde..e0e90c7c 100644 --- a/tests/models/test_models_lli_rank.py +++ b/tests/models/test_models_lli_rank.py @@ -1,7 +1,7 @@ import asyncio import unittest -from memoryscope.models.llama_index_rank_model import LlamaIndexRankModel +from memoryscope.core.models.llama_index_rank_model import LlamaIndexRankModel class TestLLIReRank(unittest.TestCase): diff --git a/tests/operations/test_interface.py b/tests/operations/test_interface.py deleted file mode 100644 index d89efb93..00000000 --- a/tests/operations/test_interface.py +++ /dev/null @@ -1,18 +0,0 @@ -from memoryscope.cli import MemoryScope -from memoryscope.scheme.message import Message - -ms = MemoryScope().load_config("config/demo_config_no_stream.yaml") -memory_service = ms.default_service -memory_chat = ms.default_chat_handle - -# new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name="我", content="我的爱好是弹琴并且喜欢看电影。") -# memory_service.add_messages(new_message) - -res: Message = memory_chat.chat_with_memory(query="我的爱好是弹琴。", remember_response=True) -print(res.message.content) - -res: Message = memory_chat.chat_with_memory(query="昨天弹出一个光粒,消灭了星系0x4be。", remember_response=True) -print(res.message.content) - -res: Message = memory_chat.chat_with_memory(query="今天弹出一个二向箔,消灭了星系0xa2e。", remember_response=True) -print(res.message.content) diff --git a/tests/operations/init_test.py b/tests/other/init_test.py similarity index 100% rename from tests/operations/init_test.py rename to tests/other/init_test.py diff --git a/tests/other/read_prompt.yaml b/tests/other/read_prompt.yaml new file mode 100644 index 00000000..efb0e419 --- /dev/null +++ b/tests/other/read_prompt.yaml @@ -0,0 +1,3 @@ +a: + cn: c + en: e \ No newline at end of file diff --git a/tests/other/read_yaml.py b/tests/other/read_yaml.py new file mode 100644 index 00000000..ac0513f4 --- /dev/null +++ b/tests/other/read_yaml.py @@ -0,0 +1,11 @@ +import sys + +sys.path.append(".") # noqa: E402 + +from memoryscope.core.utils.prompt_handler import PromptHandler + +if __name__ == "__main__": + file_path: str = __file__ + print(file_path) + handler = PromptHandler(__file__, language="cn", prompt_file="read_prompt", ) + print(handler.prompt_dict) diff --git a/tests/other/test_attr.py b/tests/other/test_attr.py new file mode 100644 index 00000000..eaf2fbf8 --- /dev/null +++ b/tests/other/test_attr.py @@ -0,0 +1,15 @@ +class MyClass: + def __init__(self): + self.existing_attribute = "I exist" + + def do(self, name: str, **kwargs): + print("do %s %s" % (name, kwargs)) + + def __getattr__(self, name): + return lambda **kwargs: self.do(name, **kwargs) + + +# 创建类的实例 +obj = MyClass() + +obj.haha(a=1, b=2) diff --git a/tests/other/test_cli.py b/tests/other/test_cli.py new file mode 100644 index 00000000..89a99f17 --- /dev/null +++ b/tests/other/test_cli.py @@ -0,0 +1,14 @@ +import fire + + +class CLI: + def run(self, **kwargs): + """ + 打印传入的 kwargs + """ + for key, value in kwargs.items(): + print(f"{key}: {value}") + + +if __name__ == '__main__': + fire.Fire(CLI().run) diff --git a/tests/storages/test_storages_lli_es.py b/tests/storages/test_storages_lli_es.py index 791286aa..a72d1aab 100644 --- a/tests/storages/test_storages_lli_es.py +++ b/tests/storages/test_storages_lli_es.py @@ -1,8 +1,8 @@ import unittest -from memoryscope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore class TestLlamaIndexElasticSearchStore(unittest.TestCase): diff --git a/tests/storages/test_storages_lli_synces.py b/tests/storages/test_storages_lli_synces.py index 0bb8b554..f407a389 100644 --- a/tests/storages/test_storages_lli_synces.py +++ b/tests/storages/test_storages_lli_synces.py @@ -1,8 +1,8 @@ import unittest -from memoryscope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.models.llama_index_embedding_model import LlamaIndexEmbeddingModel +from memoryscope.core.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore from memoryscope.scheme.memory_node import MemoryNode -from memoryscope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore class TestLlamaIndexElasticSearchStore(unittest.TestCase): @@ -164,7 +164,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): meta_data={"5": "5"}, timestamp=13 )) - + def test_retrieve(self): filter_dict = { "timestamp": 12, @@ -172,12 +172,11 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): # "score_rank": 0, } - res = self.es_store.retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15) print(len(res)) print(res) - def test_retrieve_wo_query(self,): + def test_retrieve_wo_query(self, ): filter_dict = { "memory_id": "bbb456", } @@ -185,6 +184,5 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): print(len(res)) print(res) - def tearDown(self): self.es_store.close() diff --git a/tests/worker/test_workers_cn.py b/tests/worker/test_workers_cn.py index a8114b35..77cad3ff 100644 --- a/tests/worker/test_workers_cn.py +++ b/tests/worker/test_workers_cn.py @@ -1,44 +1,51 @@ import datetime import unittest -from memoryscope.cli import MemoryScope from memoryscope.constants.common_constants import CHAT_MESSAGES, NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, \ - MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES + MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES, \ + MEMORYSCOPE_CONTEXT +from memoryscope.core.config.arguments import Arguments +from memoryscope.core.memoryscope import MemoryScope +from memoryscope.core.utils.tool_functions import init_instance_by_config +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode from memoryscope.scheme.message import Message -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger -from memoryscope.utils.tool_functions import init_instance_by_config class TestWorkersCn(unittest.TestCase): """Tests for LLIEmbedding""" def setUp(self): - datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S') - self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True) - - ms = MemoryScope() - ms.load_config("config/demo_config_cn.yaml") - ms.init_global_content_by_config() + arguments = Arguments( + language="cn", + memory_chat_class="api_memory_chat", + generation_backend="dashscope_generation", + generation_model="qwen-max", + embedding_backend="dashscope_embedding", + embedding_model="text-embedding-v2", + use_dummy_ranker=False, + rank_backend="dashscope_rank", + rank_model="gte-rerank", + ) + self.ms = MemoryScope(arguments=arguments) + config = self.ms.dump_config() + self.ms.logger.info(f"config=\n{config}") def tearDown(self): - self.logger.close() + self.ms.close() @unittest.skip def test_extract_time(self): name = "extract_time" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) query = "明天我去上海出差" query_timestamp = int(datetime.datetime.now().timestamp()) @@ -53,13 +60,12 @@ class TestWorkersCn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜"), @@ -79,13 +85,12 @@ class TestWorkersCn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="你知道北京哪里的海鲜最新鲜吗"), @@ -115,13 +120,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜"), @@ -140,7 +144,7 @@ class TestWorkersCn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -149,13 +153,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。"), @@ -172,7 +175,7 @@ class TestWorkersCn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -181,13 +184,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation_with_time" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="去年我们一起合作了因果推断技术"), @@ -202,7 +204,7 @@ class TestWorkersCn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_WITH_TIME_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_WITH_TIME_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -211,43 +213,42 @@ class TestWorkersCn(unittest.TestCase): name = "contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"), MemoryNode(user_name="AI", target_name="用户", content="用户在阿里巴巴工作"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result1 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作"), MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result2 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作或有工作经验。"), MemoryNode(user_name="AI", target_name="用户", content="用户跳槽至openai工作。"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result3 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="我喜欢吃西瓜"), @@ -257,10 +258,10 @@ class TestWorkersCn(unittest.TestCase): MemoryNode(user_name="AI", target_name="用户", content="我爱吃苹果和香蕉"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result4 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) worker.logger.info(f"result1={result1}") worker.logger.info(f"result2={result2}") @@ -272,13 +273,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_reflection_subject" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"), @@ -293,62 +293,60 @@ class TestWorkersCn(unittest.TestCase): MemoryNode(content="用户想知道维持广泛社交关系的方法。"), ] - worker.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes) - worker.memory_handler.set_memories(INSIGHT_NODES, []) + worker.memory_manager.set_memories(NOT_REFLECTED_NODES, nodes) + worker.memory_manager.set_memories(INSIGHT_NODES, []) worker.run() - result = [node.key for node in worker.memory_handler.get_memories(INSIGHT_NODES)] + result = [node.key for node in worker.memory_manager.get_memories(INSIGHT_NODES)] result = "\n".join(result) worker.logger.info(f"result.get_reflection={result}") return worker @unittest.skip def test_update_insight_worker(self): - reflection_worker = self.test_get_reflection_subject.__wrapped__(self) + reflection_worker: MemoryBaseWorker = self.test_get_reflection_subject.__wrapped__(self) name = "update_insight" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, context=reflection_worker.context, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="用户喜欢打王者荣耀"), ] - worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(INSIGHT_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(INSIGHT_NODES)] result = "\n".join(result) worker.logger.info(f"result.update_insight={result}") - @unittest.skip + # @unittest.skip def test_long_contra_repeat_worker(self): name = "long_contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"), MemoryNode(content="用户在北京工作,感到压力大,寻求放松方式。"), MemoryNode(content="用户在上海工作。"), ] - worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) worker.unit_test_flag = True worker.run() - result = [node.content for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result.long_contra_repeat={result}") diff --git a/tests/worker/test_workers_en.py b/tests/worker/test_workers_en.py index 6b1c012a..ff421253 100644 --- a/tests/worker/test_workers_en.py +++ b/tests/worker/test_workers_en.py @@ -1,44 +1,51 @@ import datetime import unittest -from memoryscope.cli import MemoryScope from memoryscope.constants.common_constants import CHAT_MESSAGES, NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, \ - MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES + MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES, \ + MEMORYSCOPE_CONTEXT +from memoryscope.core.config.arguments import Arguments +from memoryscope.core.memoryscope import MemoryScope +from memoryscope.core.utils.tool_functions import init_instance_by_config +from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker from memoryscope.enumeration.message_role_enum import MessageRoleEnum -from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker from memoryscope.scheme.memory_node import MemoryNode from memoryscope.scheme.message import Message -from memoryscope.utils.global_context import G_CONTEXT -from memoryscope.utils.logger import Logger -from memoryscope.utils.tool_functions import init_instance_by_config class TestWorkersEn(unittest.TestCase): """Tests for LLIEmbedding""" def setUp(self): - datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S') - self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True) - - ms = MemoryScope() - ms.load_config("config/demo_config_en.yaml") - ms.init_global_content_by_config() + arguments = Arguments( + language="en", + memory_chat_class="api_memory_chat", + generation_backend="dashscope_generation", + generation_model="qwen-max", + embedding_backend="dashscope_embedding", + embedding_model="text-embedding-v2", + use_dummy_ranker=False, + rank_backend="dashscope_rank", + rank_model="gte-rerank", + ) + self.ms = MemoryScope(arguments=arguments) + config = self.ms.dump_config() + self.ms.logger.info(f"config=\n{config}") def tearDown(self): - self.logger.close() + self.ms.close() @unittest.skip def test_extract_time(self): name = "extract_time" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) query = "I will be on a business trip to Shanghai tomorrow." query_timestamp = int(datetime.datetime.now().timestamp()) @@ -53,13 +60,12 @@ class TestWorkersEn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="I love to eat Sichuan cuisine."), @@ -80,13 +86,12 @@ class TestWorkersEn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="Do you know where the freshest seafood is in Beijing?"), @@ -129,13 +134,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) # FIXME Does the appearance of 'am' indicate the presence of a time keyword? chat_messages = [ @@ -150,7 +154,7 @@ class TestWorkersEn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -159,13 +163,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, @@ -193,7 +196,7 @@ class TestWorkersEn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -202,13 +205,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation_with_time" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, @@ -226,7 +228,7 @@ class TestWorkersEn(unittest.TestCase): worker.set_context(CHAT_MESSAGES, chat_messages) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_WITH_TIME_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_WITH_TIME_NODES)] result = "\n".join(result) worker.logger.info(f"result={result}") @@ -235,43 +237,42 @@ class TestWorkersEn(unittest.TestCase): name = "contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="User is working in Meituan"), MemoryNode(user_name="AI", target_name="用户", content="User works at Alibaba"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result1 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="User works at JD.com"), MemoryNode(user_name="AI", target_name="用户", content="Users working in Meituan"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result2 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="User works at JD.com"), MemoryNode(user_name="AI", target_name="用户", content="User is working in Meituan"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result3 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="I like to eat watermelon"), @@ -279,10 +280,10 @@ class TestWorkersEn(unittest.TestCase): MemoryNode(user_name="AI", target_name="用户", content="I don't like watermelon"), ] - worker.memory_handler.set_memories(NEW_OBS_NODES, nodes) + worker.memory_manager.set_memories(NEW_OBS_NODES, nodes) worker.run() result4 = "\n".join([" ".join([node.content, node.store_status, node.action_status]) - for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]) + for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]) worker.logger.info(f"result1={result1}") worker.logger.info(f"result2={result2}") @@ -294,13 +295,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_reflection_subject" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="Users are interested in strategy games and looking for new challenges."), @@ -316,37 +316,36 @@ class TestWorkersEn(unittest.TestCase): MemoryNode(content="Users want to know how to maintain extensive social relationships."), ] - worker.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes) - worker.memory_handler.set_memories(INSIGHT_NODES, []) + worker.memory_manager.set_memories(NOT_REFLECTED_NODES, nodes) + worker.memory_manager.set_memories(INSIGHT_NODES, []) worker.run() - result = [node.key for node in worker.memory_handler.get_memories(INSIGHT_NODES)] + result = [node.key for node in worker.memory_manager.get_memories(INSIGHT_NODES)] result = "\n".join(result) worker.logger.info(f"result.get_reflection={result}") return worker @unittest.skip def test_update_insight_worker(self): - reflection_worker = self.test_get_reflection_subject.__wrapped__(self) + reflection_worker: MemoryBaseWorker = self.test_get_reflection_subject.__wrapped__(self) name = "update_insight" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, context=reflection_worker.context, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="Users like to play King of Glory"), ] - worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) worker.run() - result = [node.content for node in worker.memory_handler.get_memories(INSIGHT_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(INSIGHT_NODES)] result = "\n".join(result) worker.logger.info(f"result.update_insight={result}") @@ -355,23 +354,22 @@ class TestWorkersEn(unittest.TestCase): name = "long_contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=G_CONTEXT.worker_config[name], - suffix_name="worker", + config=self.ms.context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={}, + context={MEMORYSCOPE_CONTEXT: self.ms.context}, context_lock=None, - thread_pool=G_CONTEXT.thread_pool) + thread_pool=self.ms.context.thread_pool) nodes = [ MemoryNode(content="Users are interested in strategy games and looking for new challenges."), MemoryNode(content="The user works in Beijing, feels stressed, and is looking for ways to relax."), MemoryNode(content="User works in Shanghai."), ] - worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes) + worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes) worker.unit_test_flag = True worker.run() - result = [node.content for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)] + result = [node.content for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)] result = "\n".join(result) worker.logger.info(f"result.long_contra_repeat={result}")