mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(gemini): count Anthropic messages, system and tools on /v1/messages/count_tokens and fall back on provider errors
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
860bc7811d
commit
d9b612e1e2
9 changed files with 416 additions and 6 deletions
|
|
@ -495,17 +495,43 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
|
|||
import copy
|
||||
|
||||
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
|
||||
from litellm.llms.gemini.count_tokens.transformation import build_count_tokens_payload
|
||||
|
||||
if contents is None and not messages:
|
||||
return None
|
||||
|
||||
deployment = deployment or {}
|
||||
count_tokens_params_request: Final = copy.deepcopy(deployment.get("litellm_params", {}))
|
||||
payload: Final = (
|
||||
build_count_tokens_payload(model=model_to_use, messages=messages, system=system, tools=tools)
|
||||
if contents is None
|
||||
else None
|
||||
)
|
||||
count_tokens_params: Final = {
|
||||
"model": model_to_use,
|
||||
"contents": contents,
|
||||
"contents": payload.contents if payload is not None else contents,
|
||||
**(
|
||||
{"system_instruction": payload.system_instruction}
|
||||
if payload is not None and payload.system_instruction is not None
|
||||
else {}
|
||||
),
|
||||
**({"tools": payload.tools} if payload is not None and payload.tools is not None else {}),
|
||||
}
|
||||
count_tokens_params_request.update(count_tokens_params)
|
||||
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
**count_tokens_params_request,
|
||||
)
|
||||
try:
|
||||
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
**count_tokens_params_request,
|
||||
)
|
||||
except (litellm.APIError, litellm.APIConnectionError) as e:
|
||||
return TokenCountResponse(
|
||||
total_tokens=0,
|
||||
request_model=request_model,
|
||||
model_used=model_to_use,
|
||||
tokenizer_type="gemini_api",
|
||||
error=True,
|
||||
error_message=e.message,
|
||||
status_code=e.status_code,
|
||||
)
|
||||
|
||||
if result is not None:
|
||||
return TokenCountResponse(
|
||||
|
|
|
|||
|
|
@ -4,6 +4,8 @@ import httpx
|
|||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.vertex_ai import SystemInstructions, Tools
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -51,7 +53,7 @@ class GoogleAIStudioTokenCounter:
|
|||
"""
|
||||
Construct the URL for the Google Gen AI Studio countTokens endpoint.
|
||||
"""
|
||||
base_url: Final = api_base or "https://generativelanguage.googleapis.com"
|
||||
base_url: Final = api_base or get_secret_str("GEMINI_API_BASE") or "https://generativelanguage.googleapis.com"
|
||||
return f"{base_url}/v1beta/models/{model}:countTokens"
|
||||
|
||||
async def validate_environment(
|
||||
|
|
@ -84,6 +86,8 @@ class GoogleAIStudioTokenCounter:
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
system_instruction: SystemInstructions | None = None,
|
||||
tools: list[Tools] | None = None,
|
||||
**kwargs: object,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
|
|
@ -130,7 +134,18 @@ class GoogleAIStudioTokenCounter:
|
|||
|
||||
# Prepare request body - clean up contents to remove unsupported fields
|
||||
cleaned_contents: Final = self._clean_contents_for_gemini_api(contents)
|
||||
request_body: Final = {"contents": cleaned_contents}
|
||||
request_body: Final = (
|
||||
{"contents": cleaned_contents}
|
||||
if system_instruction is None and tools is None
|
||||
else {
|
||||
"generateContentRequest": {
|
||||
"model": f"models/{model}",
|
||||
"contents": cleaned_contents,
|
||||
**({"systemInstruction": system_instruction} if system_instruction is not None else {}),
|
||||
**({"tools": tools} if tools is not None else {}),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
async_httpx_client: Final = get_async_httpx_client(
|
||||
llm_provider=LlmProviders.GEMINI,
|
||||
|
|
|
|||
61
litellm/llms/gemini/count_tokens/transformation.py
Normal file
61
litellm/llms/gemini/count_tokens/transformation.py
Normal file
|
|
@ -0,0 +1,61 @@
|
|||
"""Translate an Anthropic /v1/messages/count_tokens request into a Gemini
|
||||
countTokens payload (contents + systemInstruction + tools)."""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, cast
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
|
||||
LiteLLMAnthropicMessagesAdapter,
|
||||
)
|
||||
from litellm.llms.vertex_ai.gemini.transformation import (
|
||||
_gemini_convert_messages_with_history, # pyright: ignore[reportPrivateUsage] # shared helper already used by gemini/chat, context_caching, and vertex_and_google_ai_studio_gemini
|
||||
_transform_system_message,
|
||||
)
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig
|
||||
from litellm.types.llms.anthropic import AnthropicMessagesRequest
|
||||
from litellm.types.llms.vertex_ai import ContentType, SystemInstructions, Tools
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GeminiCountTokensPayload:
|
||||
contents: list[ContentType]
|
||||
system_instruction: SystemInstructions | None
|
||||
tools: list[Tools] | None
|
||||
|
||||
|
||||
def build_count_tokens_payload(
|
||||
model: str,
|
||||
messages: Sequence[Mapping[str, object]],
|
||||
system: object | None,
|
||||
tools: Sequence[Mapping[str, object]] | None,
|
||||
) -> GeminiCountTokensPayload:
|
||||
anthropic_request: Final[AnthropicMessagesRequest] = cast(
|
||||
AnthropicMessagesRequest, # cast-ok: untrusted client payload, adapter reads the anthropic-shape keys only
|
||||
{
|
||||
"model": model,
|
||||
"messages": list(messages),
|
||||
**({"system": system} if system else {}),
|
||||
**({"tools": list(tools)} if tools else {}),
|
||||
},
|
||||
)
|
||||
openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
|
||||
anthropic_request, custom_llm_provider="gemini"
|
||||
)
|
||||
system_instruction, remaining_messages = _transform_system_message(
|
||||
supports_system_message=True,
|
||||
messages=list(openai_request["messages"]),
|
||||
)
|
||||
contents: Final = _gemini_convert_messages_with_history(
|
||||
messages=remaining_messages,
|
||||
model=model,
|
||||
custom_llm_provider="gemini",
|
||||
)
|
||||
openai_tools: Final = openai_request.get("tools")
|
||||
return GeminiCountTokensPayload(
|
||||
contents=contents,
|
||||
system_instruction=system_instruction,
|
||||
tools=VertexGeminiConfig()._map_function(value=[dict(tool) for tool in openai_tools], optional_params={})
|
||||
if openai_tools
|
||||
else None,
|
||||
)
|
||||
93
tests/e2e/llm_translation/test_messages_count_tokens_e2e.py
Normal file
93
tests/e2e/llm_translation/test_messages_count_tokens_e2e.py
Normal file
|
|
@ -0,0 +1,93 @@
|
|||
"""Live e2e: POST /v1/messages/count_tokens against a gemini/ deployment.
|
||||
|
||||
Registers a fresh gemini/gemini-2.5-flash deployment via /model/new (deleted on
|
||||
teardown) and drives the endpoint through the shared transport, since no
|
||||
official provider SDK covers this route.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from models import (
|
||||
AnthropicCustomTool,
|
||||
ChatMessage,
|
||||
CountTokensBody,
|
||||
JsonSchemaProperty,
|
||||
LiteLLMParamsBody,
|
||||
ToolInputSchema,
|
||||
)
|
||||
from proxy_client import ProxyClient
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
BACKEND_MODEL = "gemini/gemini-2.5-flash"
|
||||
GEMINI_API_KEY = "os.environ/GEMINI_API_KEY"
|
||||
|
||||
WEATHER_TOOL = AnthropicCustomTool(
|
||||
name="get_weather",
|
||||
description="Get the current weather for a city.",
|
||||
input_schema=ToolInputSchema(
|
||||
properties={"city": JsonSchemaProperty(type="string")},
|
||||
required=["city"],
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _provision(proxy: ProxyClient, resources: ResourceManager) -> str:
|
||||
model_name = f"e2e-count-tokens-gemini-{unique_marker()}"
|
||||
model_id = proxy.create_model(
|
||||
model_name,
|
||||
LiteLLMParamsBody(model=BACKEND_MODEL, api_key=GEMINI_API_KEY),
|
||||
)
|
||||
resources.defer(lambda: proxy.delete_model(model_id))
|
||||
return model_name
|
||||
|
||||
|
||||
class TestMessagesCountTokens:
|
||||
def test_count_tokens_gemini_returns_input_tokens(
|
||||
self, proxy: ProxyClient, resources: ResourceManager, scoped_key: str
|
||||
) -> None:
|
||||
model = _provision(proxy, resources)
|
||||
|
||||
response = unwrap(
|
||||
proxy.count_tokens(
|
||||
scoped_key,
|
||||
CountTokensBody(
|
||||
model=model,
|
||||
messages=[ChatMessage(role="user", content=f"hello world {unique_marker()}")],
|
||||
),
|
||||
)
|
||||
)
|
||||
assert response.input_tokens > 0, f"input_tokens not positive: {response.input_tokens}"
|
||||
|
||||
def test_count_tokens_gemini_with_system_and_tools(
|
||||
self, proxy: ProxyClient, resources: ResourceManager, scoped_key: str
|
||||
) -> None:
|
||||
model = _provision(proxy, resources)
|
||||
prompt = f"hello world {unique_marker()}"
|
||||
|
||||
plain = unwrap(
|
||||
proxy.count_tokens(
|
||||
scoped_key,
|
||||
CountTokensBody(
|
||||
model=model,
|
||||
messages=[ChatMessage(role="user", content=prompt)],
|
||||
),
|
||||
)
|
||||
)
|
||||
with_system_and_tools = unwrap(
|
||||
proxy.count_tokens(
|
||||
scoped_key,
|
||||
CountTokensBody(
|
||||
model=model,
|
||||
messages=[ChatMessage(role="user", content=prompt)],
|
||||
system="You are a helpful assistant",
|
||||
tools=[WEATHER_TOOL],
|
||||
),
|
||||
)
|
||||
)
|
||||
assert with_system_and_tools.input_tokens > plain.input_tokens, (
|
||||
f"system + tools count {with_system_and_tools.input_tokens} did not exceed "
|
||||
f"the plain message count {plain.input_tokens}; the extra prompt was not counted"
|
||||
)
|
||||
|
|
@ -574,6 +574,8 @@ class CountTokensBody(BaseModel):
|
|||
|
||||
model: str
|
||||
messages: list[ChatMessage]
|
||||
system: str | None = None
|
||||
tools: list[AnthropicTool] | None = None
|
||||
|
||||
|
||||
class AnthropicMessagesResponse(BaseModel):
|
||||
|
|
|
|||
0
tests/test_litellm/llms/gemini/count_tokens/__init__.py
Normal file
0
tests/test_litellm/llms/gemini/count_tokens/__init__.py
Normal file
64
tests/test_litellm/llms/gemini/count_tokens/test_handler.py
Normal file
64
tests/test_litellm/llms/gemini/count_tokens/test_handler.py
Normal file
|
|
@ -0,0 +1,64 @@
|
|||
import json
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
|
||||
|
||||
COUNT_TOKENS_URL = "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash:countTokens"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acount_tokens_sends_generate_content_request_when_system_or_tools_present(monkeypatch):
|
||||
recorded: list[httpx.Request] = []
|
||||
|
||||
def _handler(request: httpx.Request) -> httpx.Response:
|
||||
recorded.append(request)
|
||||
return httpx.Response(200, json={"totalTokens": 42})
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.llms.gemini.count_tokens.handler.get_async_httpx_client",
|
||||
lambda **kwargs: httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
|
||||
)
|
||||
|
||||
result = await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
model="gemini-2.5-flash",
|
||||
contents=[{"role": "user", "parts": [{"text": "hello world"}]}],
|
||||
api_key="test-key",
|
||||
system_instruction={"parts": [{"text": "You are a helpful assistant"}]},
|
||||
tools=[{"function_declarations": [{"name": "get_weather"}]}],
|
||||
)
|
||||
|
||||
assert result == {"totalTokens": 42}
|
||||
request = recorded[-1]
|
||||
assert request.url == COUNT_TOKENS_URL
|
||||
body = json.loads(request.content)
|
||||
assert "contents" not in body
|
||||
generate_content_request = body["generateContentRequest"]
|
||||
assert generate_content_request["model"] == "models/gemini-2.5-flash"
|
||||
assert generate_content_request["contents"] == [{"role": "user", "parts": [{"text": "hello world"}]}]
|
||||
assert generate_content_request["systemInstruction"] == {"parts": [{"text": "You are a helpful assistant"}]}
|
||||
assert generate_content_request["tools"][0]["function_declarations"][0]["name"] == "get_weather"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acount_tokens_keeps_contents_body_without_system_or_tools(monkeypatch):
|
||||
recorded: list[httpx.Request] = []
|
||||
|
||||
def _handler(request: httpx.Request) -> httpx.Response:
|
||||
recorded.append(request)
|
||||
return httpx.Response(200, json={"totalTokens": 4})
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.llms.gemini.count_tokens.handler.get_async_httpx_client",
|
||||
lambda **kwargs: httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
|
||||
)
|
||||
|
||||
await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
model="gemini-2.5-flash",
|
||||
contents=[{"role": "user", "parts": [{"text": "hi"}]}],
|
||||
api_key="test-key",
|
||||
)
|
||||
|
||||
body = json.loads(recorded[-1].content)
|
||||
assert body == {"contents": [{"role": "user", "parts": [{"text": "hi"}]}]}
|
||||
|
|
@ -0,0 +1,62 @@
|
|||
from litellm.llms.gemini.count_tokens.transformation import build_count_tokens_payload
|
||||
|
||||
|
||||
def test_build_count_tokens_payload_translates_anthropic_request():
|
||||
payload = build_count_tokens_payload(
|
||||
model="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "hello world"}],
|
||||
system="You are a helpful assistant",
|
||||
tools=[
|
||||
{
|
||||
"name": "get_weather",
|
||||
"description": "Get the current weather for a city.",
|
||||
"input_schema": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
"required": ["city"],
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
assert payload.contents
|
||||
assert payload.contents[0]["parts"][0].get("text") == "hello world"
|
||||
assert payload.system_instruction is not None
|
||||
assert payload.system_instruction["parts"][0].get("text") == "You are a helpful assistant"
|
||||
assert payload.tools is not None
|
||||
function_declarations = payload.tools[0]["function_declarations"]
|
||||
assert function_declarations[0]["name"] == "get_weather"
|
||||
assert function_declarations[0].get("parameters", {}).get("required") == ["city"]
|
||||
|
||||
|
||||
def test_build_count_tokens_payload_without_system_or_tools():
|
||||
payload = build_count_tokens_payload(
|
||||
model="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
system=None,
|
||||
tools=None,
|
||||
)
|
||||
|
||||
assert payload.contents
|
||||
assert payload.system_instruction is None
|
||||
assert payload.tools is None
|
||||
|
||||
|
||||
def test_build_count_tokens_payload_passes_openai_tools_through():
|
||||
payload = build_count_tokens_payload(
|
||||
model="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
system=None,
|
||||
tools=[
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}},
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
assert payload.tools is not None
|
||||
assert payload.tools[0]["function_declarations"][0]["name"] == "get_weather"
|
||||
|
|
@ -161,6 +161,93 @@ class TestGoogleAIStudioTokenCounter:
|
|||
model=model_to_use, contents=contents
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_translates_anthropic_messages_system_and_tools(self):
|
||||
"""Anthropic-format messages are converted to gemini contents/system/tools
|
||||
before hitting the countTokens endpoint."""
|
||||
token_counter = GoogleAIStudioTokenCounter()
|
||||
|
||||
with patch(
|
||||
"litellm.llms.gemini.count_tokens.handler.GoogleAIStudioTokenCounter.acount_tokens",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_acount_tokens:
|
||||
mock_acount_tokens.return_value = {"totalTokens": 12}
|
||||
|
||||
result = await token_counter.count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "hello world"}],
|
||||
contents=None,
|
||||
deployment=None,
|
||||
request_model="gemini/gemini-2.5-flash",
|
||||
tools=[
|
||||
{
|
||||
"name": "get_weather",
|
||||
"description": "Get the current weather for a city.",
|
||||
"input_schema": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
"required": ["city"],
|
||||
},
|
||||
}
|
||||
],
|
||||
system="You are a helpful assistant",
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.total_tokens == 12
|
||||
kwargs = mock_acount_tokens.call_args.kwargs
|
||||
assert kwargs["contents"]
|
||||
assert kwargs["contents"][0]["parts"][0].get("text") == "hello world"
|
||||
assert kwargs["system_instruction"]["parts"][0].get("text") == "You are a helpful assistant"
|
||||
assert kwargs["tools"][0]["function_declarations"][0]["name"] == "get_weather"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_provider_error_returns_error_response(self):
|
||||
"""A provider APIError must surface as an error TokenCountResponse so the
|
||||
proxy falls back to the local tokenizer instead of 500ing."""
|
||||
import litellm
|
||||
|
||||
token_counter = GoogleAIStudioTokenCounter()
|
||||
|
||||
with patch(
|
||||
"litellm.llms.gemini.count_tokens.handler.GoogleAIStudioTokenCounter.acount_tokens",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_acount_tokens:
|
||||
mock_acount_tokens.side_effect = litellm.APIError(
|
||||
status_code=400,
|
||||
message="Google Gen AI Studio API error: 400",
|
||||
llm_provider="gemini",
|
||||
model="gemini-2.5-flash",
|
||||
)
|
||||
|
||||
result = await token_counter.count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "hello world"}],
|
||||
contents=None,
|
||||
deployment=None,
|
||||
request_model="gemini/gemini-2.5-flash",
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.error is True
|
||||
assert result.status_code == 400
|
||||
assert result.total_tokens == 0
|
||||
assert result.error_message is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_returns_none_without_contents_or_messages(self):
|
||||
token_counter = GoogleAIStudioTokenCounter()
|
||||
|
||||
result = await token_counter.count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=None,
|
||||
contents=None,
|
||||
deployment=None,
|
||||
request_model="gemini/gemini-2.5-flash",
|
||||
)
|
||||
|
||||
assert result is None
|
||||
|
||||
def test_clean_contents_for_gemini_api_removes_id_field(self):
|
||||
"""Test that _clean_contents_for_gemini_api removes unsupported 'id' field from function responses"""
|
||||
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue