[dev] rename MEMORY_HANDLER to MEMORY_MANAGER

This commit is contained in:
jinli.yl 2024-07-27 01:49:20 +08:00
parent 6aaa26742e
commit 590cbf65c4
23 changed files with 177 additions and 161 deletions

View file

@ -5,11 +5,13 @@
WORKFLOW_NAME = "workflow_name"
MEMORYSCOPE_CONTEXT = "memoryscope_context"
RESULT = "result"
CHAT_MESSAGES = "chat_messages"
MEMORY_HANDLER = "memory_handler"
MEMORY_MANAGER = "memory_manager"
CHAT_KWARGS = "chat_kwargs"

View file

@ -1,7 +1,6 @@
import time
from typing import List
from memoryscope.constants.common_constants import CHAT_KWARGS, RESULT, CHAT_MESSAGES
from memoryscope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE
from memoryscope.memory.operation.base_workflow import BaseWorkflow
@ -54,17 +53,14 @@ class BackendOperation(BaseWorkflow, BaseOperation):
Returns:
Any: The result obtained after executing the workflow.
"""
self.context.clear()
# Add additional arguments to the context
kwargs.update(**self.kwargs)
self.context[CHAT_KWARGS] = kwargs
# Include the most recent messages in the operation context
self.context[CHAT_MESSAGES] = self.chat_messages
# prepare kwargs
workflow_kwargs = {
CHAT_MESSAGES: self.chat_messages,
CHAT_KWARGS: {**kwargs, **self.kwargs},
}
# Execute the workflow with the prepared context
self.run_workflow()
self.run_workflow(**workflow_kwargs)
# Retrieve the result from the context after workflow execution
return self.context.get(RESULT)
@ -114,8 +110,8 @@ class BackendOperation(BaseWorkflow, BaseOperation):
"""
if not self._loop_switch:
self._loop_switch = True
self._backend_task = G_CONTEXT.thread_pool.submit(self._loop_operation)
self.logger.info(f"start operation={operation.name}...")
self._backend_task = self.thread_pool.submit(self._loop_operation)
self.logger.info(f"start operation={self.name}...")
def stop_operation_backend(self, wait_task_end: bool = False):
"""
@ -128,5 +124,3 @@ class BackendOperation(BaseWorkflow, BaseOperation):
self.logger.info(f"stop operation={self.name}...")
else:
self.logger.info(f"send stop signal to operation={self.name}...")

View file

@ -4,7 +4,7 @@ from concurrent.futures import ThreadPoolExecutor, as_completed
from itertools import zip_longest
from typing import Dict, Any, List
from memoryscope.constants.common_constants import WORKFLOW_NAME
from memoryscope.constants.common_constants import WORKFLOW_NAME, MEMORYSCOPE_CONTEXT
from memoryscope.memory.worker.base_worker import BaseWorker
from memoryscope.memoryscope_context import MemoryscopeContext
from memoryscope.utils.logger import Logger
@ -22,6 +22,7 @@ class BaseWorkflow(object):
self.name: str = name
self.memoryscope_context: MemoryscopeContext = memoryscope_context
self.thread_pool: ThreadPoolExecutor = self.memoryscope_context.thread_pool
self.workflow: str = workflow
self.kwargs = kwargs
@ -133,12 +134,11 @@ class BaseWorkflow(object):
self.worker_dict[name] = init_instance_by_config(
config=self.memoryscope_context.worker_conf_dict[name],
suffix_name="worker",
name=name,
is_multi_thread=is_backend or self.worker_dict[name],
context=self.context,
context_lock=self.context_lock,
thread_pool=self.memoryscope_context.thread_pool,
thread_pool=self.thread_pool,
**kwargs)
def _run_sub_workflow(self, worker_list: List[str]) -> bool:
@ -150,7 +150,7 @@ class BaseWorkflow(object):
return False
return True
def run_workflow(self):
def run_workflow(self, **kwargs):
"""
Executes the workflow by orchestrating the steps defined in `self.workflow_worker_list`.
This method supports both sequential and parallel execution of sub-workflows based on the structure
@ -159,9 +159,18 @@ class BaseWorkflow(object):
If a workflow part consists of a single item, it is executed sequentially. For parts with multiple items,
they are submitted for parallel execution using a thread pool. The workflow will stop if any sub-workflow
returns False.
Args:
**kwargs: Additional keyword arguments to be passed to context.
"""
with Timer(f"workflow.{self.name}", time_log_type="wrap"):
self.context[WORKFLOW_NAME] = self.name
self.context.clear()
self.context.update({
WORKFLOW_NAME: self.name,
MEMORYSCOPE_CONTEXT: self.memoryscope_context,
**kwargs,
})
# Iterate over each part of the workflow
for workflow_part in self.workflow_worker_list:
@ -174,7 +183,7 @@ class BaseWorkflow(object):
t_list = []
# Submit tasks to the thread pool
for sub_workflow in workflow_part:
t_list.append(G_CONTEXT.thread_pool.submit(self._run_sub_workflow, sub_workflow))
t_list.append(self.thread_pool.submit(self._run_sub_workflow, sub_workflow))
# Check results; if any task returns False, stop the workflow
flag = True

View file

@ -3,10 +3,10 @@ from memoryscope.enumeration.message_role_enum import MessageRoleEnum
from memoryscope.memory.operation.backend_operation import BackendOperation
class SummaryObservationOp(BackendOperation):
class ConsolidateOperation(BackendOperation):
def __init__(self, **kwargs):
super(SummaryObservationOp, self).__init__(**kwargs)
super(ConsolidateOperation, self).__init__(**kwargs)
self.message_lock = kwargs.get("message_lock", None)
self.contextual_msg_min_count: int = kwargs.get("contextual_msg_min_count", 0)
@ -43,17 +43,14 @@ class SummaryObservationOp(BackendOperation):
f"contextual_msg_min_count({self.contextual_msg_min_count}), skip.")
return
self.context.clear()
# Add additional arguments to the context
kwargs.update(**self.kwargs)
self.context[CHAT_KWARGS] = kwargs
# Include the most recent messages in the operation context
self.context[CHAT_MESSAGES] = chat_messages
# prepare kwargs
workflow_kwargs = {
CHAT_MESSAGES: chat_messages,
CHAT_KWARGS: {**kwargs, **self.kwargs},
}
# Execute the workflow with the prepared context
self.run_workflow()
self.run_workflow(**workflow_kwargs)
# Retrieve the result from the context after workflow execution
result = self.context.get(RESULT)

View file

@ -39,17 +39,15 @@ class FrontendOperation(BaseWorkflow, BaseOperation):
Returns:
Any: The result obtained from executing the workflow.
"""
self.context.clear()
# Include the most recent messages in the operation context
self.context[CHAT_MESSAGES] = self.chat_messages
# Add additional arguments to the context
kwargs.update(**self.kwargs)
self.context[CHAT_KWARGS] = kwargs
# prepare kwargs
workflow_kwargs = {
CHAT_MESSAGES: self.chat_messages,
CHAT_KWARGS: {**kwargs, **self.kwargs},
}
# Execute the workflow with the prepared context
self.run_workflow()
self.run_workflow(**workflow_kwargs)
# Retrieve the result from the context after workflow execution
return self.context.get(RESULT)

