Refactor code to use workflow context consistently

This commit is contained in:
青轩 2024-08-14 11:34:32 +08:00
parent 39b23c178c
commit ceffb56a53
19 changed files with 78 additions and 82 deletions

View file

@ -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]}"

View file

@ -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:

View file

@ -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)

View file

@ -15,12 +15,6 @@ from .tool_functions import (
cosine_similarity
)
def get_context():
from memoryscope import MemoryscopeContext
return MemoryscopeContext()
__all__ = [
"DatetimeHandler",
"Logger",

View file

@ -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):
"""

View file

@ -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))

View file

@ -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

View file

@ -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}")

View file

@ -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)

View file

@ -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))

View file

@ -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)

View file

@ -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)

View file

@ -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)

View file

@ -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!")

View file

@ -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))

View file

@ -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]:
"""

View file

@ -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])
)
)

View file

@ -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)]

View file

@ -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)]