Coverage for open_webui/utils/context_compaction.py: 19%

236 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 05:07 +0000

1from __future__ import annotations 

2 

3import logging 

4from typing import Any 

5 

6from fastapi.responses import JSONResponse 

7from open_webui.models.chats import Chats 

8from open_webui.models.config import Config 

9from open_webui.utils.chat_id import is_saved_chat_id 

10from open_webui.utils.json_codec import JSONCodec 

11from open_webui.utils.misc import get_content_from_message, get_last_user_message, get_message_list 

12from open_webui.utils.payload import apply_params_to_form_data 

13from open_webui.utils.task import ( 

14 prompt_template, 

15 prompt_variables_template, 

16 replace_messages_variable, 

17 replace_prompt_variable, 

18) 

19 

20log = logging.getLogger(__name__) 

21 

22DEFAULT_CONTEXT_COMPACTION_PROMPT = """### Task: 

23Summarize the conversation history that will be compacted out of the active chat context. 

24 

25### Instructions: 

26- Preserve key decisions, user preferences, and constraints. 

27- Preserve files, artifacts, tool results, and code changes that matter going forward. 

28- Preserve the current task state, unresolved questions, and next steps. 

29- Be factual and specific. Do not invent details. 

30- Keep the summary concise, but complete enough for the assistant to continue without the removed messages. 

31 

32### Previous Summary: 

33{{PREVIOUS_SUMMARY}} 

34 

35### Messages Being Compacted: 

36{{COMPACTED_MESSAGES}} 

37 

38### Recent Messages Kept In Context: 

39{{RECENT_MESSAGES}}""" 

40 

41 

42async def compact_messages_for_request( 

43 request, 

44 user, 

45 messages: list[dict], 

46 metadata: dict, 

47 model_id: str, 

48 models: dict, 

49 system_prompt: str = '', 

50) -> tuple[list[dict], str | None, bool]: 

51 config = await _load_config() 

52 if not config['enable']: 

53 return messages, None, False 

54 

55 system_messages = [messages[0]] if messages and messages[0].get('role') == 'system' else [] 

56 messages = messages[1:] if system_messages else messages 

57 

58 messages, previous_summary = _apply_latest_summary_checkpoint(messages) 

59 token_threshold = _resolve_token_threshold(config['token_threshold'], config['token_cap'], metadata) 

60 if not _exceeds_token_threshold(messages, system_prompt, previous_summary, token_threshold) or len(messages) <= 3: 

61 return [*system_messages, *messages], previous_summary, False 

62 

63 boundary = _find_compaction_boundary(messages, config['retention_percentage']) 

64 compacted_messages = messages[:boundary] 

65 recent_messages = messages[boundary:] 

66 if not compacted_messages or not recent_messages: 

67 return [*system_messages, *messages], previous_summary, False 

68 

69 event_emitter = None 

70 if metadata.get('chat_id') and metadata.get('message_id'): 

71 from open_webui.socket.main import get_event_emitter 

72 

73 event_emitter = await get_event_emitter(metadata) 

74 

75 if event_emitter: 

76 await event_emitter( 

77 { 

78 'type': 'context_compaction', 

79 'data': { 

80 'action': 'context_compaction', 

81 'description': 'Compacting context', 

82 'done': False, 

83 }, 

84 } 

85 ) 

86 

87 try: 

88 summary = await _generate_summary( 

89 request, 

90 user, 

91 model_id, 

92 models, 

93 compacted_messages, 

94 recent_messages, 

95 previous_summary, 

96 config['prompt_template'], 

97 ) 

98 except Exception: 

99 if event_emitter: 

100 await event_emitter( 

101 { 

102 'type': 'context_compaction', 

103 'data': { 

104 'action': 'context_compaction', 

105 'description': 'Context compaction failed', 

106 'done': True, 

107 'error': True, 

108 }, 

109 } 

110 ) 

111 raise 

112 

113 chat_id = metadata.get('chat_id') 

114 checkpoint_message_id = ( 

115 recent_messages[0].get('id') or metadata.get('user_message_id') or metadata.get('message_id') 

116 ) 

117 if is_saved_chat_id(chat_id) and checkpoint_message_id: 

118 await Chats.upsert_message_to_chat_by_id_and_message_id( 

119 chat_id, 

120 checkpoint_message_id, 

121 {'contextSummary': summary}, 

122 touch=False, 

123 ) 

124 

125 log.info( 

126 'Compacted chat context for chat=%s checkpoint=%s response=%s dropped=%d kept=%d summary_chars=%d', 

127 chat_id, 

128 checkpoint_message_id, 

129 metadata.get('message_id'), 

130 len(compacted_messages), 

131 len(recent_messages), 

132 len(summary), 

133 ) 