View file

@ -76,7 +76,7 @@ class MemoryScopeService(BaseMemoryService):
name=name,
chat_messages=self.chat_messages,
message_lock=self.message_lock,
context=self.context,
memoryscope_context=self.context,
contextual_msg_max_count=self.contextual_msg_max_count,
contextual_msg_min_count=self.contextual_msg_min_count)

View file

@ -43,13 +43,13 @@ class ContraRepeatWorker(MemoryBaseWorker):
6. Updates the status of nodes accordingly.
7. Persists the changes back to memory storage.
"""
all_obs_nodes: List[MemoryNode] = self.memory_handler.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES])
all_obs_nodes: List[MemoryNode] = self.memory_manager.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES])
if not all_obs_nodes:
self.logger.info("all_obs_nodes is empty!")
# self.continue_run = False
return
today_obs_nodes: List[MemoryNode] = self.memory_handler.get_memories(TODAY_NODES)
today_obs_nodes: List[MemoryNode] = self.memory_manager.get_memories(TODAY_NODES)
if today_obs_nodes:
all_obs_nodes.extend(today_obs_nodes)
@ -121,4 +121,4 @@ class ContraRepeatWorker(MemoryBaseWorker):
merge_obs_nodes.append(node)
# save context
self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes, log_repeat=False)
self.memory_manager.set_memories(MERGE_OBS_NODES, merge_obs_nodes, log_repeat=False)

View file

@ -42,11 +42,11 @@ class GetObservationWorker(MemoryBaseWorker):
MemoryTypeEnum.CONVERSATION.value: message.content,
TIME_INFER: time_infer,
"keywords": keywords,
**{k: str(v) for k, v in dt_handler.dt_info_dict.items()},
**{k: str(v) for k, v in dt_handler.get_dt_info_dict(self.language).items()},
}
if time_infer:
dt_info_dict = DatetimeHandler.extract_date_parts(input_string=time_infer)
dt_info_dict = DatetimeHandler.extract_date_parts(input_string=time_infer, language=self.language)
meta_data.update({f"event_{k}": str(v) for k, v in dt_info_dict.items()})
obs_content = (f"{obs_content} ({self.get_language_value(TIME_INFER_WORD)}"
f"{self.get_language_value(COLON_WORD)} {time_infer})")
@ -68,7 +68,7 @@ class GetObservationWorker(MemoryBaseWorker):
"""
filter_messages = []
for msg in self.chat_messages:
if not DatetimeHandler.has_time_word(query=msg.content):
if not DatetimeHandler.has_time_word(query=msg.content, language=self.language):
filter_messages.append(msg)
self.logger.info(f"after filter_messages.size from {len(self.chat_messages)} to {len(filter_messages)}")
@ -184,4 +184,4 @@ class GetObservationWorker(MemoryBaseWorker):
keywords=keywords))
# Stores the extracted and structured observations in the conversation memory
self.memory_handler.set_memories(self.OBS_STORE_KEY, new_obs_nodes)
self.memory_manager.set_memories(self.OBS_STORE_KEY, new_obs_nodes)

