From 67cfa917c12004d4e662d0f8879779b75ad6629a Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sat, 10 Oct 2026 23:02:39 +0400 Subject: [PATCH] refac --- .../open_webui/utils/context_compaction.py | 36 +++++++++-- backend/open_webui/utils/middleware.py | 64 +++++++++++++------ 2 files changed, 75 insertions(+), 25 deletions(-) diff --git a/backend/open_webui/utils/context_compaction.py b/backend/open_webui/utils/context_compaction.py index d9cb7879f7..19d44f417d 100644 --- a/backend/open_webui/utils/context_compaction.py +++ b/backend/open_webui/utils/context_compaction.py @@ -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) diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index e8a12e3e35..bc251652ae 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -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,