[dev] modify op name & service base func

This commit is contained in:
jinli.yl 2024-07-13 23:02:29 +08:00
parent 4c91212b16
commit c1957af6af
17 changed files with 122 additions and 142 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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