mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-05 02:41:34 +00:00
Merge 231ee6b6a0 into 4ef7e35b88
This commit is contained in:
commit
a5d096f262
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]
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue