mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
61994a6704
commit
bb90256ba1
2 changed files with 49 additions and 8 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"}],
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue