diff --git a/backend/open_webui/utils/context_compaction.py b/backend/open_webui/utils/context_compaction.py index 3509b00a12..b38a4eb346 100644 --- a/backend/open_webui/utils/context_compaction.py +++ b/backend/open_webui/utils/context_compaction.py @@ -259,9 +259,25 @@ async def get_chat_context_usage(chat: Any, model_id: str | None = None) -> dict for idx in range(len(messages) - 1, -1, -1): usage = messages[idx].get('usage') or (messages[idx].get('info') or {}).get('usage') - input_tokens = (usage or {}).get('prompt_tokens') or (usage or {}).get('input_tokens') - if isinstance(usage, dict) and input_tokens: - tokens = int(input_tokens or 0) + int(usage.get('completion_tokens') or usage.get('output_tokens') or 0) + if isinstance(usage, dict) and ( + tokens := ( + int( + usage.get('prompt_tokens') + or usage.get('input_tokens') + or usage.get('prompt_eval_count') + or usage.get('prompt_n') + or 0 + ) + + int( + usage.get('completion_tokens') + or usage.get('output_tokens') + or usage.get('eval_count') + or usage.get('predicted_n') + or 0 + ) + + int(usage.get('cache_n') or 0) + ) + ): tokens += _estimate_messages_tokens(messages[idx + 1 :]) return _build_context_usage(tokens, threshold) @@ -300,11 +316,26 @@ def _exceeds_token_threshold(messages: list[dict], system_prompt: str, summary: 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 (usage.get('prompt_tokens') or usage.get('input_tokens')): - total = int(usage.get('prompt_tokens') or usage.get('input_tokens') or 0) + int( - usage.get('completion_tokens') or usage.get('output_tokens') or 0 + if isinstance(usage, dict) and ( + tokens := ( + int( + usage.get('prompt_tokens') + or usage.get('input_tokens') + or usage.get('prompt_eval_count') + or usage.get('prompt_n') + or 0 + ) + + int( + usage.get('completion_tokens') + or usage.get('output_tokens') + or usage.get('eval_count') + or usage.get('predicted_n') + or 0 + ) + + int(usage.get('cache_n') or 0) ) - return total + _estimate_messages_tokens(messages[idx + 1 :]) > threshold + ): + 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 diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index 699730da8c..e2973d9a12 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -2035,7 +2035,7 @@ async def load_messages_from_db(chat_id: str, message_id: str) -> Optional[list[ return None return [ - {k: v for k, v in msg.items() if k in ('id', 'role', 'content', 'output', 'files', 'contextSummary')} + {k: v for k, v in msg.items() if k in ('id', 'role', 'content', 'output', 'files', 'contextSummary', 'usage')} for msg in db_messages ] @@ -2095,6 +2095,7 @@ def strip_compaction_fields(messages: list[dict]) -> list[dict]: clean = dict(message) clean.pop('contextSummary', None) clean.pop('context_summary', None) + clean.pop('usage', None) clean.pop('id', None) stripped.append(clean) return stripped @@ -2290,7 +2291,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): { k: v for k, v in assistant_message.items() - if k in ('id', 'role', 'content', 'output', 'files', 'contextSummary') + if k in ('id', 'role', 'content', 'output', 'files', 'contextSummary', 'usage') } )