fix(anthropic): normalize images for provider token counting (#45185)

This commit is contained in:
tin-berri 2026-10-07 16:36:27 -07:00 • committed by GitHub
parent d7c6c4b80f
commit 6befb9ad7d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 202 additions and 3 deletions

View file

@ -12,6 +12,7 @@ from pydantic import JsonValue, TypeAdapter
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.llms.anthropic.common_utils import AnthropicError
from litellm.llms.anthropic.count_tokens.transformation import (
AnthropicCountTokensConfig,
@ -62,7 +63,7 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig):
verbose_logger.debug("Processing Anthropic CountTokens request for model: %s", model)
# Transform request to Anthropic format
request_body: Final = self.transform_request_to_count_tokens(
request_body: Final = await asyncify(self.transform_request_to_count_tokens)(
model=model,
messages=messages,
tools=tools,

View file

@ -13,11 +13,41 @@ from pydantic import JsonValue, TypeAdapter
from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION
from litellm.llms.anthropic.common_utils import merge_anthropic_beta_headers
from litellm.llms.anthropic.wif import resolve_anthropic_base
from litellm.types.llms.openai import ChatCompletionImageObject
_COUNT_REQUEST: Final = TypeAdapter(dict[str, JsonValue])
_IMAGE_BLOCK: Final = TypeAdapter(ChatCompletionImageObject)
COUNT_TOKEN_OPTION_NAMES: Final = ("thinking", "tool_choice", "output_config")
def _count_image(block: JsonValue) -> JsonValue:
if not isinstance(block, dict) or block.get("type") != "image_url":
return block
from litellm.litellm_core_utils.prompt_templates.factory import convert_to_anthropic_image_obj
image_block: Final = _IMAGE_BLOCK.validate_python(block)
image_url: Final = image_block["image_url"]
source: Final = convert_to_anthropic_image_obj(
openai_image_url=image_url if isinstance(image_url, str) else image_url["url"],
format=image_url.get("format") if isinstance(image_url, dict) else None,
)
image: Final = _COUNT_REQUEST.validate_python({"type": "image", "source": source})
return {**{key: value for key, value in block.items() if key not in {"type", "image_url"}}, **image}
def _count_block(block: JsonValue) -> JsonValue:
if not isinstance(block, dict) or block.get("type") != "tool_result":
return _count_image(block)
content: Final = block.get("content")
if not isinstance(content, list):
return block
return {**block, "content": [_count_image(part) for part in content]}
def _count_content(content: JsonValue) -> JsonValue:
return [_count_block(block) for block in content] if isinstance(content, list) else content
class AnthropicCountTokensConfig:
"""
Configuration and transformation logic for Anthropic CountTokens API.
@ -62,7 +92,7 @@ class AnthropicCountTokensConfig:
MappingProxyType(
{
"model": model,
"messages": messages,
"messages": [{**message, "content": _count_content(message["content"])} for message in messages],
**MappingProxyType(
{key: value for key, value in (("system", system), ("tools", tools)) if value is not None}
),

View file

@ -10,6 +10,7 @@ import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.llms.anthropic.common_utils import AnthropicError
from litellm.llms.azure_ai.anthropic.count_tokens.transformation import (
AzureAIAnthropicCountTokensConfig,
@ -59,7 +60,7 @@ class AzureAIAnthropicCountTokensHandler(AzureAIAnthropicCountTokensConfig):
verbose_logger.debug("Processing Azure AI Anthropic CountTokens request for model: %s", model)
# Transform request to Anthropic format
request_body: Final = self.transform_request_to_count_tokens(
request_body: Final = await asyncify(self.transform_request_to_count_tokens)(
model=model,
messages=messages,
tools=tools,

View file

@ -1,12 +1,113 @@
import asyncio
from base64 import b64encode
from copy import deepcopy
from threading import get_ident
from typing import Final
import httpx
import pytest
import respx
from pydantic import JsonValue, TypeAdapter
import litellm
from litellm.llms.anthropic.count_tokens.handler import AnthropicCountTokensHandler
from litellm.llms.anthropic.count_tokens.transformation import (
AnthropicCountTokensConfig,
)
from litellm.llms.azure_ai.anthropic.count_tokens.handler import (
AzureAIAnthropicCountTokensHandler,
)
from litellm.llms.azure_ai.anthropic.count_tokens.transformation import (
AzureAIAnthropicCountTokensConfig,
)
@pytest.mark.parametrize(
"config_type", (AnthropicCountTokensConfig, AzureAIAnthropicCountTokensConfig)
)
@pytest.mark.parametrize(
("image_url", "source"),
(
("data:image/png;base64,aW1hZ2U=", {"type": "base64", "media_type": "image/png", "data": "aW1hZ2U="}),
({"url": "data:image/png;base64,aW1hZ2U="}, {"type": "base64", "media_type": "image/png", "data": "aW1hZ2U="}),
(
{"url": "data:image/png;base64,aW1hZ2U=", "format": "image/jpeg", "detail": "high"},
{"type": "base64", "media_type": "image/jpeg", "data": "aW1hZ2U="},
),
),
)
def test_count_translates_openai_images_without_mutating_input(
config_type: type[AnthropicCountTokensConfig],
image_url: str | dict[str, JsonValue],
source: dict[str, JsonValue],
) -> None:
cache_control: Final[dict[str, JsonValue]] = {"type": "ephemeral"}
messages: Final[list[dict[str, JsonValue]]] = [{
"role": "user", "content": [
{"type": "text", "text": "Count this image"},
{"type": "image_url", "image_url": image_url, "cache_control": cache_control},
],
}]
original: Final = deepcopy(messages)
result: Final = config_type().transform_request_to_count_tokens(
model="claude-opus-5-5", messages=messages
)
assert result == {
"model": "claude-opus-5-5", "messages": [{
"role": "user", "content": [
{"type": "text", "text": "Count this image"},
{"type": "image", "source": source, "cache_control": cache_control},
],
}],
}
assert messages == original
@pytest.mark.parametrize(
"config_type", (AnthropicCountTokensConfig, AzureAIAnthropicCountTokensConfig)
)
def test_count_normalizes_nested_tool_images_and_preserves_native_fields(
config_type: type[AnthropicCountTokensConfig],
) -> None:
openai_image: Final[dict[str, JsonValue]] = {
"type": "image_url", "image_url": {"url": "data:image/png;base64,aW1hZ2U="}
}
native_image: Final[dict[str, JsonValue]] = {
"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "aW1hZ2U="}
}
assistant: Final[dict[str, JsonValue]] = {"role": "assistant", "content": [
{"type": "thinking", "thinking": "inspect screenshot", "signature": "fixture-signature"},
{"type": "tool_use", "id": "read-1", "name": "Read", "input": {"content": [openai_image]}},
]}
tool_result: Final[dict[str, JsonValue]] = {
"type": "tool_result", "tool_use_id": "read-1", "is_error": False,
"content": [{"type": "text", "text": "Screenshot"}, native_image, openai_image],
"cache_control": {"type": "ephemeral"},
}
text_result: Final[dict[str, JsonValue]] = {"type": "tool_result", "tool_use_id": "read-2", "content": "done"}
messages: Final[list[dict[str, JsonValue]]] = [
assistant, {"role": "user", "content": [native_image, tool_result, text_result]}
]
tools: Final[list[dict[str, JsonValue]]] = [{
"name": "Read", "input_schema": {"type": "object", "examples": [openai_image]}
}]
system: Final[JsonValue] = [{"type": "text", "text": "policy", "cache_control": {"type": "ephemeral"}}]
options: Final[dict[str, JsonValue]] = {
"thinking": {"type": "adaptive"}, "tool_choice": {"type": "auto"}, "output_config": {"effort": "high"}
}
original: Final = deepcopy((messages, tools, system, options))
result: Final = config_type().transform_request_to_count_tokens(
model="claude-opus-5-5", messages=messages, tools=tools, system=system, optional_params=options
)
assert result == {
"model": "claude-opus-5-5", "system": system, "tools": tools, **options,
"messages": [assistant, {"role": "user", "content": [native_image, {
**tool_result, "content": [{"type": "text", "text": "Screenshot"}, native_image, native_image]
}, text_result]}],
}
assert (messages, tools, system, options) == original
def test_transform_basic_request():
@ -162,3 +263,69 @@ async def test_handler_posts_to_count_tokens_path_under_deployment_api_base(http
assert route.called
assert result == {"input_tokens": 7}
@pytest.mark.asyncio
@pytest.mark.usefixtures("httpx_transport_clients")
@pytest.mark.parametrize(
"handler_type", (AnthropicCountTokensHandler, AzureAIAnthropicCountTokensHandler)
)
@pytest.mark.parametrize("scheme", ("http", "https"))
@pytest.mark.parametrize("dict_url", (False, True))
async def test_remote_image_fetch_keeps_counting_handler_event_loop_responsive(
handler_type: type[AnthropicCountTokensHandler] | type[AzureAIAnthropicCountTokensHandler],
scheme: str,
dict_url: bool,
) -> None:
loop: Final = asyncio.get_running_loop()
loop_thread: Final = get_ident()
witness: Final = asyncio.Event()
image_bytes: Final = b"\x89PNG\r\n\x1a\ncount-image"
image_url: Final = f"{scheme}://1.1.1.1/{handler_type.__name__}-{dict_url}.png"
model: Final = "claude-opus-5-5"
api_base: Final = "https://gateway.example/anthropic"
image: Final[dict[str, JsonValue]] = {
"type": "image_url", "image_url": {"url": image_url} if dict_url else image_url,
"cache_control": {"type": "ephemeral"},
}
messages: Final[list[dict[str, JsonValue]]] = [{
"role": "user", "content": [image, {"type": "tool_result", "tool_use_id": "read-1", "content": [image]}]
}]
original: Final = deepcopy(messages)
async def run_witness() -> None:
witness.set()
def image_response(_request: httpx.Request) -> httpx.Response:
assert get_ident() != loop_thread, "image fetch blocked the counting handler's event loop"
asyncio.run_coroutine_threadsafe(run_witness(), loop).result(timeout=5)
return httpx.Response(200, content=image_bytes, headers={"Content-Type": "image/png"})
with respx.mock:
image_route: Final = respx.get(image_url).mock(side_effect=image_response)
count_route: Final = respx.post(f"{api_base}/v1/messages/count_tokens").mock(
return_value=httpx.Response(200, json={"input_tokens": 7})
)
handler: Final = handler_type()
result: Final = await (
handler.handle_count_tokens_request(
model=model, messages=messages, api_base=api_base, auth_header={"x-api-key": "test-key"}
) if isinstance(handler, AnthropicCountTokensHandler) else handler.handle_count_tokens_request(
model=model, messages=messages, api_base=api_base, api_key="test-key"
)
)
assert witness.is_set()
assert image_route.call_count == count_route.call_count == 1
assert result == {"input_tokens": 7}
native_image: Final[dict[str, JsonValue]] = {
"type": "image", "source": {
"type": "base64", "media_type": "image/png", "data": b64encode(image_bytes).decode()
}, "cache_control": {"type": "ephemeral"},
}
assert TypeAdapter(dict[str, JsonValue]).validate_json(count_route.calls.last.request.content) == {
"model": model, "messages": [{"role": "user", "content": [
native_image, {"type": "tool_result", "tool_use_id": "read-1", "content": [native_image]}
]}],
}
assert messages == original