134 

135 if event_emitter: 

136 await event_emitter( 

137 { 

138 'type': 'context_compaction', 

139 'data': { 

140 'action': 'context_compaction', 

141 'description': 'Context compacted', 

142 'done': True, 

143 }, 

144 } 

145 ) 

146 

147 return [*system_messages, *recent_messages], summary, True 

148 

149 

150async def compact_chat_branch(request, user, chat: Any, model_id: str, models: dict) -> dict: 

151 config = await _load_config() 

152 if not config['enable']: 152 ↛ 155line 152 didn't jump to line 155 because the condition on line 152 was always true

153 return {'ok': True, 'compacted': False, 'reason': 'disabled'} 

154 

155 chat_data = chat.chat or {} 

156 history = chat_data.get('history') or {} 

157 current_id = getattr(chat, 'current_message_id', None) or history.get('currentId') 

158 if not current_id: 

159 current_id = chat_data.get('currentId') or chat_data.get('branchPointMessageId') 

160 if not current_id and isinstance(chat_data.get('messages'), list) and chat_data['messages']: 

161 current_id = chat_data['messages'][-1].get('id') 

162 if not current_id: 

163 return {'ok': True, 'compacted': False, 'reason': 'empty'} 

164 

165 messages_map = await Chats.get_messages_map_by_chat_id(chat.id) 

166 if not messages_map: 

167 messages_map = history.get('messages') or {} 

168 

169 messages, previous_summary = _apply_latest_summary_checkpoint(get_message_list(messages_map, current_id)) 

170 compacted_messages = messages[:-1] 

171 recent_messages = messages[-1:] 

172 if not compacted_messages or not recent_messages: 

173 return {'ok': True, 'compacted': False, 'reason': 'too_short'} 

174 

175 summary = await _generate_summary( 

176 request, 

177 user, 

178 model_id, 

179 models, 

180 compacted_messages, 

181 recent_messages, 

182 previous_summary, 

183 config['prompt_template'], 

184 ) 

185 await Chats.upsert_message_to_chat_by_id_and_message_id( 

186 chat.id, current_id, {'contextSummary': summary}, touch=False 

187 ) 

188 

189 return { 

190 'ok': True, 

191 'compacted': True, 

192 'dropped_messages': len(compacted_messages), 

193 'kept_messages': len(recent_messages), 

194 'summary_chars': len(summary), 

195 } 

196 

197 

198async def _load_config() -> dict: 

199 values = await Config.get_many( 

200 'chat.context_compaction.enable', 

201 'chat.context_compaction.token_threshold', 

202 'chat.context_compaction.token_cap', 

203 'chat.context_compaction.retention_percentage', 

204 'chat.context_compaction.prompt_template', 

205 ) 

206 token_threshold = _parse_positive_int(values.get('chat.context_compaction.token_threshold')) or 80000 

207 return { 

208 'enable': bool(values.get('chat.context_compaction.enable', False)), 

209 'token_threshold': token_threshold, 

210 'token_cap': _parse_positive_int(values.get('chat.context_compaction.token_cap')) or token_threshold, 

211 'retention_percentage': _clamp_retention_percentage(values.get('chat.context_compaction.retention_percentage')), 

212 'prompt_template': values.get('chat.context_compaction.prompt_template', '') or '', 

213 } 

214 

215 

216def _parse_positive_int(value: Any) -> int | None: 

217 try: 

218 parsed = int(value) 

219 except (TypeError, ValueError): 

220 return None 

221 return parsed if parsed > 0 else None 

222 

223 

224def _clamp_retention_percentage(value: Any) -> int: 

225 try: 

226 parsed = int(value) 

227 except (TypeError, ValueError): 

228 parsed = 40 

229 return min(50, max(10, parsed)) 

230 

231 

232def _resolve_token_threshold(global_threshold: int, global_cap: int, metadata: dict) -> int: 

233 configured_threshold = _parse_positive_int((metadata.get('params') or {}).get('compact_token_threshold')) 

234 return min(configured_threshold or global_threshold, global_cap) 

235 

236 

237def _usage_token_count(usage: dict) -> int: 

238 prompt_tokens = int(usage.get('prompt_tokens') or usage.get('prompt_eval_count') or 0) 

239 if not prompt_tokens and (usage.get('prompt_n') is not None or usage.get('cache_n') is not None): 

240 prompt_tokens = int(usage.get('prompt_n') or 0) + int(usage.get('cache_n') or 0) 

241 if not prompt_tokens: 

