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