mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
[dev] update cli memory chat
This commit is contained in:
parent
6b3fb6bbc9
commit
aeeb6012da
7 changed files with 24 additions and 9 deletions
|
|
@ -21,9 +21,9 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
}
|
||||
|
||||
def __init__(self, memory_service: str, generation_model: str, **kwargs):
|
||||
super().__init__(**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:
|
||||
|
|
@ -43,7 +43,9 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
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])
|
||||
all_prompt_list = [system_prompt, memory_prompt]
|
||||
all_prompt_list.extend(related_memories)
|
||||
system_prompt = "\n".join([x.strip() for x in all_prompt_list])
|
||||
return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created)
|
||||
|
||||
def chat_with_memory(self, query: str):
|
||||
|
|
@ -53,7 +55,7 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
|
||||
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.do_operation("read_memory")
|
||||
related_memories: List[str] = self.memory_service.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)
|
||||
|
||||
|
|
|
|||
|
|
@ -20,3 +20,6 @@ class BaseOperation(metaclass=ABCMeta):
|
|||
|
||||
def run_operation_backend(self):
|
||||
pass
|
||||
|
||||
def stop_operation_backend(self):
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ class SummaryMemory(BaseWorkflow, BaseOperation):
|
|||
|
||||
self._operation_status_run: bool = False
|
||||
self._loop_switch: bool = False
|
||||
self._run_thread = None
|
||||
|
||||
def init_workflow(self):
|
||||
self.init_workers()
|
||||
|
|
@ -44,4 +45,7 @@ class SummaryMemory(BaseWorkflow, BaseOperation):
|
|||
def run_operation_backend(self):
|
||||
if not self._loop_switch:
|
||||
self._loop_switch = True
|
||||
return G_CONTEXT.thread_pool.submit(self._loop_operation)
|
||||
self._run_thread = G_CONTEXT.thread_pool.submit(self._loop_operation)
|
||||
|
||||
def stop_operation_backend(self):
|
||||
self._loop_switch = False
|
||||
|
|
|
|||
|
|
@ -73,3 +73,6 @@ class WriteMemory(BaseOperation, BaseWorkflow):
|
|||
if not self._loop_switch:
|
||||
self._loop_switch = True
|
||||
return G_CONTEXT.thread_pool.submit(self._loop_operation)
|
||||
|
||||
def stop_operation_backend(self):
|
||||
self._loop_switch = False
|
||||
|
|
|
|||
|
|
@ -6,7 +6,8 @@ from memory_scope.utils.logger import Logger
|
|||
|
||||
|
||||
class BaseMemoryService(metaclass=ABCMeta):
|
||||
def __init__(self, **kwargs):
|
||||
def __init__(self, read_memory_key: str = "read_memory", **kwargs):
|
||||
self.read_memory_key: str = read_memory_key
|
||||
self.logger = Logger.get_logger()
|
||||
self.kwargs = kwargs
|
||||
|
||||
|
|
@ -23,3 +24,6 @@ class BaseMemoryService(metaclass=ABCMeta):
|
|||
@abstractmethod
|
||||
def get_op_description_dict(self) -> Dict[str, str]:
|
||||
pass
|
||||
|
||||
def read_memory(self):
|
||||
return self.do_operation(self.read_memory_key)
|
||||
|
|
|
|||
|
|
@ -43,15 +43,14 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
self.logger.info(f"----- worker_{self.name}_end cost={t.cost_str}-----")
|
||||
|
||||
def get_context(self, key: str, default=None):
|
||||
return self.context_dict.get(key, default)
|
||||
return self.context.get(key, default)
|
||||
|
||||
def set_context(self, key: str, value: Any):
|
||||
if self.is_multi_thread:
|
||||
with self.context_lock:
|
||||
self.context_dict[key] = value
|
||||
else:
|
||||
self.context_dict[key] = value
|
||||
self.context[key] = value
|
||||
|
||||
def __getattr__(self, key):
|
||||
# raise exception if not exists
|
||||
return self.kwargs[key]
|
||||
|
|
|
|||
|
|
@ -5,4 +5,4 @@ from memory_scope.memory.worker.base_worker import BaseWorker
|
|||
class DummyWorker(BaseWorker):
|
||||
def _run(self):
|
||||
self.set_context(RESULT, ["test 123"])
|
||||
self.logger.info("enter dummy worker!")
|
||||
self.logger.info("enter dummy worker!")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue