This commit is contained in:
Timothy Jaeryang Baek
2026-08-17 00:57:57 -07:00
parent 017075a2d7
commit 0b27fa5e87
2 changed files with 22 additions and 45 deletions
+19 -38
View File
@@ -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)
+3 -7
View File
@@ -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