diff --git a/config/config.yaml b/config/config.yaml index bad2cac2..9c3c9906 100644 --- a/config/config.yaml +++ b/config/config.yaml @@ -5,7 +5,7 @@ global_config: open_ai_apikey: memory_chat: cli_memory_chat: - class: chat.cli_memory_chat + class: chat_v2.cli_memory_chat memory_service: memory_chat_service generation_model: dashscope_generation memory_service: @@ -17,20 +17,25 @@ memory_service: read_user_message: class: memory.operation.read_memory workflow: dummy + description: "read session messages of the user" read_memory: class: memory.operation.read_memory workflow: dummy + description: "read related memories of the user" list_memory: class: memory.operation.read_memory workflow: dummy + description: "read all memories of the user" write_memory: class: memory.operation.write_memory workflow: dummy + description: "write observation memories of the user" interval_time: 60 contextual_msg_count: 6 summary_memory: class: memory.operation.summary_memory workflow: dummy + description: "summary observation memories of the user" interval_time: 300 models: dashscope_generation: diff --git a/memory_scope/chat/memory_chat.py b/memory_scope/chat/memory_chat.py index 859758de..73f63fb6 100644 --- a/memory_scope/chat/memory_chat.py +++ b/memory_scope/chat/memory_chat.py @@ -29,17 +29,7 @@ class MemoryChat(BaseMemoryChat): ] return self._generation_model - @staticmethod - def get_system_prompt(related_memories: List[str], time_created: int) -> Message: - system_prompt = SYSTEM_PROMPT[GLOBAL_CONTEXT.language] - if related_memories: - memory_prompt = MEMORY_PROMPT[GLOBAL_CONTEXT.language] - system_prompt = "\n".join([system_prompt, memory_prompt] + related_memories) - return Message( - role=MessageRoleEnum.SYSTEM, - content=system_prompt.strip(), - time_created=time_created, - ) + def chat_with_memory(self, query: str): query = query.strip() diff --git a/memory_scope/chat_v2/base_memory_chat.py b/memory_scope/chat_v2/base_memory_chat.py index b5bda713..f647cb98 100644 --- a/memory_scope/chat_v2/base_memory_chat.py +++ b/memory_scope/chat_v2/base_memory_chat.py @@ -2,9 +2,6 @@ from abc import ABCMeta, abstractmethod class BaseMemoryChat(metaclass=ABCMeta): - def __init__(self, memory_service: str, **kwargs): - self.kwargs = kwargs - @abstractmethod def chat_with_memory(self, query: str): diff --git a/memory_scope/chat_v2/base_memory_service.py b/memory_scope/chat_v2/base_memory_service.py deleted file mode 100644 index 9cf3fd76..00000000 --- a/memory_scope/chat_v2/base_memory_service.py +++ /dev/null @@ -1,4 +0,0 @@ -class BaseMemoryService(object): - def __init__(self, **kwargs): - - self.kwargs = kwargs diff --git a/memory_scope/chat_v2/cli_memory_chat.py b/memory_scope/chat_v2/cli_memory_chat.py index 1444edcd..f9e60d8a 100644 --- a/memory_scope/chat_v2/cli_memory_chat.py +++ b/memory_scope/chat_v2/cli_memory_chat.py @@ -1,83 +1,124 @@ import datetime +import time +from typing import Dict, List import questionary from rich.console import Console -from .memory_chat import MemoryChat -from enumeration.message_role_enum import MessageRoleEnum -from scheme.message import Message +from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat +from memory_scope.chat_v2.global_context import G_CONTEXT +from memory_scope.enumeration.message_role_enum import MessageRoleEnum +from memory_scope.memory.service.base_memory_service import BaseMemoryService +from memory_scope.models.base_model import BaseModel +from memory_scope.prompts.memory_chat_prompt import MEMORY_PROMPT, SYSTEM_PROMPT +from memory_scope.scheme.message import Message -class CliMemoryChat(MemoryChat): - +class CliMemoryChat(BaseMemoryChat): USER_COMMANDS = { - "/exit": "exit the CLI", - "/memory": "print the current contents of agent memory", - "/retrieve": "retrieve related memory", - "/log": "log chat progress", - # TODO add more commands + "exit": "exit the CLI", + "help": "get cli commands help", } - def chat_with_memory(self, query): # for testing + def __init__(self, memory_service: str, generation_model: str, **kwargs): + self._memory_service: BaseMemoryService | str = memory_service + self._generation_model: BaseModel | str = generation_model + self.kwargs: dict = kwargs + + @property + def memory_service(self) -> BaseMemoryService: + if isinstance(self._memory_service, str): + self._memory_service = G_CONTEXT.memory_service_dict[self._memory_service] + self._memory_service.prepare_service() + return self._memory_service + + @property + def generation_model(self) -> BaseModel: + if isinstance(self._generation_model, str): + self._generation_model = G_CONTEXT.model_dict[self._generation_model] + return self._generation_model + + @staticmethod + def get_system_prompt(related_memories: List[str], time_created: int) -> Message: + system_prompt = SYSTEM_PROMPT[G_CONTEXT.language] + if related_memories: + memory_prompt = MEMORY_PROMPT[G_CONTEXT.language] + system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt] + related_memories]) + return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created) + + def chat_with_memory(self, query: str): query = query.strip() if not query: return time_created = int(datetime.datetime.now().timestamp()) - message = Message( - role=MessageRoleEnum.USER, content=query, time_created=time_created - ) - messages = [message] - return self.generation_model.call(messages=messages, stream=True) + new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created) + related_memories: List[str] = self.memory_service.do_operation("read_memory") + system_message: Message = self.get_system_prompt(related_memories, time_created) + return self.generation_model.call(messages=[system_message, new_message], stream=True) - def retrieve_all(self): # for testing - return "memory 1. 2. 3." - def run(self): - console = Console() +def run(self): + op_description_dict: Dict[str, str] = self.memory_service.get_op_description_dict() + self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()}) + + console = Console() + while True: + query = questionary.text( + "Please enter your message or command:", + multiline=False, + qmark=">", + ).ask() + + query: str = query.rstrip() + + if query == "": + console.print("Empty input received. Please try again!") + continue + + # handle cli / commands with memory ops + if query.startswith("/"): + query_split = query.lstrip("/").lower().split(" ") + query = query_split[0] + args = query_split[1:] + if query == "exit": + break + elif query == "help": + questionary.print("CLI commands", "bold") + for cmd, desc in self.USER_COMMANDS.items(): + questionary.print(cmd, "bold") + questionary.print(f" {desc}") + elif query in op_description_dict: + if not args: + result = self.memory_service.do_operation(op_name=query) + questionary.print(result) + + elif args[0].isdigit(): + refresh_time = int(args[0]) + while True: + time.sleep(refresh_time) + result = self.memory_service.do_operation(op_name=query) + questionary.print(result) + else: + console.print("unknown command received. Please try again!") + else: + console.print("unknown command received. Please try again!") + continue + while True: - query = questionary.text( - "Enter your message or command:", - multiline=False, - qmark=">", - ).ask() - - query = query.rstrip() - - if query == "": - console.print("Empty input received. Try again!") - continue - - # Handle CLI commands - if query.startswith("/"): - if query.lower() == "/exit": + try: + # with console.status("[bold cyan]Thinking..."): + for msg in self.chat_with_memory(query=query): + console.print(msg.delta, end="") + console.print() + break + except KeyboardInterrupt: + console.print("User interrupt occurred.") + retry = questionary.confirm("Retry chat_with_memory()?").ask() + if not retry: break - elif query.lower() == "/memory": - console.print(self.memory_service.retrieve_all()) - elif query.lower() == "/help": - questionary.print("CLI commands", "bold") - for cmd, desc in self.USER_COMMANDS.items(): - questionary.print(cmd, "bold") - questionary.print(f" {desc}") - - continue - - while True: - try: - # with console.status("[bold cyan]Thinking..."): - for msg in self.chat_with_memory(query=query): - console.print(msg.delta, end="") - console.print() + except Exception as e: + console.print(f"An exception occurred when running chat_with_memory(): {e}") + retry = questionary.confirm("Retry chat_with_memory()?").ask() + if not retry: break - except KeyboardInterrupt: - console.print("User interrupt occurred.") - retry = questionary.confirm("Retry chat_with_memory()?").ask() - if not retry: - break - except Exception as e: - console.print( - f"An exception occurred when running chat_with_memory(): {e}" - ) - retry = questionary.confirm("Retry chat_with_memory()?").ask() - if not retry: - break diff --git a/memory_scope/chat_v2/memory_chat.py b/memory_scope/chat_v2/memory_chat.py deleted file mode 100644 index 859758de..00000000 --- a/memory_scope/chat_v2/memory_chat.py +++ /dev/null @@ -1,67 +0,0 @@ -import datetime -from typing import List - -from .base_memory_chat import BaseMemoryChat -from .global_context import GLOBAL_CONTEXT -from enumeration.message_role_enum import MessageRoleEnum -from models.base_model import BaseModel -from prompts.memory_chat_prompt import SYSTEM_PROMPT, MEMORY_PROMPT -from scheme.message import Message -from .memory_service import MemoryService - - -class MemoryChat(BaseMemoryChat): - - def __init__(self, generation_model: str, history_msg_count: int, chat_name: str, **kwargs): - super().__init__(**kwargs) - self.memory_service = MemoryService(chat_name=chat_name, **kwargs) - self.generation_model_name: str = generation_model - self.history_msg_count: int = history_msg_count - - self._generation_model: BaseModel | None = None - self.history_message_list: List[Message] = [] - - @property - def generation_model(self): - if self._generation_model is None: - self._generation_model = GLOBAL_CONTEXT.model_dict[ - self.generation_model_name - ] - return self._generation_model - - @staticmethod - def get_system_prompt(related_memories: List[str], time_created: int) -> Message: - system_prompt = SYSTEM_PROMPT[GLOBAL_CONTEXT.language] - if related_memories: - memory_prompt = MEMORY_PROMPT[GLOBAL_CONTEXT.language] - system_prompt = "\n".join([system_prompt, memory_prompt] + related_memories) - return Message( - role=MessageRoleEnum.SYSTEM, - content=system_prompt.strip(), - time_created=time_created, - ) - - def chat_with_memory(self, query: str): - query = query.strip() - if not query: - return - - time_created = int(datetime.datetime.now().timestamp()) - new_message: Message = Message( - role=MessageRoleEnum.USER, content=query, time_created=time_created - ) - related_memories: List[str] = self.memory_service.retrieve(message=new_message) - system_message = self.get_system_prompt(related_memories, time_created) - self.history_message_list.append(new_message) - self.history_message_list = self.history_message_list[-self.history_msg_count :] - all_messages = [system_message] + self.history_message_list - # TODO at xian zhe - return self.generation_model.call(messages=all_messages, stream=True) - - def run(self): - self.memory_service.start_memory_backend() - while True: - query = input("wait for input:") - if query in ["stop", "停止"]: - break - self.chat_with_memory(query=query) diff --git a/memory_scope/chat_v2/memory_service.py b/memory_scope/chat_v2/memory_service.py deleted file mode 100644 index e9fca97a..00000000 --- a/memory_scope/chat_v2/memory_service.py +++ /dev/null @@ -1,70 +0,0 @@ -from constants.common_constants import RELATED_MEMORIES -from enumeration.memory_method_enum import MemoryMethodEnum -from scheme.message import Message -from utils.pipeline import Pipeline -from .base_memory_service import BaseMemoryService - - -class MemoryService(BaseMemoryService): - def __init__( - self, - chat_name: str, - retrieve_pipeline: str, - retrieve_all_pipeline: str, - summary_short_pipeline: str, - summary_long_pipeline: str, - summary_short_interval_time: int = 60, - summary_short_minimum_count: int = 5, - summary_long_interval_time: int = 60 * 5, - summary_long_minimum_count: int = 5 * 5, - **kwargs - ): - super().__init__(**kwargs) - self.retrieve_pipeline = Pipeline( - chat_name=chat_name, - memory_method_type=MemoryMethodEnum.RETRIEVE, - pipeline_str=retrieve_pipeline, - ) - - self.retrieve_all_pipeline = Pipeline( - chat_name=chat_name, - memory_method_type=MemoryMethodEnum.RETRIEVE_ALL, - pipeline_str=retrieve_all_pipeline, - ) - - self.summary_short_pipeline = Pipeline( - chat_name=chat_name, - memory_method_type=MemoryMethodEnum.SUMMARY_SHORT, - pipeline_str=summary_short_pipeline, - loop_interval_time=summary_short_interval_time, - loop_minimum_count=summary_short_minimum_count, - ) - - self.summary_long_pipeline = Pipeline( - chat_name=chat_name, - memory_method_type=MemoryMethodEnum.SUMMARY_LONG, - pipeline_str=summary_long_pipeline, - loop_interval_time=summary_long_interval_time, - loop_minimum_count=summary_long_minimum_count, - ) - - def retrieve(self, message: Message): - self.retrieve_pipeline.submit_message(message, with_lock=False) - self.summary_short_pipeline.submit_message(message) - self.summary_long_pipeline.submit_message(message) - return self.retrieve_pipeline.run(RELATED_MEMORIES) - - def retrieve_all(self): - return self.retrieve_all_pipeline.run(RELATED_MEMORIES) - - def start_memory_backend(self): - self.summary_short_pipeline.start_loop_run() - self.summary_long_pipeline.start_loop_run() - - def get_worker_list(self) -> list: - worker_set = set() - worker_set.update(self.retrieve_pipeline.worker_set) - worker_set.update(self.retrieve_all_pipeline.worker_set) - worker_set.update(self.summary_short_pipeline.worker_set) - worker_set.update(self.summary_long_pipeline.worker_set) - return sorted(worker_set) diff --git a/memory_scope/memory/operation/base_operation.py b/memory_scope/memory/operation/base_operation.py index 70d1d103..3f531b4e 100644 --- a/memory_scope/memory/operation/base_operation.py +++ b/memory_scope/memory/operation/base_operation.py @@ -7,6 +7,10 @@ OPERATION_TYPE = Literal["frontend", "backend"] class BaseOperation(metaclass=ABCMeta): operation_type: OPERATION_TYPE = "frontend" + def __init__(self, name: str, description: str = "", **kwargs): + self.name: str = name + self.description: str = description + def init_workflow(self): pass diff --git a/memory_scope/memory/operation/read_memory.py b/memory_scope/memory/operation/read_memory.py index 8ffe24d2..89956137 100644 --- a/memory_scope/memory/operation/read_memory.py +++ b/memory_scope/memory/operation/read_memory.py @@ -9,8 +9,15 @@ from memory_scope.scheme.message import Message class ReadMemory(BaseWorkflow, BaseOperation): operation_type: OPERATION_TYPE = "frontend" - def __init__(self, chat_messages: List[Message], his_msg_count: int = 0, contextual_msg_count: int = 0, **kwargs): - super().__init__(**kwargs) + def __init__(self, + name: str, + description: str, + chat_messages: List[Message], + his_msg_count: int = 0, + contextual_msg_count: int = 0, + **kwargs): + super().__init__(name=name, **kwargs) + BaseOperation.__init__(self, name=name, description=description) self.chat_messages: List[Message] = chat_messages self.his_msg_count: int = his_msg_count self.contextual_msg_count: int = contextual_msg_count diff --git a/memory_scope/memory/operation/summary_memory.py b/memory_scope/memory/operation/summary_memory.py index 3deba4ba..30dc96c8 100644 --- a/memory_scope/memory/operation/summary_memory.py +++ b/memory_scope/memory/operation/summary_memory.py @@ -1,6 +1,7 @@ import time from memory_scope.chat_v2.global_context import G_CONTEXT +from memory_scope.constants.common_constants import RESULT from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE from memory_scope.memory.operation.base_workflow import BaseWorkflow @@ -8,8 +9,13 @@ from memory_scope.memory.operation.base_workflow import BaseWorkflow class SummaryMemory(BaseWorkflow, BaseOperation): operation_type: OPERATION_TYPE = "backend" - def __init__(self, interval_time: int = 300, **kwargs): - super().__init__(**kwargs) + def __init__(self, + name: str, + description: str, + interval_time: int = 300, + **kwargs): + super().__init__(name=name, **kwargs) + BaseOperation.__init__(self, name=name, description=description) self.interval_time: int = interval_time @@ -25,8 +31,10 @@ class SummaryMemory(BaseWorkflow, BaseOperation): self._operation_status_run = True self.run_workflow() + result = self.context.get(RESULT) self.context.clear() self._operation_status_run = False + return result def _loop_operation(self): while self._loop_switch: diff --git a/memory_scope/memory/operation/write_memory.py b/memory_scope/memory/operation/write_memory.py index 9d6dce17..1e9070b7 100644 --- a/memory_scope/memory/operation/write_memory.py +++ b/memory_scope/memory/operation/write_memory.py @@ -2,7 +2,7 @@ import time from typing import List from memory_scope.chat_v2.global_context import G_CONTEXT -from memory_scope.constants.common_constants import CHAT_MESSAGES +from memory_scope.constants.common_constants import CHAT_MESSAGES, RESULT from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE from memory_scope.memory.operation.base_workflow import BaseWorkflow from memory_scope.scheme.message import Message @@ -12,13 +12,17 @@ class WriteMemory(BaseOperation, BaseWorkflow): operation_type: OPERATION_TYPE = "backend" def __init__(self, + name: str, + description: str, chat_messages: List[Message], his_msg_count: int = 0, message_lock=None, interval_time: int = 60, contextual_msg_count: int = 6, **kwargs): - super().__init__(**kwargs) + + super().__init__(name=name, **kwargs) + BaseOperation.__init__(self, name=name, description=description) self.chat_messages: List[Message] = chat_messages self.his_msg_count: int = his_msg_count @@ -54,9 +58,11 @@ class WriteMemory(BaseOperation, BaseWorkflow): max_count = not_memorized_size + self.his_msg_count self.context[CHAT_MESSAGES] = [x.copy() for x in self.chat_messages[-max_count:]] self.run_workflow() + result = self.context.get(RESULT) self.context.clear() self.set_memorized() self._operation_status_run = False + return result def _loop_operation(self): while self._loop_switch: diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index 79299265..e2d2a452 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -1,5 +1,7 @@ from abc import ABCMeta, abstractmethod +from typing import List, Dict +from memory_scope.scheme.message import Message from memory_scope.utils.logger import Logger @@ -8,6 +10,16 @@ class BaseMemoryService(metaclass=ABCMeta): self.logger = Logger.get_logger() self.kwargs = kwargs + def submit_messages(self, messages: List[Message] | Message): + pass + + def prepare_service(self): + pass + @abstractmethod def do_operation(self, op_name: str): pass + + @abstractmethod + def get_op_description_dict(self) -> Dict[str, str]: + pass diff --git a/memory_scope/memory/service/chat_memory_service.py b/memory_scope/memory/service/chat_memory_service.py index f822c0dd..5bb9895c 100644 --- a/memory_scope/memory/service/chat_memory_service.py +++ b/memory_scope/memory/service/chat_memory_service.py @@ -59,3 +59,6 @@ class ChatMemoryService(BaseMemoryService): self.logger.warning(f"op_name={op_name} is not inited!") return return self.op_dict[op_name].run_operation() + + def get_op_description_dict(self) -> Dict[str, str]: + return {k: v.description for k, v in self.op_dict.items()}