diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index b5b7c964..53cfa650 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -11,7 +11,8 @@ 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 -from ..models.model_response import ModelResponse, ModelResponseGen +from memory_scope.scheme.model_response import ModelResponse, ModelResponseGen +from memory_scope.utils.logger import Logger class CliMemoryChat(BaseMemoryChat): @@ -27,6 +28,8 @@ class CliMemoryChat(BaseMemoryChat): self.stream: bool = stream self.kwargs: dict = kwargs + self.logger = Logger.get_logger() + @property def memory_service(self) -> BaseMemoryService: if isinstance(self._memory_service, str): @@ -66,75 +69,84 @@ class CliMemoryChat(BaseMemoryChat): self.memory_service.add_messages(result.message) + def process_commands(self, query: str) -> bool: + continue_run = True + query_split = query.lstrip("/").lower().split(" ") + query = query_split[0] + args = query_split[1:] + if query == "exit": + self.memory_service.stop_service() + continue_run = False + + 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 == "stream": + self.stream = bool(args[0]) + questionary.print(f"stream: {self.stream}") + + elif query in self.memory_service.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, flush=True) + + else: + questionary.print("unknown command received. Please try again!") + + else: + questionary.print("unknown command received. Please try again!") + + return continue_run + def run(self): self.USER_COMMANDS.update({f"/{k}": v for k, v in self.memory_service.op_description_dict.items()}) while True: - query = questionary.text( - "Please enter your message or command:", - multiline=False, - qmark=">", - ).ask() + try: + query = questionary.text( + message="Please enter your message or command:", + multiline=False, + qmark=">", + ).ask() + query: str = query.strip() - query: str = query.rstrip() + if query == "": + questionary.print("Empty input received. Please try again!") + continue - if query == "": - 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") - print(f" {desc}") - elif query == "stream": - questionary.print(f"stream: {self.stream}") - self.stream = ~self.stream - elif query in self.memory_service.op_description_dict: - if not args: - result = self.memory_service.do_operation(op_name=query) - print(result) - - elif args[0].isdigit(): - refresh_time = int(args[0]) - try: - while True: - time.sleep(refresh_time) - result = self.memory_service.do_operation(op_name=query) - print(result, flush=True) - except KeyboardInterrupt: - print("stop refresh!") + # handle cli / commands with memory ops + if query.startswith("/"): + if self.process_commands(query=query): + continue else: - print("unknown command received. Please try again!") - else: - print("unknown command received. Please try again!") - continue + break - while True: - # try: if self.stream: for msg in self.chat_with_memory(query=query): - print(msg.delta, end="") - print() + questionary.print(msg.delta, end="") + questionary.print("") else: msg = self.chat_with_memory(query=query) - print(msg.message.content) - break - # except KeyboardInterrupt: - # questionary.print("User interrupt occurred.") - # retry = questionary.confirm("Retry chat_with_memory()?").ask() - # if not retry: - # break - # except Exception as e: - # questionary.print(f"An exception occurred when running chat_with_memory(): {e}") - # # retry = questionary.confirm("Retry chat_with_memory()?").ask() - # # if not retry: - # # break - # raise e + questionary.print(msg.message.content) + + except KeyboardInterrupt: + questionary.print("User interrupt occurred.") + is_exit = questionary.confirm("continue exit?").ask() + if is_exit: + self.memory_service.stop_service() + break + + except Exception as e: + questionary.print(f"An exception occurred when running cli memory chat. args={e.args}") + self.logger.exception(f"An exception occurred when running cli memory chat. args={e.args}") + continue diff --git a/memory_scope/memory/service/base_memory_service.py b/memory_scope/memory/service/base_memory_service.py index 2a8e78bc..41a8c141 100644 --- a/memory_scope/memory/service/base_memory_service.py +++ b/memory_scope/memory/service/base_memory_service.py @@ -47,3 +47,6 @@ class BaseMemoryService(metaclass=ABCMeta): def read_memory(self): assert self.read_memory_key in self._operation_dict, f"op={self.read_memory_key} is not inited!" return self.do_operation(self.read_memory_key) + + def stop_service(self): + pass diff --git a/tests/models/test_models_lli_embedding.py b/tests/models/test_models_lli_embedding.py index 5bd3f28b..c71236f5 100644 --- a/tests/models/test_models_lli_embedding.py +++ b/tests/models/test_models_lli_embedding.py @@ -8,6 +8,7 @@ import unittest from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel from memory_scope.utils.logger import Logger + class TestLLIEmbedding(unittest.TestCase): """Tests for LlamaIndexEmbeddingModel""" @@ -18,17 +19,20 @@ class TestLLIEmbedding(unittest.TestCase): "clazz": "models.base_embedding_model" } self.emb = LlamaIndexEmbeddingModel(**config) + print() self.logger = Logger.get_logger() def test_single_embedding(self): text = "您吃了吗?" result = self.emb.call(text=text) - self.logger.info(result) + self.logger.info(result.m_type) + self.logger.info(len(result.embedding_results)) def test_batch_embedding(self): texts = ["您吃了吗?", "吃了吗您?"] result = self.emb.call(text=texts) + print() self.logger.info(result) def test_async_embedding(self): @@ -36,4 +40,5 @@ class TestLLIEmbedding(unittest.TestCase): "吃了吗您?"] # 调用异步函数并等待其结果 result = asyncio.run(self.emb.async_call(text=texts)) + print() self.logger.info(result)