mirror of
https://github.com/open-webui/open-webui.git
synced 2026-08-13 01:02:25 -06:00
refac
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user