mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
[dev] modify op name & service base func
This commit is contained in:
parent
4c91212b16
commit
c1957af6af
17 changed files with 122 additions and 142 deletions
|
|
@ -39,13 +39,13 @@ memory_service:
|
|||
description: "delete all long-term memory"
|
||||
|
||||
write_memory:
|
||||
class: memory.operation.write_memory
|
||||
class: memory.operation.write_memory_op
|
||||
workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory
|
||||
description: "write observation memory of the user"
|
||||
interval_time: 5
|
||||
|
||||
summary_memory:
|
||||
class: memory.operation.summary_memory
|
||||
class: memory.operation.backend_operation
|
||||
workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory
|
||||
description: "summary observation memory of the user"
|
||||
interval_time: 30
|
||||
|
|
|
|||
|
|
@ -110,7 +110,8 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
if self._memory_service not in G_CONTEXT.memory_service_dict:
|
||||
raise ValueError("Missing declaration of memory_service in yaml configuration: " + self._memory_service)
|
||||
self._memory_service = G_CONTEXT.memory_service_dict[self._memory_service]
|
||||
self._memory_service.start_service() # ⭐ Initialize and start the memory service
|
||||
self._memory_service.init_service()
|
||||
self._memory_service.start_backend_service()
|
||||
return self._memory_service
|
||||
|
||||
@property
|
||||
|
|
@ -230,7 +231,7 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
command, kwargs = self.parse_query_command(query)
|
||||
|
||||
if command == "exit":
|
||||
self.memory_service.stop_service()
|
||||
self.memory_service.stop_backend_service()
|
||||
continue_run = False
|
||||
|
||||
elif command == "clear":
|
||||
|
|
@ -250,16 +251,23 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
refresh_time = kwargs.pop("refresh_time", "")
|
||||
if refresh_time and refresh_time.isdigit():
|
||||
refresh_time = int(refresh_time)
|
||||
self.memory_service.stop_backend_service()
|
||||
while True:
|
||||
result = self.memory_service.do_operation(op_name=command, **kwargs)
|
||||
os.system("clear")
|
||||
self.print_logo()
|
||||
questionary.print(result)
|
||||
if result:
|
||||
questionary.print(result)
|
||||
else:
|
||||
questionary.print(f"command={command} result is empty! kwargs={kwargs}")
|
||||
time.sleep(refresh_time)
|
||||
|
||||
else:
|
||||
result = self.memory_service.do_operation(op_name=command, **kwargs)
|
||||
questionary.print(result)
|
||||
if result:
|
||||
questionary.print(result)
|
||||
else:
|
||||
questionary.print(f"command={command} result is empty! kwargs={kwargs}")
|
||||
|
||||
else:
|
||||
questionary.print(f"Unknown command={command} received.")
|
||||
|
|
@ -297,6 +305,7 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
questionary.print(f"{self.assistant_name}: ", end="", style="bold")
|
||||
|
||||
# Fetch and display AI's response, with support for streaming
|
||||
self.memory_service.start_backend_service()
|
||||
if self.stream:
|
||||
model_response = None
|
||||
for model_response in self.chat_with_memory(query=query):
|
||||
|
|
@ -316,7 +325,7 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
questionary.print("User interrupt occurred.")
|
||||
is_exit = questionary.confirm("Continue exit?").unsafe_ask()
|
||||
if is_exit:
|
||||
self.memory_service.stop_service()
|
||||
self.memory_service.stop_backend_service()
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -17,18 +17,16 @@ class BaseOperation(metaclass=ABCMeta):
|
|||
|
||||
operation_type: OPERATION_TYPE = "frontend"
|
||||
|
||||
def __init__(self, name: str, description: str = "", **kwargs):
|
||||
def __init__(self, name: str, description: str = ""):
|
||||
"""
|
||||
Initializes a new instance of the BaseOperation.
|
||||
|
||||
Args:
|
||||
name (str): The name identifying the operation.
|
||||
description (str): An optional description detailing the operation's purpose or behavior.
|
||||
**kwargs: Arbitrary keyword arguments for custom settings or parameters.
|
||||
"""
|
||||
self.name: str = name
|
||||
self.description: str = description
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
def init_workflow(self, **kwargs):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -146,6 +146,7 @@ class BaseWorkflow(object):
|
|||
worker = self.worker_dict[name]
|
||||
worker.run()
|
||||
if not worker.continue_run:
|
||||
self.logger.warning(f"worker={worker.name} stop workflow!")
|
||||
return False
|
||||
return True
|
||||
|
||||
|
|
@ -159,7 +160,7 @@ class BaseWorkflow(object):
|
|||
they are submitted for parallel execution using a thread pool. The workflow will stop if any sub-workflow
|
||||
returns False.
|
||||
"""
|
||||
with Timer(self.name, time_log_type="wrap"):
|
||||
with Timer(f"workflow.{self.name}", time_log_type="wrap"):
|
||||
self.context[WORKFLOW_NAME] = self.name
|
||||
|
||||
# Iterate over each part of the workflow
|
||||
|
|
|
|||
|
|
@ -30,6 +30,9 @@ class WriteMemoryOp(BackendOperation):
|
|||
Any: The result obtained from running the workflow.
|
||||
"""
|
||||
|
||||
if not self.chat_messages:
|
||||
return
|
||||
|
||||
# Use shallow copy to prevent adding new messages.
|
||||
chat_messages = self.chat_messages.copy()
|
||||
|
||||
|
|
|
|||
|
|
@ -41,38 +41,10 @@ class BaseMemoryService(metaclass=ABCMeta):
|
|||
self.logger = Logger.get_logger()
|
||||
self.kwargs = kwargs
|
||||
|
||||
@abstractmethod
|
||||
def _init_operation(self, memory_operations: Dict[str, dict]):
|
||||
"""
|
||||
Initializes the memory operations with a given dictionary of operations.
|
||||
|
||||
This method is to be implemented by subclasses to set up or configure
|
||||
the memory operations based on the provided dictionary.
|
||||
|
||||
Args:
|
||||
memory_operations (Dict[str, dict]): A dictionary containing configuration
|
||||
details for each memory operation.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: This exception is raised to indicate that the method
|
||||
needs to be overridden in the subclass.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def add_messages(self, messages: List[Message] | Message):
|
||||
raise NotImplementedError
|
||||
|
||||
def start_service(self, **kwargs):
|
||||
"""
|
||||
This method is intended to initiate the service with provided keyword arguments,
|
||||
preparing the necessary resources for executing operations.
|
||||
|
||||
Args:
|
||||
**kwargs: Additional keyword arguments used to configure the service upon startup.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def do_operation(self, op_name: str, **kwargs):
|
||||
"""
|
||||
|
|
@ -123,9 +95,12 @@ class BaseMemoryService(metaclass=ABCMeta):
|
|||
assert self.read_message_key in self._operation_dict, f"op={self.read_message_key} is not inited!"
|
||||
return self.do_operation(self.read_message_key)
|
||||
|
||||
def stop_service(self):
|
||||
"""
|
||||
Placeholder method to stop the service.
|
||||
Intended to be overridden by subclasses to define specific shutdown logic.
|
||||
"""
|
||||
@abstractmethod
|
||||
def init_service(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def start_backend_service(self):
|
||||
pass
|
||||
|
||||
def stop_backend_service(self):
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from typing import List, Dict
|
||||
from typing import List
|
||||
|
||||
from memory_scope.memory.operation.base_operation import BaseOperation
|
||||
from memory_scope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.utils.tool_functions import init_instance_by_config
|
||||
|
|
@ -12,35 +13,6 @@ class ChatMemoryService(BaseMemoryService):
|
|||
self.contextual_msg_count: int = contextual_msg_count
|
||||
assert self.history_msg_count >= self.contextual_msg_count
|
||||
|
||||
self._init_operation(self.memory_operations)
|
||||
|
||||
def _init_operation(self, memory_operations: Dict[str, dict]):
|
||||
"""
|
||||
Initializes memory operations based on the provided configuration dictionary.
|
||||
Ensures that each operation is only initialized once by checking for duplicates.
|
||||
|
||||
Args:
|
||||
memory_operations (Dict[str, dict]): A dictionary where keys are operation names
|
||||
and values are configuration dictionaries for each operation.
|
||||
|
||||
Note:
|
||||
Logs a warning if an attempt is made to initialize an operation with a name that already exists.
|
||||
Logs an info message upon successful initialization of each operation.
|
||||
"""
|
||||
for name, operation_config in memory_operations.items():
|
||||
if name in self._operation_dict:
|
||||
self.logger.warning(f"memory operation={name} is repeated!")
|
||||
continue
|
||||
|
||||
# ⭐ Initialize operation instance by its config
|
||||
self._operation_dict[name] = init_instance_by_config(
|
||||
config=operation_config,
|
||||
name=name,
|
||||
chat_messages=self.chat_messages,
|
||||
message_lock=self.message_lock,
|
||||
contextual_msg_count=self.contextual_msg_count)
|
||||
self.logger.info(f"service={self.__class__.__name__} init operation={name}")
|
||||
|
||||
def add_messages(self, messages: List[Message] | Message):
|
||||
"""
|
||||
Adds a single message or a list of messages to the chat history, ensuring the message list
|
||||
|
|
@ -56,28 +28,16 @@ class ChatMemoryService(BaseMemoryService):
|
|||
|
||||
# Sort the messages by their creation time to maintain chronological order
|
||||
messages = sorted(messages, key=lambda x: x.time_created)
|
||||
|
||||
|
||||
# Append the sorted messages to the chat history
|
||||
self.chat_messages.extend(messages)
|
||||
|
||||
|
||||
# If the chat history exceeds the allowed message count, remove the oldest messages
|
||||
if len(self.chat_messages) > self.history_msg_count:
|
||||
gap_size = len(self.chat_messages) - self.history_msg_count
|
||||
for _ in range(gap_size):
|
||||
self.chat_messages.pop(0)
|
||||
|
||||
def start_service(self, **kwargs):
|
||||
"""
|
||||
Initializes and starts backend operations defined in `_operation_dict`.
|
||||
|
||||
Args:
|
||||
**kwargs: Additional keyword arguments passed to each operation's initialization.
|
||||
"""
|
||||
for _, operation in self._operation_dict.items():
|
||||
operation.init_workflow(**kwargs) # Initialize workflow for each operation
|
||||
if operation.operation_type == "backend":
|
||||
operation.run_operation_backend() # Run backend operations
|
||||
|
||||
def do_operation(self, op_name: str, **kwargs):
|
||||
"""
|
||||
Executes a specific operation by its name with provided keyword arguments.
|
||||
|
|
@ -97,10 +57,37 @@ class ChatMemoryService(BaseMemoryService):
|
|||
return
|
||||
return self._operation_dict[op_name].run_operation(**kwargs) # Execute the operation
|
||||
|
||||
def stop_service(self):
|
||||
def init_service(self):
|
||||
for name, operation_config in self.memory_operations.items():
|
||||
if name in self._operation_dict:
|
||||
self.logger.warning(f"memory operation={name} is repeated!")
|
||||
continue
|
||||
|
||||
# ⭐ Initialize operation instance by its config
|
||||
operation: BaseOperation = init_instance_by_config(
|
||||
config=operation_config,
|
||||
name=name,
|
||||
chat_messages=self.chat_messages,
|
||||
message_lock=self.message_lock,
|
||||
contextual_msg_count=self.contextual_msg_count)
|
||||
operation.init_workflow() # Initialize workflow for each operation
|
||||
|
||||
self._operation_dict[name] = operation
|
||||
self.logger.info(f"service={self.__class__.__name__} init operation={name}")
|
||||
|
||||
def start_backend_service(self):
|
||||
for _, operation in self._operation_dict.items():
|
||||
if operation.operation_type == "backend":
|
||||
# Run backend operations
|
||||
operation.run_operation_backend()
|
||||
self.logger.info(f"start operation={operation.name}...")
|
||||
|
||||
def stop_backend_service(self):
|
||||
"""
|
||||
Stops all backend operations that are currently running.
|
||||
"""
|
||||
for _, operation in self._operation_dict.items():
|
||||
if operation.operation_type == "backend":
|
||||
operation.stop_operation_backend() # Stop backend operations
|
||||
# Stop backend operations
|
||||
operation.stop_operation_backend()
|
||||
self.logger.info(f"stop operation={operation.name}...")
|
||||
|
|
|
|||
|
|
@ -62,6 +62,10 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
|
||||
def run(self):
|
||||
with Timer(f"worker.{self.name}", time_log_type="wrap"):
|
||||
self.continue_run = True
|
||||
self.async_task_list.clear()
|
||||
self.thread_task_list.clear()
|
||||
|
||||
if self.raise_exception:
|
||||
self._run()
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.constants.common_constants import RETRIEVE_MEMORY_NODES, RESULT
|
||||
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memory_scope.enumeration.store_status_enum import StoreStatusEnum
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.utils.datetime_handler import DatetimeHandler
|
||||
|
|
@ -25,7 +25,7 @@ class PrintMemoryWorker(MemoryBaseWorker):
|
|||
dt_handler = DatetimeHandler(node.timestamp)
|
||||
dt = dt_handler.datetime_format("%Y%m%d-%H:%M:%S")
|
||||
line = f"{dt} {node.content}"
|
||||
if MemoryNodeStatus(node.status) is MemoryNodeStatus.EXPIRED:
|
||||
if StoreStatusEnum(node.store_status) is StoreStatusEnum.EXPIRED:
|
||||
if node.content in expired_content_set:
|
||||
continue
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -114,12 +114,12 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
|
|||
self.logger.info(f"memory_node_list.size={len(memory_node_list)}")
|
||||
|
||||
if not memory_node_list:
|
||||
self.continue_run = False
|
||||
return
|
||||
|
||||
memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True)
|
||||
for node in memory_node_list:
|
||||
node.action_status = ActionStatusEnum.NONE
|
||||
self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} "
|
||||
f"type={node.memory_type} status={node.status}")
|
||||
self.memory_handler.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list)
|
||||
self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} 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)
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS,
|
|||
from memory_scope.memory.worker.base_worker import BaseWorker
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.storage.base_memory_store import BaseMemoryStore
|
||||
from memory_scope.storage.base_monitor import BaseMonitor
|
||||
from memory_scope.utils.global_context import G_CONTEXT
|
||||
from memory_scope.utils.memory_handler import MemoryHandler
|
||||
|
|
@ -36,6 +37,8 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
self._embedding_model: BaseModel | str = embedding_model
|
||||
self._generation_model: BaseModel | str = generation_model
|
||||
self._rank_model: BaseModel | str = rank_model
|
||||
|
||||
self._memory_store: BaseMemoryStore | None = None
|
||||
self._monitor: BaseMonitor | None = None
|
||||
|
||||
self._user_name: str | None = None
|
||||
|
|
@ -105,10 +108,10 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
return self._rank_model
|
||||
|
||||
@property
|
||||
def memory_handler(self) -> 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)
|
||||
def memory_store(self) -> BaseMemoryStore:
|
||||
if self._memory_store is None:
|
||||
self._memory_store = G_CONTEXT.memory_store
|
||||
return self._memory_store
|
||||
|
||||
@property
|
||||
def monitor(self) -> BaseMonitor:
|
||||
|
|
@ -160,6 +163,12 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
self._prompt_handler = PromptHandler(self.FILE_PATH, **self.kwargs)
|
||||
return self._prompt_handler
|
||||
|
||||
@property
|
||||
def memory_handler(self) -> 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)
|
||||
|
||||
def __getattr__(self, key: str):
|
||||
"""
|
||||
Custom attribute access to directly retrieve values from kwargs.
|
||||
|
|
|
|||
|
|
@ -22,7 +22,7 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
|
||||
def retrieve_similar_content(self, node: MemoryNode) -> (MemoryNode, List[MemoryNode]):
|
||||
"""
|
||||
Retrieves memory nodes with content similar to the given node, filtering by user, target, status, and memory type.
|
||||
Retrieves memory nodes with content similar to the given node, filtering by user/target/status/memory_type.
|
||||
Only returns nodes whose similarity score meets or exceeds the predefined threshold.
|
||||
|
||||
Args:
|
||||
|
|
@ -139,7 +139,8 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
node.store_status = StoreStatusEnum.EXPIRED.value
|
||||
|
||||
merge_obs_nodes.append(node)
|
||||
self.logger.info(f"after_long_contra_repeat: {node.content} {node.status}")
|
||||
self.logger.info(f"after_long_contra_repeat: {node.content} store_status={node.store_status} "
|
||||
f"action_status={node.action_status}")
|
||||
|
||||
# save context
|
||||
self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes)
|
||||
|
|
|
|||
|
|
@ -74,7 +74,7 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
# parse text
|
||||
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__)
|
||||
if len(idx_merge_obs_list) <= 0:
|
||||
self.add_run_info("idx_merge_obs_list is empty!")
|
||||
self.logger.warning("idx_merge_obs_list is empty!")
|
||||
return
|
||||
|
||||
# add merged obs
|
||||
|
|
|
|||
|
|
@ -76,7 +76,6 @@ class InfoFilterWorker(MemoryBaseWorker):
|
|||
info_score_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__)
|
||||
if len(info_score_list) != len(info_messages):
|
||||
self.logger.warning(f"score_size != messages_size, {len(info_score_list)} vs {len(info_messages)}")
|
||||
self.continue_run = False
|
||||
return
|
||||
|
||||
# filter messages
|
||||
|
|
|
|||
|
|
@ -5,7 +5,6 @@ from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
|
|||
from memory_scope.enumeration.store_status_enum import StoreStatusEnum
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.utils.datetime_handler import DatetimeHandler
|
||||
from memory_scope.utils.timer import timer
|
||||
|
||||
|
|
@ -63,24 +62,18 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
self.memory_handler.set_memories(INSIGHT_NODES, nodes)
|
||||
|
||||
@timer
|
||||
def retrieve_today_memory(self):
|
||||
if not self.today_obs_top_k:
|
||||
def retrieve_today_memory(self, query: str, dt: str):
|
||||
if not self.retrieve_today_top_k:
|
||||
return
|
||||
|
||||
if not self.chat_messages:
|
||||
self.logger.warning("chat_messages is empty!")
|
||||
return
|
||||
|
||||
message: Message = self.chat_messages[-1]
|
||||
dt_handler = DatetimeHandler(message.time_created)
|
||||
filter_dict = {
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"store_status": StoreStatusEnum.VALID.value,
|
||||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"dt": dt_handler.datetime_format(),
|
||||
"dt": dt,
|
||||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=message.content,
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_today_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
|
|
@ -95,9 +88,14 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
This method serves as the controller for data retrieval operations, enhancing efficiency
|
||||
by handling tasks concurrently.
|
||||
"""
|
||||
mock_query = "-" # Placeholder query
|
||||
self.submit_thread_task(self.retrieve_not_reflected_memory, query=mock_query)
|
||||
self.submit_thread_task(self.retrieve_not_updated_memory, query=mock_query)
|
||||
self.submit_thread_task(self.retrieve_insight_memory, query=mock_query)
|
||||
self.submit_thread_task(self.retrieve_today_memory)
|
||||
self.gather_thread_result() # Waits for all submitted tasks to complete
|
||||
|
||||
# Placeholder query
|
||||
query = "-"
|
||||
dt = DatetimeHandler().datetime_format()
|
||||
self.submit_thread_task(self.retrieve_not_reflected_memory, query=query)
|
||||
self.submit_thread_task(self.retrieve_not_updated_memory, query=query)
|
||||
self.submit_thread_task(self.retrieve_insight_memory, query=query)
|
||||
self.submit_thread_task(self.retrieve_today_memory, query=query, dt=dt)
|
||||
|
||||
# Waits for all submitted tasks to complete
|
||||
self.gather_thread_result()
|
||||
|
|
|
|||
|
|
@ -23,4 +23,4 @@ class StoreMemoryWorker(MemoryBaseWorker):
|
|||
timestamp=dt_handler.timestamp)
|
||||
self.memory_handler.update_memories(nodes=node)
|
||||
else:
|
||||
self.memory_handler.update_memories(self.store_key)
|
||||
self.memory_handler.update_memories(key=self.store_key)
|
||||
|
|
|
|||
|
|
@ -37,7 +37,7 @@ class MemoryHandler(object):
|
|||
self._id_memory_dict.clear()
|
||||
self._key_id_dict.clear()
|
||||
|
||||
def add_memories(self, nodes: MemoryNode | List[MemoryNode], log_repeat: bool = True):
|
||||
def set_memories(self, key: str, nodes: MemoryNode | List[MemoryNode], log_repeat: bool = True):
|
||||
if nodes is None:
|
||||
nodes = []
|
||||
elif isinstance(nodes, MemoryNode):
|
||||
|
|
@ -47,15 +47,13 @@ class MemoryHandler(object):
|
|||
if node.memory_id in self._id_memory_dict:
|
||||
if log_repeat:
|
||||
self.logger.warning(f"repeated_id memory id={node.memory_id} content={node.content} "
|
||||
f"status={node.status}")
|
||||
f"store_status={node.store_status} action_status={node.action_status}")
|
||||
continue
|
||||
|
||||
self._id_memory_dict[node.memory_id] = node
|
||||
self.logger.info(f"add to memory context memory id={node.memory_id} content={node.content} "
|
||||
f"status={node.status}")
|
||||
f"store_status={node.store_status} action_status={node.action_status}")
|
||||
|
||||
def set_memories(self, key: str, nodes: MemoryNode | List[MemoryNode], log_repeat: bool = True):
|
||||
self.add_memories(nodes=nodes, log_repeat=log_repeat)
|
||||
self._key_id_dict[key] = [n.memory_id for n in nodes]
|
||||
|
||||
def get_memories(self, keys: str | List[str]) -> List[MemoryNode]:
|
||||
|
|
@ -78,25 +76,23 @@ class MemoryHandler(object):
|
|||
keys = [keys]
|
||||
|
||||
for key in keys:
|
||||
if key not in self._key_id_dict:
|
||||
continue
|
||||
memory_ids: List[str] = self._key_id_dict[key]
|
||||
if memory_ids:
|
||||
memories.update({x: self._id_memory_dict[x] for x in memory_ids})
|
||||
return list(memories.values())
|
||||
|
||||
def update_memories(self, keys: str | List[str] = None, nodes: MemoryNode | List[MemoryNode] = None):
|
||||
if keys is None:
|
||||
keys = []
|
||||
elif keys == "all":
|
||||
keys = list(self._id_memory_dict.keys())
|
||||
elif isinstance(keys, str):
|
||||
keys = [keys]
|
||||
|
||||
# combine keys
|
||||
def update_memories(self, key: str = "", nodes: MemoryNode | List[MemoryNode] = None):
|
||||
ids: Set[str] = set()
|
||||
for key in keys:
|
||||
t_ids: List[str] = self._key_id_dict[key]
|
||||
if t_ids:
|
||||
ids.update(t_ids)
|
||||
|
||||
if key == "all":
|
||||
ids.update(self._id_memory_dict.keys())
|
||||
elif key:
|
||||
for k in key.split(","):
|
||||
t_ids: List[str] = self._key_id_dict.get(k.strip())
|
||||
if t_ids:
|
||||
ids.update(t_ids)
|
||||
|
||||
# Remove and collect nodes by IDs
|
||||
update_nodes: List[MemoryNode] = [self._id_memory_dict.pop(_) for _ in ids]
|
||||
|
|
@ -148,9 +144,9 @@ class MemoryHandler(object):
|
|||
if modified_memories:
|
||||
for n in modified_memories:
|
||||
n.action_status = ActionStatusEnum.NONE.value
|
||||
self.memory_store.batch_update(c_modified_memories, update_embedding=False)
|
||||
self.memory_store.batch_update(modified_memories, update_embedding=False)
|
||||
|
||||
# set memories expired
|
||||
delete_memories = [n for n in nodes if n.action_status == ActionStatusEnum.DELETE]
|
||||
if delete_memories:
|
||||
self.memory_store.batch_delete(nodes)
|
||||
self.memory_store.batch_delete(delete_memories)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue