mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
[dev] rename MEMORY_HANDLER to MEMORY_MANAGER
This commit is contained in:
parent
6aaa26742e
commit
590cbf65c4
23 changed files with 177 additions and 161 deletions
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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}...")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)())
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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] = []
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue