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:
@@ -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
|
||||
|
||||
@@ -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')
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user