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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1from __future__ import annotations
3import logging
4from typing import Any
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)
20log = logging.getLogger(__name__)
22DEFAULT_CONTEXT_COMPACTION_PROMPT = """### Task:
23Summarize the conversation history that will be compacted out of the active chat context.
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.
32### Previous Summary:
33{{PREVIOUS_SUMMARY}}
35### Messages Being Compacted:
36{{COMPACTED_MESSAGES}}
38### Recent Messages Kept In Context:
39{{RECENT_MESSAGES}}"""
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
55 system_messages = [messages[0]] if messages and messages[0].get('role') == 'system' else []
56 messages = messages[1:] if system_messages else messages
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
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
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
73 event_emitter = await get_event_emitter(metadata)
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 )
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
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 )
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 )
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 )
147 return [*system_messages, *recent_messages], summary, True
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'}
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'}
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 {}
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'}
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 )
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 }
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 }
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
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))
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)
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)
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
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
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
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
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)
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)
286 tokens = _estimate_tokens(previous_summary or '') + _estimate_messages_tokens(messages)
287 return _build_context_usage(tokens, threshold)
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 }
300def _apply_latest_summary_checkpoint(messages: list[dict]) -> tuple[list[dict], str | None]:
301 summary = None
302 summary_idx = None
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
310 if summary_idx is None:
311 return messages, None
312 return messages[summary_idx:], summary
315def _exceeds_token_threshold(messages: list[dict], system_prompt: str, summary: str | None, threshold: int) -> bool:
316 if threshold <= 0:
317 return False
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
324 estimated = _estimate_tokens(system_prompt) + _estimate_tokens(summary or '') + _estimate_messages_tokens(messages)
325 return estimated > threshold
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)
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
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')
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)
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 }
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 }
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
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]
398def _response_text(response: Any) -> str:
399 if isinstance(response, list) and len(response) == 1:
400 response = response[0]
402 if isinstance(response, JSONResponse):
403 try:
404 response = JSONCodec.loads(response.body.decode('utf-8', 'replace'))
405 except Exception:
406 return ''
408 if not isinstance(response, dict):
409 return ''
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 ''
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)
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)
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
446def _estimate_tokens(value: Any) -> int:
447 if value is None:
448 return 0
450 if not isinstance(value, str):
451 try:
452 value = JSONCodec.dumps(value, ensure_ascii=False)
453 except Exception:
454 value = str(value)
456 if not value:
457 return 0
459 return max(1, len(value) // 4)