View file

@ -36,7 +36,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
"""
dt_handler = DatetimeHandler()
# Prepare metadata with current datetime info
meta_data = {k: str(v) for k, v in dt_handler.dt_info_dict.items()}
meta_data = {k: str(v) for k, v in dt_handler.get_dt_info_dict.items()}
return MemoryNode(user_name=self.user_name,
target_name=self.target_name,
@ -58,8 +58,8 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
- Parsing the model's responses for new insight keys.
- Creating new insight nodes and updating the memory status accordingly.
"""
not_reflected_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_REFLECTED_NODES)
insight_nodes: List[MemoryNode] = self.memory_handler.get_memories(INSIGHT_NODES)
not_reflected_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_REFLECTED_NODES)
insight_nodes: List[MemoryNode] = self.memory_manager.get_memories(INSIGHT_NODES)
# Count unaudited nodes
not_reflected_count = len(not_reflected_nodes)
@ -104,7 +104,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
new_insight_keys = ResponseTextParser(response.message.content, self.__class__.__name__).parse_v2()
if new_insight_keys:
for insight_key in new_insight_keys:
self.memory_handler.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key))
self.memory_manager.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key))
# Mark unaudited nodes as reflected
for node in not_reflected_nodes:

View file

@ -33,7 +33,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
}
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_not_reflected_top_k,
filter_dict=filter_dict)
self.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes)
self.memory_manager.set_memories(NOT_REFLECTED_NODES, nodes)
@timer
def retrieve_not_updated_memory(self):
@ -52,7 +52,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
}
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_not_updated_top_k,
filter_dict=filter_dict)
self.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
self.memory_manager.set_memories(NOT_UPDATED_NODES, nodes)
@timer
def retrieve_insight_memory(self):
@ -70,7 +70,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
}
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_insight_top_k,
filter_dict=filter_dict)
self.memory_handler.set_memories(INSIGHT_NODES, nodes)
self.memory_manager.set_memories(INSIGHT_NODES, nodes)
@timer
def retrieve_today_memory(self, dt: str):
@ -93,7 +93,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_today_top_k,
filter_dict=filter_dict)
self.memory_handler.set_memories(TODAY_NODES, nodes)
self.memory_manager.set_memories(TODAY_NODES, nodes)
def _run(self):
"""

View file

@ -63,7 +63,7 @@ class LongContraRepeatWorker(MemoryBaseWorker):
The process helps in maintaining conversation coherence by resolving contradictions and redundancies.
"""
not_updated_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_UPDATED_NODES)
not_updated_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_UPDATED_NODES)
for node in not_updated_nodes:
self.submit_thread_task(fn=self.retrieve_similar_content, node=node)
@ -157,4 +157,4 @@ class LongContraRepeatWorker(MemoryBaseWorker):
f"action_status={node.action_status}")
# save context
self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes)
self.memory_manager.set_memories(MERGE_OBS_NODES, merge_obs_nodes)

View file

