mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-09 22:31:05 +00:00
update es store logging system
This commit is contained in:
commit
d0a9ea729d
30 changed files with 230 additions and 76 deletions
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -143,6 +143,7 @@ docs/sphinx_doc/build/
|
|||
*runs/
|
||||
memoryscope.db
|
||||
tmp*.json
|
||||
tmp*.py
|
||||
cradle*
|
||||
|
||||
# sphinx docs
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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,16 @@ 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 = ""
|
||||
|
||||
print_workflow_dynamic: bool = False
|
||||
|
||||
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.")
|
||||
|
|
|
|||
|
|
@ -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,10 +34,11 @@ 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
|
||||
self.logger = Logger.get_logger()
|
||||
self.logger = Logger.get_logger(Logger.append_timestamp("base_model"))
|
||||
|
||||
@property
|
||||
def model(self):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -28,15 +28,20 @@ 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()
|
||||
self.logger: Logger = Logger.get_logger(Logger.append_timestamp("workflow"))
|
||||
|
||||
if self.workflow:
|
||||
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.
|
||||
|
|
@ -132,11 +137,12 @@ 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,
|
||||
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,
|
||||
|
|
@ -165,21 +171,31 @@ class BaseWorkflow(object):
|
|||
**kwargs: Additional keyword arguments to be passed to context.
|
||||
"""
|
||||
with Timer(f"workflow.{self.name}", time_log_type="wrap"):
|
||||
self.context.clear()
|
||||
|
||||
self.context.update({WORKFLOW_NAME: self.name, **kwargs})
|
||||
|
||||
log_buf = f"Operation: {self.name}"
|
||||
self.logger.info(log_buf)
|
||||
self.workflow_print_console(log_buf, style="bold red")
|
||||
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 workflow_part in self.workflow_worker_list:
|
||||
for index, workflow_part in enumerate(self.workflow_worker_list):
|
||||
# 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]}"
|
||||
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
|
||||
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):
|
||||
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)
|
||||
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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -1,12 +1,21 @@
|
|||
import logging
|
||||
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
|
||||
|
||||
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):
|
||||
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):
|
||||
"""
|
||||
|
|
@ -63,6 +72,49 @@ 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):
|
||||
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 = []
|
||||
buf.append('\n')
|
||||
buf.append(f"LM Input:\n")
|
||||
for chat_message in message.meta_data['data']['messages']:
|
||||
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 = []
|
||||
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.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):
|
||||
"""
|
||||
Adds a file handler to the logger which logs messages to a rotating file.
|
||||
|
|
@ -76,6 +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")
|
||||
|
||||
# Instantiate a rotating file handler with specified parameters
|
||||
file_handler = RotatingFileHandler(
|
||||
|
|
@ -153,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
|
||||
|
|
@ -182,3 +235,8 @@ class Logger(logging.Logger):
|
|||
LOGGER_DICT[name] = Logger(name=name, **kwargs)
|
||||
|
||||
return LOGGER_DICT[name]
|
||||
|
||||
@staticmethod
|
||||
def append_timestamp(name: str) -> str:
|
||||
from memoryscope.core.memoryscope_context import get_ms_context
|
||||
return f"{name}_{get_ms_context().memory_scope_uuid}"
|
||||
9
memoryscope/core/utils/singleton.py
Normal file
9
memoryscope/core/utils/singleton.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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!")
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
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]:
|
||||
"""
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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, workerflow_name: str ="default_worker"):
|
||||
self.memoryscope_context: MemoryscopeContext = memoryscope_context
|
||||
|
||||
self._memory_store: BaseMemoryStore | None = None
|
||||
|
|
@ -24,7 +24,10 @@ 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"))
|
||||
|
||||
self.workerflow_name = workerflow_name
|
||||
|
||||
|
||||
@property
|
||||
def memory_store(self) -> BaseMemoryStore:
|
||||
|
|
@ -95,7 +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._key_id_dict[key] = [n.memory_id for n in nodes]
|
||||
if nodes:
|
||||
self.logger.info(
|
||||
self.logger.wrap_in_box(
|
||||
'\n'.join([f"workerflow_name: {self.workerflow_name} | memory_type:{node.memory_type} | content:{node.content}" for node in nodes])
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def get_memories(self, keys: str | List[str]) -> List[MemoryNode]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -14,4 +14,5 @@ dashscope~=1.19.1
|
|||
elasticsearch~=8.14.0
|
||||
pyyaml~=6.0.1
|
||||
ray~=2.31.0
|
||||
numpy~=1.26.4
|
||||
numpy~=1.26.4
|
||||
rich
|
||||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue