fix(gemini): count system and tools with native contents, map unexpected provider errors, and inject http client in tests

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-23 14:31:36 +00:00
parent 3bc391c16e
commit cdd3d2f930
4 changed files with 112 additions and 73 deletions

View file

@ -491,6 +491,7 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
request_model: str = "",
tools: list[dict[str, object]] | None = None,
system: object | None = None,
client: httpx.AsyncClient | None = None,
) -> TokenCountResponse | None:
import copy
@ -507,23 +508,26 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
if contents is None
else None
)
system_instruction: Final = payload.system_instruction if payload is not None else system
gemini_tools: Final = payload.tools if payload is not None else tools
count_tokens_params: Final = {
"model": model_to_use,
"contents": payload.contents if payload is not None else contents,
**(
{"system_instruction": payload.system_instruction} # mutable-ok: kwargs dict for acount_tokens
if payload is not None and payload.system_instruction is not None
{"system_instruction": system_instruction} # mutable-ok: kwargs dict for acount_tokens
if system_instruction is not None
else {} # mutable-ok: kwargs dict for acount_tokens
),
**(
{"tools": payload.tools} # mutable-ok: kwargs dict for acount_tokens
if payload is not None and payload.tools is not None
{"tools": gemini_tools} # mutable-ok: kwargs dict for acount_tokens
if gemini_tools is not None
else {} # mutable-ok: kwargs dict for acount_tokens
),
}
count_tokens_params_request.update(count_tokens_params)
try:
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
client=client,
**count_tokens_params_request,
)
except (litellm.APIError, litellm.APIConnectionError) as e:

View file

@ -88,6 +88,7 @@ class GoogleAIStudioTokenCounter:
timeout: float | httpx.Timeout | None = None,
system_instruction: SystemInstructions | None = None,
tools: list[Tools] | None = None,
client: httpx.AsyncClient | None = None,
**kwargs: object,
) -> dict[str, Any]:
"""
@ -159,7 +160,7 @@ class GoogleAIStudioTokenCounter:
}
)
async_httpx_client: Final = get_async_httpx_client(
async_httpx_client: Final = client or get_async_httpx_client(
llm_provider=LlmProviders.GEMINI,
)
@ -185,5 +186,9 @@ class GoogleAIStudioTokenCounter:
error_msg = f"Request to Google Gen AI Studio failed: {e}"
raise litellm.APIConnectionError(message=error_msg, llm_provider="gemini", model=model) from e
except Exception as e:
error_msg = f"Unexpected error during token counting: {e}"
raise Exception(error_msg) from e
raise litellm.APIError(
message=f"Unexpected error during token counting: {e}",
llm_provider="gemini",
model=model,
status_code=500,
) from e

View file

@ -9,17 +9,14 @@ COUNT_TOKENS_URL = "https://generativelanguage.googleapis.com/v1beta/models/gemi
@pytest.mark.asyncio
async def test_acount_tokens_sends_generate_content_request_when_system_or_tools_present(monkeypatch):
async def test_acount_tokens_sends_generate_content_request_when_system_or_tools_present():
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)),
)
client = httpx.AsyncClient(transport=httpx.MockTransport(_handler))
result = await GoogleAIStudioTokenCounter().acount_tokens(
model="gemini-2.5-flash",
@ -27,6 +24,7 @@ async def test_acount_tokens_sends_generate_content_request_when_system_or_tools
api_key="test-key",
system_instruction={"parts": [{"text": "You are a helpful assistant"}]},
tools=[{"function_declarations": [{"name": "get_weather"}]}],
client=client,
)
assert result == {"totalTokens": 42}
@ -42,22 +40,20 @@ async def test_acount_tokens_sends_generate_content_request_when_system_or_tools
@pytest.mark.asyncio
async def test_acount_tokens_keeps_contents_body_without_system_or_tools(monkeypatch):
async def test_acount_tokens_keeps_contents_body_without_system_or_tools():
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)),
)
client = 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",
client=client,
)
body = json.loads(recorded[-1].content)

View file

@ -1,3 +1,4 @@
import json
from unittest.mock import AsyncMock, patch
import pytest
@ -158,81 +159,114 @@ class TestGoogleAIStudioTokenCounter:
# Verify the mock was called correctly
mock_acount_tokens.assert_called_once_with(
model=model_to_use, contents=contents
model=model_to_use, contents=contents, client=None
)
@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."""
import httpx
recorded: list = []
def _handler(request):
recorded.append(request)
return httpx.Response(200, json={"totalTokens": 12})
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={"litellm_params": {"api_key": "test-key", "api_base": "https://gemini.example.test"}},
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",
client=httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
)
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
body = json.loads(recorded[-1].content)
generate_content_request = body["generateContentRequest"]
assert generate_content_request["contents"]
assert generate_content_request["contents"][0]["parts"][0].get("text") == "hello world"
assert generate_content_request["systemInstruction"]["parts"][0].get("text") == "You are a helpful assistant"
assert generate_content_request["tools"][0]["function_declarations"][0]["name"] == "get_weather"
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_passes_system_and_tools_with_native_contents(self):
"""A request that already carries gemini contents still counts the
caller-supplied system instruction and tools."""
import httpx
recorded: list = []
def _handler(request):
recorded.append(request)
return httpx.Response(200, json={"totalTokens": 20})
token_counter = GoogleAIStudioTokenCounter()
result = await token_counter.count_tokens(
model_to_use="gemini-2.5-flash",
messages=None,
contents=[{"role": "user", "parts": [{"text": "hello world"}]}],
deployment={"litellm_params": {"api_key": "test-key", "api_base": "https://gemini.example.test"}},
request_model="gemini/gemini-2.5-flash",
system={"parts": [{"text": "You are a helpful assistant"}]},
tools=[{"function_declarations": [{"name": "get_weather"}]}],
client=httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
)
assert result is not None
assert result.total_tokens == 20
body = json.loads(recorded[-1].content)
generate_content_request = body["generateContentRequest"]
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_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
import httpx
def _handler(request):
return httpx.Response(
400,
json={"error": {"code": 400, "message": "bad request", "status": "INVALID_ARGUMENT"}},
)
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={"litellm_params": {"api_key": "test-key", "api_base": "https://gemini.example.test"}},
request_model="gemini/gemini-2.5-flash",
client=httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
)
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
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):