@ -95,7 +95,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
content = f"{key}{self.get_language_value(COLON_WORD)} {insight_value}"
insight_node.content = content
insight_node.value = insight_value
insight_node.meta_data.update({k: str(v) for k, v in dt_handler.dt_info_dict.items()})
insight_node.meta_data.update({k: str(v) for k, v in dt_handler.get_dt_info_dict.items()})
insight_node.timestamp = dt_handler.timestamp
insight_node.dt = dt_handler.datetime_format()
if insight_node.action_status == ActionStatusEnum.NONE.value:
@ -175,9 +175,9 @@ class UpdateInsightWorker(MemoryBaseWorker):
6. Gather the results of all update tasks.
7. Mark processed nodes as updated in memory.
"""
insight_nodes: List[MemoryNode] = self.memory_handler.get_memories(INSIGHT_NODES)
not_updated_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_UPDATED_NODES)
not_reflected_nodes: List[MemoryNode] = self.memory_handler.get_memories(keys=[NOT_REFLECTED_NODES,
insight_nodes: List[MemoryNode] = self.memory_manager.get_memories(INSIGHT_NODES)
not_updated_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_UPDATED_NODES)
not_reflected_nodes: List[MemoryNode] = self.memory_manager.get_memories(keys=[NOT_REFLECTED_NODES,
NOT_UPDATED_NODES])
if not insight_nodes:
@ -216,7 +216,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
# delete empty nodes
empty_nodes = [n for n in insight_nodes if not n.content.strip()]
self.memory_handler.delete_memories(empty_nodes)
self.memory_manager.delete_memories(empty_nodes)
for node in not_updated_nodes:
node.obs_updated = 1

View file

@ -46,7 +46,7 @@ class UpdateMemoryWorker(MemoryBaseWorker):
if not self.memory_key:
return
return self.memory_handler.get_memories(keys=self.memory_key)
return self.memory_manager.get_memories(keys=self.memory_key)
def delete_all(self):
"""
@ -55,7 +55,7 @@ class UpdateMemoryWorker(MemoryBaseWorker):
Returns:
List[MemoryNode]: A list of all MemoryNode objects marked for deletion.
"""
nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all")
nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all")
for node in nodes:
node.action_status = ActionStatusEnum.DELETE.value
self.logger.info(f"delete_all.size={len(nodes)}")
@ -74,7 +74,7 @@ class UpdateMemoryWorker(MemoryBaseWorker):
return
i = 0
nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all")
nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all")
for node in nodes:
if node.content == query:
i += 1
@ -88,7 +88,7 @@ class UpdateMemoryWorker(MemoryBaseWorker):
return
i = 0
nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all")
nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all")
for node in nodes:
if node.memory_id == memory_id:
i += 1
@ -109,4 +109,4 @@ class UpdateMemoryWorker(MemoryBaseWorker):
if not hasattr(self, method):
self.logger.info(f"method={method} is missing!")
return
self.memory_handler.update_memories(nodes=getattr(self, method)())
self.memory_manager.update_memories(nodes=getattr(self, method)())

View file

@ -62,7 +62,7 @@ class FuseRerankWorker(MemoryBaseWorker):
"""
# Parse input parameters from the worker's context
extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT)
memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RANKED_MEMORY_NODES)
memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RANKED_MEMORY_NODES)
# Check if memory nodes are available; warn and return if not
if not memory_node_list:

View file

@ -22,7 +22,7 @@ class PrintMemoryWorker(MemoryBaseWorker):
3. Set the formatted string back into the worker's context
"""
# get long-term memory
memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RETRIEVE_MEMORY_NODES)
memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RETRIEVE_MEMORY_NODES)
memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True)
observation_memory_list: List[str] = []

View file

@ -139,4 +139,4 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
self.logger.info(f"recall_stage: content={node.content} score={node.score_rerank} type={node.memory_type} "
f"store_status={node.store_status} action_status={node.action_status}")
self.memory_handler.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list)
self.memory_manager.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list)

View file

@ -29,7 +29,7 @@ class SemanticRankWorker(MemoryBaseWorker):
"""
# query
query, _ = self.get_context(QUERY_WITH_TS)
memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RETRIEVE_MEMORY_NODES)
memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RETRIEVE_MEMORY_NODES)
if not memory_node_list:
self.logger.warning("Retrieve memory nodes is empty!")
return
@ -58,4 +58,4 @@ class SemanticRankWorker(MemoryBaseWorker):
self.logger.info(f"Rank stage: Content={node.content}, Score={node.score_rank}")
# save ranked nodes back to memory
self.memory_handler.set_memories(RANKED_MEMORY_NODES, memory_node_list, log_repeat=False)
self.memory_manager.set_memories(RANKED_MEMORY_NODES, memory_node_list, log_repeat=False)

