[dev] update cli memory chat

This commit is contained in:
jinli.yl 2024-06-27 11:14:16 +08:00
parent 6b3fb6bbc9
commit aeeb6012da
7 changed files with 24 additions and 9 deletions

View file

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

View file

@ -20,3 +20,6 @@ class BaseOperation(metaclass=ABCMeta):
def run_operation_backend(self):
pass
def stop_operation_backend(self):
pass

View file

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

View file

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

View file

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

View file

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

View file

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