mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
3bc391c16e
commit
cdd3d2f930
4 changed files with 112 additions and 73 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue