aboutsummaryrefslogtreecommitdiff
path: root/supergenerator.py
diff options
context:
space:
mode:
Diffstat (limited to 'supergenerator.py')
-rw-r--r--supergenerator.py48
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" \