mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(anthropic): normalize images for provider token counting (#45185)
This commit is contained in:
parent
d7c6c4b80f
commit
6befb9ad7d
4 changed files with 202 additions and 3 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue