mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-05 02:41:34 +00:00
feat: run context compaction inside the tool call loop
Long agentic runs died mid-turn on the provider's context limit. Compaction only ran once, right before the first model call of a turn, so every tool call and result after that grew the request unchecked until the provider rejected it. Clicking "Continue Chat" appeared to fix it only because that starts a new turn. The tool loop now checks the threshold before each continuation request and compacts when it is crossed. Turn-start compaction cuts only at a user message and an agentic run has none after its first, so this gets its own boundary that cuts on any non-tool message, which keeps every tool call together with its result. The user's request is carried forward when the cut would drop it, and the summary travels in its own system message so no existing prompt is rewritten. Simulated 10-iteration run at an 8000 token threshold: the raw request grows from 3059 to 30338 tokens while what is actually sent stays under 6583. Closes #27599
This commit is contained in:
parent
e9a0164690
commit
231ee6b6a0
2 changed files with 168 additions and 2 deletions
|
|
@ -8,7 +8,12 @@ from open_webui.models.chats import Chats
|
|||
from open_webui.models.config import Config
|
||||
from open_webui.utils.chat_id import is_saved_chat_id
|
||||
from open_webui.utils.json_codec import JSONCodec
|
||||
from open_webui.utils.misc import get_content_from_message, get_last_user_message, get_message_list
|
||||
from open_webui.utils.misc import (
|
||||
get_content_from_message,
|
||||
get_last_user_message,
|
||||
get_last_user_message_item,
|
||||
get_message_list,
|
||||
)
|
||||
from open_webui.utils.payload import apply_params_to_form_data
|
||||
from open_webui.utils.task import (
|
||||
prompt_template,
|
||||
|
|
@ -147,6 +152,108 @@ async def compact_messages_for_request(
|
|||
return [*system_messages, *recent_messages], summary, True
|
||||
|
||||
|
||||
class ToolLoopCompactor:
|
||||
"""Compacts the native tool call loop's own history when it crosses the token threshold."""
|
||||
|
||||
def __init__(self, request, user, metadata: dict, messages: list[dict]) -> None:
|
||||
self._request = request
|
||||
self._user = user
|
||||
self._metadata = metadata
|
||||
self._summary: str | None = None
|
||||
self._config: dict | None = None
|
||||
self._stopped = False
|
||||
self._dropped_count = 0
|
||||
self._sent_count = len(messages) - (1 if messages and messages[0].get('role') == 'system' else 0)
|
||||
self._user_message = get_last_user_message_item(messages)
|
||||
self._preserved_user: list[dict] = []
|
||||
|
||||
async def apply(self, messages: list[dict], usage: dict | None, model_id: str) -> list[dict]:
|
||||
if self._config is None:
|
||||
self._config = await _load_config()
|
||||
if not self._config['enable']:
|
||||
return messages
|
||||
|
||||
system_messages = [messages[0]] if messages and messages[0].get('role') == 'system' else []
|
||||
system_prompt = (get_content_from_message(system_messages[0]) or '') if system_messages else ''
|
||||
history = messages[len(system_messages) :]
|
||||
|
||||
# An outside insert can shift the cut, so walk it back onto a whole call/result block.
|
||||
while self._dropped_count < len(history) and history[self._dropped_count].get('role') == 'tool':
|
||||
self._dropped_count += 1
|
||||
history = history[self._dropped_count :]
|
||||
|
||||
if not self._stopped and self._exceeds_threshold(history, system_prompt, usage):
|
||||
try:
|
||||
history = await self._summarize(history, model_id)
|
||||
except Exception:
|
||||
self._stopped = True
|
||||
log.exception('Tool loop context compaction failed; keeping the history trimmed so far')
|
||||
|
||||
self._sent_count = len(history)
|
||||
summary_messages = (
|
||||
[{'role': 'system', 'content': f'[CONVERSATION SUMMARY]\n{self._summary}'}] if self._summary else []
|
||||
)
|
||||
return [*system_messages, *summary_messages, *self._preserved_user, *history]
|
||||
|
||||
def _exceeds_threshold(self, history: list[dict], system_prompt: str, usage: dict | None) -> bool:
|
||||
if len(history) <= 3:
|
||||
return False
|
||||
|
||||
threshold = _resolve_token_threshold(self._config['token_threshold'], self._config['token_cap'], self._metadata)
|
||||
|
||||
reported_tokens = _usage_token_count(usage or {})
|
||||
if reported_tokens:
|
||||
# The assistant turn that follows what was sent is already inside completion_tokens.
|
||||
return reported_tokens + _estimate_messages_tokens(history[self._sent_count + 1 :]) > threshold
|
||||
|
||||
estimated = (
|
||||
_estimate_tokens(system_prompt)
|
||||
+ _estimate_tokens(self._summary or '')
|
||||
+ _estimate_messages_tokens(self._preserved_user)
|
||||
+ _estimate_messages_tokens(history)
|
||||
)
|
||||
return estimated > threshold
|
||||
|
||||
async def _summarize(self, history: list[dict], model_id: str) -> list[dict]:
|
||||
boundary = _find_tool_loop_boundary(history, self._config['retention_percentage'])
|
||||
compacted_messages, recent_messages = history[:boundary], history[boundary:]
|
||||
drops_only_the_user_message = all(message is self._user_message for message in compacted_messages)
|
||||
if drops_only_the_user_message:
|
||||
return history
|
||||
|
||||
await _emit_compaction_status(self._metadata, 'Compacting context', done=False)
|
||||
try:
|
||||
self._summary = await _generate_summary(
|
||||
self._request,
|
||||
self._user,
|
||||
model_id,
|
||||
_get_compaction_models(self._request),
|
||||
_describe_tool_calls(compacted_messages),
|
||||
_describe_tool_calls(recent_messages),
|
||||
self._summary,
|
||||
self._config['prompt_template'],
|
||||
)
|
||||
except Exception:
|
||||
await _emit_compaction_status(self._metadata, 'Context compaction failed', done=True, error=True)
|
||||
raise
|
||||
|
||||
self._dropped_count += boundary
|
||||
keeps_user_message = any(message is self._user_message for message in recent_messages)
|
||||
if self._user_message and not self._preserved_user and not keeps_user_message:
|
||||
self._preserved_user = [self._user_message]
|
||||
|
||||
log.info(
|
||||
'Compacted tool loop context for chat=%s response=%s dropped=%d kept=%d summary_chars=%d',
|
||||
self._metadata.get('chat_id'),
|
||||
self._metadata.get('message_id'),
|
||||
len(compacted_messages),
|
||||
len(recent_messages),
|
||||
len(self._summary or ''),
|
||||
)
|
||||
await _emit_compaction_status(self._metadata, 'Context compacted', done=True)
|
||||
return recent_messages
|
||||
|
||||
|
||||
async def compact_chat_branch(request, user, chat: Any, model_id: str, models: dict) -> dict:
|
||||
config = await _load_config()
|
||||
if not config['enable']:
|
||||
|
|
@ -297,6 +404,31 @@ def _build_context_usage(tokens: int, threshold: int) -> dict:
|
|||
}
|
||||
|
||||
|
||||
def _get_compaction_models(request) -> dict:
|
||||
"""Return the model registry with any direct-connection model from request.state merged in."""
|
||||
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
|
||||
return {**dict(request.app.state.MODELS.items()), request.state.model['id']: request.state.model}
|
||||
return request.app.state.MODELS
|
||||
|
||||
|
||||
async def _emit_compaction_status(metadata: dict, description: str, done: bool, error: bool = False) -> None:
|
||||
if not (metadata.get('chat_id') and metadata.get('message_id')):
|
||||
return
|
||||
|
||||
from open_webui.socket.main import get_event_emitter
|
||||
|
||||
data = {'action': 'context_compaction', 'description': description, 'done': done}
|
||||
if error:
|
||||
data['error'] = True
|
||||
|
||||
try:
|
||||
event_emitter = await get_event_emitter(metadata)
|
||||
if event_emitter:
|
||||
await event_emitter({'type': 'context_compaction', 'data': data})
|
||||
except Exception:
|
||||
log.debug('Could not emit context compaction status')
|
||||
|
||||
|
||||
def _apply_latest_summary_checkpoint(messages: list[dict]) -> tuple[list[dict], str | None]:
|
||||
summary = None
|
||||
summary_idx = None
|
||||
|
|
@ -333,6 +465,14 @@ def _find_compaction_boundary(messages: list[dict], retention_percentage: int =
|
|||
return next((idx for idx in reversed(boundaries) if idx <= target), 0)
|
||||
|
||||
|
||||
def _find_tool_loop_boundary(messages: list[dict], retention_percentage: int) -> int:
|
||||
keep_count = max(2, len(messages) * retention_percentage // 100)
|
||||
target = max(1, len(messages) - keep_count)
|
||||
# Cutting on a tool result orphans it from its call and the provider rejects the request.
|
||||
boundaries = [idx for idx, message in enumerate(messages) if idx and message.get('role') != 'tool']
|
||||
return next((idx for idx in reversed(boundaries) if idx <= target), 0)
|
||||
|
||||
|
||||
async def _generate_summary(
|
||||
request,
|
||||
user,
|
||||
|
|
@ -395,6 +535,24 @@ async def _generate_summary(
|
|||
return '\n'.join(parts)[:4000]
|
||||
|
||||
|
||||
def _describe_tool_calls(messages: list[dict]) -> list[dict]:
|
||||
"""Fold each tool call's name and arguments into its message content."""
|
||||
described = []
|
||||
for message in messages:
|
||||
tool_calls = message.get('tool_calls') or []
|
||||
if not tool_calls:
|
||||
described.append(message)
|
||||
continue
|
||||
|
||||
calls = ', '.join(
|
||||
f'{(call.get("function") or {}).get("name", "")}({(call.get("function") or {}).get("arguments", "")})'
|
||||
for call in tool_calls
|
||||
)
|
||||
content = get_content_from_message(message) or ''
|
||||
described.append({**message, 'content': f'{content}\n[TOOL CALLS] {calls}'.strip()})
|
||||
return described
|
||||
|
||||
|
||||
def _response_text(response: Any) -> str:
|
||||
if isinstance(response, list) and len(response) == 1:
|
||||
response = response[0]
|
||||
|
|
|
|||
|
|
@ -87,7 +87,7 @@ from open_webui.utils.ask_user import stage_ask_user_tool_calls
|
|||
from open_webui.utils.chat import generate_chat_completion
|
||||
from open_webui.utils.chat_id import is_saved_chat_id
|
||||
from open_webui.utils.code_interpreter import execute_code_jupyter
|
||||
from open_webui.utils.context_compaction import compact_messages_for_request
|
||||
from open_webui.utils.context_compaction import ToolLoopCompactor, compact_messages_for_request
|
||||
from open_webui.utils.files import (
|
||||
convert_markdown_base64_images,
|
||||
get_file_url_from_base64,
|
||||
|
|
@ -5724,6 +5724,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
|
||||
tool_loop_compactor = ToolLoopCompactor(request, user, metadata, form_data['messages'])
|
||||
user_message = get_last_user_message(form_data['messages'])
|
||||
|
||||
# Check if citations are enabled for this model
|
||||
|
|
@ -6158,6 +6159,13 @@ async def streaming_chat_response_handler(response, ctx):
|
|||
}
|
||||
)
|
||||
|
||||
try:
|
||||
new_form_data['messages'] = await tool_loop_compactor.apply(
|
||||
new_form_data['messages'], usage, model_id
|
||||
)
|
||||
except Exception:
|
||||
log.exception('Tool loop compaction failed; continuing with full tool history')
|
||||
|
||||
new_form_data = await convert_url_images_to_base64(new_form_data, user=user)
|
||||
|
||||
if filter_functions:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue