refac
This commit is contained in:
@@ -235,6 +235,23 @@ def _resolve_token_threshold(global_threshold: int, global_cap: int, metadata: d
|
||||
return min(configured_threshold or global_threshold, global_cap)
|
||||
|
||||
|
||||
def _usage_token_count(usage: dict) -> int:
|
||||
prompt_tokens = int(usage.get('prompt_tokens') or usage.get('prompt_eval_count') or 0)
|
||||
if not prompt_tokens and (usage.get('prompt_n') is not None or usage.get('cache_n') is not None):
|
||||
prompt_tokens = int(usage.get('prompt_n') or 0) + int(usage.get('cache_n') or 0)
|
||||
if not prompt_tokens:
|
||||
prompt_tokens = int(usage.get('input_tokens') or 0)
|
||||
|
||||
completion_tokens = int(
|
||||
usage.get('completion_tokens')
|
||||
or usage.get('output_tokens')
|
||||
or usage.get('eval_count')
|
||||
or usage.get('predicted_n')
|
||||
or 0
|
||||
)
|
||||
return prompt_tokens + completion_tokens
|
||||
|
||||
|
||||
async def get_chat_context_usage(chat: Any, model_id: str | None = None) -> dict | None:
|
||||
chat_data = chat.chat or {}
|
||||
history = chat_data.get('history') or {}
|
||||
@@ -263,25 +280,7 @@ 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')
|
||||
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)
|
||||
)
|
||||
):
|
||||
if isinstance(usage, dict) and (tokens := _usage_token_count(usage)):
|
||||
tokens += _estimate_messages_tokens(messages[idx + 1 :])
|
||||
return _build_context_usage(tokens, threshold)
|
||||
|
||||
@@ -320,25 +319,7 @@ 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 (
|
||||
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)
|
||||
)
|
||||
):
|
||||
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)
|
||||
|
||||
@@ -24,13 +24,9 @@ def normalize_usage(usage: dict) -> dict:
|
||||
return {}
|
||||
|
||||
# Map various field names to standard names
|
||||
input_tokens = (
|
||||
usage.get('input_tokens') # Already standard
|
||||
or usage.get('prompt_tokens') # OpenAI
|
||||
or usage.get('prompt_eval_count') # Ollama
|
||||
or usage.get('prompt_n') # llama.cpp
|
||||
or 0
|
||||
)
|
||||
input_tokens = usage.get('input_tokens') or usage.get('prompt_tokens') or usage.get('prompt_eval_count')
|
||||
if input_tokens is None:
|
||||
input_tokens = int(usage.get('prompt_n') or 0) + int(usage.get('cache_n') or 0)
|
||||
|
||||
output_tokens = (
|
||||
usage.get('output_tokens') # Already standard
|
||||
|
||||
Reference in New Issue
Block a user