mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-07 08:26:06 +00:00
[dev] add operation description
This commit is contained in:
parent
3959360fa0
commit
900ba7b9eb
13 changed files with 157 additions and 225 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -1,4 +0,0 @@
|
|||
class BaseMemoryService(object):
|
||||
def __init__(self, **kwargs):
|
||||
|
||||
self.kwargs = kwargs
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue