From a5e50e25729bb21468c89fc82517ad948cbca0b0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E8=BD=A9?= Date: Mon, 12 Aug 2024 15:37:18 +0800 Subject: [PATCH 01/12] update log system --- .gitignore | 1 + memoryscope/core/models/base_model.py | 2 +- .../models/llama_index_embedding_model.py | 5 +++ .../models/llama_index_generation_model.py | 6 ++- .../core/models/llama_index_rank_model.py | 6 +++ memoryscope/core/utils/logger.py | 40 +++++++++++++++++++ 6 files changed, 58 insertions(+), 2 deletions(-) diff --git a/.gitignore b/.gitignore index 9120669e..717e5fec 100644 --- a/.gitignore +++ b/.gitignore @@ -143,6 +143,7 @@ docs/sphinx_doc/build/ *runs/ memoryscope.db tmp*.json +tmp*.py cradle* # sphinx docs diff --git a/memoryscope/core/models/base_model.py b/memoryscope/core/models/base_model.py index 07cbae69..69bb5c3f 100644 --- a/memoryscope/core/models/base_model.py +++ b/memoryscope/core/models/base_model.py @@ -35,7 +35,7 @@ class BaseModel(metaclass=ABCMeta): self.kwargs: dict = kwargs self._model: Any = None - self.logger = Logger.get_logger() + self.logger = Logger.get_logger(Logger.append_timestamp("base_model")) @property def model(self): diff --git a/memoryscope/core/models/llama_index_embedding_model.py b/memoryscope/core/models/llama_index_embedding_model.py index ed91cbc6..742b9558 100644 --- a/memoryscope/core/models/llama_index_embedding_model.py +++ b/memoryscope/core/models/llama_index_embedding_model.py @@ -15,6 +15,10 @@ class LlamaIndexEmbeddingModel(BaseModel): """ m_type: ModelEnum = ModelEnum.EMBEDDING_MODEL + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.logger = self.logger.get_logger(self.logger.append_timestamp("llama_index_embedding_model")) + @classmethod def register_model(cls, model_name: str, model_class: type): """ @@ -34,6 +38,7 @@ class LlamaIndexEmbeddingModel(BaseModel): if isinstance(text, str): text = [text] model_response.meta_data["data"] = dict(texts=text) + self.logger.info("Embedding Model:\n" + text[0]) def after_call(self, model_response: ModelResponse, **kwargs) -> ModelResponse: embeddings = model_response.raw diff --git a/memoryscope/core/models/llama_index_generation_model.py b/memoryscope/core/models/llama_index_generation_model.py index 60f9d8d8..9b3bb45a 100644 --- a/memoryscope/core/models/llama_index_generation_model.py +++ b/memoryscope/core/models/llama_index_generation_model.py @@ -25,6 +25,10 @@ class LlamaIndexGenerationModel(BaseModel): MODEL_REGISTRY.register("dashscope_generation", DashScope) MODEL_REGISTRY.register("openai_generation", OpenAI) + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.logger = self.logger.get_logger(self.logger.append_timestamp("llama_index_generation_model")) + def before_call(self, model_response: ModelResponse, **kwargs): """ Prepares the input data before making a call to the language model. @@ -77,7 +81,7 @@ class LlamaIndexGenerationModel(BaseModel): model_response.message.content = call_result.message.content else: raise NotImplementedError - + self.logger.info(self.logger.format_chat_message(model_response)) return model_response def _call(self, model_response: ModelResponse, stream: bool = False, **kwargs): diff --git a/memoryscope/core/models/llama_index_rank_model.py b/memoryscope/core/models/llama_index_rank_model.py index 26a6c79c..6e4c4245 100644 --- a/memoryscope/core/models/llama_index_rank_model.py +++ b/memoryscope/core/models/llama_index_rank_model.py @@ -19,6 +19,10 @@ class LlamaIndexRankModel(BaseModel): m_type: ModelEnum = ModelEnum.RANK_MODEL MODEL_REGISTRY.register("dashscope_rank", DashScopeRerank) + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.logger = self.logger.get_logger(self.logger.append_timestamp("llama_index_rank_model")) def before_call(self, model_response: ModelResponse, **kwargs): """ @@ -65,6 +69,8 @@ class LlamaIndexRankModel(BaseModel): text = node.node.text idx = documents_map[text] model_response.rank_scores[idx] = node.score + + self.logger.info(self.logger.format_rank_message(model_response)) return model_response def _call(self, model_response: ModelResponse, **kwargs): diff --git a/memoryscope/core/utils/logger.py b/memoryscope/core/utils/logger.py index d97d58b6..3caa5c7b 100644 --- a/memoryscope/core/utils/logger.py +++ b/memoryscope/core/utils/logger.py @@ -1,4 +1,5 @@ import logging +from datetime import datetime from logging.handlers import RotatingFileHandler from pathlib import Path @@ -63,6 +64,40 @@ class Logger(logging.Logger): self.info(f"logger={name} is inited.") # Logs an initialization message + def format_chat_message(self, message): + buf = '\n' + buf += f"++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++\n" + buf += f"LM Input:\n" + for chat_message in message.meta_data['data']['messages']: + buf += chat_message.content + buf += '\n' + buf += f"--------------------------------------------------------------\n" + buf += f"LM Output:\n" + buf += message.message.content + buf += '\n' + buf += f"++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++\n" + buf += '\n' + return buf + + def format_rank_message(self, model_response): + buf = '\n' + buf += f"++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++\n" + buf += f"Query Input:\n" + buf += model_response.meta_data['data']['query_str'] + buf += '\n' + buf += f"--------------------------------------------------------------\n" + buf += f"Rank:\n" + rank = 0 + for index, score in model_response.rank_scores.items(): + rank += 1 + node = model_response.meta_data['data']['nodes'][index] + node_text = node.text + buf += f"Score {score} | Rank {rank} | {node_text}\n" + buf += '\n' + buf += f"++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++\n" + buf += '\n' + return buf + def _add_file_handler(self): """ Adds a file handler to the logger which logs messages to a rotating file. @@ -76,6 +111,7 @@ class Logger(logging.Logger): file_path = Path().joinpath(self.dir_path, f"{self.name}.{self.file_type}") file_path.parent.mkdir(exist_ok=True) # Ensure the directory exists file_name = file_path.as_posix() # Get the absolute path as a string + print(f"[{self.name}] Registering Logger to file at: ", file_name) # Instantiate a rotating file handler with specified parameters file_handler = RotatingFileHandler( @@ -182,3 +218,7 @@ class Logger(logging.Logger): LOGGER_DICT[name] = Logger(name=name, **kwargs) return LOGGER_DICT[name] + + @staticmethod + def append_timestamp(name: str) -> str: + return f"{name}_{datetime.now().strftime(r'%Y%m%d_%H%M%S')}" \ No newline at end of file From 799c2381e4906d103b94049df61876fc4ab7b692 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E8=BD=A9?= Date: Mon, 12 Aug 2024 16:10:47 +0800 Subject: [PATCH 02/12] log work flow --- memoryscope/core/operation/base_workflow.py | 16 +++++++++++----- 1 file changed, 11 insertions(+), 5 deletions(-) diff --git a/memoryscope/core/operation/base_workflow.py b/memoryscope/core/operation/base_workflow.py index ea11d8d9..de3c41fc 100644 --- a/memoryscope/core/operation/base_workflow.py +++ b/memoryscope/core/operation/base_workflow.py @@ -31,7 +31,7 @@ class BaseWorkflow(object): self.context: Dict[str, Any] = {} self.context_lock = threading.Lock() - self.logger: Logger = Logger.get_logger() + self.logger: Logger = Logger.get_logger(Logger.append_timestamp("workflow")) if self.workflow: self.workflow_worker_list = self._parse_workflow() @@ -165,21 +165,27 @@ class BaseWorkflow(object): **kwargs: Additional keyword arguments to be passed to context. """ with Timer(f"workflow.{self.name}", time_log_type="wrap"): + self.logger.info(f"\n\n\n++++++++++++++++++++++++ [{self.name}] ++++++++++++++++++++++++") + self.context.clear() - self.context.update({WORKFLOW_NAME: self.name, **kwargs}) - + n_stage = len(self.workflow_worker_list) # Iterate over each part of the workflow - for workflow_part in self.workflow_worker_list: + for index, workflow_part in enumerate(self.workflow_worker_list): # Sequential execution for single-item parts if len(workflow_part) == 1: + self.logger.info(f"\n-----------------------------------------------------") + self.logger.info(f"sequential execution ({self.name}) | {index+1}/{n_stage}: {workflow_part[0]}") if not self._run_sub_workflow(workflow_part[0]): break # Parallel execution for multi-item parts else: t_list = [] # Submit tasks to the thread pool - for sub_workflow in workflow_part: + n_sub_stage = len(workflow_part) + for sub_index, sub_workflow in enumerate(workflow_part): + self.logger.info(f"\n-----------------------------------------------------") + self.logger.info(f"parallel sequential execution ({self.name}) | {index+1}/{n_stage} | {sub_index+1}/{n_sub_stage}: {str(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 From 7e0f85d538bb46b9be72ebeec6e647441e31aae9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E8=BD=A9?= Date: Mon, 12 Aug 2024 19:23:24 +0800 Subject: [PATCH 03/12] improve logging --- memoryscope/core/operation/base_workflow.py | 4 ++- memoryscope/core/utils/logger.py | 35 +++++++++++++++++++-- memoryscope/core/worker/memory_manager.py | 7 +++-- 3 files changed, 41 insertions(+), 5 deletions(-) diff --git a/memoryscope/core/operation/base_workflow.py b/memoryscope/core/operation/base_workflow.py index de3c41fc..b8670395 100644 --- a/memoryscope/core/operation/base_workflow.py +++ b/memoryscope/core/operation/base_workflow.py @@ -132,6 +132,7 @@ class BaseWorkflow(object): if name not in self.memoryscope_context.worker_conf_dict: raise RuntimeError(f"worker={name} is not exists in worker config!") + # note: shared context object in all workers self.worker_dict[name] = init_instance_by_config( config=self.memoryscope_context.worker_conf_dict[name], name=name, @@ -165,13 +166,14 @@ class BaseWorkflow(object): **kwargs: Additional keyword arguments to be passed to context. """ with Timer(f"workflow.{self.name}", time_log_type="wrap"): - self.logger.info(f"\n\n\n++++++++++++++++++++++++ [{self.name}] ++++++++++++++++++++++++") + self.logger.info(f"\n\n\n++++++++++++++++++++++++ [Operation: {self.name}] ++++++++++++++++++++++++") self.context.clear() self.context.update({WORKFLOW_NAME: self.name, **kwargs}) n_stage = len(self.workflow_worker_list) # Iterate over each part of the workflow for index, workflow_part in enumerate(self.workflow_worker_list): + self.logger.info(self.logger.format_current_context(self.context)) # Sequential execution for single-item parts if len(workflow_part) == 1: self.logger.info(f"\n-----------------------------------------------------") diff --git a/memoryscope/core/utils/logger.py b/memoryscope/core/utils/logger.py index 3caa5c7b..6a533f26 100644 --- a/memoryscope/core/utils/logger.py +++ b/memoryscope/core/utils/logger.py @@ -2,6 +2,7 @@ import logging from datetime import datetime from logging.handlers import RotatingFileHandler from pathlib import Path +from rich.console import Console LOG_FORMAT = "%(asctime)s %(levelname)s %(threadName)s %(module)s:%(lineno)d] %(message)s" DATE_FORMAT = "%Y-%m-%d %H:%M:%S" @@ -64,6 +65,36 @@ class Logger(logging.Logger): self.info(f"logger={name} is inited.") # Logs an initialization message + def format_current_context(self, context): + from rich.panel import Panel + from rich.text import Text + import pprint + pp = pprint.PrettyPrinter() + pretty_string = pp.pformat(context) + + def rich2text(rich_table): + console = Console(width=150) + with console.capture() as capture: + console.print(rich_table) + return '\n' + str(Text.from_ansi(capture.get())) + + return rich2text(Panel(pretty_string, width=128)) + + def format_current_memory(self, context): + from rich.panel import Panel + from rich.text import Text + import pprint + pp = pprint.PrettyPrinter() + pretty_string = pp.pformat(context) + + def rich2text(rich_table): + console = Console(width=150) + with console.capture() as capture: + console.print(rich_table) + return '\n' + str(Text.from_ansi(capture.get())) + + return rich2text(Panel(pretty_string, width=128)) + def format_chat_message(self, message): buf = '\n' buf += f"++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++\n" @@ -111,7 +142,7 @@ class Logger(logging.Logger): file_path = Path().joinpath(self.dir_path, f"{self.name}.{self.file_type}") file_path.parent.mkdir(exist_ok=True) # Ensure the directory exists file_name = file_path.as_posix() # Get the absolute path as a string - print(f"[{self.name}] Registering Logger to file at: ", file_name) + Console().print(f"[{self.name}] Registering Logger to file at: " + file_name, style="bold blue") # Instantiate a rotating file handler with specified parameters file_handler = RotatingFileHandler( @@ -221,4 +252,4 @@ class Logger(logging.Logger): @staticmethod def append_timestamp(name: str) -> str: - return f"{name}_{datetime.now().strftime(r'%Y%m%d_%H%M%S')}" \ No newline at end of file + return f"{name}_{datetime.now().strftime(r'%Y%m%d_%H%M')}" \ No newline at end of file diff --git a/memoryscope/core/worker/memory_manager.py b/memoryscope/core/worker/memory_manager.py index 936912a3..7d577993 100644 --- a/memoryscope/core/worker/memory_manager.py +++ b/memoryscope/core/worker/memory_manager.py @@ -24,7 +24,7 @@ class MemoryManager(object): # dict: key -> memory_id self._key_id_dict: Dict[str, List[str]] = {} - self.logger = Logger.get_logger() + self.logger = Logger.get_logger(Logger.append_timestamp("memory_manager")) @property def memory_store(self) -> BaseMemoryStore: @@ -95,7 +95,10 @@ class MemoryManager(object): self.logger.info(f"add to memory context memory id={node.memory_id} content={node.content} " f"store_status={node.store_status} action_status={node.action_status}") - self._key_id_dict[key] = [n.memory_id for n in nodes] + self.logger.info(self.logger.format_current_memory( + [f"{node.memory_type} | {node.content}" for node in nodes] + )) + def get_memories(self, keys: str | List[str]) -> List[MemoryNode]: """ From eacf6b5f85e0a141137fbb9e4e8efb4e8bc084f7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E8=BD=A9?= Date: Mon, 12 Aug 2024 19:34:52 +0800 Subject: [PATCH 04/12] revise memory logging --- memoryscope/core/utils/logger.py | 33 ++++++------------- memoryscope/core/worker/memory_base_worker.py | 2 +- memoryscope/core/worker/memory_manager.py | 16 ++++++--- 3 files changed, 22 insertions(+), 29 deletions(-) diff --git a/memoryscope/core/utils/logger.py b/memoryscope/core/utils/logger.py index 6a533f26..92d814bd 100644 --- a/memoryscope/core/utils/logger.py +++ b/memoryscope/core/utils/logger.py @@ -3,12 +3,20 @@ from datetime import datetime from logging.handlers import RotatingFileHandler from pathlib import Path from rich.console import Console +from rich.panel import Panel LOG_FORMAT = "%(asctime)s %(levelname)s %(threadName)s %(module)s:%(lineno)d] %(message)s" DATE_FORMAT = "%Y-%m-%d %H:%M:%S" LOGGER_DICT = {} +def rich2text(rich_table): + from rich.text import Text + console = Console(width=150) + with console.capture() as capture: + console.print(rich_table) + return '\n' + str(Text.from_ansi(capture.get())) + class Logger(logging.Logger): """ @@ -66,34 +74,13 @@ class Logger(logging.Logger): self.info(f"logger={name} is inited.") # Logs an initialization message def format_current_context(self, context): - from rich.panel import Panel - from rich.text import Text import pprint pp = pprint.PrettyPrinter() pretty_string = pp.pformat(context) - - def rich2text(rich_table): - console = Console(width=150) - with console.capture() as capture: - console.print(rich_table) - return '\n' + str(Text.from_ansi(capture.get())) - return rich2text(Panel(pretty_string, width=128)) - def format_current_memory(self, context): - from rich.panel import Panel - from rich.text import Text - import pprint - pp = pprint.PrettyPrinter() - pretty_string = pp.pformat(context) - - def rich2text(rich_table): - console = Console(width=150) - with console.capture() as capture: - console.print(rich_table) - return '\n' + str(Text.from_ansi(capture.get())) - - return rich2text(Panel(pretty_string, width=128)) + def wrap_in_box(self, context): + return rich2text(Panel(context, width=128)) def format_chat_message(self, message): buf = '\n' diff --git a/memoryscope/core/worker/memory_base_worker.py b/memoryscope/core/worker/memory_base_worker.py index 3d8fc0d7..53a72642 100644 --- a/memoryscope/core/worker/memory_base_worker.py +++ b/memoryscope/core/worker/memory_base_worker.py @@ -204,7 +204,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): MemoryHandler: An instance of MemoryHandler. """ if not self.has_content(MEMORY_MANAGER): - self.set_context(MEMORY_MANAGER, MemoryManager(self.memoryscope_context)) + self.set_context(MEMORY_MANAGER, MemoryManager(self.memoryscope_context, worker_name=self.name)) return self.get_context(MEMORY_MANAGER) def get_language_value(self, languages: dict | List[dict]) -> Any | List[Any]: diff --git a/memoryscope/core/worker/memory_manager.py b/memoryscope/core/worker/memory_manager.py index 7d577993..601b567b 100644 --- a/memoryscope/core/worker/memory_manager.py +++ b/memoryscope/core/worker/memory_manager.py @@ -13,7 +13,7 @@ class MemoryManager(object): The `MemoryHandler` class manages memory nodes with memory store. """ - def __init__(self, memoryscope_context: MemoryscopeContext): + def __init__(self, memoryscope_context: MemoryscopeContext, worker_name: str ="default_worker"): self.memoryscope_context: MemoryscopeContext = memoryscope_context self._memory_store: BaseMemoryStore | None = None @@ -26,6 +26,9 @@ class MemoryManager(object): self.logger = Logger.get_logger(Logger.append_timestamp("memory_manager")) + self.worker_name = worker_name + + @property def memory_store(self) -> BaseMemoryStore: """ @@ -95,10 +98,13 @@ class MemoryManager(object): self.logger.info(f"add to memory context memory id={node.memory_id} content={node.content} " f"store_status={node.store_status} action_status={node.action_status}") - self.logger.info(self.logger.format_current_memory( - [f"{node.memory_type} | {node.content}" for node in nodes] - )) - + if nodes: + self.logger.info( + self.logger.wrap_in_box( + '\n'.join([f"worker_name: {self.worker_name} | memory_type:{node.memory_type} | content:{node.content}" for node in nodes]) + ) + ) + def get_memories(self, keys: str | List[str]) -> List[MemoryNode]: """ From 22683d29ffae2981db63d7cbf811cc21b466633c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E8=BD=A9?= Date: Tue, 13 Aug 2024 15:25:13 +0800 Subject: [PATCH 05/12] context passing --- memoryscope/core/config/arguments.py | 1 - memoryscope/core/memoryscope.py | 5 +++++ memoryscope/core/memoryscope_context.py | 14 +++++++++++++- memoryscope/core/models/base_model.py | 3 +++ memoryscope/core/utils/__init__.py | 6 ++++++ memoryscope/core/utils/logger.py | 11 ++++------- memoryscope/core/utils/singleton.py | 9 +++++++++ 7 files changed, 40 insertions(+), 9 deletions(-) create mode 100644 memoryscope/core/utils/singleton.py diff --git a/memoryscope/core/config/arguments.py b/memoryscope/core/config/arguments.py index 92ccf9e1..f409f411 100644 --- a/memoryscope/core/config/arguments.py +++ b/memoryscope/core/config/arguments.py @@ -1,7 +1,6 @@ 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"}) diff --git a/memoryscope/core/memoryscope.py b/memoryscope/core/memoryscope.py index 0e5219bd..d673d0ea 100644 --- a/memoryscope/core/memoryscope.py +++ b/memoryscope/core/memoryscope.py @@ -1,4 +1,5 @@ from concurrent.futures import ThreadPoolExecutor +from datetime import datetime from memoryscope.core.chat.base_memory_chat import BaseMemoryChat from memoryscope.core.config.config_manager import ConfigManager @@ -32,6 +33,10 @@ class MemoryScope(ConfigManager): self.logger.warning("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.") + self.context.memory_scope_uuid = datetime.now().strftime(global_conf["logger_name_time_suffix"]) + + # set context_initialized + self.context.context_initialized = True # init memory_chat memory_chat_conf_dict = self.config["memory_chat"] diff --git a/memoryscope/core/memoryscope_context.py b/memoryscope/core/memoryscope_context.py index ac8d89fb..4de426b8 100644 --- a/memoryscope/core/memoryscope_context.py +++ b/memoryscope/core/memoryscope_context.py @@ -2,8 +2,9 @@ from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass, field from memoryscope.enumeration.language_enum import LanguageEnum +from memoryscope.core.utils.singleton import singleton - +@singleton @dataclass class MemoryscopeContext(object): """ @@ -27,3 +28,14 @@ class MemoryscopeContext(object): worker_conf_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> worker_conf"}) meta_data: dict = field(default_factory=lambda: {}) + + memory_scope_uuid: str = "" + + context_initialized: bool = False + +def get_ms_context(): + ms_context = MemoryscopeContext() + if ms_context.context_initialized: + return ms_context + else: + raise RuntimeError("MemoryscopeContext is not initialized yet. Please initialize it first.") diff --git a/memoryscope/core/models/base_model.py b/memoryscope/core/models/base_model.py index 69bb5c3f..4fc6befe 100644 --- a/memoryscope/core/models/base_model.py +++ b/memoryscope/core/models/base_model.py @@ -8,6 +8,8 @@ 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.core.memoryscope_context import MemoryscopeContext +from memoryscope.core.memoryscope_context import get_ms_context MODEL_REGISTRY = Registry("models") @@ -32,6 +34,7 @@ class BaseModel(metaclass=ABCMeta): self.retry_interval: float = retry_interval self.kwargs_filter: bool = kwargs_filter self.raise_exception: bool = raise_exception + self.context: MemoryscopeContext = get_ms_context() self.kwargs: dict = kwargs self._model: Any = None diff --git a/memoryscope/core/utils/__init__.py b/memoryscope/core/utils/__init__.py index 903b05a3..ab3e278c 100644 --- a/memoryscope/core/utils/__init__.py +++ b/memoryscope/core/utils/__init__.py @@ -15,6 +15,12 @@ from .tool_functions import ( cosine_similarity ) + +def get_context(): + from memoryscope import MemoryscopeContext + return MemoryscopeContext() + + __all__ = [ "DatetimeHandler", "Logger", diff --git a/memoryscope/core/utils/logger.py b/memoryscope/core/utils/logger.py index 92d814bd..2c9d64be 100644 --- a/memoryscope/core/utils/logger.py +++ b/memoryscope/core/utils/logger.py @@ -84,7 +84,6 @@ class Logger(logging.Logger): def format_chat_message(self, message): buf = '\n' - buf += f"++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++\n" buf += f"LM Input:\n" for chat_message in message.meta_data['data']['messages']: buf += chat_message.content @@ -93,13 +92,11 @@ class Logger(logging.Logger): buf += f"LM Output:\n" buf += message.message.content buf += '\n' - buf += f"++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++\n" buf += '\n' - return buf + return self.wrap_in_box(buf) def format_rank_message(self, model_response): buf = '\n' - buf += f"++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++\n" buf += f"Query Input:\n" buf += model_response.meta_data['data']['query_str'] buf += '\n' @@ -112,9 +109,8 @@ class Logger(logging.Logger): node_text = node.text buf += f"Score {score} | Rank {rank} | {node_text}\n" buf += '\n' - buf += f"++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++\n" buf += '\n' - return buf + return self.wrap_in_box(buf) def _add_file_handler(self): """ @@ -239,4 +235,5 @@ class Logger(logging.Logger): @staticmethod def append_timestamp(name: str) -> str: - return f"{name}_{datetime.now().strftime(r'%Y%m%d_%H%M')}" \ No newline at end of file + from memoryscope.core.memoryscope_context import get_ms_context + return f"{name}_{get_ms_context().memory_scope_uuid}" \ No newline at end of file diff --git a/memoryscope/core/utils/singleton.py b/memoryscope/core/utils/singleton.py new file mode 100644 index 00000000..b767cd1e --- /dev/null +++ b/memoryscope/core/utils/singleton.py @@ -0,0 +1,9 @@ +def singleton(cls): + _instance = {} + + def _singleton(*args, **kargs): + if cls not in _instance: + _instance[cls] = cls(*args, **kargs) + return _instance[cls] + + return _singleton \ No newline at end of file From 24bea8dae65a2872a4eebbebb97f55b6a8d5939f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E8=BD=A9?= Date: Tue, 13 Aug 2024 16:03:22 +0800 Subject: [PATCH 06/12] update es store logging system --- .../core/storage/llama_index_es_memory_store.py | 17 ++++++++++++++--- memoryscope/core/utils/logger.py | 11 +++++++---- 2 files changed, 21 insertions(+), 7 deletions(-) diff --git a/memoryscope/core/storage/llama_index_es_memory_store.py b/memoryscope/core/storage/llama_index_es_memory_store.py index f4142103..43e09b3f 100644 --- a/memoryscope/core/storage/llama_index_es_memory_store.py +++ b/memoryscope/core/storage/llama_index_es_memory_store.py @@ -37,7 +37,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store, embed_model=self.embedding_model.model) - self.logger = Logger.get_logger() + self.logger = Logger.get_logger(Logger.append_timestamp("es_memory_store")) def retrieve_memories(self, query: str = "", @@ -65,7 +65,11 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): text_nodes = retriever.retrieve(query) if text_nodes and text_nodes[0].embedding: self.emb_dims = len(text_nodes[0].embedding) - + self.logger.log_dictionary_info({ + "action": "retrieve_memories", + "query": query, + "text_nodes": [f"ID: {n.node_id} |Text: {n.text}" for n in text_nodes] + }) return [self._text_node_2_memory_node(n) for n in text_nodes] async def a_retrieve_memories(self, @@ -115,14 +119,21 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): def insert(self, node: MemoryNode): self.index.insert_nodes([self._memory_node_2_text_node(node)]) + self.logger.log_dictionary_info({ + "action": "insert", + "node": f"ID: {node.memory_id} | Text: {node.content} | Key: {node.key} | Type: {node.memory_type}" + }) def delete(self, node: MemoryNode): + self.logger.log_dictionary_info({ + "action": "delete", + "id": node.memory_id, + }) return self.es_store.delete(node.memory_id) def update(self, node: MemoryNode, update_embedding: bool = True): if update_embedding: node.vector = [] - self.delete(node) self.insert(node) diff --git a/memoryscope/core/utils/logger.py b/memoryscope/core/utils/logger.py index 2c9d64be..9f7ff446 100644 --- a/memoryscope/core/utils/logger.py +++ b/memoryscope/core/utils/logger.py @@ -73,15 +73,18 @@ class Logger(logging.Logger): self.info(f"logger={name} is inited.") # Logs an initialization message + def log_dictionary_info(self, dictionary): + self.info(self.format_current_context(dictionary)) + def format_current_context(self, context): import pprint pp = pprint.PrettyPrinter() pretty_string = pp.pformat(context) return rich2text(Panel(pretty_string, width=128)) - + def wrap_in_box(self, context): return rich2text(Panel(context, width=128)) - + def format_chat_message(self, message): buf = '\n' buf += f"LM Input:\n" @@ -111,7 +114,7 @@ class Logger(logging.Logger): buf += '\n' buf += '\n' return self.wrap_in_box(buf) - + def _add_file_handler(self): """ Adds a file handler to the logger which logs messages to a rotating file. @@ -125,7 +128,7 @@ class Logger(logging.Logger): file_path = Path().joinpath(self.dir_path, f"{self.name}.{self.file_type}") file_path.parent.mkdir(exist_ok=True) # Ensure the directory exists file_name = file_path.as_posix() # Get the absolute path as a string - Console().print(f"[{self.name}] Registering Logger to file at: " + file_name, style="bold blue") + Console().print(f"[{self.name}] Registering logger to file at: " + file_name, style="bold blue") # Instantiate a rotating file handler with specified parameters file_handler = RotatingFileHandler( From 5f6a52215f51f1ee51bdfb3b2fff0d115e5142c0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E8=BD=A9?= Date: Wed, 14 Aug 2024 10:20:52 +0800 Subject: [PATCH 07/12] improve es logger --- memoryscope/core/operation/base_workflow.py | 16 +++++++------- .../storage/llama_index_sync_elasticsearch.py | 21 +++++++++++++++---- memoryscope/core/worker/memory_base_worker.py | 4 ++-- requirements.txt | 3 ++- 4 files changed, 29 insertions(+), 15 deletions(-) diff --git a/memoryscope/core/operation/base_workflow.py b/memoryscope/core/operation/base_workflow.py index b8670395..26efa62f 100644 --- a/memoryscope/core/operation/base_workflow.py +++ b/memoryscope/core/operation/base_workflow.py @@ -3,6 +3,7 @@ import threading from concurrent.futures import ThreadPoolExecutor, as_completed from itertools import zip_longest from typing import Dict, Any, List +from rich.console import Console from memoryscope.constants.common_constants import WORKFLOW_NAME from memoryscope.core.memoryscope_context import MemoryscopeContext @@ -11,7 +12,6 @@ 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, @@ -166,18 +166,18 @@ class BaseWorkflow(object): **kwargs: Additional keyword arguments to be passed to context. """ with Timer(f"workflow.{self.name}", time_log_type="wrap"): - self.logger.info(f"\n\n\n++++++++++++++++++++++++ [Operation: {self.name}] ++++++++++++++++++++++++") - + log_buf = f"Operation: {self.name}" + self.logger.info(log_buf); Console().print(log_buf, style="bold red") self.context.clear() self.context.update({WORKFLOW_NAME: self.name, **kwargs}) n_stage = len(self.workflow_worker_list) # Iterate over each part of the workflow for index, workflow_part in enumerate(self.workflow_worker_list): - self.logger.info(self.logger.format_current_context(self.context)) + # self.logger.info(self.logger.format_current_context(self.context)) # Sequential execution for single-item parts if len(workflow_part) == 1: - self.logger.info(f"\n-----------------------------------------------------") - self.logger.info(f"sequential execution ({self.name}) | {index+1}/{n_stage}: {workflow_part[0]}") + log_buf = f"\t- Operation: {self.name} | {index+1}/{n_stage}: {workflow_part[0]}" + self.logger.info(log_buf); Console().print(log_buf, style="bold red") if not self._run_sub_workflow(workflow_part[0]): break # Parallel execution for multi-item parts @@ -186,8 +186,8 @@ class BaseWorkflow(object): # Submit tasks to the thread pool n_sub_stage = len(workflow_part) for sub_index, sub_workflow in enumerate(workflow_part): - self.logger.info(f"\n-----------------------------------------------------") - self.logger.info(f"parallel sequential execution ({self.name}) | {index+1}/{n_stage} | {sub_index+1}/{n_sub_stage}: {str(sub_workflow)}") + log_buf = f"\t- Operation: {self.name} | {index+1}/{n_stage} | sub workflow {sub_index+1}/{n_sub_stage}: {str(sub_workflow)}" + self.logger.info(log_buf); Console().print(log_buf, style="red") t_list.append(self.thread_pool.submit(self._run_sub_workflow, sub_workflow)) # Check results; if any task returns False, stop the workflow diff --git a/memoryscope/core/storage/llama_index_sync_elasticsearch.py b/memoryscope/core/storage/llama_index_sync_elasticsearch.py index 87c0d97e..ed468416 100644 --- a/memoryscope/core/storage/llama_index_sync_elasticsearch.py +++ b/memoryscope/core/storage/llama_index_sync_elasticsearch.py @@ -1,10 +1,10 @@ """Elasticsearch vector store.""" -from logging import getLogger from typing import Any, Callable, Dict, List, Literal, Optional, Union, cast import nest_asyncio import numpy as np +from memoryscope.core.utils.logger import Logger from elasticsearch import AsyncElasticsearch, Elasticsearch from elasticsearch.helpers.vectorstore import ( AsyncBM25Strategy, @@ -30,8 +30,6 @@ from llama_index.vector_stores.elasticsearch.utils import ( get_user_agent, ) -logger = getLogger(__name__) - DISTANCE_STRATEGIES = Literal[ "COSINE", "DOT_PRODUCT", @@ -366,6 +364,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore): batch_size: int = 200 distance_strategy: Optional[DISTANCE_STRATEGIES] = "COSINE" retrieval_strategy: AsyncRetrievalStrategy + logger: Logger = None _store = PrivateAttr() @@ -431,6 +430,8 @@ class SyncElasticsearchStore(BasePydanticVectorStore): retrieval_strategy=retrieval_strategy, ) + self.logger = Logger.get_logger(Logger.append_timestamp("elastic_search")) + @property def client(self) -> Any: """ @@ -471,6 +472,10 @@ class SyncElasticsearchStore(BasePydanticVectorStore): Note: This method delegates the actual operation to the `sync_add` method. """ + self.logger.log_dictionary_info({ + "action": "add", + "node_count": len(nodes), + }) return self.sync_add(nodes, create_index_if_not_exists=create_index_if_not_exists) def sync_add( @@ -550,6 +555,10 @@ class SyncElasticsearchStore(BasePydanticVectorStore): This method internally calls a synchronous delete method (`sync_delete`) to execute the deletion operation against Elasticsearch. """ + self.logger.log_dictionary_info({ + "action": "delete", + "id": ref_doc_id, + }) return self.sync_delete(ref_doc_id, **delete_kwargs) def sync_delete(self, ref_doc_id: str, **delete_kwargs: Any) -> None: @@ -604,6 +613,10 @@ class SyncElasticsearchStore(BasePydanticVectorStore): Exception: If an error occurs during the Elasticsearch query execution. """ + self.logger.log_dictionary_info({ + "action": "query", + "query": query.query_str, + }) return self.sync_query(query, custom_query, es_filter, **kwargs) def sync_query( @@ -673,7 +686,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore): node.embedding = embedding except Exception: # Legacy support for old metadata format - logger.warning( + self.logger.warning( f"Could not parse metadata from hit {hit['_source']['metadata']}" ) node_info = source.get("node_info") diff --git a/memoryscope/core/worker/memory_base_worker.py b/memoryscope/core/worker/memory_base_worker.py index 53a72642..6ca3f37a 100644 --- a/memoryscope/core/worker/memory_base_worker.py +++ b/memoryscope/core/worker/memory_base_worker.py @@ -246,9 +246,9 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): system_message = Message(role=MessageRoleEnum.SYSTEM.value, content=system_content) if concat_system_prompt: - user_content_list = [system_content, few_shot, user_query] + user_content_list = [system_content, '\n', few_shot, '\n', user_query] else: - user_content_list = [few_shot, user_query] + user_content_list = [few_shot, '\n', user_query] user_message = Message(role=MessageRoleEnum.USER.value, content="\n".join([x.strip() for x in user_content_list])) return [system_message, user_message] diff --git a/requirements.txt b/requirements.txt index 5fc77936..c8617f72 100644 --- a/requirements.txt +++ b/requirements.txt @@ -14,4 +14,5 @@ dashscope~=1.19.1 elasticsearch~=8.14.0 pyyaml~=6.0.1 ray~=2.31.0 -numpy~=1.26.4 \ No newline at end of file +numpy~=1.26.4 +rich \ No newline at end of file From bc8bb78efb34480574bb0fdef52533501ca8b846 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E8=BD=A9?= Date: Wed, 14 Aug 2024 10:29:02 +0800 Subject: [PATCH 08/12] bug fix --- memoryscope/core/memoryscope_context.py | 2 ++ memoryscope/core/operation/base_workflow.py | 14 +++++++++++--- 2 files changed, 13 insertions(+), 3 deletions(-) diff --git a/memoryscope/core/memoryscope_context.py b/memoryscope/core/memoryscope_context.py index 4de426b8..99a163ae 100644 --- a/memoryscope/core/memoryscope_context.py +++ b/memoryscope/core/memoryscope_context.py @@ -31,6 +31,8 @@ class MemoryscopeContext(object): memory_scope_uuid: str = "" + print_workflow_dynamic: bool = False + context_initialized: bool = False def get_ms_context(): diff --git a/memoryscope/core/operation/base_workflow.py b/memoryscope/core/operation/base_workflow.py index 26efa62f..4f3d1783 100644 --- a/memoryscope/core/operation/base_workflow.py +++ b/memoryscope/core/operation/base_workflow.py @@ -37,6 +37,11 @@ class BaseWorkflow(object): self.workflow_worker_list = self._parse_workflow() self._print_workflow() + def workflow_print_console(self, *args, **kwargs): + if self.memoryscope_context.print_workflow_dynamic: + Console().print(*args, **kwargs) + return + def _parse_workflow(self): """ Parses the workflow string to configure worker threads and organizes them into execution order. @@ -167,7 +172,8 @@ class BaseWorkflow(object): """ with Timer(f"workflow.{self.name}", time_log_type="wrap"): log_buf = f"Operation: {self.name}" - self.logger.info(log_buf); Console().print(log_buf, style="bold red") + self.logger.info(log_buf) + self.workflow_print_console(log_buf, style="bold red") self.context.clear() self.context.update({WORKFLOW_NAME: self.name, **kwargs}) n_stage = len(self.workflow_worker_list) @@ -177,7 +183,8 @@ class BaseWorkflow(object): # Sequential execution for single-item parts if len(workflow_part) == 1: log_buf = f"\t- Operation: {self.name} | {index+1}/{n_stage}: {workflow_part[0]}" - self.logger.info(log_buf); Console().print(log_buf, style="bold red") + self.logger.info(log_buf) + self.workflow_print_console(log_buf, style="bold red") if not self._run_sub_workflow(workflow_part[0]): break # Parallel execution for multi-item parts @@ -187,7 +194,8 @@ class BaseWorkflow(object): n_sub_stage = len(workflow_part) for sub_index, sub_workflow in enumerate(workflow_part): log_buf = f"\t- Operation: {self.name} | {index+1}/{n_stage} | sub workflow {sub_index+1}/{n_sub_stage}: {str(sub_workflow)}" - self.logger.info(log_buf); Console().print(log_buf, style="red") + self.logger.info(log_buf) + self.workflow_print_console(log_buf, style="red") t_list.append(self.thread_pool.submit(self._run_sub_workflow, sub_workflow)) # Check results; if any task returns False, stop the workflow From bf6dfd3911182826993c371e47087faf74c278af Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E8=BD=A9?= Date: Wed, 14 Aug 2024 10:44:37 +0800 Subject: [PATCH 09/12] refactor: reorganize logger imports and methods --- memoryscope/core/utils/logger.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/memoryscope/core/utils/logger.py b/memoryscope/core/utils/logger.py index 9f7ff446..0accdc9b 100644 --- a/memoryscope/core/utils/logger.py +++ b/memoryscope/core/utils/logger.py @@ -4,6 +4,7 @@ from logging.handlers import RotatingFileHandler from pathlib import Path from rich.console import Console from rich.panel import Panel +from rich.text import Text LOG_FORMAT = "%(asctime)s %(levelname)s %(threadName)s %(module)s:%(lineno)d] %(message)s" DATE_FORMAT = "%Y-%m-%d %H:%M:%S" @@ -11,7 +12,6 @@ DATE_FORMAT = "%Y-%m-%d %H:%M:%S" LOGGER_DICT = {} def rich2text(rich_table): - from rich.text import Text console = Console(width=150) with console.capture() as capture: console.print(rich_table) @@ -206,7 +206,7 @@ class Logger(logging.Logger): if extra is None: extra = {} if self.trace_id: - extra["trace_id"] = self.trace_id # ⭐ Include trace_id from the logger in the log record extra data + extra["trace_id"] = self.trace_id # тнР Include trace_id from the logger in the log record extra data return super().makeRecord(name, level, fn, lno, msg, args, exc_info, func, extra, sinfo) @classmethod From c66bea14ab23f8b97c80d97257714d094667cc80 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E8=BD=A9?= Date: Wed, 14 Aug 2024 11:02:35 +0800 Subject: [PATCH 10/12] import fix --- memoryscope/core/utils/logger.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/memoryscope/core/utils/logger.py b/memoryscope/core/utils/logger.py index 0accdc9b..49d42645 100644 --- a/memoryscope/core/utils/logger.py +++ b/memoryscope/core/utils/logger.py @@ -5,6 +5,7 @@ from pathlib import Path from rich.console import Console from rich.panel import Panel from rich.text import Text +import pprint LOG_FORMAT = "%(asctime)s %(levelname)s %(threadName)s %(module)s:%(lineno)d] %(message)s" DATE_FORMAT = "%Y-%m-%d %H:%M:%S" @@ -17,7 +18,6 @@ def rich2text(rich_table): console.print(rich_table) return '\n' + str(Text.from_ansi(capture.get())) - class Logger(logging.Logger): """ The `Logger` class handle the stream of information or errors in activities. @@ -77,7 +77,6 @@ class Logger(logging.Logger): self.info(self.format_current_context(dictionary)) def format_current_context(self, context): - import pprint pp = pprint.PrettyPrinter() pretty_string = pp.pformat(context) return rich2text(Panel(pretty_string, width=128)) @@ -206,7 +205,7 @@ class Logger(logging.Logger): if extra is None: extra = {} if self.trace_id: - extra["trace_id"] = self.trace_id # тнР Include trace_id from the logger in the log record extra data + extra["trace_id"] = self.trace_id # Include trace_id from the logger in the log record extra data return super().makeRecord(name, level, fn, lno, msg, args, exc_info, func, extra, sinfo) @classmethod From 39b23c178c0012173cfca9482f2487305bebec9e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E8=BD=A9?= Date: Wed, 14 Aug 2024 11:03:31 +0800 Subject: [PATCH 11/12] rotate import order --- memoryscope/core/utils/logger.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/memoryscope/core/utils/logger.py b/memoryscope/core/utils/logger.py index 49d42645..2209d632 100644 --- a/memoryscope/core/utils/logger.py +++ b/memoryscope/core/utils/logger.py @@ -1,11 +1,10 @@ import logging -from datetime import datetime +import pprint from logging.handlers import RotatingFileHandler from pathlib import Path from rich.console import Console from rich.panel import Panel from rich.text import Text -import pprint LOG_FORMAT = "%(asctime)s %(levelname)s %(threadName)s %(module)s:%(lineno)d] %(message)s" DATE_FORMAT = "%Y-%m-%d %H:%M:%S" From ceffb56a53b180410b62b85bfd5f8f35de0874f7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E8=BD=A9?= Date: Wed, 14 Aug 2024 11:34:32 +0800 Subject: [PATCH 12/12] Refactor code to use workflow context consistently --- memoryscope/core/operation/base_workflow.py | 10 ++--- .../core/operation/consolidate_memory_op.py | 2 +- .../core/operation/frontend_operation.py | 2 +- memoryscope/core/utils/__init__.py | 6 --- memoryscope/core/utils/logger.py | 42 ++++++++++--------- .../worker/backend/update_memory_worker.py | 2 +- memoryscope/core/worker/base_worker.py | 14 +++---- memoryscope/core/worker/dummy_worker.py | 6 +-- .../worker/frontend/extract_time_worker.py | 4 +- .../worker/frontend/fuse_rerank_worker.py | 4 +- .../worker/frontend/print_memory_worker.py | 2 +- .../worker/frontend/read_message_worker.py | 2 +- .../worker/frontend/retrieve_memory_worker.py | 2 +- .../worker/frontend/semantic_rank_worker.py | 2 +- .../core/worker/frontend/set_query_worker.py | 2 +- memoryscope/core/worker/memory_base_worker.py | 24 +++++------ memoryscope/core/worker/memory_manager.py | 6 +-- tests/worker/test_workers_cn.py | 14 +++---- tests/worker/test_workers_en.py | 14 +++---- 19 files changed, 78 insertions(+), 82 deletions(-) diff --git a/memoryscope/core/operation/base_workflow.py b/memoryscope/core/operation/base_workflow.py index 4f3d1783..e1fd1495 100644 --- a/memoryscope/core/operation/base_workflow.py +++ b/memoryscope/core/operation/base_workflow.py @@ -28,7 +28,7 @@ class BaseWorkflow(object): self.workflow_worker_list: List[List[List[str]]] = [] self.worker_dict: Dict[str, BaseWorker | bool] = {} - self.context: Dict[str, Any] = {} + self.workflow_context: Dict[str, Any] = {} self.context_lock = threading.Lock() self.logger: Logger = Logger.get_logger(Logger.append_timestamp("workflow")) @@ -142,7 +142,7 @@ class BaseWorkflow(object): config=self.memoryscope_context.worker_conf_dict[name], name=name, is_multi_thread=is_backend or self.worker_dict[name], - context=self.context, + context=self.workflow_context, memoryscope_context=self.memoryscope_context, context_lock=self.context_lock, thread_pool=self.thread_pool, @@ -174,12 +174,12 @@ class BaseWorkflow(object): log_buf = f"Operation: {self.name}" self.logger.info(log_buf) self.workflow_print_console(log_buf, style="bold red") - self.context.clear() - self.context.update({WORKFLOW_NAME: self.name, **kwargs}) + self.workflow_context.clear() + self.workflow_context.update({WORKFLOW_NAME: self.name, **kwargs}) n_stage = len(self.workflow_worker_list) # Iterate over each part of the workflow for index, workflow_part in enumerate(self.workflow_worker_list): - # self.logger.info(self.logger.format_current_context(self.context)) + # self.logger.info(self.logger.format_current_context(self.workflow_context)) # Sequential execution for single-item parts if len(workflow_part) == 1: log_buf = f"\t- Operation: {self.name} | {index+1}/{n_stage}: {workflow_part[0]}" diff --git a/memoryscope/core/operation/consolidate_memory_op.py b/memoryscope/core/operation/consolidate_memory_op.py index 5fb691f5..d10bb7d3 100644 --- a/memoryscope/core/operation/consolidate_memory_op.py +++ b/memoryscope/core/operation/consolidate_memory_op.py @@ -70,7 +70,7 @@ class ConsolidateMemoryOp(BackendOperation): self.run_workflow(**workflow_kwargs) # Retrieve the result from the context after workflow execution - result = self.context.get(RESULT) + result = self.workflow_context.get(RESULT) # set message memorized with self.message_lock: diff --git a/memoryscope/core/operation/frontend_operation.py b/memoryscope/core/operation/frontend_operation.py index d74212e5..ea9b4b9f 100644 --- a/memoryscope/core/operation/frontend_operation.py +++ b/memoryscope/core/operation/frontend_operation.py @@ -58,4 +58,4 @@ class FrontendOperation(BaseWorkflow, BaseOperation): self.run_workflow(**workflow_kwargs) # Retrieve the result from the context after workflow execution - return self.context.get(RESULT) + return self.workflow_context.get(RESULT) diff --git a/memoryscope/core/utils/__init__.py b/memoryscope/core/utils/__init__.py index ab3e278c..903b05a3 100644 --- a/memoryscope/core/utils/__init__.py +++ b/memoryscope/core/utils/__init__.py @@ -15,12 +15,6 @@ from .tool_functions import ( cosine_similarity ) - -def get_context(): - from memoryscope import MemoryscopeContext - return MemoryscopeContext() - - __all__ = [ "DatetimeHandler", "Logger", diff --git a/memoryscope/core/utils/logger.py b/memoryscope/core/utils/logger.py index 2209d632..8c3bf251 100644 --- a/memoryscope/core/utils/logger.py +++ b/memoryscope/core/utils/logger.py @@ -84,34 +84,36 @@ class Logger(logging.Logger): return rich2text(Panel(context, width=128)) def format_chat_message(self, message): - buf = '\n' - buf += f"LM Input:\n" + buf = [] + buf.append('\n') + buf.append(f"LM Input:\n") for chat_message in message.meta_data['data']['messages']: - buf += chat_message.content - buf += '\n' - buf += f"--------------------------------------------------------------\n" - buf += f"LM Output:\n" - buf += message.message.content - buf += '\n' - buf += '\n' - return self.wrap_in_box(buf) + buf.append(chat_message.content) + buf.append('\n') + buf.append(f"--------------------------------------------------------------\n") + buf.append(f"LM Output:\n") + buf.append(message.message.content) + buf.append('\n') + buf.append('\n') + return self.wrap_in_box(''.join(buf)) def format_rank_message(self, model_response): - buf = '\n' - buf += f"Query Input:\n" - buf += model_response.meta_data['data']['query_str'] - buf += '\n' - buf += f"--------------------------------------------------------------\n" - buf += f"Rank:\n" + buf = [] + buf.append('\n') + buf.append(f"Query Input:\n") + buf.append(model_response.meta_data['data']['query_str']) + buf.append('\n') + buf.append(f"--------------------------------------------------------------\n") + buf.append(f"Rank:\n") rank = 0 for index, score in model_response.rank_scores.items(): rank += 1 node = model_response.meta_data['data']['nodes'][index] node_text = node.text - buf += f"Score {score} | Rank {rank} | {node_text}\n" - buf += '\n' - buf += '\n' - return self.wrap_in_box(buf) + buf.append(f"Score {score} | Rank {rank} | {node_text}\n") + buf.append('\n') + buf.append('\n') + return self.wrap_in_box(''.join(buf)) def _add_file_handler(self): """ diff --git a/memoryscope/core/worker/backend/update_memory_worker.py b/memoryscope/core/worker/backend/update_memory_worker.py index 0ca13438..914602eb 100644 --- a/memoryscope/core/worker/backend/update_memory_worker.py +++ b/memoryscope/core/worker/backend/update_memory_worker.py @@ -116,4 +116,4 @@ class UpdateMemoryWorker(MemoryBaseWorker): for action, nodes in updated_nodes.items(): for node in nodes: line.append(f"{action} {node.memory_type}: {node.content} ({node.store_status})") - self.set_context(RESULT, "\n".join(line)) + self.set_workflow_context(RESULT, "\n".join(line)) diff --git a/memoryscope/core/worker/base_worker.py b/memoryscope/core/worker/base_worker.py index 6a07c26e..2b2eaf45 100644 --- a/memoryscope/core/worker/base_worker.py +++ b/memoryscope/core/worker/base_worker.py @@ -37,7 +37,7 @@ class BaseWorker(metaclass=ABCMeta): """ self.name: str = name - self.context: Dict[str, Any] = context + self.workflow_context: Dict[str, Any] = context self.memoryscope_context: MemoryscopeContext = memoryscope_context self.context_lock = context_lock self.raise_exception: bool = raise_exception @@ -164,7 +164,7 @@ class BaseWorker(metaclass=ABCMeta): except Exception as e: self.logger.exception(f"run {self.name} failed! args={e.args}") - def get_context(self, key: str, default=None): + def get_workflow_context(self, key: str, default=None): """ Retrieves a value from the shared context. @@ -175,9 +175,9 @@ class BaseWorker(metaclass=ABCMeta): Returns: The value from the context or the default value. """ - return self.context.get(key, default) + return self.workflow_context.get(key, default) - def set_context(self, key: str, value: Any): + def set_workflow_context(self, key: str, value: Any): """ Sets a value in the shared context. @@ -187,9 +187,9 @@ class BaseWorker(metaclass=ABCMeta): """ if self.is_multi_thread: with self.context_lock: - self.context[key] = value + self.workflow_context[key] = value else: - self.context[key] = value + self.workflow_context[key] = value def has_content(self, key: str): """ @@ -201,4 +201,4 @@ class BaseWorker(metaclass=ABCMeta): Returns: bool: True if the key is in the context, otherwise False. """ - return key in self.context + return key in self.workflow_context diff --git a/memoryscope/core/worker/dummy_worker.py b/memoryscope/core/worker/dummy_worker.py index a5f4b201..d0619bb8 100644 --- a/memoryscope/core/worker/dummy_worker.py +++ b/memoryscope/core/worker/dummy_worker.py @@ -12,11 +12,11 @@ class DummyWorker(MemoryBaseWorker): This method utilizes the BaseWorker's capabilities to interact with the workflow context. """ - workflow_name = self.get_context(WORKFLOW_NAME) - chat_kwargs = self.get_context(CHAT_KWARGS) + workflow_name = self.get_workflow_context(WORKFLOW_NAME) + chat_kwargs = self.get_workflow_context(CHAT_KWARGS) self.logger.info(f"Entering workflow={workflow_name}.dummy_worker!") # Records the current timestamp as an integer ts = int(datetime.datetime.now().timestamp()) # Retrieves the current file's path file_path = __file__ - self.set_context(RESULT, f"test {workflow_name} kwargs={chat_kwargs} file_path={file_path} \nts={ts}") + self.set_workflow_context(RESULT, f"test {workflow_name} kwargs={chat_kwargs} file_path={file_path} \nts={ts}") diff --git a/memoryscope/core/worker/frontend/extract_time_worker.py b/memoryscope/core/worker/frontend/extract_time_worker.py index cff18073..70e1ba00 100644 --- a/memoryscope/core/worker/frontend/extract_time_worker.py +++ b/memoryscope/core/worker/frontend/extract_time_worker.py @@ -29,7 +29,7 @@ class ExtractTimeWorker(MemoryBaseWorker): The response is parsed for time-related data using regex, translated via a language-specific key map, and the resulting time data is stored in the shared context. """ - query, query_timestamp = self.get_context(QUERY_WITH_TS) + query, query_timestamp = self.get_workflow_context(QUERY_WITH_TS) # Identify if the query contains datetime keywords contain_datetime = DatetimeHandler.has_time_word(query, self.language) @@ -62,4 +62,4 @@ class ExtractTimeWorker(MemoryBaseWorker): if key in key_map.keys(): extract_time_dict[key_map[key]] = value self.logger.info(f"response_text={response_text} matches={matches} filters={extract_time_dict}") - self.set_context(EXTRACT_TIME_DICT, extract_time_dict) + self.set_workflow_context(EXTRACT_TIME_DICT, extract_time_dict) diff --git a/memoryscope/core/worker/frontend/fuse_rerank_worker.py b/memoryscope/core/worker/frontend/fuse_rerank_worker.py index 176f0167..b137f354 100644 --- a/memoryscope/core/worker/frontend/fuse_rerank_worker.py +++ b/memoryscope/core/worker/frontend/fuse_rerank_worker.py @@ -61,7 +61,7 @@ class FuseRerankWorker(MemoryBaseWorker): 5. Logs reranking details and formats the final list of memories for output. """ # Parse input parameters from the worker's context - extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT) + extract_time_dict: Dict[str, str] = self.get_workflow_context(EXTRACT_TIME_DICT) memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RANKED_MEMORY_NODES) # Check if memory nodes are available; warn and return if not @@ -106,4 +106,4 @@ class FuseRerankWorker(MemoryBaseWorker): memories.append(f"[{datetime} {weekday}] {node.content}") # Set the final list of formatted memories back into the worker's context - self.set_context(RESULT, "\n".join(memories)) + self.set_workflow_context(RESULT, "\n".join(memories)) diff --git a/memoryscope/core/worker/frontend/print_memory_worker.py b/memoryscope/core/worker/frontend/print_memory_worker.py index 00da5855..7421614d 100644 --- a/memoryscope/core/worker/frontend/print_memory_worker.py +++ b/memoryscope/core/worker/frontend/print_memory_worker.py @@ -63,4 +63,4 @@ class PrintMemoryWorker(MemoryBaseWorker): observation_memory="\n".join(observation_memory_list), insight_memory="\n".join(insight_memory_list), expired_memory="\n".join(expired_memory_list)).strip() - self.set_context(RESULT, result) + self.set_workflow_context(RESULT, result) diff --git a/memoryscope/core/worker/frontend/read_message_worker.py b/memoryscope/core/worker/frontend/read_message_worker.py index 62ad8b18..2f378eef 100644 --- a/memoryscope/core/worker/frontend/read_message_worker.py +++ b/memoryscope/core/worker/frontend/read_message_worker.py @@ -37,4 +37,4 @@ class ReadMessageWorker(MemoryBaseWorker): for messages in chat_messages_not_memorized[-contextual_msg_max_count:]: chat_message_scatter.extend(messages) chat_message_scatter.sort(key=lambda _: _.time_created) - self.set_context(RESULT, chat_message_scatter) + self.set_workflow_context(RESULT, chat_message_scatter) diff --git a/memoryscope/core/worker/frontend/retrieve_memory_worker.py b/memoryscope/core/worker/frontend/retrieve_memory_worker.py index c38ae4b3..a539c7f4 100644 --- a/memoryscope/core/worker/frontend/retrieve_memory_worker.py +++ b/memoryscope/core/worker/frontend/retrieve_memory_worker.py @@ -119,7 +119,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker): 6. Logs detailed information about each memory node. 7. Stores the processed memory nodes for further use. """ - query, _ = self.get_context(QUERY_WITH_TS) + query, _ = self.get_workflow_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) diff --git a/memoryscope/core/worker/frontend/semantic_rank_worker.py b/memoryscope/core/worker/frontend/semantic_rank_worker.py index 4cd8d8bd..894785fd 100644 --- a/memoryscope/core/worker/frontend/semantic_rank_worker.py +++ b/memoryscope/core/worker/frontend/semantic_rank_worker.py @@ -32,7 +32,7 @@ class SemanticRankWorker(MemoryBaseWorker): appropriate warnings are logged. """ # query - query, _ = self.get_context(QUERY_WITH_TS) + query, _ = self.get_workflow_context(QUERY_WITH_TS) 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!") diff --git a/memoryscope/core/worker/frontend/set_query_worker.py b/memoryscope/core/worker/frontend/set_query_worker.py index 0a586110..fe541c11 100644 --- a/memoryscope/core/worker/frontend/set_query_worker.py +++ b/memoryscope/core/worker/frontend/set_query_worker.py @@ -36,4 +36,4 @@ class SetQueryWorker(MemoryBaseWorker): timestamp = _timestamp # Store the determined query and its timestamp in the context - self.set_context(QUERY_WITH_TS, (query, timestamp)) + self.set_workflow_context(QUERY_WITH_TS, (query, timestamp)) diff --git a/memoryscope/core/worker/memory_base_worker.py b/memoryscope/core/worker/memory_base_worker.py index 6ca3f37a..1293a933 100644 --- a/memoryscope/core/worker/memory_base_worker.py +++ b/memoryscope/core/worker/memory_base_worker.py @@ -54,7 +54,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): Returns: List[Message]: List of chat messages. """ - return self.get_context(CHAT_MESSAGES) + return self.get_workflow_context(CHAT_MESSAGES) @property def chat_messages_scatter(self) -> List[Message]: @@ -64,7 +64,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): Returns: List[Message]: List of chat messages. """ - result = self.get_context(CHAT_MESSAGES_SCATTER) + result = self.get_workflow_context(CHAT_MESSAGES_SCATTER) if not result: if isinstance(self.chat_messages[0], list): @@ -73,13 +73,13 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): if messages: chat_messages.extend(messages) chat_messages.sort(key=lambda _: _.time_created) - self.set_context(CHAT_MESSAGES_SCATTER, chat_messages) + self.set_workflow_context(CHAT_MESSAGES_SCATTER, chat_messages) else: assert isinstance(self.chat_messages[0], Message) - self.set_context(CHAT_MESSAGES_SCATTER, self.chat_messages) + self.set_workflow_context(CHAT_MESSAGES_SCATTER, self.chat_messages) - return self.get_context(CHAT_MESSAGES_SCATTER) + return self.get_workflow_context(CHAT_MESSAGES_SCATTER) @chat_messages_scatter.setter def chat_messages_scatter(self, value: List[Message]): @@ -87,7 +87,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): Set the chat messages with the new value. """ - self.set_context(CHAT_MESSAGES_SCATTER, value) + self.set_workflow_context(CHAT_MESSAGES_SCATTER, value) @property def chat_kwargs(self) -> Dict[str, Any]: @@ -100,19 +100,19 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): Returns: Dict[str, str]: A dictionary containing the chat keyword arguments. """ - return self.get_context(CHAT_KWARGS) + return self.get_workflow_context(CHAT_KWARGS) @property def user_name(self) -> str: - return self.get_context(USER_NAME) + return self.get_workflow_context(USER_NAME) @property def target_name(self) -> str: - return self.get_context(TARGET_NAME) + return self.get_workflow_context(TARGET_NAME) @property def workflow_name(self) -> str: - return self.get_context(WORKFLOW_NAME) + return self.get_workflow_context(WORKFLOW_NAME) @property def language(self) -> LanguageEnum: @@ -204,8 +204,8 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): MemoryHandler: An instance of MemoryHandler. """ if not self.has_content(MEMORY_MANAGER): - self.set_context(MEMORY_MANAGER, MemoryManager(self.memoryscope_context, worker_name=self.name)) - return self.get_context(MEMORY_MANAGER) + self.set_workflow_context(MEMORY_MANAGER, MemoryManager(self.memoryscope_context, workerflow_name=self.workflow_name)) + return self.get_workflow_context(MEMORY_MANAGER) def get_language_value(self, languages: dict | List[dict]) -> Any | List[Any]: """ diff --git a/memoryscope/core/worker/memory_manager.py b/memoryscope/core/worker/memory_manager.py index 601b567b..bf053a67 100644 --- a/memoryscope/core/worker/memory_manager.py +++ b/memoryscope/core/worker/memory_manager.py @@ -13,7 +13,7 @@ class MemoryManager(object): The `MemoryHandler` class manages memory nodes with memory store. """ - def __init__(self, memoryscope_context: MemoryscopeContext, worker_name: str ="default_worker"): + def __init__(self, memoryscope_context: MemoryscopeContext, workerflow_name: str ="default_worker"): self.memoryscope_context: MemoryscopeContext = memoryscope_context self._memory_store: BaseMemoryStore | None = None @@ -26,7 +26,7 @@ class MemoryManager(object): self.logger = Logger.get_logger(Logger.append_timestamp("memory_manager")) - self.worker_name = worker_name + self.workerflow_name = workerflow_name @property @@ -101,7 +101,7 @@ class MemoryManager(object): if nodes: self.logger.info( self.logger.wrap_in_box( - '\n'.join([f"worker_name: {self.worker_name} | memory_type:{node.memory_type} | content:{node.content}" for node in nodes]) + '\n'.join([f"workerflow_name: {self.workerflow_name} | memory_type:{node.memory_type} | content:{node.content}" for node in nodes]) ) ) diff --git a/tests/worker/test_workers_cn.py b/tests/worker/test_workers_cn.py index bb6cd49c..b7f23f2d 100644 --- a/tests/worker/test_workers_cn.py +++ b/tests/worker/test_workers_cn.py @@ -52,10 +52,10 @@ class TestWorkersCn(unittest.TestCase): query = "明天我去上海出差" query_timestamp = int(datetime.datetime.now().timestamp()) - worker.set_context(QUERY_WITH_TS, (query, query_timestamp)) + worker.set_workflow_context(QUERY_WITH_TS, (query, query_timestamp)) worker.run() - result = worker.get_context(EXTRACT_TIME_DICT) + result = worker.get_workflow_context(EXTRACT_TIME_DICT) worker.logger.info(f"result={result}") # @unittest.skip @@ -85,7 +85,7 @@ class TestWorkersCn(unittest.TestCase): role_name=self.arguments.human_name), ] - worker.set_context(CHAT_MESSAGES_SCATTER, chat_messages) + worker.set_workflow_context(CHAT_MESSAGES_SCATTER, chat_messages) worker.run() result = [msg.content for msg in worker.chat_messages_scatter] @@ -133,7 +133,7 @@ class TestWorkersCn(unittest.TestCase): role_name=self.arguments.human_name), ] - worker.set_context(CHAT_MESSAGES_SCATTER, chat_messages) + worker.set_workflow_context(CHAT_MESSAGES_SCATTER, chat_messages) worker.run() result = [msg.content for msg in worker.chat_messages_scatter] @@ -167,7 +167,7 @@ class TestWorkersCn(unittest.TestCase): # Message(role=MessageRoleEnum.USER.value, content="我在一家叫京东的公司干活"), # ] - worker.set_context(CHAT_MESSAGES_SCATTER, chat_messages) + worker.set_workflow_context(CHAT_MESSAGES_SCATTER, chat_messages) worker.run() result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] @@ -198,7 +198,7 @@ class TestWorkersCn(unittest.TestCase): Message(role=MessageRoleEnum.USER.value, content="最后一个问题,你知道怎么才能维持广泛的社交关系吗?", role_name=self.arguments.human_name), ] - worker.set_context(CHAT_MESSAGES_SCATTER, chat_messages) + worker.set_workflow_context(CHAT_MESSAGES_SCATTER, chat_messages) worker.run() result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] @@ -227,7 +227,7 @@ class TestWorkersCn(unittest.TestCase): Message(role=MessageRoleEnum.USER.value, content="明天是我生日", role_name=self.arguments.human_name), ] - worker.set_context(CHAT_MESSAGES_SCATTER, chat_messages) + worker.set_workflow_context(CHAT_MESSAGES_SCATTER, chat_messages) worker.run() result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_WITH_TIME_NODES)] diff --git a/tests/worker/test_workers_en.py b/tests/worker/test_workers_en.py index 6616afc0..92f81cd3 100644 --- a/tests/worker/test_workers_en.py +++ b/tests/worker/test_workers_en.py @@ -50,10 +50,10 @@ class TestWorkersEn(unittest.TestCase): query = "I will be on a business trip to Shanghai tomorrow." query_timestamp = int(datetime.datetime.now().timestamp()) - worker.set_context(QUERY_WITH_TS, (query, query_timestamp)) + worker.set_workflow_context(QUERY_WITH_TS, (query, query_timestamp)) worker.run() - result = worker.get_context(EXTRACT_TIME_DICT) + result = worker.get_workflow_context(EXTRACT_TIME_DICT) worker.logger.info(f"result={result}") @unittest.skip @@ -75,7 +75,7 @@ class TestWorkersEn(unittest.TestCase): content="I'm going to take the college entrance examination tomorrow."), ] - worker.set_context(CHAT_MESSAGES, chat_messages) + worker.set_workflow_context(CHAT_MESSAGES, chat_messages) worker.run() result = [msg.content for msg in worker.chat_messages_scatter] @@ -123,7 +123,7 @@ class TestWorkersEn(unittest.TestCase): content="Last question, do you know how to maintain extensive social relationships?"), ] - worker.set_context(CHAT_MESSAGES, chat_messages) + worker.set_workflow_context(CHAT_MESSAGES, chat_messages) worker.run() result = [msg.content for msg in worker.chat_messages_scatter] @@ -152,7 +152,7 @@ class TestWorkersEn(unittest.TestCase): Message(role=MessageRoleEnum.USER.value, content="I work for a company called JD.com"), ] - worker.set_context(CHAT_MESSAGES, chat_messages) + worker.set_workflow_context(CHAT_MESSAGES, chat_messages) worker.run() result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] @@ -194,7 +194,7 @@ class TestWorkersEn(unittest.TestCase): content="Last question, do you know how to maintain extensive social relationships?"), ] - worker.set_context(CHAT_MESSAGES, chat_messages) + worker.set_workflow_context(CHAT_MESSAGES, chat_messages) worker.run() result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)] @@ -226,7 +226,7 @@ class TestWorkersEn(unittest.TestCase): Message(role=MessageRoleEnum.USER.value, content="Tomorrow is my birthday."), ] - worker.set_context(CHAT_MESSAGES, chat_messages) + worker.set_workflow_context(CHAT_MESSAGES, chat_messages) worker.run() result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_WITH_TIME_NODES)]