242 prompt_tokens = int(usage.get('input_tokens') or 0) 

243 

244 completion_tokens = int( 

245 usage.get('completion_tokens') 

246 or usage.get('output_tokens') 

247 or usage.get('eval_count') 

248 or usage.get('predicted_n') 

249 or 0 

250 ) 

251 return prompt_tokens + completion_tokens 

252 

253 

254async def get_chat_context_usage(chat: Any, model_id: str | None = None) -> dict | None: 

255 chat_data = chat.chat or {} 

256 history = chat_data.get('history') or {} 

257 current_id = getattr(chat, 'current_message_id', None) or history.get('currentId') 

258 if not current_id: 

259 current_id = chat_data.get('currentId') or chat_data.get('branchPointMessageId') 

260 if not current_id and isinstance(chat_data.get('messages'), list) and chat_data['messages']: 260 ↛ 261line 260 didn't jump to line 261 because the condition on line 260 was never true

261 current_id = chat_data['messages'][-1].get('id') 

262 if not current_id: 

263 return None 

264 

265 messages_map = await Chats.get_messages_map_by_chat_id(chat.id) 

266 messages = get_message_list(messages_map or history.get('messages') or {}, current_id) 

267 if not messages: 267 ↛ 268line 267 didn't jump to line 268 because the condition on line 267 was never true

268 return None 

269 

270 config = await _load_config() 

271 if not config['enable']: 271 ↛ 274line 271 didn't jump to line 274 because the condition on line 271 was always true

272 return None 

273 

274 params = ((chat.chat or {}).get('params') or {}).copy() 

275 if model_id: 

276 params['model'] = model_id 

277 threshold = _resolve_token_threshold(config['token_threshold'], config['token_cap'], {'params': params}) 

278 messages, previous_summary = _apply_latest_summary_checkpoint(messages) 

279 

280 for idx in range(len(messages) - 1, -1, -1): 

281 usage = messages[idx].get('usage') or (messages[idx].get('info') or {}).get('usage') 

282 if isinstance(usage, dict) and (tokens := _usage_token_count(usage)): 

283 tokens += _estimate_messages_tokens(messages[idx + 1 :]) 

284 return _build_context_usage(tokens, threshold) 

285 

286 tokens = _estimate_tokens(previous_summary or '') + _estimate_messages_tokens(messages) 

287 return _build_context_usage(tokens, threshold) 

288 

289 

290def _build_context_usage(tokens: int, threshold: int) -> dict: 

291 return { 

292 'tokens': tokens, 

293 'estimated_tokens': tokens, 

294 'threshold': threshold, 

295 'percent': round((tokens / threshold) * 100) if threshold > 0 else 0, 

296 'source': 'estimated', 

297 } 

298 

299 

300def _apply_latest_summary_checkpoint(messages: list[dict]) -> tuple[list[dict], str | None]: 

301 summary = None 

302 summary_idx = None 

303 

304 for idx, message in enumerate(messages): 

305 value = message.get('contextSummary') or message.get('context_summary') 

306 if isinstance(value, str) and value.strip(): 

307 summary = value 

308 summary_idx = idx 

309 

310 if summary_idx is None: 

311 return messages, None 

312 return messages[summary_idx:], summary 

313 

314 

315def _exceeds_token_threshold(messages: list[dict], system_prompt: str, summary: str | None, threshold: int) -> bool: 

316 if threshold <= 0: 

317 return False 

318 

319 for idx in range(len(messages) - 1, -1, -1): 

320 usage = messages[idx].get('usage') or (messages[idx].get('info') or {}).get('usage') 

321 if isinstance(usage, dict) and (tokens := _usage_token_count(usage)): 

322 return tokens + _estimate_messages_tokens(messages[idx + 1 :]) > threshold 

323 

324 estimated = _estimate_tokens(system_prompt) + _estimate_tokens(summary or '') + _estimate_messages_tokens(messages) 

325 return estimated > threshold 

326 

327 

328def _find_compaction_boundary(messages: list[dict], retention_percentage: int = 40) -> int: 

329 retention_percentage = _clamp_retention_percentage(retention_percentage) 

