mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-05 08:06:15 +00:00
[dev] rename worker name
This commit is contained in:
parent
c1957af6af
commit
0edaaa4256
23 changed files with 179 additions and 153 deletions
|
|
@ -20,33 +20,33 @@ memory_service:
|
|||
|
||||
read_memory:
|
||||
class: memory.operation.frontend_operation
|
||||
workflow: set_query,[extract_time|retrieve_memory1,semantic_rank],fuse_rerank
|
||||
workflow: set_query,[extract_time|retrieve_obs_ins,semantic_rank],fuse_rerank
|
||||
description: "read long-term memory"
|
||||
|
||||
list_memory:
|
||||
class: memory.operation.frontend_operation
|
||||
workflow: set_query,retrieve_memory2,print_memory
|
||||
workflow: set_query,retrieve_top_memory,print_memory
|
||||
description: "read all long-term memory of the user"
|
||||
|
||||
delete_memory:
|
||||
class: memory.operation.frontend_operation
|
||||
workflow: set_query,retrieve_memory3,delete_memory
|
||||
workflow: set_query,retrieve_all_memory,delete_memory
|
||||
description: "delete all long-term memory"
|
||||
|
||||
add_memory:
|
||||
class: memory.operation.frontend_operation
|
||||
workflow: store_memory
|
||||
description: "delete all long-term memory"
|
||||
workflow: add_memory
|
||||
description: "add a single observation"
|
||||
|
||||
write_memory:
|
||||
class: memory.operation.write_memory_op
|
||||
workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory
|
||||
workflow: info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat,store_memory
|
||||
description: "write observation memory of the user"
|
||||
interval_time: 5
|
||||
|
||||
summary_memory:
|
||||
class: memory.operation.backend_operation
|
||||
workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory
|
||||
workflow: load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat,store_memory
|
||||
description: "summary observation memory of the user"
|
||||
interval_time: 30
|
||||
|
||||
|
|
@ -60,15 +60,13 @@ worker:
|
|||
class: memory.worker.frontend.read_message_worker
|
||||
set_query:
|
||||
class: memory.worker.frontend.set_query_worker
|
||||
retrieve_memory1:
|
||||
retrieve_obs_ins:
|
||||
class: memory.worker.frontend.retrieve_memory_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_pf_top_k: 100
|
||||
retrieve_expired_top_k: 0
|
||||
retrieve_ins_top_k: 100
|
||||
extract_time:
|
||||
class: memory.worker.frontend.extract_time_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_top_k: 1
|
||||
semantic_rank:
|
||||
class: memory.worker.frontend.semantic_rank_worker
|
||||
rank_model: dashscope_rank
|
||||
|
|
@ -82,75 +80,59 @@ worker:
|
|||
insight: 2.0
|
||||
fuse_time_ratio: 2.0
|
||||
fuse_rerank_top_k: 10
|
||||
retrieve_memory2:
|
||||
retrieve_top_memory:
|
||||
class: memory.worker.frontend.retrieve_memory_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_pf_top_k: 100
|
||||
retrieve_ins_top_k: 100
|
||||
retrieve_expired_top_k: 100
|
||||
print_memory:
|
||||
class: memory.worker.frontend.print_memory_worker
|
||||
retrieve_memory3:
|
||||
retrieve_all_memory:
|
||||
class: memory.worker.frontend.retrieve_memory_worker
|
||||
retrieve_obs_top_k: 10000
|
||||
retrieve_ins_pf_top_k: 10000
|
||||
retrieve_ins_top_k: 10000
|
||||
retrieve_expired_top_k: 10000
|
||||
delete_memory:
|
||||
class: memory.worker.frontend.update_status_worker
|
||||
class: memory.worker.frontend.update_memory_worker
|
||||
method: modify_action_status
|
||||
expired_action: delete
|
||||
valid_action_dict:
|
||||
obs_customized: delete
|
||||
insight: delete
|
||||
observation: delete
|
||||
add_memory:
|
||||
class: memory.worker.frontend.update_memory_worker
|
||||
method: from_query
|
||||
info_filter:
|
||||
class: memory.worker.write.info_filter_worker
|
||||
generation_model: dashscope_generation
|
||||
preserved_scores: 2,3
|
||||
info_filter_msg_max_size: 200
|
||||
generation_model_top_k: 1
|
||||
load_memory1:
|
||||
load_today_memory:
|
||||
class: memory.worker.write.load_memory_worker
|
||||
retrieve_not_reflected_top_k: 0
|
||||
retrieve_not_updated_top_k: 0
|
||||
retrieve_insight_top_k: 0
|
||||
retrieve_today_top_k: 100
|
||||
get_observation:
|
||||
class: memory.worker.write.get_observation_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_top_k: 1
|
||||
get_observation_with_time:
|
||||
class: memory.worker.write.get_observation_with_time_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_top_k: 1
|
||||
contra_repeat:
|
||||
class: memory.worker.write.contra_repeat_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_top_k: 1
|
||||
retrieve_top_k: 30
|
||||
contra_repeat_max_count: 50
|
||||
store_memory:
|
||||
class: memory.worker.write.store_memory_worker
|
||||
store_key: all
|
||||
load_memory2:
|
||||
class: memory.worker.write.update_memory_worker
|
||||
method: from_memory_key
|
||||
memory_key: all
|
||||
load_obs_and_insight:
|
||||
class: memory.worker.write.load_memory_worker
|
||||
retrieve_not_reflected_top_k: 100
|
||||
retrieve_not_updated_top_k: 100
|
||||
retrieve_insight_top_k: 100
|
||||
retrieve_today_top_k: 0
|
||||
get_reflection_subject:
|
||||
class: memory.worker.summary.get_reflection_subject_worker
|
||||
retrieve_top_k: 100
|
||||
reflect_obs_cnt_threshold: 10
|
||||
generation_model_top_k: 1
|
||||
update_insight:
|
||||
class: memory.worker.summary.update_insight_worker
|
||||
update_insight_threshold: 0.1
|
||||
generation_model_top_k: 1
|
||||
update_insight_max_thread: 10
|
||||
long_contra_repeat:
|
||||
class: memory.worker.summary.long_contra_repeat_worker
|
||||
long_contra_repeat_top_k: 2
|
||||
long_contra_repeat_threshold: 0.1
|
||||
generation_model_top_k: 1
|
||||
|
||||
models:
|
||||
dashscope_generation:
|
||||
|
|
|
|||
|
|
@ -230,6 +230,10 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
continue_run = True
|
||||
command, kwargs = self.parse_query_command(query)
|
||||
|
||||
# Print prompt for AI's response
|
||||
questionary.print("> ", end="", style="fg:yellow")
|
||||
questionary.print(f"{self.assistant_name}: ", end="", style="bold")
|
||||
|
||||
if command == "exit":
|
||||
self.memory_service.stop_backend_service()
|
||||
continue_run = False
|
||||
|
|
@ -257,6 +261,8 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
os.system("clear")
|
||||
self.print_logo()
|
||||
if result:
|
||||
if isinstance(result, list):
|
||||
result = "\n".join([str(x) for x in result])
|
||||
questionary.print(result)
|
||||
else:
|
||||
questionary.print(f"command={command} result is empty! kwargs={kwargs}")
|
||||
|
|
@ -265,6 +271,8 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
else:
|
||||
result = self.memory_service.do_operation(op_name=command, **kwargs)
|
||||
if result:
|
||||
if isinstance(result, list):
|
||||
result = "\n".join([str(x) for x in result])
|
||||
questionary.print(result)
|
||||
else:
|
||||
questionary.print(f"command={command} result is empty! kwargs={kwargs}")
|
||||
|
|
|
|||
|
|
@ -96,7 +96,7 @@ class BaseMemoryService(metaclass=ABCMeta):
|
|||
return self.do_operation(self.read_message_key)
|
||||
|
||||
@abstractmethod
|
||||
def init_service(self):
|
||||
def init_service(self, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
def start_backend_service(self):
|
||||
|
|
|
|||
|
|
@ -57,7 +57,7 @@ class ChatMemoryService(BaseMemoryService):
|
|||
return
|
||||
return self._operation_dict[op_name].run_operation(**kwargs) # Execute the operation
|
||||
|
||||
def init_service(self):
|
||||
def init_service(self, **kwargs):
|
||||
for name, operation_config in self.memory_operations.items():
|
||||
if name in self._operation_dict:
|
||||
self.logger.warning(f"memory operation={name} is repeated!")
|
||||
|
|
@ -70,7 +70,7 @@ class ChatMemoryService(BaseMemoryService):
|
|||
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
|
||||
operation.init_workflow(**kwargs) # Initialize workflow for each operation
|
||||
|
||||
self._operation_dict[name] = operation
|
||||
self.logger.info(f"service={self.__class__.__name__} init operation={name}")
|
||||
|
|
|
|||
|
|
@ -24,13 +24,17 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
self.raise_exception: bool = raise_exception
|
||||
self.is_multi_thread: bool = is_multi_thread
|
||||
self.thread_pool: ThreadPoolExecutor = thread_pool
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
self.continue_run: bool = True
|
||||
self.async_task_list: list = []
|
||||
self.thread_task_list: list = []
|
||||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
self._parse_params(**kwargs)
|
||||
|
||||
def _parse_params(self, **kwargs):
|
||||
pass
|
||||
|
||||
def submit_async_task(self, fn, *args, **kwargs):
|
||||
if self.is_multi_thread:
|
||||
raise RuntimeError(f"async_task is not allowed in multi_thread condition")
|
||||
|
|
|
|||
|
|
@ -1,10 +1,10 @@
|
|||
import datetime
|
||||
|
||||
from memory_scope.constants.common_constants import RESULT, WORKFLOW_NAME, CHAT_KWARGS
|
||||
from memory_scope.memory.worker.base_worker import BaseWorker
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class DummyWorker(BaseWorker):
|
||||
class DummyWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
"""
|
||||
Executes the dummy worker's run logic by logging workflow entry, capturing the current timestamp,
|
||||
|
|
|
|||
|
|
@ -18,6 +18,9 @@ class ExtractTimeWorker(MemoryBaseWorker):
|
|||
EXTRACT_TIME_PATTERN = r'-\s*(\S+):(\d+)'
|
||||
FILE_PATH: str = __file__
|
||||
|
||||
def _parse_params(self, **kwargs):
|
||||
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
|
||||
|
||||
def _run(self):
|
||||
"""
|
||||
Executes the primary logic of identifying and extracting time data from an LLM's response.
|
||||
|
|
|
|||
|
|
@ -8,6 +8,12 @@ from memory_scope.utils.datetime_handler import DatetimeHandler
|
|||
|
||||
class FuseRerankWorker(MemoryBaseWorker):
|
||||
|
||||
def _parse_params(self, **kwargs):
|
||||
self.fuse_score_threshold: float = kwargs.get("fuse_score_threshold", 0.1)
|
||||
self.fuse_ratio_dict: Dict[str, float] = kwargs.get("fuse_ratio_dict", {})
|
||||
self.fuse_time_ratio: float = kwargs.get("fuse_time_ratio", 2.0)
|
||||
self.fuse_rerank_top_k: int = kwargs.get("fuse_rerank_top_k", 10)
|
||||
|
||||
@staticmethod
|
||||
def match_node_time(extract_time_dict: Dict[str, str], node: MemoryNode):
|
||||
if extract_time_dict:
|
||||
|
|
@ -65,6 +71,8 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
continue
|
||||
|
||||
# Calculate type-based adjustment factor
|
||||
if node.memory_type not in self.fuse_ratio_dict:
|
||||
self.logger.warning(f"{node.memory_type} 'factor is not configured!")
|
||||
type_ratio: float = self.fuse_ratio_dict.get(node.memory_type, 0.1)
|
||||
|
||||
# Determine time relevance adjustment factor
|
||||
|
|
|
|||
|
|
@ -16,6 +16,11 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
|
|||
facilitating efficient memory retrieval operations within a given scope.
|
||||
"""
|
||||
|
||||
def _parse_params(self, **kwargs):
|
||||
self.retrieve_obs_top_k: int = kwargs.get("retrieve_obs_top_k", 0)
|
||||
self.retrieve_ins_top_k: int = kwargs.get("retrieve_ins_top_k", 0)
|
||||
self.retrieve_expired_top_k: int = kwargs.get("retrieve_expired_top_k", 0)
|
||||
|
||||
@timer
|
||||
def retrieve_from_observation(self, query: str) -> List[MemoryNode]:
|
||||
"""
|
||||
|
|
@ -45,7 +50,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
|
|||
filter_dict=filter_dict)
|
||||
|
||||
@timer
|
||||
def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]:
|
||||
def retrieve_from_insight(self, query: str) -> List[MemoryNode]:
|
||||
"""
|
||||
Retrieves memories marked as insights from the database based on a query, filtered by user, target,
|
||||
and set to active status.
|
||||
|
|
@ -58,7 +63,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
|
|||
limited by 'retrieve_ins_pf_top_k'.
|
||||
Returns an empty list if 'retrieve_ins_pf_top_k' is not set.
|
||||
"""
|
||||
if not self.retrieve_ins_pf_top_k:
|
||||
if not self.retrieve_ins_top_k:
|
||||
return []
|
||||
|
||||
filter_dict = {
|
||||
|
|
@ -69,7 +74,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
|
|||
}
|
||||
# ⭐ Retrieve insights matching the query, filtered, and limited by top_k
|
||||
return self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_ins_pf_top_k,
|
||||
top_k=self.retrieve_ins_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
@timer
|
||||
|
|
@ -104,7 +109,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
|
|||
"""
|
||||
query, _ = self.get_context(QUERY_WITH_TS)
|
||||
self.submit_thread_task(self.retrieve_from_observation, query=query)
|
||||
self.submit_thread_task(self.retrieve_from_insight_and_profile, query=query)
|
||||
self.submit_thread_task(self.retrieve_from_insight, query=query)
|
||||
self.submit_thread_task(self.retrieve_expired_memory, query=query)
|
||||
|
||||
memory_node_list: List[MemoryNode] = []
|
||||
|
|
|
|||
|
|
@ -24,16 +24,20 @@ class SetQueryWorker(MemoryBaseWorker):
|
|||
query = "_" # Default query value
|
||||
query_timestamp = int(datetime.datetime.now().timestamp()) # Current timestamp as default
|
||||
|
||||
# Check if a specific 'query' has been provided via chat kwargs
|
||||
if "query" in self.chat_kwargs:
|
||||
# Check if a specific 'query' has been provided via chat kwargs
|
||||
query = self.chat_kwargs["query"]
|
||||
if not query:
|
||||
query = ""
|
||||
query = query.strip()
|
||||
|
||||
# If no explicit query is given, use the content of the latest chat message
|
||||
elif self.chat_messages:
|
||||
message = self.chat_messages[-1]
|
||||
assert message.role == MessageRoleEnum.USER.value
|
||||
query = message.content
|
||||
query_timestamp = message.time_created
|
||||
# If no explicit query is given, use the content of the latest chat message
|
||||
chat_messages = [msg for msg in self.chat_messages if msg.role == MessageRoleEnum.USER.value]
|
||||
if chat_messages:
|
||||
message = chat_messages[-1]
|
||||
query = message.content
|
||||
query_timestamp = message.time_created
|
||||
|
||||
# Store the determined query and its timestamp in the context
|
||||
self.set_context(QUERY_WITH_TS, (query, query_timestamp))
|
||||
|
|
|
|||
|
|
@ -1,24 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.constants.common_constants import RETRIEVE_MEMORY_NODES
|
||||
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
|
||||
|
||||
|
||||
class UpdateStatusWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
expired_action = self.expired_action
|
||||
valid_action_dict: dict = self.valid_action_dict
|
||||
memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
if not memory_node_list:
|
||||
return
|
||||
|
||||
for node in memory_node_list:
|
||||
if node.store_status == StoreStatusEnum.EXPIRED.value:
|
||||
node.action_status = expired_action
|
||||
|
||||
elif node.memory_type in valid_action_dict:
|
||||
node.action_status = valid_action_dict[node.memory_type]
|
||||
|
||||
self.memory_handler.update_memories(nodes=memory_node_list)
|
||||
|
|
@ -169,18 +169,6 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
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.
|
||||
|
||||
Args:
|
||||
key (str): The attribute key to look up in kwargs.
|
||||
|
||||
Returns:
|
||||
Any: The value associated with the key in kwargs.
|
||||
"""
|
||||
return self.kwargs[key]
|
||||
|
||||
@staticmethod
|
||||
def get_language_value(languages: dict | list[dict]) -> Any | list[Any]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -19,6 +19,11 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
|
|||
"""
|
||||
FILE_PATH: str = __file__
|
||||
|
||||
def _parse_params(self, **kwargs):
|
||||
self.reflect_obs_cnt_threshold: int = kwargs.get("reflect_obs_cnt_threshold", 10)
|
||||
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
|
||||
self.reflect_num_questions: int = kwargs.get("reflect_num_questions", 5)
|
||||
|
||||
def new_insight_node(self, insight_key: str) -> MemoryNode:
|
||||
"""
|
||||
Creates a new MemoryNode for an insight with the given key, enriched with current datetime metadata.
|
||||
|
|
|
|||
|
|
@ -20,6 +20,11 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
"""
|
||||
FILE_PATH: str = __file__
|
||||
|
||||
def _parse_params(self, **kwargs):
|
||||
self.long_contra_repeat_top_k: int = kwargs.get("long_contra_repeat_top_k", 2)
|
||||
self.long_contra_repeat_threshold: float = kwargs.get("long_contra_repeat_threshold", 0.1)
|
||||
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
|
||||
|
||||
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/memory_type.
|
||||
|
|
|
|||
|
|
@ -19,6 +19,11 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
"""
|
||||
FILE_PATH: str = __file__
|
||||
|
||||
def _parse_params(self, **kwargs):
|
||||
self.update_insight_threshold: float = kwargs.get("update_insight_threshold", 0.1)
|
||||
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
|
||||
self.update_insight_max_count: int = kwargs.get("update_insight_max_count", 10)
|
||||
|
||||
def filter_obs_nodes(self,
|
||||
insight_node: MemoryNode,
|
||||
obs_nodes: List[MemoryNode]) -> (MemoryNode, List[MemoryNode], float):
|
||||
|
|
@ -178,7 +183,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
if not filtered_nodes:
|
||||
continue
|
||||
result_list.append(result)
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)[: self.update_insight_max_thread]
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)[: self.update_insight_max_count]
|
||||
|
||||
# Submit tasks to update insights for the top nodes
|
||||
for insight_node, filtered_nodes, _ in result_sorted:
|
||||
|
|
|
|||
|
|
@ -23,6 +23,11 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
"""
|
||||
FILE_PATH: str = __file__
|
||||
|
||||
def _parse_params(self, **kwargs):
|
||||
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
|
||||
self.retrieve_top_k: int = kwargs.get("retrieve_top_k", 30)
|
||||
self.contra_repeat_max_count: int = kwargs.get("contra_repeat_max_count", 50)
|
||||
|
||||
def _run(self):
|
||||
"""
|
||||
Executes the primary routine of the ContraRepeatWorker which involves fetching memory nodes,
|
||||
|
|
|
|||
|
|
@ -16,6 +16,9 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
FILE_PATH: str = __file__
|
||||
OBS_STORE_KEY: str = NEW_OBS_NODES
|
||||
|
||||
def _parse_params(self, **kwargs):
|
||||
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
|
||||
|
||||
def add_observation(self, message: Message, time_infer: str, obs_content: str, keywords: str):
|
||||
dt_handler = DatetimeHandler(dt=message.time_created)
|
||||
|
||||
|
|
|
|||
|
|
@ -17,6 +17,11 @@ class InfoFilterWorker(MemoryBaseWorker):
|
|||
"""
|
||||
FILE_PATH: str = __file__
|
||||
|
||||
def _parse_params(self, **kwargs):
|
||||
self.preserved_scores: str = kwargs.get("preserved_scores", "2,3")
|
||||
self.info_filter_msg_max_size: int = kwargs.get("info_filter_msg_max_size", 200)
|
||||
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
|
||||
|
||||
def _run(self):
|
||||
"""
|
||||
Filters user messages in the chat, generates a prompt incorporating these messages,
|
||||
|
|
|
|||
|
|
@ -10,6 +10,11 @@ from memory_scope.utils.timer import timer
|
|||
|
||||
|
||||
class LoadMemoryWorker(MemoryBaseWorker):
|
||||
def _parse_params(self, **kwargs):
|
||||
self.retrieve_not_reflected_top_k: int = kwargs.get("retrieve_not_reflected_top_k", 0)
|
||||
self.retrieve_not_updated_top_k: int = kwargs.get("retrieve_not_updated_top_k", 0)
|
||||
self.retrieve_insight_top_k: int = kwargs.get("retrieve_insight_top_k", 0)
|
||||
self.retrieve_today_top_k: int = kwargs.get("retrieve_today_top_k", 0)
|
||||
|
||||
@timer
|
||||
def retrieve_not_reflected_memory(self, query: str):
|
||||
|
|
|
|||
|
|
@ -1,26 +0,0 @@
|
|||
from memory_scope.enumeration.action_status_enum import ActionStatusEnum
|
||||
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
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
|
||||
|
||||
|
||||
class StoreMemoryWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
if "query" in self.chat_kwargs:
|
||||
query = self.chat_kwargs["query"]
|
||||
query = query.strip()
|
||||
if not query:
|
||||
return
|
||||
|
||||
dt_handler = DatetimeHandler()
|
||||
node = MemoryNode(user_name=self.user_name,
|
||||
target_name=self.target_name,
|
||||
content=query,
|
||||
memory_type=MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
action_status=ActionStatusEnum.NEW.value,
|
||||
timestamp=dt_handler.timestamp)
|
||||
self.memory_handler.update_memories(nodes=node)
|
||||
else:
|
||||
self.memory_handler.update_memories(key=self.store_key)
|
||||
60
memory_scope/memory/worker/write/update_memory_worker.py
Normal file
60
memory_scope/memory/worker/write/update_memory_worker.py
Normal file
|
|
@ -0,0 +1,60 @@
|
|||
from typing import Dict, List
|
||||
|
||||
from memory_scope.enumeration.action_status_enum import ActionStatusEnum
|
||||
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
|
||||
|
||||
|
||||
class UpdateMemoryWorker(MemoryBaseWorker):
|
||||
|
||||
def _parse_params(self, **kwargs):
|
||||
self.method: str = kwargs.get("method", "")
|
||||
self.memory_key: str = kwargs.get("memory_key", "")
|
||||
self.expired_action: str = kwargs.get("expired_action", "")
|
||||
self.valid_action_dict: Dict[str, str] = kwargs.get("expired_action", {})
|
||||
|
||||
def from_query(self):
|
||||
if "query" not in self.chat_kwargs:
|
||||
return
|
||||
|
||||
query = self.chat_kwargs["query"].strip()
|
||||
if not query:
|
||||
return
|
||||
|
||||
dt_handler = DatetimeHandler()
|
||||
node = MemoryNode(user_name=self.user_name,
|
||||
target_name=self.target_name,
|
||||
content=query,
|
||||
memory_type=MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
action_status=ActionStatusEnum.NEW.value,
|
||||
timestamp=dt_handler.timestamp)
|
||||
return [node]
|
||||
|
||||
def from_memory_key(self):
|
||||
if not self.memory_key:
|
||||
return
|
||||
|
||||
return self.memory_handler.get_memories(keys=self.memory_key)
|
||||
|
||||
def modify_action_status(self):
|
||||
nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all")
|
||||
for node in nodes:
|
||||
if self.expired_action and node.store_status == StoreStatusEnum.EXPIRED.value:
|
||||
node.action_status = self.expired_action
|
||||
|
||||
elif node.memory_type in self.valid_action_dict:
|
||||
action_status = self.valid_action_dict[node.memory_type]
|
||||
if action_status:
|
||||
node.action_status = action_status
|
||||
|
||||
return nodes
|
||||
|
||||
def _run(self):
|
||||
method = self.method.strip()
|
||||
if not hasattr(self, method):
|
||||
self.logger.info(f"method={method} is missing!")
|
||||
return
|
||||
self.memory_handler.update_memories(nodes=getattr(self, method)())
|
||||
|
|
@ -22,7 +22,7 @@ class Message(BaseModel):
|
|||
|
||||
content: str = Field(..., description="The primary content of the message")
|
||||
|
||||
time_created: int = Field(int(datetime.datetime.now().timestamp()),
|
||||
time_created: int = Field(default_factory=lambda: int(datetime.datetime.now().timestamp()),
|
||||
description="Timestamp marking the message creation time")
|
||||
|
||||
memorized: bool = Field(False, description="Indicates if the message is flagged for memory retention")
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Dict, List, Set
|
||||
from typing import Dict, List
|
||||
|
||||
from memory_scope.enumeration.action_status_enum import ActionStatusEnum
|
||||
from memory_scope.enumeration.store_status_enum import StoreStatusEnum
|
||||
|
|
@ -57,53 +57,34 @@ class MemoryHandler(object):
|
|||
self._key_id_dict[key] = [n.memory_id for n in nodes]
|
||||
|
||||
def get_memories(self, keys: str | List[str]) -> List[MemoryNode]:
|
||||
"""
|
||||
Retrieves memory nodes associated with the given keys.
|
||||
|
||||
This method accepts a single key or a list of keys. For each key, it fetches the
|
||||
associated memory IDs from the context. If memory IDs are found, they are used to
|
||||
collect the corresponding MemoryNode objects from the contex_memory_dict. The final
|
||||
result is a list of unique MemoryNode values, avoiding duplicates.
|
||||
|
||||
Args:
|
||||
keys (str | List[str]): The key or list of keys to retrieve memories for.
|
||||
|
||||
Returns:
|
||||
List[MemoryNode]: A list of MemoryNode objects associated with the input keys.
|
||||
"""
|
||||
memories: Dict[str, MemoryNode] = {}
|
||||
|
||||
if isinstance(keys, str):
|
||||
keys = [keys]
|
||||
|
||||
for key in keys:
|
||||
if key not in self._key_id_dict:
|
||||
if key == "all":
|
||||
memories.update(self._id_memory_dict)
|
||||
break
|
||||
elif key not in self._key_id_dict:
|
||||
continue
|
||||
memory_ids: List[str] = self._key_id_dict[key]
|
||||
|
||||
memory_ids: List[str] = self._key_id_dict.get(key.strip())
|
||||
if memory_ids:
|
||||
memories.update({x: self._id_memory_dict[x] for x in memory_ids})
|
||||
|
||||
return list(memories.values())
|
||||
|
||||
def update_memories(self, key: str = "", nodes: MemoryNode | List[MemoryNode] = None):
|
||||
ids: Set[str] = set()
|
||||
|
||||
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]
|
||||
def update_memories(self, keys: str = "", nodes: MemoryNode | List[MemoryNode] = None):
|
||||
update_memories: Dict[str, MemoryNode] = {n.memory_id: n for n in self.get_memories(keys=keys)}
|
||||
|
||||
if nodes is not None:
|
||||
if isinstance(nodes, MemoryNode):
|
||||
nodes = [nodes]
|
||||
update_nodes.extend(nodes)
|
||||
update_memories.update({n.memory_id: n for n in nodes})
|
||||
|
||||
# Save collected nodes to memory store
|
||||
self._update_memories(update_nodes)
|
||||
self._update_memories(list(update_memories.values()))
|
||||
|
||||
def _update_memories(self, nodes: List[MemoryNode]):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue