diff --git a/backend/open_webui/functions.py b/backend/open_webui/functions.py index 1e032759ea..a3f99bb182 100644 --- a/backend/open_webui/functions.py +++ b/backend/open_webui/functions.py @@ -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) diff --git a/backend/open_webui/routers/ollama.py b/backend/open_webui/routers/ollama.py index 8311fee5d4..01fbf10f4f 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -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: diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index ab2eec9527..b85a8f83b3 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -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: diff --git a/backend/open_webui/utils/automations.py b/backend/open_webui/utils/automations.py index 95b931e320..05955d54d9 100644 --- a/backend/open_webui/utils/automations.py +++ b/backend/open_webui/utils/automations.py @@ -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') diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index ff0edd166e..31f5deb2f8 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -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, diff --git a/backend/open_webui/utils/payload.py b/backend/open_webui/utils/payload.py index 7cf4fe4de3..63063f4983 100644 --- a/backend/open_webui/utils/payload.py +++ b/backend/open_webui/utils/payload.py @@ -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', [])) diff --git a/backend/open_webui/utils/task.py b/backend/open_webui/utils/task.py index 15213b1f03..dd5e3af72e 100644 --- a/backend/open_webui/utils/task.py +++ b/backend/open_webui/utils/task.py @@ -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