View file

@ -1,14 +1,16 @@
from abc import ABCMeta
from typing import List, Dict, Any
from memoryscope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, MEMORY_HANDLER
from memoryscope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, MEMORYSCOPE_CONTEXT, \
WORKFLOW_NAME, MEMORY_MANAGER
from memoryscope.enumeration.language_enum import LanguageEnum
from memoryscope.memory.worker.base_worker import BaseWorker
from memoryscope.memory.worker.memory_manager import MemoryManager
from memoryscope.memoryscope_context import MemoryscopeContext
from memoryscope.models.base_model import BaseModel
from memoryscope.scheme.message import Message
from memoryscope.storage.base_memory_store import BaseMemoryStore
from memoryscope.storage.base_monitor import BaseMonitor
from memoryscope.utils.global_context import G_CONTEXT
from memoryscope.utils.memory_handler import MemoryHandler
from memoryscope.utils.prompt_handler import PromptHandler
@ -77,6 +79,18 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
"""
return self.get_context(CHAT_KWARGS)
@property
def workflow_name(self) -> str:
return self.get_context(WORKFLOW_NAME)
@property
def memoryscope_context(self) -> MemoryscopeContext:
return self.get_context(MEMORYSCOPE_CONTEXT)
@property
def language(self) -> LanguageEnum:
return self.memoryscope_context.language
@property
def embedding_model(self) -> BaseModel:
"""
@ -87,8 +101,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
BaseModel: The embedding model used for converting text into vector representations.
"""
if isinstance(self._embedding_model, str):
self._embedding_model = G_CONTEXT.model_conf_dict[self._embedding_model]
# ⭐ Retrieve the actual model instance when the attribute is a string reference
self._embedding_model = self.memoryscope_context.model_dict[self._embedding_model]
return self._embedding_model
@property
@ -101,8 +114,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
BaseModel: The model used for text generation.
"""
if isinstance(self._generation_model, str):
self._generation_model = G_CONTEXT.model_conf_dict[self._generation_model]
# ⭐ Retrieve the model instance if currently a string reference
self._generation_model = self.memoryscope_context.model_dict[self._generation_model]
return self._generation_model
@property
@ -115,7 +127,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
BaseModel: The rank model instance used for ranking tasks.
"""
if isinstance(self._rank_model, str):
self._rank_model = G_CONTEXT.model_conf_dict[self._rank_model] # Fetch model instance if string reference
self._rank_model = self.memoryscope_context.model_dict[self._rank_model]
return self._rank_model
@property
@ -128,7 +140,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
BaseMemoryStore: The memory store instance used for inserting, updating, retrieving and deleting operations.
"""
if self._memory_store is None:
self._memory_store = G_CONTEXT.memory_store_conf
self._memory_store = self.memoryscope_context.memory_store
return self._memory_store
@property
@ -141,7 +153,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
BaseMonitor: The monitoring component instance.
"""
if self._monitor is None:
self._monitor = G_CONTEXT.monitor_conf
self._monitor = self.memoryscope_context.monitor
return self._monitor
@property
@ -154,7 +166,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
str: The name of the assistant.
"""
if self._user_name is None:
self._user_name = G_CONTEXT.meta_data["assistant_name"]
self._user_name = self.memoryscope_context.meta_data["assistant_name"]
return self._user_name
@property
@ -166,7 +178,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
str: The readable name of the human.
"""
if self._target_name is None:
self._target_name = G_CONTEXT.meta_data["human_name"]
self._target_name = self.memoryscope_context.meta_data["human_name"]
return self._target_name
@property
@ -182,19 +194,18 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
return self._prompt_handler
@property
def memory_handler(self) -> MemoryHandler:
def memory_manager(self) -> MemoryManager:
"""
Lazily initializes and returns the MemoryHandler instance.
Returns:
MemoryHandler: An instance of MemoryHandler.
"""
if not self.has_content(MEMORY_HANDLER):
self.set_context(MEMORY_HANDLER, MemoryHandler()) # Initialize the memory handler if not present
return self.get_context(MEMORY_HANDLER)
if not self.has_content(MEMORY_MANAGER):
self.set_context(MEMORY_MANAGER, MemoryManager(self.memoryscope_context))
return self.get_context(MEMORY_MANAGER)
@staticmethod
def get_language_value(languages: dict | List[dict]) -> Any | List[Any]:
def get_language_value(self, languages: dict | List[dict]) -> Any | List[Any]:
"""
Retrieves the value(s) corresponding to the current language context.
@ -205,5 +216,5 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
Any | list[Any]: The value or list of values matching the current language setting.
"""
if isinstance(languages, list):
return [x[G_CONTEXT.language] for x in languages]
return languages[G_CONTEXT.language]
return [x[self.language] for x in languages]
return languages[self.language]

View file

@ -2,21 +2,20 @@ from typing import Dict, List
from memoryscope.enumeration.action_status_enum import ActionStatusEnum
from memoryscope.enumeration.store_status_enum import StoreStatusEnum
from memoryscope.memoryscope_context import MemoryscopeContext
from memoryscope.scheme.memory_node import MemoryNode
from memoryscope.storage.base_memory_store import BaseMemoryStore
from memoryscope.utils.global_context import G_CONTEXT
from memoryscope.utils.logger import Logger
class MemoryHandler(object):
class MemoryManager(object):
"""
The `MemoryHandler` class manages memory nodes with memory store.
"""
def __init__(self):
"""
Initializes the MemoryHandler.
"""
def __init__(self, memoryscope_context: MemoryscopeContext):
self.memoryscope_context: MemoryscopeContext = memoryscope_context
self._memory_store: BaseMemoryStore | None = None
# dict: memory_id -> MemoryNode
@ -36,7 +35,7 @@ class MemoryHandler(object):
BaseMemoryStore: The memory store instance associated with this worker.
"""
if self._memory_store is None:
self._memory_store = G_CONTEXT.memory_store_conf
self._memory_store = self.memoryscope_context.memory_store
return self._memory_store
def clear(self):

