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

1import logging 

2import math 

3import re 

4import uuid 

5from datetime import datetime 

6from typing import Any, Optional 

7 

8from open_webui.config import DEFAULT_RAG_TEMPLATE 

9from open_webui.utils.misc import get_last_user_message, get_messages_content 

10 

11log = logging.getLogger(__name__) 

12 

13 

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 

26 

27 return task_model_id 

28 

29 

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 

34 

35 

36async def prompt_template(template: str, user: Optional[Any] = None) -> str: 

37 USER_VARIABLES = {} 

38 

39 if user: 

40 if hasattr(user, 'model_dump'): 

41 user = user.model_dump() 

42 

43 if isinstance(user, dict): 

44 user_info = user.get('info', {}) or {} 

45 birth_date = user.get('date_of_birth') 

46 age = None 

47 

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') 

53 

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 

58 

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 

66 

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 

71 

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 } 

82 

83 # Get the current date 

84 current_date = datetime.now() 

85 

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') 

90 

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) 

95 

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', '')) 

104 

105 return template 

106 

107 

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) 

114 

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 '' 

129 

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 

134 

135 

136def truncate_content(content: str, max_chars: int, mode: str = 'middletruncate') -> str: 

137 """Truncate a string to max_chars using the specified mode. 

138 

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 '' 

146 

147 if not content or len(content) <= max_chars: 

148 return content 

149 

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) :]}' 

157 

158 

159def apply_content_filter(messages: list[dict], filter_str: str) -> list[dict]: 

160 """Apply a content filter to each message's content. 

161 

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 

168 

169 mode = parts[0].lower() 

170 try: 

171 max_chars = int(parts[1]) 

172 except ValueError: 

173 return messages 

174 

175 if mode not in ('middletruncate', 'start', 'end'): 

176 return messages 

177 

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 

195 

196 

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) 

212 

213 # If messages is None, handle it as an empty list 

214 if messages is None: 

215 return '' 

216 

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 

238 

239 # Apply content filter if present 

240 if content_filter: 

241 selected = apply_content_filter(selected, content_filter) 

242 

243 return get_messages_content(selected) 

244 

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 ) 

256 

257 return template 

258 

259 

260# {{prompt:middletruncate:8000}} 

261 

262 

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 

268 

269 template = await prompt_template(template) 

270 

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.") 

273 

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 ) 

280 

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]')) 

286 

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}}')) 

291 

292 template = template.replace('[context]', context) 

293 template = template.replace('{{CONTEXT}}', context) 

294 

295 template = template.replace('[query]', query) 

296 template = template.replace('{{QUERY}}', query) 

297 

298 for query_placeholder, original_placeholder in query_placeholders: 

299 template = template.replace(query_placeholder, original_placeholder) 

300 

301 return template 

302 

303 

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) 

308 

309 template = await prompt_template(template, user) 

310 

311 return template 

312 

313 

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) 

318 

319 template = await prompt_template(template, user) 

320 return template 

321 

322 

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) 

327 

328 template = await prompt_template(template, user) 

329 return template 

330 

331 

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) 

336 

337 template = await prompt_template(template, user) 

338 return template 

339 

340 

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) 

344 

345 return template 

346 

347 

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) 

358 

359 template = await prompt_template(template, user) 

360 return template 

361 

362 

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) 

367 

368 template = await prompt_template(template, user) 

369 return template 

370 

371 

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) 

378 

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 '' 

393 

394 template = re.sub( 

395 r'{{prompt}}|{{prompt:start:(\d+)}}|{{prompt:end:(\d+)}}|{{prompt:middletruncate:(\d+)}}', 

396 replacement_function, 

397 template, 

398 ) 

399 

400 responses = [f'"""{response}"""' for response in responses] 

401 responses = '\n\n'.join(responses) 

402 

403 template = template.replace('{{responses}}', responses) 

404 return template 

405 

406 

407def tools_function_calling_generation_template(template: str, tools_specs: str) -> str: 

408 template = template.replace('{{TOOLS}}', tools_specs) 

409 return template