This commit is contained in:
Timothy Jaeryang Baek 2026-10-10 23:02:39 +04:00
parent 930f8c3640
commit 67cfa917c1
2 changed files with 75 additions and 25 deletions

View file

@ -59,6 +59,7 @@ async def compact_messages_for_request(
messages = messages[1:] if system_messages else messages
messages, previous_summary = _apply_latest_summary_checkpoint(messages)
system_prompt = system_prompt or (get_content_from_message(system_messages[0]) if system_messages else '')
token_threshold = _resolve_token_threshold(config['token_threshold'], config['token_cap'], metadata)
if not _exceeds_token_threshold(messages, system_prompt, previous_summary, token_threshold) or len(messages) <= 3:
return [*system_messages, *messages], previous_summary, False
@ -117,7 +118,8 @@ async def compact_messages_for_request(
checkpoint_message_id = (
recent_messages[0].get('id') or metadata.get('user_message_id') or metadata.get('message_id')
)
if is_saved_chat_id(chat_id) and checkpoint_message_id:
# Only whole-turn user boundaries correspond to persisted chat checkpoints.
if recent_messages[0].get('role') == 'user' and is_saved_chat_id(chat_id) and checkpoint_message_id:
await Chats.upsert_message_to_chat_by_id_and_message_id(
chat_id,
checkpoint_message_id,
@ -319,10 +321,12 @@ def _exceeds_token_threshold(messages: list[dict], system_prompt: str, summary:
if threshold <= 0:
return False
for idx in range(len(messages) - 1, -1, -1):
usage = messages[idx].get('usage') or (messages[idx].get('info') or {}).get('usage')
if isinstance(usage, dict) and (tokens := _usage_token_count(usage)):
return tokens + _estimate_messages_tokens(messages[idx + 1 :]) > threshold
# Expanded tool histories include fresh results and rebuilt prompts absent from prior usage.
if not any(message.get('role') == 'tool' for message in messages):
for idx in range(len(messages) - 1, -1, -1):
usage = messages[idx].get('usage') or (messages[idx].get('info') or {}).get('usage')
if isinstance(usage, dict) and (tokens := _usage_token_count(usage)):
return tokens + _estimate_messages_tokens(messages[idx + 1 :]) > threshold
estimated = _estimate_tokens(system_prompt) + _estimate_tokens(summary or '') + _estimate_messages_tokens(messages)
return estimated > threshold
@ -332,7 +336,15 @@ def _find_compaction_boundary(messages: list[dict], retention_percentage: int =
retention_percentage = _clamp_retention_percentage(retention_percentage)
keep_count = max(2, len(messages) * retention_percentage // 100)
target = max(1, len(messages) - keep_count)
boundaries = [idx for idx, message in enumerate(messages) if message.get('role') == 'user'][1:]
if any(message.get('role') == 'tool' for message in messages):
# Keep each call/result batch and its following image messages together.
boundaries = [
idx
for idx, message in enumerate(messages)
if message.get('role') == 'assistant' and message.get('tool_calls')
][1:]
else:
boundaries = [idx for idx, message in enumerate(messages) if message.get('role') == 'user'][1:]
return next((idx for idx in reversed(boundaries) if idx <= target), 0)
@ -348,6 +360,9 @@ async def _generate_summary(
) -> str:
from open_webui.utils.chat import generate_chat_completion
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
models = {**dict(models.items()), request.state.model['id']: request.state.model}
task_config = await Config.get_many(
'task.model.params',
'chat.context_compaction.model',
@ -358,6 +373,15 @@ async def _generate_summary(
raise ValueError('No available model for context compaction')
summary_prompt_template = summary_prompt_template.strip() or DEFAULT_CONTEXT_COMPACTION_PROMPT
# The text-only prompt formatter otherwise omits tool names and arguments.
compacted_messages, recent_messages = list(compacted_messages), list(recent_messages)
for messages in (compacted_messages, recent_messages):
for idx, message in enumerate(messages):
if message.get('tool_calls'):
messages[idx] = {
**message,
'content': f'{get_content_from_message(message) or ""}\n[TOOL CALLS] {JSONCodec.dumps(message["tool_calls"])}',
}
all_messages = [*compacted_messages, *recent_messages]
prompt = replace_prompt_variable(summary_prompt_template, get_last_user_message(all_messages) or '')
prompt = replace_messages_variable(prompt, all_messages)

View file

@ -2514,14 +2514,6 @@ async def process_chat_payload(request, form_data, user, metadata, model):
form_data['messages'].append({'role': 'user', 'content': regeneration_prompt})
if is_saved_chat_id(chat_id) and user_message_id:
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
compaction_models = {
**dict(request.app.state.MODELS.items()),
request.state.model['id']: request.state.model,
}
else:
compaction_models = request.app.state.MODELS
system_message = get_system_message(form_data.get('messages', []))
chat_system_prompt = get_content_from_message(system_message) if system_message else ''
@ -2532,7 +2524,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
form_data.get('messages', []),
metadata,
form_data.get('model'),
compaction_models,
request.app.state.MODELS,
chat_system_prompt,
)
if context_summary:
@ -5906,6 +5898,7 @@ async def streaming_chat_response_handler(response, ctx):
tool_call_sources = [] # Track citation sources from tool results
all_tool_call_sources = [] # Accumulated sources across all iterations
user_message = get_last_user_message(form_data['messages'])
messages = list(form_data['messages'])
# Check if citations are enabled for this model
citations_enabled = (model.get('info', {}).get('meta', {}).get('capabilities') or {}).get(
@ -6267,13 +6260,19 @@ async def streaming_chat_response_handler(response, ctx):
image_urls.append(part.get('image_url', ''))
message['content'] = ''.join(text_parts)
new_form_data['messages'] = [
*form_data['messages'],
# Refresh citation context without restoring compacted history.
user_prompt = get_last_user_message_item(form_data['messages'])
for message in messages:
if message.get('contextSummary') and message.get('role') == 'user' and user_prompt:
message.update(user_prompt)
messages = [
*[message for message in form_data['messages'] if message.get('role') == 'system'],
*[message for message in messages if message.get('role') != 'system'],
*tool_messages,
]
if image_urls:
new_form_data['messages'].append(
messages.append(
{
'role': 'user',
'content': [
@ -6286,6 +6285,39 @@ async def streaming_chat_response_handler(response, ctx):
}
)
try:
messages, context_summary, compacted = await compact_messages_for_request(
request, user, messages, metadata, model_id, request.app.state.MODELS
)
if compacted:
if user_prompt is not None:
messages = [message for message in messages if message != user_prompt]
messages.insert(
1 if messages and messages[0].get('role') == 'system' else 0,
dict(user_prompt),
)
checkpoint = next(
idx for idx, message in enumerate(messages) if message.get('role') != 'system'
)
messages[checkpoint] = {**messages[checkpoint], 'contextSummary': context_summary}
except Exception:
log.exception('Tool loop context compaction failed; keeping the current context')
# Use the existing checkpoint format internally and strip it from the provider payload.
new_form_data['messages'] = process_messages_with_output(messages)
context_summary = next(
(message['contextSummary'] for message in messages if message.get('contextSummary')),
None,
)
if context_summary:
new_form_data['messages'].insert(
1
if new_form_data['messages']
and new_form_data['messages'][0].get('role') == 'system'
else 0,
{'role': 'system', 'content': f'[CONVERSATION SUMMARY]\n{context_summary}'},
)
new_form_data = await convert_url_images_to_base64(new_form_data, user=user)
if filter_functions:
@ -6314,8 +6346,6 @@ async def streaming_chat_response_handler(response, ctx):
# keeps indices aligned. The display prefix
# ensures the UI shows tool history during
# streaming.
continued_output = prior_output
round_output = output
prior_output = list(full_output())
# Trim the trailing empty placeholder message
# so it doesn't persist as a ghost item once
@ -6328,13 +6358,9 @@ async def streaming_chat_response_handler(response, ctx):
msg_parts = prior_output[-1].get('content', [])
if not msg_parts or (len(msg_parts) == 1 and not msg_parts[0].get('text', '').strip()):
prior_output.pop()
round_output = round_output[:-1]
output = []
output_start = len(prior_output)
await stream_body_handler(res, new_form_data)
# A continued reply's earlier items are already in form_data['messages']
output = [*round_output, *output]
prior_output = continued_output
elif getattr(res, 'status_code', 200) >= 400:
await emit_message_error(get_message_error_content(get_response_error_detail(res)))
break
@ -6503,7 +6529,7 @@ async def streaming_chat_response_handler(response, ctx):
'stream': True,
'metadata': metadata,
'messages': [
*form_data['messages'],
*messages,
*convert_output_to_messages(
output,
raw=True,