Coverage for open_webui/utils/task.py: 8%
253 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1import logging
2import math
3import re
4import uuid
5from datetime import datetime
6from typing import Any, Optional
8from open_webui.config import DEFAULT_RAG_TEMPLATE
9from open_webui.utils.misc import get_last_user_message, get_messages_content
11log = logging.getLogger(__name__)
14# Let the right tool be given for the work at hand,
15# not the one that flatters, but the one that serves.
16def get_task_model_id(default_model_id: str, task_model: str, task_model_external: str, models) -> str:
17 # Set the task model
18 task_model_id = default_model_id
19 # Check if the user has a custom task model and use that model
20 if models.get(task_model_id, {}).get('connection_type') == 'local':
21 if task_model and task_model in models:
22 task_model_id = task_model
23 else:
24 if task_model_external and task_model_external in models:
25 task_model_id = task_model_external
27 return task_model_id
30def prompt_variables_template(template: str, variables: dict[str, str]) -> str:
31 for variable, value in variables.items():
32 template = template.replace(variable, value)
33 return template
36async def prompt_template(template: str, user: Optional[Any] = None) -> str:
37 USER_VARIABLES = {}
39 if user:
40 if hasattr(user, 'model_dump'):
41 user = user.model_dump()
43 if isinstance(user, dict):
44 user_info = user.get('info', {}) or {}
45 birth_date = user.get('date_of_birth')
46 age = None
48 if birth_date:
49 try:
50 # If birth_date is str, convert to datetime
51 if isinstance(birth_date, str):
52 birth_date = datetime.strptime(birth_date, '%Y-%m-%d')
54 today = datetime.now()
55 age = today.year - birth_date.year - ((today.month, today.day) < (birth_date.month, birth_date.day))
56 except Exception as e:
57 pass
59 # Resolve user groups from DB only when the template uses {{USER_GROUPS}}
60 groups = ''
61 if '{{USER_GROUPS}}' in template:
62 user_id = user.get('id')
63 if user_id:
64 try:
65 from open_webui.models.groups import Groups
67 user_groups = await Groups.get_groups_by_member_id(user_id)
68 groups = ', '.join(g.name for g in user_groups)
69 except Exception:
70 pass
72 USER_VARIABLES = {
73 'name': str(user.get('name')),
74 'email': str(user.get('email')),
75 'location': str(user_info.get('location')),
76 'bio': str(user.get('bio')),
77 'gender': str(user.get('gender')),
78 'birth_date': str(birth_date),
79 'age': str(age),
80 'groups': groups,
81 }
83 # Get the current date
84 current_date = datetime.now()
86 # Format the date to YYYY-MM-DD
87 formatted_date = current_date.strftime('%Y-%m-%d')
88 formatted_time = current_date.strftime('%I:%M:%S %p')
89 formatted_weekday = current_date.strftime('%A')
91 template = template.replace('{{CURRENT_DATE}}', formatted_date)
92 template = template.replace('{{CURRENT_TIME}}', formatted_time)
93 template = template.replace('{{CURRENT_DATETIME}}', f'{formatted_date} {formatted_time}')
94 template = template.replace('{{CURRENT_WEEKDAY}}', formatted_weekday)
96 template = template.replace('{{USER_NAME}}', USER_VARIABLES.get('name', 'Unknown'))
97 template = template.replace('{{USER_EMAIL}}', USER_VARIABLES.get('email', 'Unknown'))
98 template = template.replace('{{USER_BIO}}', USER_VARIABLES.get('bio', 'Unknown'))
99 template = template.replace('{{USER_GENDER}}', USER_VARIABLES.get('gender', 'Unknown'))
100 template = template.replace('{{USER_BIRTH_DATE}}', USER_VARIABLES.get('birth_date', 'Unknown'))
101 template = template.replace('{{USER_AGE}}', str(USER_VARIABLES.get('age', 'Unknown')))
102 template = template.replace('{{USER_LOCATION}}', USER_VARIABLES.get('location', 'Unknown'))
103 template = template.replace('{{USER_GROUPS}}', USER_VARIABLES.get('groups', ''))
105 return template
108def replace_prompt_variable(template: str, prompt: str) -> str:
109 def replacement_function(match):
110 full_match = match.group(0).lower() # Normalize to lowercase for consistent handling
111 start_length = match.group(1)
112 end_length = match.group(2)
113 middle_length = match.group(3)
115 if full_match == '{{prompt}}':
116 return prompt
117 elif start_length is not None:
118 return prompt[: int(start_length)]
119 elif end_length is not None:
120 return prompt[-int(end_length) :]
121 elif middle_length is not None:
122 middle_length = int(middle_length)
123 if len(prompt) <= middle_length:
124 return prompt
125 start = prompt[: math.ceil(middle_length / 2)]
126 end = prompt[-math.floor(middle_length / 2) :]
127 return f'{start}...{end}'
128 return ''
130 # Updated regex pattern to make it case-insensitive with the `(?i)` flag
131 pattern = r'(?i){{prompt}}|{{prompt:start:(\d+)}}|{{prompt:end:(\d+)}}|{{prompt:middletruncate:(\d+)}}'
132 template = re.sub(pattern, replacement_function, template)
133 return template
136def truncate_content(content: str, max_chars: int, mode: str = 'middletruncate') -> str:
137 """Truncate a string to max_chars using the specified mode.
139 Modes:
140 - middletruncate: keep beginning and end, join with '...'
141 - start: keep first max_chars characters
142 - end: keep last max_chars characters
143 """
144 if max_chars <= 0:
145 return ''
147 if not content or len(content) <= max_chars:
148 return content
150 if mode == 'start':
151 return content[:max_chars]
152 elif mode == 'end':
153 return content[-max_chars:]
154 else: # middletruncate
155 half = max_chars // 2
156 return f'{content[:half]}...{content[-(max_chars - half) :]}'
159def apply_content_filter(messages: list[dict], filter_str: str) -> list[dict]:
160 """Apply a content filter to each message's content.
162 filter_str is like 'middletruncate:500', 'start:200', or 'end:200'.
163 Returns a new list with truncated content (original messages are not mutated).
164 """
165 parts = filter_str.split(':')
166 if len(parts) != 2:
167 return messages
169 mode = parts[0].lower()
170 try:
171 max_chars = int(parts[1])
172 except ValueError:
173 return messages
175 if mode not in ('middletruncate', 'start', 'end'):
176 return messages
178 result = []
179 for msg in messages:
180 new_msg = dict(msg)
181 if isinstance(new_msg.get('content'), str):
182 new_msg['content'] = truncate_content(new_msg['content'], max_chars, mode)
183 elif isinstance(new_msg.get('content'), list):
184 new_content = []
185 for item in new_msg['content']:
186 if isinstance(item, dict) and item.get('type') == 'text':
187 new_item = dict(item)
188 new_item['text'] = truncate_content(item.get('text', ''), max_chars, mode)
189 new_content.append(new_item)
190 else:
191 new_content.append(item)
192 new_msg['content'] = new_content
193 result.append(new_msg)
194 return result
197def replace_messages_variable(
198 template: str, messages: Optional[list[dict]] = None, variable_name: str = 'MESSAGES'
199) -> str:
200 def replacement_function(match):
201 # Groups: (1) filter for bare MESSAGES
202 # (2) START count, (3) filter for START
203 # (4) END count, (5) filter for END
204 # (6) MIDDLE count,(7) filter for MIDDLE
205 bare_filter = match.group(1)
206 start_length = match.group(2)
207 start_filter = match.group(3)
208 end_length = match.group(4)
209 end_filter = match.group(5)
210 middle_length = match.group(6)
211 middle_filter = match.group(7)
213 # If messages is None, handle it as an empty list
214 if messages is None:
215 return ''
217 # Select messages based on the variant
218 if start_length is not None:
219 selected = messages[: int(start_length)]
220 content_filter = start_filter
221 elif end_length is not None:
222 selected = messages[-int(end_length) :]
223 content_filter = end_filter
224 elif middle_length is not None:
225 mid = int(middle_length)
226 if len(messages) <= mid:
227 selected = messages
228 else:
229 half = mid // 2
230 start_msgs = messages[:half]
231 end_msgs = messages[-half:] if mid % 2 == 0 else messages[-(half + 1) :]
232 selected = start_msgs + end_msgs
233 content_filter = middle_filter
234 else:
235 # Bare {{MESSAGES}} or {{MESSAGES|filter}}
236 selected = messages
237 content_filter = bare_filter
239 # Apply content filter if present
240 if content_filter:
241 selected = apply_content_filter(selected, content_filter)
243 return get_messages_content(selected)
245 variable_pattern = re.escape(variable_name)
246 template = re.sub(
247 r'(?:'
248 rf'\{{\{{{variable_pattern}(?:\|(\w+:\d+))?\}}\}}'
249 rf'|\{{\{{{variable_pattern}:START:(\d+)(?:\|(\w+:\d+))?\}}\}}'
250 rf'|\{{\{{{variable_pattern}:END:(\d+)(?:\|(\w+:\d+))?\}}\}}'
251 rf'|\{{\{{{variable_pattern}:MIDDLETRUNCATE:(\d+)(?:\|(\w+:\d+))?\}}\}}'
252 r')',
253 replacement_function,
254 template,
255 )
257 return template
260# {{prompt:middletruncate:8000}}
263# Let the context given here not distort the question,
264# but illuminate it, so that the answer serves the one who asked.
265async def rag_template(template: str, context: str, query: str):
266 if template.strip() == '':
267 template = DEFAULT_RAG_TEMPLATE
269 template = await prompt_template(template)
271 if '[context]' not in template and '{{CONTEXT}}' not in template:
272 log.debug("WARNING: The RAG template does not contain the '[context]' or '{{CONTEXT}}' placeholder.")
274 if '<context>' in context and '</context>' in context:
275 log.debug(
276 'WARNING: Potential prompt injection attack: the RAG '
277 "context contains '<context>' and '</context>'. This might be "
278 'nothing, or the user might be trying to hack something.'
279 )
281 query_placeholders = []
282 if '[query]' in context:
283 query_placeholder = '{{QUERY' + str(uuid.uuid4()) + '}}'
284 template = template.replace('[query]', query_placeholder)
285 query_placeholders.append((query_placeholder, '[query]'))
287 if '{{QUERY}}' in context:
288 query_placeholder = '{{QUERY' + str(uuid.uuid4()) + '}}'
289 template = template.replace('{{QUERY}}', query_placeholder)
290 query_placeholders.append((query_placeholder, '{{QUERY}}'))
292 template = template.replace('[context]', context)
293 template = template.replace('{{CONTEXT}}', context)
295 template = template.replace('[query]', query)
296 template = template.replace('{{QUERY}}', query)
298 for query_placeholder, original_placeholder in query_placeholders:
299 template = template.replace(query_placeholder, original_placeholder)
301 return template
304async def title_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
305 prompt = get_last_user_message(messages)
306 template = replace_prompt_variable(template, prompt)
307 template = replace_messages_variable(template, messages)
309 template = await prompt_template(template, user)
311 return template
314async def follow_up_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
315 prompt = get_last_user_message(messages)
316 template = replace_prompt_variable(template, prompt)
317 template = replace_messages_variable(template, messages)
319 template = await prompt_template(template, user)
320 return template
323async def tags_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
324 prompt = get_last_user_message(messages)
325 template = replace_prompt_variable(template, prompt)
326 template = replace_messages_variable(template, messages)
328 template = await prompt_template(template, user)
329 return template
332async def image_prompt_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
333 prompt = get_last_user_message(messages)
334 template = replace_prompt_variable(template, prompt)
335 template = replace_messages_variable(template, messages)
337 template = await prompt_template(template, user)
338 return template
341async def emoji_generation_template(template: str, prompt: str, user: Optional[Any] = None) -> str:
342 template = replace_prompt_variable(template, prompt)
343 template = await prompt_template(template, user)
345 return template
348async def autocomplete_generation_template(
349 template: str,
350 prompt: str,
351 messages: Optional[list[dict]] = None,
352 type: Optional[str] = None,
353 user: Optional[Any] = None,
354) -> str:
355 template = template.replace('{{TYPE}}', type if type else '')
356 template = replace_prompt_variable(template, prompt)
357 template = replace_messages_variable(template, messages)
359 template = await prompt_template(template, user)
360 return template
363async def query_generation_template(template: str, messages: list[dict], user: Optional[Any] = None) -> str:
364 prompt = get_last_user_message(messages)
365 template = replace_prompt_variable(template, prompt)
366 template = replace_messages_variable(template, messages)
368 template = await prompt_template(template, user)
369 return template
372def moa_response_generation_template(template: str, prompt: str, responses: list[str]) -> str:
373 def replacement_function(match):
374 full_match = match.group(0)
375 start_length = match.group(1)
376 end_length = match.group(2)
377 middle_length = match.group(3)
379 if full_match == '{{prompt}}':
380 return prompt
381 elif start_length is not None:
382 return prompt[: int(start_length)]
383 elif end_length is not None:
384 return prompt[-int(end_length) :]
385 elif middle_length is not None:
386 middle_length = int(middle_length)
387 if len(prompt) <= middle_length:
388 return prompt
389 start = prompt[: math.ceil(middle_length / 2)]
390 end = prompt[-math.floor(middle_length / 2) :]
391 return f'{start}...{end}'
392 return ''
394 template = re.sub(
395 r'{{prompt}}|{{prompt:start:(\d+)}}|{{prompt:end:(\d+)}}|{{prompt:middletruncate:(\d+)}}',
396 replacement_function,
397 template,
398 )
400 responses = [f'"""{response}"""' for response in responses]
401 responses = '\n\n'.join(responses)
403 template = template.replace('{{responses}}', responses)
404 return template
407def tools_function_calling_generation_template(template: str, tools_specs: str) -> str:
408 template = template.replace('{{TOOLS}}', tools_specs)
409 return template