mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-11 03:38:02 +00:00
refac
This commit is contained in:
parent
930f8c3640
commit
67cfa917c1
2 changed files with 75 additions and 25 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue