refac
This commit is contained in:
@@ -284,7 +284,7 @@ async def generate_function_chat_completion(request, form_data, user, models: di
|
||||
if params:
|
||||
system = params.pop('system', None)
|
||||
form_data = apply_model_params_to_body_openai(params, form_data)
|
||||
form_data = apply_system_prompt_to_body(system, form_data, metadata, user)
|
||||
form_data = await apply_system_prompt_to_body(system, form_data, metadata, user)
|
||||
|
||||
pipe_id = get_pipe_id(form_data)
|
||||
function_module = await get_function_module_by_id(request, pipe_id)
|
||||
|
||||
@@ -1127,7 +1127,7 @@ async def generate_chat_completion(
|
||||
|
||||
payload = apply_model_params_to_body_ollama(params, payload)
|
||||
if not bypass_system_prompt:
|
||||
payload = apply_system_prompt_to_body(system, payload, metadata, user)
|
||||
payload = await apply_system_prompt_to_body(system, payload, metadata, user)
|
||||
|
||||
await check_model_access(user, model_info, bypass_filter)
|
||||
else:
|
||||
@@ -1282,7 +1282,7 @@ async def generate_openai_chat_completion(
|
||||
system = params.pop('system', None)
|
||||
|
||||
payload = apply_model_params_to_body_openai(params, payload)
|
||||
payload = apply_system_prompt_to_body(system, payload, metadata, user)
|
||||
payload = await apply_system_prompt_to_body(system, payload, metadata, user)
|
||||
|
||||
await check_model_access(user, model_info)
|
||||
else:
|
||||
|
||||
@@ -1120,7 +1120,7 @@ async def generate_chat_completion(
|
||||
|
||||
payload = apply_model_params_to_body_openai(params, payload)
|
||||
if not bypass_system_prompt:
|
||||
payload = apply_system_prompt_to_body(system, payload, metadata, user)
|
||||
payload = await apply_system_prompt_to_body(system, payload, metadata, user)
|
||||
|
||||
await check_model_access(user, model_info, bypass_filter)
|
||||
else:
|
||||
|
||||
@@ -360,7 +360,7 @@ async def execute_automation(app, automation: AutomationModel) -> None:
|
||||
await _record_run(automation.id, 'error', error='User not found')
|
||||
return
|
||||
|
||||
prompt = prompt_template(automation.data['prompt'], user)
|
||||
prompt = await prompt_template(automation.data['prompt'], user)
|
||||
model_id = automation.data['model_id']
|
||||
terminal_config = automation.data.get('terminal')
|
||||
|
||||
|
||||
@@ -969,7 +969,7 @@ def get_source_context(sources: list, source_ids: dict = None, include_content:
|
||||
return context_string
|
||||
|
||||
|
||||
def apply_source_context_to_messages(
|
||||
async def apply_source_context_to_messages(
|
||||
request: Request,
|
||||
messages: list,
|
||||
sources: list,
|
||||
@@ -995,13 +995,13 @@ def apply_source_context_to_messages(
|
||||
|
||||
if RAG_SYSTEM_CONTEXT:
|
||||
return add_or_update_system_message(
|
||||
rag_template(request.app.state.config.RAG_TEMPLATE, context, user_message),
|
||||
await rag_template(request.app.state.config.RAG_TEMPLATE, context, user_message),
|
||||
messages,
|
||||
append=True,
|
||||
)
|
||||
else:
|
||||
return add_or_update_user_message(
|
||||
rag_template(request.app.state.config.RAG_TEMPLATE, context, user_message),
|
||||
await rag_template(request.app.state.config.RAG_TEMPLATE, context, user_message),
|
||||
messages,
|
||||
append=False,
|
||||
)
|
||||
@@ -2310,7 +2310,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
|
||||
system_message = get_system_message(form_data.get('messages', []))
|
||||
if system_message: # Chat Controls/User Settings
|
||||
try:
|
||||
form_data = apply_system_prompt_to_body(
|
||||
form_data = await apply_system_prompt_to_body(
|
||||
system_message.get('content'), form_data, metadata, user, replace=True
|
||||
) # Required to handle system prompt variables
|
||||
except Exception:
|
||||
@@ -2368,7 +2368,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
|
||||
|
||||
if folder and folder.data:
|
||||
if 'system_prompt' in folder.data:
|
||||
form_data = apply_system_prompt_to_body(folder.data['system_prompt'], form_data, metadata, user)
|
||||
form_data = await apply_system_prompt_to_body(folder.data['system_prompt'], form_data, metadata, user)
|
||||
if 'files' in folder.data:
|
||||
if metadata.get('params', {}).get('function_calling') != 'native':
|
||||
form_data['files'] = [
|
||||
@@ -2854,7 +2854,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
|
||||
|
||||
# If context is not empty, insert it into the messages
|
||||
if sources and prompt:
|
||||
form_data['messages'] = apply_source_context_to_messages(request, form_data['messages'], sources, prompt)
|
||||
form_data['messages'] = await apply_source_context_to_messages(request, form_data['messages'], sources, prompt)
|
||||
|
||||
# If there are citations, add them to the data_items
|
||||
sources = [
|
||||
@@ -3012,6 +3012,7 @@ async def background_tasks_handler(ctx):
|
||||
tasks = ctx['tasks']
|
||||
event_emitter = ctx['event_emitter']
|
||||
|
||||
|
||||
message = None
|
||||
messages = []
|
||||
|
||||
@@ -4676,7 +4677,7 @@ async def streaming_chat_response_handler(response, ctx):
|
||||
)
|
||||
source_context = source_context.strip()
|
||||
if source_context:
|
||||
rag_content = rag_template(
|
||||
rag_content = await rag_template(
|
||||
request.app.state.config.RAG_TEMPLATE,
|
||||
source_context,
|
||||
user_message,
|
||||
|
||||
@@ -13,7 +13,7 @@ import json
|
||||
# What goes out cannot be taken back. Let it be shaped
|
||||
# well before it leaves this place.
|
||||
# inplace function: form_data is modified
|
||||
def apply_system_prompt_to_body(
|
||||
async def apply_system_prompt_to_body(
|
||||
system: Optional[str],
|
||||
form_data: dict,
|
||||
metadata: Optional[dict] = None,
|
||||
@@ -30,7 +30,7 @@ def apply_system_prompt_to_body(
|
||||
system = prompt_variables_template(system, variables)
|
||||
|
||||
# Legacy (API Usage)
|
||||
system = prompt_template(system, user)
|
||||
system = await prompt_template(system, user)
|
||||
|
||||
if replace:
|
||||
form_data['messages'] = replace_system_message_content(system, form_data.get('messages', []))
|
||||
|
||||
@@ -35,7 +35,7 @@ def prompt_variables_template(template: str, variables: dict[str, str]) -> str:
|
||||
return template
|
||||
|
||||
|
||||
def prompt_template(template: str, user: Optional[Any] = None) -> str:
|
||||
async def prompt_template(template: str, user: Optional[Any] = None) -> str:
|
||||
USER_VARIABLES = {}
|
||||
|
||||
if user:
|
||||
@@ -58,6 +58,19 @@ def prompt_template(template: str, user: Optional[Any] = None) -> str:
|
||||
except Exception as e:
|
||||
pass
|
||||
|
||||
# Resolve user groups from DB only when the template uses {{USER_GROUPS}}
|
||||
groups = ''
|
||||
if '{{USER_GROUPS}}' in template:
|
||||
user_id = user.get('id')
|
||||
if user_id:
|
||||
try:
|
||||
from open_webui.models.groups import Groups
|
||||
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id)
|
||||
groups = ', '.join(g.name for g in user_groups)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
USER_VARIABLES = {
|
||||
'name': str(user.get('name')),
|
||||
'email': str(user.get('email')),
|
||||
@@ -66,6 +79,7 @@ def prompt_template(template: str, user: Optional[Any] = None) -> str:
|
||||
'gender': str(user.get('gender')),
|
||||
'birth_date': str(birth_date),
|
||||
'age': str(age),
|
||||
'groups': groups,
|
||||
}
|
||||
|
||||
# Get the current date
|
||||
@@ -88,6 +102,7 @@ def prompt_template(template: str, user: Optional[Any] = None) -> str:
|
||||
template = template.replace('{{USER_BIRTH_DATE}}', USER_VARIABLES.get('birth_date', 'Unknown'))
|
||||
template = template.replace('{{USER_AGE}}', str(USER_VARIABLES.get('age', 'Unknown')))
|
||||
template = template.replace('{{USER_LOCATION}}', USER_VARIABLES.get('location', 'Unknown'))
|
||||
template = template.replace('{{USER_GROUPS}}', USER_VARIABLES.get('groups', ''))
|
||||
|
||||
return template
|
||||
|
||||
@@ -243,11 +258,11 @@ def replace_messages_variable(template: str, messages: Optional[list[dict]] = No
|
||||
|
||||
# Let the context given here not distort the question,
|
||||
# but illuminate it, so that the answer serves the one who asked.
|
||||
def rag_template(template: str, context: str, query: str):
|
||||
async def rag_template(template: str, context: str, query: str):
|
||||
if template.strip() == '':
|
||||
template = DEFAULT_RAG_TEMPLATE
|
||||
|
||||
template = prompt_template(template)
|
||||
template = await prompt_template(template)
|
||||
|
||||
if '[context]' not in template and '{{CONTEXT}}' not in template:
|
||||
log.debug("WARNING: The RAG template does not contain the '[context]' or '{{CONTEXT}}' placeholder.")
|
||||
@@ -282,51 +297,51 @@ def rag_template(template: str, context: str, query: str):
|
||||
return template
|
||||
|
||||
|
||||
def title_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
|
||||
async def title_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
|
||||
prompt = get_last_user_message(messages)
|
||||
template = replace_prompt_variable(template, prompt)
|
||||
template = replace_messages_variable(template, messages)
|
||||
|
||||
template = prompt_template(template, user)
|
||||
template = await prompt_template(template, user)
|
||||
|
||||
return template
|
||||
|
||||
|
||||
def follow_up_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
|
||||
async def follow_up_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
|
||||
prompt = get_last_user_message(messages)
|
||||
template = replace_prompt_variable(template, prompt)
|
||||
template = replace_messages_variable(template, messages)
|
||||
|
||||
template = prompt_template(template, user)
|
||||
template = await prompt_template(template, user)
|
||||
return template
|
||||
|
||||
|
||||
def tags_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
|
||||
async def tags_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
|
||||
prompt = get_last_user_message(messages)
|
||||
template = replace_prompt_variable(template, prompt)
|
||||
template = replace_messages_variable(template, messages)
|
||||
|
||||
template = prompt_template(template, user)
|
||||
template = await prompt_template(template, user)
|
||||
return template
|
||||
|
||||
|
||||
def image_prompt_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
|
||||
async def image_prompt_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
|
||||
prompt = get_last_user_message(messages)
|
||||
template = replace_prompt_variable(template, prompt)
|
||||
template = replace_messages_variable(template, messages)
|
||||
|
||||
template = prompt_template(template, user)
|
||||
template = await prompt_template(template, user)
|
||||
return template
|
||||
|
||||
|
||||
def emoji_generation_template(template: str, prompt: str, user: Optional[Any] = None) -> str:
|
||||
async def emoji_generation_template(template: str, prompt: str, user: Optional[Any] = None) -> str:
|
||||
template = replace_prompt_variable(template, prompt)
|
||||
template = prompt_template(template, user)
|
||||
template = await prompt_template(template, user)
|
||||
|
||||
return template
|
||||
|
||||
|
||||
def autocomplete_generation_template(
|
||||
async def autocomplete_generation_template(
|
||||
template: str,
|
||||
prompt: str,
|
||||
messages: Optional[list[dict]] = None,
|
||||
@@ -337,16 +352,16 @@ def autocomplete_generation_template(
|
||||
template = replace_prompt_variable(template, prompt)
|
||||
template = replace_messages_variable(template, messages)
|
||||
|
||||
template = prompt_template(template, user)
|
||||
template = await prompt_template(template, user)
|
||||
return template
|
||||
|
||||
|
||||
def query_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
|
||||
async def query_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
|
||||
prompt = get_last_user_message(messages)
|
||||
template = replace_prompt_variable(template, prompt)
|
||||
template = replace_messages_variable(template, messages)
|
||||
|
||||
template = prompt_template(template, user)
|
||||
template = await prompt_template(template, user)
|
||||
return template
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user