From deedb1772a6383ccbc850f3ed94827aed157fbc2 Mon Sep 17 00:00:00 2001 From: hs Date: Fri, 21 Jun 2024 17:48:14 +0800 Subject: [PATCH] cli chat stream --- memory_scope/cli.py | 120 ++++++++++++++++++++++++++++++++++++++------ 1 file changed, 104 insertions(+), 16 deletions(-) diff --git a/memory_scope/cli.py b/memory_scope/cli.py index 96219085..54f7575e 100644 --- a/memory_scope/cli.py +++ b/memory_scope/cli.py @@ -2,7 +2,10 @@ import json import os from concurrent.futures import ThreadPoolExecutor from typing import Dict, Any, List - +import questionary +from rich.console import Console +import sys +import time import fire from memory_scope.chat.base_memory_chat import BaseMemoryChat @@ -10,7 +13,77 @@ from memory_scope.chat.global_context import GLOBAL_CONTEXT from memory_scope.enumeration.language_enum import LanguageEnum from memory_scope.enumeration.model_enum import ModelEnum from memory_scope.utils.logger import Logger -from memory_scope.utils.tool_functions import complete_config_name, init_instance_by_config +from memory_scope.utils.tool_functions import ( + complete_config_name, + init_instance_by_config, +) +from memory_scope.chat.memory_chat import MemoryChat + + +class CliMemoryChat(object): # object -> MemoryChat + + 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 + } + + def chat_with_memory(self, query): # for testing + return query + + def retrieve_all(self): # for testing + return "memory 1. 2. 3." + + def run(self): + console = Console() + 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": + 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 + + continue + + while True: + try: + with console.status("[bold cyan]Thinking..."): + messages = self.chat_with_memory(query=query) + console.print(messages) + 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 class CliJob(object): @@ -26,7 +99,9 @@ class CliJob(object): def init_memory_chat(self): for chat_name in self.config["chat_list"]: memory_chat_config = self.config[chat_name] - memory_chat: BaseMemoryChat = init_instance_by_config(memory_chat_config, chat_name=chat_name) + memory_chat: BaseMemoryChat = init_instance_by_config( + memory_chat_config, chat_name=chat_name + ) GLOBAL_CONTEXT.memory_chat_dict[chat_name] = memory_chat for worker_name in memory_chat.memory_service.get_worker_list(): @@ -41,15 +116,20 @@ class CliJob(object): if not model_name or model_name in GLOBAL_CONTEXT.model_dict: return - with open(os.path.join(self.config_base_dir, "model", complete_config_name(model_name))) as f: + with open( + os.path.join( + self.config_base_dir, "model", complete_config_name(model_name) + ) + ) as f: model_config = json.load(f) GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config(model_config) def init_workers(self): - """ load worker config & init workers - """ + """load worker config & init workers""" worker_config_name: str = self.config["workers"] - with open(os.path.join(self.config_base_dir, complete_config_name(worker_config_name))) as f: + with open( + os.path.join(self.config_base_dir, complete_config_name(worker_config_name)) + ) as f: worker_config_dict = json.load(f) for worker_name, worker_config in worker_config_dict.items(): @@ -60,10 +140,13 @@ class CliJob(object): for chat_name in chat_name_list: if chat_name not in GLOBAL_CONTEXT.worker_dict: GLOBAL_CONTEXT.worker_dict[chat_name] = {} - GLOBAL_CONTEXT.worker_dict[chat_name][worker_name] = init_instance_by_config( - worker_config, - suffix_name="worker", - **GLOBAL_CONTEXT.global_configs) + GLOBAL_CONTEXT.worker_dict[chat_name][worker_name] = ( + init_instance_by_config( + worker_config, + suffix_name="worker", + **GLOBAL_CONTEXT.global_configs, + ) + ) self.init_model(worker_config.get(ModelEnum.EMBEDDING_MODEL.value)) self.init_model(worker_config.get(ModelEnum.GENERATION_MODEL.value)) @@ -71,10 +154,13 @@ class CliJob(object): @staticmethod def set_global_config(): - """ TODO set global_configs & set apikey into env - """ - GLOBAL_CONTEXT.language = LanguageEnum(GLOBAL_CONTEXT.global_configs["language"]) - GLOBAL_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(GLOBAL_CONTEXT.global_configs["max_workers"])) + """TODO set global_configs & set apikey into env""" + GLOBAL_CONTEXT.language = LanguageEnum( + GLOBAL_CONTEXT.global_configs["language"] + ) + GLOBAL_CONTEXT.thread_pool = ThreadPoolExecutor( + max_workers=int(GLOBAL_CONTEXT.global_configs["max_workers"]) + ) def init_global_content_by_config(self): with open(complete_config_name(self.config_path)) as f: @@ -86,7 +172,9 @@ class CliJob(object): self.init_memory_chat() self.init_workers() - GLOBAL_CONTEXT.vector_store = init_instance_by_config(self.config["vector_store"]) + GLOBAL_CONTEXT.vector_store = init_instance_by_config( + self.config["vector_store"] + ) GLOBAL_CONTEXT.monitor = init_instance_by_config(self.config["monitor"]) @staticmethod