330 keep_count = max(2, len(messages) * retention_percentage // 100) 

331 target = max(1, len(messages) - keep_count) 

332 boundaries = [idx for idx, message in enumerate(messages) if message.get('role') == 'user'][1:] 

333 return next((idx for idx in reversed(boundaries) if idx <= target), 0) 

334 

335 

336async def _generate_summary( 

337 request, 

338 user, 

339 model_id: str, 

340 models: dict, 

341 compacted_messages: list[dict], 

342 recent_messages: list[dict], 

343 previous_summary: str | None, 

344 summary_prompt_template: str, 

345) -> str: 

346 from open_webui.utils.chat import generate_chat_completion 

347 

348 task_config = await Config.get_many( 

349 'task.model.params', 

350 'chat.context_compaction.model', 

351 ) 

352 context_compaction_model = task_config.get('chat.context_compaction.model') 

353 task_model_id = context_compaction_model if context_compaction_model in models else model_id 

354 if task_model_id not in models: 

355 raise ValueError('No available model for context compaction') 

356 

357 summary_prompt_template = summary_prompt_template.strip() or DEFAULT_CONTEXT_COMPACTION_PROMPT 

358 all_messages = [*compacted_messages, *recent_messages] 

359 prompt = replace_prompt_variable(summary_prompt_template, get_last_user_message(all_messages) or '') 

360 prompt = replace_messages_variable(prompt, all_messages) 

361 prompt = replace_messages_variable(prompt, compacted_messages, 'COMPACTED_MESSAGES') 

362 prompt = replace_messages_variable(prompt, recent_messages, 'RECENT_MESSAGES') 

363 prompt = prompt_variables_template(prompt, {'{{PREVIOUS_SUMMARY}}': previous_summary or ''}) 

364 prompt = await prompt_template(prompt, user) 

365 

366 task_model_params = task_config.get('task.model.params') or {} 

367 if not isinstance(task_model_params, dict): 

368 task_model_params = {} 

369 task_model_params = {key: value for key, value in task_model_params.items() if value is not None and value != ''} 

370 task_model_params = task_model_params or { 

371 'max_tokens': models[task_model_id].get('info', {}).get('params', {}).get('max_tokens', 1000) 

372 } 

373 

374 payload = { 

375 'model': task_model_id, 

376 'messages': [{'role': 'user', 'content': prompt}], 

377 'stream': False, 

378 'metadata': { 

379 **(request.state.metadata if hasattr(request.state, 'metadata') else {}), 

380 'task': 'context_compaction', 

381 }, 

382 } 

383 

384 payload = apply_params_to_form_data(payload, models[task_model_id], task_model_params) 

385 response = await generate_chat_completion(request, form_data=payload, user=user) 

386 summary = _response_text(response).strip() 

387 if summary: 

388 return summary 

389 

390 parts = [previous_summary] if previous_summary else [] 

391 for message in compacted_messages: 

392 content = get_content_from_message(message) 

393 if content: 

394 parts.append(f'- {message.get("role", "unknown")}: {content[:500]}') 

395 return '\n'.join(parts)[:4000] 

396 

397 

398def _response_text(response: Any) -> str: 

399 if isinstance(response, list) and len(response) == 1: 

400 response = response[0] 

401 

402 if isinstance(response, JSONResponse): 

403 try: 

404 response = JSONCodec.loads(response.body.decode('utf-8', 'replace')) 

405 except Exception: 

406 return '' 

407 

408 if not isinstance(response, dict): 

409 return '' 

410 

411 choices = response.get('choices') or [] 

412 if choices: 

413 message = choices[0].get('message') or {} 

414 return message.get('content') or message.get('reasoning_content') or '' 

415 

416 parts = [] 

417 for item in response.get('output') or []: 

418 for content in item.get('content') or []: 

419 if isinstance(content, dict): 

420 parts.append(content.get('text') or content.get('content') or '') 

421 return '\n'.join(part for part in parts if part) 

422 

423 

424def _estimate_messages_tokens(messages: list[dict]) -> int: 

425 total = 0 

426 for message in messages: 

427 total += 4 

428 content = message.get('content') 

429 if isinstance(content, list): 

430 for item in content: 

431 if not isinstance(item, dict): 

432 total += _estimate_tokens(item) 

433 elif item.get('type') in {'image', 'image_url'}: 

434 total += 1000 

435 else: 

436 total += _estimate_tokens(item.get('text') or item.get('content') or item) 

437 else: 

438 total += _estimate_tokens(content) 

439 

440 total += _estimate_tokens(message.get('output')) 

441 total += _estimate_tokens(message.get('tool_calls')) 

442 total += _estimate_tokens(message.get('files')) 

443 return total 

444 

445 

446def _estimate_tokens(value: Any) -> int: 

447 if value is None: 

448 return 0 

449 

450 if not isinstance(value, str): 

451 try: 

452 value = JSONCodec.dumps(value, ensure_ascii=False) 

453 except Exception: 

454 value = str(value) 

455 

456 if not value: 

457 return 0 

458 

459 return max(1, len(value) // 4)