diff --git a/backend/open_webui/utils/context_compaction.py b/backend/open_webui/utils/context_compaction.py index 1a9b5e9f96..1fc5296c56 100644 --- a/backend/open_webui/utils/context_compaction.py +++ b/backend/open_webui/utils/context_compaction.py @@ -260,9 +260,14 @@ 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('input_tokens') or (usage or {}).get('prompt_tokens') + 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('output_tokens') or usage.get('completion_tokens') or 0) + tokens = int(input_tokens or 0) + int( + usage.get('completion_tokens') or usage.get('output_tokens') or 0 + ) tokens += _estimate_messages_tokens(messages[idx + 1 :]) return _build_context_usage(tokens, threshold) @@ -301,8 +306,10 @@ 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('input_tokens'): - total = int(usage.get('input_tokens') or 0) + int(usage.get('output_tokens') or 0) + 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 + ) return total + _estimate_messages_tokens(messages[idx + 1 :]) > threshold estimated = _estimate_tokens(system_prompt) + _estimate_tokens(summary or '') + _estimate_messages_tokens(messages) diff --git a/backend/open_webui/utils/response.py b/backend/open_webui/utils/response.py index de3ede0a07..8b0214f6c8 100644 --- a/backend/open_webui/utils/response.py +++ b/backend/open_webui/utils/response.py @@ -55,8 +55,6 @@ USAGE_TOKEN_KEYS = { 'input_tokens', 'output_tokens', 'total_tokens', - 'prompt_tokens', - 'completion_tokens', } USAGE_COST_KEYS = { @@ -105,7 +103,7 @@ def merge_usage(current: dict | None, incoming: dict | None) -> dict: """ Merge usage payloads from multiple model calls into one cumulative usage dict. - Token fields are additive; non-numeric metadata keeps the latest provider value. + Canonical token fields are additive; provider aliases keep the latest value. """ current_usage = normalize_usage(current or {}) if current else {} incoming_usage = normalize_usage(incoming or {}) if incoming else {} @@ -133,6 +131,17 @@ def merge_usage(current: dict | None, incoming: dict | None) -> dict: incoming_usage.get(key) if isinstance(incoming_usage.get(key), dict) else {}, ) + result['prompt_tokens'] = ( + incoming_usage.get('prompt_tokens') + or incoming_usage.get('input_tokens') + or current_usage.get('prompt_tokens', 0) + ) + result['completion_tokens'] = ( + incoming_usage.get('completion_tokens') + or incoming_usage.get('output_tokens') + or current_usage.get('completion_tokens', 0) + ) + return result