This commit is contained in:
Classic298 2026-10-01 01:50:18 +08:00 • committed by GitHub
commit a5d096f262
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 168 additions and 2 deletions

View file

@ -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]

View file

@ -88,7 +88,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,
@ -5826,6 +5826,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
@ -6271,6 +6272,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: