fix(proxy): count contents text in the local token_counter fallback

a contents-only request that hits a provider error used to raise
ValueError -> 500 on the local fallback; translate text parts to
chat messages so the fallback counts them
This commit is contained in:
shrey kharbanda 2026-09-24 03:19:00 +00:00
parent 61994a6704
commit bb90256ba1
2 changed files with 49 additions and 8 deletions

View file

@ -13459,6 +13459,26 @@ def _system_message(system: object) -> ChatCompletionSystemMessage | None:
return message
def _contents_as_messages(contents: object) -> tuple[Mapping[str, object], ...] | None:
"""Approximate gemini contents as chat messages for the local fallback
tokenizer; only text parts are countable locally."""
if not isinstance(contents, list):
return None
messages: Final = tuple(
{ # mutable-ok: transient chat-shaped message for the local tokenizer
"role": "assistant" if content.get("role") == "model" else "user",
"content": "\n".join(
part["text"]
for part in content.get("parts", ())
if isinstance(part, Mapping) and isinstance(part.get("text"), str)
),
}
for content in contents
if isinstance(content, Mapping)
)
return tuple(message for message in messages if message["content"]) or None
@router.post(
"/utils/token_counter",
tags=["llm utils"],
@ -13561,7 +13581,8 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False)
tokenizer_used: Final = str(_tokenizer_used["type"])
system_message: Final = _system_message(system)
typed_messages: Final = cast( # cast-ok: request messages are raw chat-shaped dicts that token_counter normalizes
Sequence[AllMessageValues] | None, messages
Sequence[AllMessageValues] | None,
messages if messages is not None else _contents_as_messages(contents),
)
counted_messages: Final = (
typed_messages if typed_messages is None or system_message is None else (system_message, *typed_messages)

View file

@ -84,9 +84,7 @@ def test_token_counter_counts_off_the_event_loop(client, auth_as, patched_token_
assert counted_off_loop == [True]
def test_token_counter_missing_input_returns_400(
client, auth_as, patched_token_counter
):
def test_token_counter_missing_input_returns_400(client, auth_as, patched_token_counter):
"""Pins ``POST /utils/token_counter`` (error: missing input)."""
with auth_as():
response = client.post("/utils/token_counter", json={"model": "gpt-4"})
@ -118,9 +116,7 @@ def patched_supported_params(monkeypatch):
def test_supported_openai_params_happy_path(client, auth_as, patched_supported_params):
"""Pins ``GET /utils/supported_openai_params``."""
with auth_as():
response = client.get(
"/utils/supported_openai_params", params={"model": "gpt-4"}
)
response = client.get("/utils/supported_openai_params", params={"model": "gpt-4"})
assert response.status_code == 200
assert normalize(response.json()) == {
"supported_openai_params": ["max_tokens", "temperature", "top_p"],
@ -359,7 +355,10 @@ def test_token_counter_fallback_counts_tools_system_and_anthropic_blocks(client,
"content": [
{"type": "text", "text": "What is in this file?"},
{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "iVBORw0KGgo="}},
{"type": "document", "source": {"type": "base64", "media_type": "application/pdf", "data": "JVBERi0xLjQK"}},
{
"type": "document",
"source": {"type": "base64", "media_type": "application/pdf", "data": "JVBERi0xLjQK"},
},
],
}
]
@ -405,3 +404,24 @@ def test_token_counter_fallback_prompt_with_tools_does_not_500(client, auth_as,
assert response.status_code == 200, response.text
assert response.json()["total_tokens"] == litellm.token_counter(model="claude-fable-5", text=prompt)
def test_token_counter_contents_only_request_counts_text_parts(client, auth_as, monkeypatch):
"""Regression: a contents-only request (google countTokens shape) reaches the
local fallback as translated messages instead of raising ValueError -> 500."""
monkeypatch.setattr(proxy_server, "llm_router", None)
monkeypatch.setattr(litellm, "disable_token_counter", False, raising=False)
contents = [
{"role": "user", "parts": [{"text": "hello world"}, {"inline_data": {"data": "abc"}}]},
{"role": "model", "parts": [{"text": "hi there"}]},
{"role": "user", "parts": [{"inline_data": {"data": "abc"}}]},
]
with auth_as():
response = client.post("/utils/token_counter", json={"model": "claude-fable-5", "contents": contents})
assert response.status_code == 200, response.text
assert response.json()["total_tokens"] == litellm.token_counter(
model="claude-fable-5",
messages=[{"role": "user", "content": "hello world"}, {"role": "assistant", "content": "hi there"}],
)