View file

@ -20,9 +20,6 @@ class DummyGenerationModel(BaseModel):
m_type: ModelEnum = ModelEnum.GENERATION_MODEL
class DummyModel:
"""
An inner class representing the dummy model placeholder.
"""
pass
MODEL_REGISTRY.register("dummy_generation", DummyModel)
@ -79,12 +76,12 @@ class DummyGenerationModel(BaseModel):
for delta in call_result:
model_response.message.content += delta
model_response.delta = delta
time.sleep(0.1) # ⭐ Introduce a delay to simulate streaming
time.sleep(0.1)
yield model_response
return gen()
else:
model_response.message.content = "".join(call_result) # ⭐ Concatenate results for non-streaming
model_response.message.content = "".join(call_result)
return model_response
def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen:

View file

@ -1,8 +1,9 @@
import datetime
import re
from typing import List
from memoryscope.constants.language_constants import WEEKDAYS, DATATIME_WORD_LIST, MONTH_DICT
from memoryscope.utils.global_context import G_CONTEXT
from memoryscope.enumeration.language_enum import LanguageEnum
from memoryscope.utils.logger import Logger
@ -40,7 +41,7 @@ class DatetimeHandler(object):
self._dt_info_dict: dict | None = None
def _parse_dt_info(self):
def _parse_dt_info(self, language: LanguageEnum):
"""
Parses the datetime object (_dt) into a dictionary containing detailed date and time components,
including language-specific weekday representation.
@ -52,17 +53,16 @@ class DatetimeHandler(object):
"""
return {
"year": self._dt.year,
"month": MONTH_DICT[G_CONTEXT.language][self._dt.month - 1],
"month": MONTH_DICT[language][self._dt.month - 1],
"day": self._dt.day,
"hour": self._dt.hour,
"minute": self._dt.minute,
"second": self._dt.second,
"week": self._dt.isocalendar().week,
"weekday": WEEKDAYS[G_CONTEXT.language][self._dt.isocalendar().weekday - 1],
"weekday": WEEKDAYS[language][self._dt.isocalendar().weekday - 1],
}
@property
def dt_info_dict(self):
def get_dt_info_dict(self, language: LanguageEnum):
"""
Property method to get the dictionary containing parsed datetime information.
If None, initialize using `_parse_dt_info`.
@ -71,7 +71,7 @@ class DatetimeHandler(object):
dict: A dictionary with parsed datetime information.
"""
if self._dt_info_dict is None:
self._dt_info_dict = self._parse_dt_info()
self._dt_info_dict = self._parse_dt_info(language=language)
return self._dt_info_dict
@classmethod
@ -207,7 +207,7 @@ class DatetimeHandler(object):
return date_info
@classmethod
def extract_date_parts(cls, input_string: str) -> dict:
def extract_date_parts(cls, input_string: str, language: LanguageEnum) -> dict:
"""
Extracts various date components from the input string based on the current language context.
@ -217,48 +217,51 @@ class DatetimeHandler(object):
Args:
input_string (str): The string containing date information to be parsed.
language (str): current language.
Returns:
dict: A dictionary containing extracted date components, or an empty dictionary if parsing fails.
"""
func_name = f"extract_date_parts_{G_CONTEXT.language.value}"
func_name = f"extract_date_parts_{language}"
if not hasattr(cls, func_name):
cls.logger.warning(f"language={G_CONTEXT.language.value} needs to complete extract_date_parts func!")
cls.logger.warning(f"language={language} needs to complete extract_date_parts func!")
return {}
return getattr(cls, func_name)(input_string=input_string)
@classmethod
def has_time_word_cn(cls, query: str) -> bool:
def has_time_word_cn(cls, query: str, datetime_word_list: List[str]) -> bool:
"""
Check if the input query contains any datetime-related words based on the cn language context.
Args:
query (str): The input string to check for datetime-related words.
datetime_word_list (list[str]): datetime keywords
Returns:
bool: True if the query contains at least one datetime-related word, False otherwise.
"""
contain_datetime = False
# TODO use re
for datetime_word in DATATIME_WORD_LIST[G_CONTEXT.language]:
for datetime_word in datetime_word_list:
if datetime_word in query:
contain_datetime = True
break
return contain_datetime
@classmethod
def has_time_word_en(cls, query: str) -> bool:
def has_time_word_en(cls, query: str, datetime_word_list: List[str]) -> bool:
"""
Check if the input query contains any datetime-related words based on the en language context.
Args:
query (str): The input string to check for datetime-related words.
datetime_word_list (list[str]): datetime keywords
Returns:
bool: True if the query contains at least one datetime-related word, False otherwise.
"""
contain_datetime = False
for datetime_word in DATATIME_WORD_LIST[G_CONTEXT.language]:
for datetime_word in datetime_word_list:
datetime_word = datetime_word.lower()
# TODO fix strip
if datetime_word in [x.strip().lower().strip(",").strip(".").strip("?").strip(":")
@ -268,12 +271,18 @@ class DatetimeHandler(object):
return contain_datetime
@classmethod
def has_time_word(cls, query: str) -> bool:
func_name = f"has_time_word_{G_CONTEXT.language.value}"
def has_time_word(cls, query: str, language: LanguageEnum) -> bool:
func_name = f"has_time_word_{language}"
if not hasattr(cls, func_name):
cls.logger.warning(f"language={G_CONTEXT.language.value} needs to complete has_time_word func!")
cls.logger.warning(f"language={language} needs to complete has_time_word function!")
return False
return getattr(cls, func_name)(query=query)
if language not in DATATIME_WORD_LIST:
cls.logger.warning(f"language={language} is missing in DATATIME_WORD_LIST!")
return False
datetime_word_list = DATATIME_WORD_LIST[language]
return getattr(cls, func_name)(query=query, datetime_word_list=datetime_word_list)
def datetime_format(self, dt_format: str = "%Y%m%d") -> str:
"""
@ -287,7 +296,7 @@ class DatetimeHandler(object):
"""
return self._dt.strftime(dt_format)
def string_format(self, string_format: str) -> str:
def string_format(self, string_format: str, language: LanguageEnum) -> str:
"""
Format the datetime information stored in the instance using a custom string format.
@ -297,7 +306,7 @@ class DatetimeHandler(object):
Returns:
str: A formatted datetime string.
"""
return string_format.format(**self.dt_info_dict)
return string_format.format(**self.get_dt_info_dict(language=language))
@property
def timestamp(self) -> int:

View file

@ -140,7 +140,7 @@ class TestWorkersCn(unittest.TestCase):
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)]
result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@ -172,7 +172,7 @@ class TestWorkersCn(unittest.TestCase):
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)]
result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@ -202,7 +202,7 @@ class TestWorkersCn(unittest.TestCase):
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_WITH_TIME_NODES)]
result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_WITH_TIME_NODES)]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@ -224,30 +224,30 @@ class TestWorkersCn(unittest.TestCase):
MemoryNode(user_name="AI", target_name="用户", content="用户在阿里巴巴工作"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.memory_manager.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result1 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)])
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作"),
MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.memory_manager.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result2 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)])
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作或有工作经验。"),
MemoryNode(user_name="AI", target_name="用户", content="用户跳槽至openai工作。"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.memory_manager.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result3 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)])
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="我喜欢吃西瓜"),
@ -257,10 +257,10 @@ class TestWorkersCn(unittest.TestCase):
MemoryNode(user_name="AI", target_name="用户", content="我爱吃苹果和香蕉"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.memory_manager.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result4 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)])
worker.logger.info(f"result1={result1}")
worker.logger.info(f"result2={result2}")
@ -293,11 +293,11 @@ class TestWorkersCn(unittest.TestCase):
MemoryNode(content="用户想知道维持广泛社交关系的方法。"),
]
worker.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes)
worker.memory_handler.set_memories(INSIGHT_NODES, [])
worker.memory_manager.set_memories(NOT_REFLECTED_NODES, nodes)
worker.memory_manager.set_memories(INSIGHT_NODES, [])
worker.run()
result = [node.key for node in worker.memory_handler.get_memories(INSIGHT_NODES)]
result = [node.key for node in worker.memory_manager.get_memories(INSIGHT_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.get_reflection={result}")
return worker
@ -320,10 +320,10 @@ class TestWorkersCn(unittest.TestCase):
nodes = [
MemoryNode(content="用户喜欢打王者荣耀"),
]
worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(INSIGHT_NODES)]
result = [node.content for node in worker.memory_manager.get_memories(INSIGHT_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.update_insight={result}")
@ -345,10 +345,10 @@ class TestWorkersCn(unittest.TestCase):
MemoryNode(content="用户在北京工作,感到压力大,寻求放松方式。"),
MemoryNode(content="用户在上海工作。"),
]
worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes)
worker.unit_test_flag = True
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]
result = [node.content for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.long_contra_repeat={result}")

View file

@ -150,7 +150,7 @@ class TestWorkersEn(unittest.TestCase):
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)]
result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@ -193,7 +193,7 @@ class TestWorkersEn(unittest.TestCase):
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)]
result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@ -226,7 +226,7 @@ class TestWorkersEn(unittest.TestCase):
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_WITH_TIME_NODES)]
result = [node.content for node in worker.memory_manager.get_memories(NEW_OBS_WITH_TIME_NODES)]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@ -248,30 +248,30 @@ class TestWorkersEn(unittest.TestCase):
MemoryNode(user_name="AI", target_name="用户", content="User works at Alibaba"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.memory_manager.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result1 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)])
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="User works at JD.com"),
MemoryNode(user_name="AI", target_name="用户", content="Users working in Meituan"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.memory_manager.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result2 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)])
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="User works at JD.com"),
MemoryNode(user_name="AI", target_name="用户", content="User is working in Meituan"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.memory_manager.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result3 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)])
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="I like to eat watermelon"),
@ -279,10 +279,10 @@ class TestWorkersEn(unittest.TestCase):
MemoryNode(user_name="AI", target_name="用户", content="I don't like watermelon"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.memory_manager.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result4 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)])
worker.logger.info(f"result1={result1}")
worker.logger.info(f"result2={result2}")
@ -316,11 +316,11 @@ class TestWorkersEn(unittest.TestCase):
MemoryNode(content="Users want to know how to maintain extensive social relationships."),
]
worker.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes)
worker.memory_handler.set_memories(INSIGHT_NODES, [])
worker.memory_manager.set_memories(NOT_REFLECTED_NODES, nodes)
worker.memory_manager.set_memories(INSIGHT_NODES, [])
worker.run()
result = [node.key for node in worker.memory_handler.get_memories(INSIGHT_NODES)]
result = [node.key for node in worker.memory_manager.get_memories(INSIGHT_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.get_reflection={result}")
return worker
@ -343,10 +343,10 @@ class TestWorkersEn(unittest.TestCase):
nodes = [
MemoryNode(content="Users like to play King of Glory"),
]
worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(INSIGHT_NODES)]
result = [node.content for node in worker.memory_manager.get_memories(INSIGHT_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.update_insight={result}")
@ -368,10 +368,10 @@ class TestWorkersEn(unittest.TestCase):
MemoryNode(content="The user works in Beijing, feels stressed, and is looking for ways to relax."),
MemoryNode(content="User works in Shanghai."),
]
worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
worker.memory_manager.set_memories(NOT_UPDATED_NODES, nodes)
worker.unit_test_flag = True
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]
result = [node.content for node in worker.memory_manager.get_memories(MERGE_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.long_contra_repeat={result}")