diff options
Diffstat (limited to 'supergenerator.py')
| -rw-r--r-- | supergenerator.py | 48 |
1 files changed, 37 insertions, 11 deletions
diff --git a/supergenerator.py b/supergenerator.py index 3695da1..6b67742 100644 --- a/supergenerator.py +++ b/supergenerator.py @@ -15,6 +15,7 @@ from config import * last_command_time = {} last_error_time = {} last_block_time = {} +alo_command = {} user_contexts = {} @@ -143,7 +144,6 @@ def get_all_user_ids(only_notify=False): return ids - def set_user_config( user_id, model_name, @@ -151,7 +151,8 @@ def set_user_config( call_mode=DEFAULT_CALL, prompt_name=DEFAULT_PROMPT, notify_mode=DEFAULT_NOTIFY, - markdown="", + markdown=DEFAULT_MARKDOWN, + prompts={}, ): with shelve.open("users") as db: db[str(user_id)] = { @@ -161,9 +162,23 @@ def set_user_config( "prompt": prompt_name, "notify": notify_mode, "markdown": markdown, + "prompts": prompts, } +def delete_user_from_db(user_id): + with shelve.open("users") as db: + if str(user_id) in db: + del db[str(user_id)] + + +def reset_user_setting(user_id, key, value): + with shelve.open("users") as db: + user = db[str(user_id)] + user[key] = value + db[str(user_id)] = user + + def ensure_user( user_id, model_name=DEFAULT_MODEL, @@ -171,7 +186,8 @@ def ensure_user( call_mode=DEFAULT_CALL, prompt_name=DEFAULT_PROMPT, notify_mode=DEFAULT_NOTIFY, - markdown="" + markdown=DEFAULT_MARKDOWN, + prompts={}, ): with shelve.open("users") as db: key = str(user_id) @@ -184,6 +200,7 @@ def ensure_user( "prompt": prompt_name, "notify": notify_mode, "markdown": markdown, + "prompts": prompts, } @@ -196,7 +213,8 @@ def get_user_config(user_id): "call": DEFAULT_CALL, "prompt": DEFAULT_PROMPT, "notify": DEFAULT_NOTIFY, - "markdown": "", + "markdown": DEFAULT_MARKDOWN, + "prompts": {}, } updated = False for key, value in defaults.items(): @@ -225,14 +243,22 @@ async def generate( settings = MODELS.get(model[0]).get(model[1]) - stream_cf = config["stream"] - system_prompt_name = config["prompt"] - if system_prompt_name in SYSTEM_PROMPTS: - system_prompt = SYSTEM_PROMPTS[system_prompt_name] + ADD_TO_PROMPT - if config["markdown"] == "new": - system_prompt += NEW_MD + stream_cf = config.get("stream", "none") + system_prompt_name = config.get("prompt", "") + user_prompts = config.get("prompts", {}) + prompts = SYSTEM_PROMPTS | user_prompts + + if system_prompt_name in prompts: + is_user_prompt = False + if system_prompt_name in SYSTEM_PROMPTS: + system_prompt = prompts[system_prompt_name]["text"] + ADD_TO_PROMPT + if config["markdown"] == "new": + system_prompt += NEW_MD + else: + system_prompt += OLD_MD else: - system_prompt += OLD_MD + is_user_prompt = True + system_prompt = prompts[system_prompt_name]["text"] else: text = "❌произошла ошибка, возможные проблемы:\n" \ f"1. у вас выбран удалённый сист промпт (`{system_prompt_name}`)\n" \ |