mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(gemini): typed countTokens request, 400 on malformed input, 502 on missing totalTokens, proxy fallback on provider exceptions
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
0a25c92f2b
commit
e3713c767f
9 changed files with 466 additions and 215 deletions
|
|
@ -3,7 +3,7 @@ import datetime
|
|||
import json
|
||||
import math
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final
|
||||
from typing import Any, Final, cast
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -497,68 +497,83 @@ class GoogleAIStudioTokenCounter(BaseTokenCounter):
|
|||
|
||||
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
|
||||
from litellm.llms.gemini.count_tokens.transformation import (
|
||||
InvalidCountTokensRequest,
|
||||
build_count_tokens_payload,
|
||||
normalize_count_tokens_tools,
|
||||
)
|
||||
from litellm.types.llms.vertex_ai import SystemInstructions
|
||||
|
||||
if contents is None and not messages:
|
||||
return None
|
||||
|
||||
deployment = deployment or {}
|
||||
count_tokens_params_request: Final = copy.deepcopy(deployment.get("litellm_params", {}))
|
||||
try:
|
||||
payload: Final = (
|
||||
build_count_tokens_payload(model=model_to_use, messages=messages, system=system, tools=tools)
|
||||
if contents is None
|
||||
else None
|
||||
)
|
||||
system_instruction: Final = (
|
||||
payload.system_instruction
|
||||
if payload is not None
|
||||
else (
|
||||
{"parts": [{"text": system}]} # mutable-ok: SystemInstructions wire shape
|
||||
if isinstance(system, str)
|
||||
else system
|
||||
)
|
||||
)
|
||||
gemini_tools: Final = payload.tools if payload is not None else normalize_count_tokens_tools(tools)
|
||||
count_tokens_params: Final = { # mutable-ok: kwargs dict for acount_tokens
|
||||
"model": model_to_use,
|
||||
"contents": payload.contents if payload is not None else contents,
|
||||
**(
|
||||
{"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": 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)
|
||||
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
client=client,
|
||||
**count_tokens_params_request,
|
||||
)
|
||||
if result is not None:
|
||||
return TokenCountResponse(
|
||||
total_tokens=result.get("totalTokens", 0),
|
||||
request_model=request_model,
|
||||
model_used=model_to_use,
|
||||
tokenizer_type="gemini_api",
|
||||
original_response=result,
|
||||
)
|
||||
return None
|
||||
except Exception as e:
|
||||
# provider counting is best-effort: translation, credential, and request
|
||||
# failures all degrade to the proxy's local-tokenizer fallback
|
||||
payload: Final = (
|
||||
build_count_tokens_payload(model=model_to_use, messages=messages, system=system, tools=tools)
|
||||
if contents is None
|
||||
else None
|
||||
)
|
||||
if isinstance(payload, InvalidCountTokensRequest):
|
||||
return TokenCountResponse(
|
||||
total_tokens=0,
|
||||
request_model=request_model,
|
||||
model_used=model_to_use,
|
||||
tokenizer_type="gemini_api",
|
||||
error=True,
|
||||
error_message=getattr(e, "message", None) or str(e),
|
||||
status_code=getattr(e, "status_code", None) or 500,
|
||||
error_message=payload.message,
|
||||
status_code=400,
|
||||
)
|
||||
system_instruction: Final[SystemInstructions | None] = (
|
||||
payload.system_instruction
|
||||
if payload is not None
|
||||
else (
|
||||
{"parts": [{"text": system}]} # mutable-ok: SystemInstructions wire shape
|
||||
if isinstance(system, str)
|
||||
else cast( # cast-ok: contents-path callers pass a Gemini-shaped systemInstruction
|
||||
"SystemInstructions | None",
|
||||
system,
|
||||
)
|
||||
)
|
||||
)
|
||||
gemini_tools: Final = payload.tools if payload is not None else normalize_count_tokens_tools(tools)
|
||||
count_tokens_params_request.update(
|
||||
{ # mutable-ok: kwargs dict for acount_tokens
|
||||
"model": model_to_use,
|
||||
"contents": payload.contents if payload is not None else contents,
|
||||
}
|
||||
)
|
||||
try:
|
||||
result: Final = await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
system_instruction=system_instruction,
|
||||
tools=gemini_tools,
|
||||
client=client,
|
||||
**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 "totalTokens" not in result:
|
||||
return TokenCountResponse(
|
||||
total_tokens=0,
|
||||
request_model=request_model,
|
||||
model_used=model_to_use,
|
||||
tokenizer_type="gemini_api",
|
||||
error=True,
|
||||
error_message="Google Gen AI Studio countTokens response has no totalTokens",
|
||||
status_code=502,
|
||||
original_response=result,
|
||||
)
|
||||
return TokenCountResponse(
|
||||
total_tokens=result["totalTokens"],
|
||||
request_model=request_model,
|
||||
model_used=model_to_use,
|
||||
tokenizer_type="gemini_api",
|
||||
original_response=result,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
|
|
@ -5,7 +6,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.llms.gemini import GeminiCountTokensRequest
|
||||
from litellm.types.llms.vertex_ai import ContentType, SystemInstructions, Tools
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -14,6 +16,41 @@ else:
|
|||
GenerateContentContentListUnionDict = Any
|
||||
|
||||
|
||||
def build_count_tokens_request(
|
||||
model: str,
|
||||
contents: Sequence[ContentType],
|
||||
system_instruction: SystemInstructions | None,
|
||||
tools: Sequence[Tools] | None,
|
||||
) -> GeminiCountTokensRequest:
|
||||
model_name: Final = f"models/{model}"
|
||||
if tools is None:
|
||||
if system_instruction is None:
|
||||
bare: Final[GeminiCountTokensRequest] = {"contents": contents}
|
||||
return bare
|
||||
with_system: Final[GeminiCountTokensRequest] = {
|
||||
"generateContentRequest": {
|
||||
"model": model_name,
|
||||
"contents": contents,
|
||||
"systemInstruction": system_instruction,
|
||||
}
|
||||
}
|
||||
return with_system
|
||||
if system_instruction is None:
|
||||
with_tools: Final[GeminiCountTokensRequest] = {
|
||||
"generateContentRequest": {"model": model_name, "contents": contents, "tools": tools}
|
||||
}
|
||||
return with_tools
|
||||
with_both: Final[GeminiCountTokensRequest] = {
|
||||
"generateContentRequest": {
|
||||
"model": model_name,
|
||||
"contents": contents,
|
||||
"systemInstruction": system_instruction,
|
||||
"tools": tools,
|
||||
}
|
||||
}
|
||||
return with_both
|
||||
|
||||
|
||||
class GoogleAIStudioTokenCounter:
|
||||
def _clean_contents_for_gemini_api(self, contents: Any) -> Any:
|
||||
"""
|
||||
|
|
@ -87,7 +124,7 @@ class GoogleAIStudioTokenCounter:
|
|||
api_base: str | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
system_instruction: SystemInstructions | None = None,
|
||||
tools: list[Tools] | None = None,
|
||||
tools: Sequence[Tools] | None = None,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
**kwargs: object,
|
||||
) -> dict[str, Any]:
|
||||
|
|
@ -133,45 +170,40 @@ class GoogleAIStudioTokenCounter:
|
|||
litellm_params=kwargs,
|
||||
)
|
||||
|
||||
cleaned_contents: Final = self._clean_contents_for_gemini_api(contents)
|
||||
request_body: Final = (
|
||||
{"contents": cleaned_contents} # mutable-ok: httpx json body takes a plain dict
|
||||
if system_instruction is None and tools is None
|
||||
else { # mutable-ok: httpx json body takes a plain dict
|
||||
"generateContentRequest": { # mutable-ok: httpx json body takes a plain dict
|
||||
"model": f"models/{model}",
|
||||
"contents": cleaned_contents,
|
||||
**(
|
||||
{ # mutable-ok: httpx json body takes a plain dict
|
||||
"systemInstruction": system_instruction,
|
||||
}
|
||||
if system_instruction is not None
|
||||
else {} # mutable-ok: httpx json body takes a plain dict
|
||||
),
|
||||
**(
|
||||
{ # mutable-ok: httpx json body takes a plain dict
|
||||
"tools": tools,
|
||||
}
|
||||
if tools is not None
|
||||
else {} # mutable-ok: httpx json body takes a plain dict
|
||||
),
|
||||
}
|
||||
}
|
||||
request_body: Final = build_count_tokens_request(
|
||||
model=model,
|
||||
contents=self._clean_contents_for_gemini_api(contents),
|
||||
system_instruction=system_instruction,
|
||||
tools=tools,
|
||||
)
|
||||
|
||||
async_httpx_client: Final = client or get_async_httpx_client(
|
||||
llm_provider=LlmProviders.GEMINI,
|
||||
)
|
||||
|
||||
response: Final = await async_httpx_client.post(url=url, headers=headers, json=request_body)
|
||||
response: Final = await async_httpx_client.post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
json=request_body, # pyright: ignore[reportArgumentType] # post() takes a bare dict; a TypedDict is one at runtime
|
||||
)
|
||||
|
||||
# Check for HTTP errors
|
||||
response.raise_for_status()
|
||||
|
||||
# Parse response
|
||||
result: Final = response.json()
|
||||
try:
|
||||
result: Final = response.json()
|
||||
except ValueError as e:
|
||||
raise litellm.APIError(
|
||||
message=f"Google Gen AI Studio API returned a non-JSON body: {response.text}",
|
||||
llm_provider="gemini",
|
||||
model=model,
|
||||
status_code=response.status_code,
|
||||
) from e
|
||||
return result
|
||||
|
||||
except litellm.APIError:
|
||||
raise
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_msg = f"Google Gen AI Studio API error: {e.response.status_code} - {e.response.text}"
|
||||
raise litellm.APIError(
|
||||
|
|
|
|||
|
|
@ -13,6 +13,8 @@ from collections.abc import Mapping, Sequence
|
|||
from dataclasses import dataclass
|
||||
from typing import Final, cast
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
|
||||
LiteLLMAnthropicMessagesAdapter,
|
||||
|
|
@ -34,6 +36,11 @@ class GeminiCountTokensPayload:
|
|||
tools: list[Tools] | None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class InvalidCountTokensRequest:
|
||||
message: str
|
||||
|
||||
|
||||
_ANTHROPIC_PART_TYPES: Final = frozenset(
|
||||
{
|
||||
"tool_use",
|
||||
|
|
@ -345,28 +352,36 @@ def _payload_from_openai_parts(
|
|||
)
|
||||
|
||||
|
||||
_ANTHROPIC_REQUEST: Final = TypeAdapter(AnthropicMessagesRequest)
|
||||
|
||||
|
||||
def _build_anthropic_payload(
|
||||
model: str,
|
||||
messages: Sequence[Mapping[str, object]],
|
||||
system: object | None,
|
||||
tools: Sequence[Mapping[str, object]] | None,
|
||||
) -> GeminiCountTokensPayload:
|
||||
) -> GeminiCountTokensPayload | InvalidCountTokensRequest:
|
||||
hosted_tools: Final = tuple(tool for tool in tools or () if _hosted_tool_type(tool) is not None)
|
||||
adapter_tools: Final = tuple(tool for tool in tools or () if _hosted_tool_type(tool) is None)
|
||||
anthropic_request: Final[AnthropicMessagesRequest] = cast( # cast-ok: adapter reads only the keys supplied
|
||||
raw_request: Final = { # mutable-ok: transient request dict for the anthropic adapter
|
||||
"model": model,
|
||||
"messages": list( # mutable-ok: adapter contract takes a list of messages
|
||||
_textify_server_side_blocks(messages)
|
||||
),
|
||||
**({"system": system} if system else {}), # mutable-ok: transient request dict for the anthropic adapter
|
||||
**(
|
||||
{"tools": list(adapter_tools)}
|
||||
if adapter_tools
|
||||
else {} # mutable-ok: transient request dict for the anthropic adapter
|
||||
),
|
||||
}
|
||||
try:
|
||||
_ANTHROPIC_REQUEST.validate_python(raw_request)
|
||||
except ValidationError as e:
|
||||
return InvalidCountTokensRequest(message=str(e))
|
||||
anthropic_request: Final = cast( # cast-ok: validated above; pydantic returns lazy Iterable validators
|
||||
AnthropicMessagesRequest,
|
||||
{ # mutable-ok: transient request dict for the anthropic adapter
|
||||
"model": model,
|
||||
"messages": list( # mutable-ok: adapter contract takes a list of messages
|
||||
_textify_server_side_blocks(messages)
|
||||
),
|
||||
**({"system": system} if system else {}), # mutable-ok: transient request dict for the anthropic adapter
|
||||
**(
|
||||
{"tools": list(adapter_tools)}
|
||||
if adapter_tools
|
||||
else {} # mutable-ok: transient request dict for the anthropic adapter
|
||||
),
|
||||
},
|
||||
raw_request,
|
||||
)
|
||||
openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
|
||||
anthropic_request, custom_llm_provider="gemini"
|
||||
|
|
@ -422,10 +437,13 @@ def build_count_tokens_payload(
|
|||
messages: Sequence[Mapping[str, object]],
|
||||
system: object | None,
|
||||
tools: Sequence[Mapping[str, object]] | None,
|
||||
) -> GeminiCountTokensPayload:
|
||||
if _has_anthropic_shape(system=system, tools=tools, messages=messages):
|
||||
return _build_anthropic_payload(model=model, messages=messages, system=system, tools=tools)
|
||||
return _build_openai_payload(model=model, messages=messages, system=system, tools=tools)
|
||||
) -> GeminiCountTokensPayload | InvalidCountTokensRequest:
|
||||
try:
|
||||
if _has_anthropic_shape(system=system, tools=tools, messages=messages):
|
||||
return _build_anthropic_payload(model=model, messages=messages, system=system, tools=tools)
|
||||
return _build_openai_payload(model=model, messages=messages, system=system, tools=tools)
|
||||
except (KeyError, TypeError, ValueError) as e:
|
||||
return InvalidCountTokensRequest(message=f"Invalid token count request: {e!r}")
|
||||
|
||||
|
||||
# Matches real inlineData blobs; a short or non-base64 `data` field (tool args,
|
||||
|
|
|
|||
|
|
@ -13436,6 +13436,20 @@ async def _try_provider_token_count(
|
|||
param="model",
|
||||
code=status_code,
|
||||
)
|
||||
except (litellm.APIError, litellm.APIConnectionError) as e:
|
||||
if litellm.disable_token_counter is True:
|
||||
raise ProxyException(
|
||||
message=e.message,
|
||||
type="token_counting_error",
|
||||
param="model",
|
||||
code=e.status_code,
|
||||
)
|
||||
verbose_proxy_logger.warning(
|
||||
"Provider token counting raised (%s): %s. Falling back to local tokenizer.",
|
||||
e.status_code,
|
||||
e.message,
|
||||
)
|
||||
return None
|
||||
if result is not None and result.error is True:
|
||||
if litellm.disable_token_counter is True:
|
||||
raise ProxyException(
|
||||
|
|
|
|||
|
|
@ -1,17 +1,34 @@
|
|||
from collections.abc import Sequence
|
||||
from enum import Enum
|
||||
from typing import Any, Literal
|
||||
|
||||
from typing_extensions import Required, TypedDict
|
||||
from typing_extensions import NotRequired, ReadOnly, Required, TypedDict
|
||||
|
||||
from .vertex_ai import (
|
||||
ContentType,
|
||||
GenerationConfig,
|
||||
HttpxBlobType,
|
||||
HttpxContentType,
|
||||
SystemInstructions,
|
||||
Tools,
|
||||
UsageMetadata,
|
||||
)
|
||||
|
||||
|
||||
class GeminiGenerateContentRequest(TypedDict):
|
||||
model: ReadOnly[str]
|
||||
contents: ReadOnly[Sequence[ContentType]]
|
||||
systemInstruction: ReadOnly[NotRequired[SystemInstructions]]
|
||||
tools: ReadOnly[NotRequired[Sequence[Tools]]]
|
||||
|
||||
|
||||
class GeminiCountTokensRequest(TypedDict, total=False):
|
||||
"""Body of models/{model}:countTokens: bare contents, or a generateContentRequest when system or tools are set."""
|
||||
|
||||
contents: ReadOnly[Sequence[ContentType]]
|
||||
generateContentRequest: ReadOnly[GeminiGenerateContentRequest]
|
||||
|
||||
|
||||
class GeminiFilesState(Enum):
|
||||
STATE_UNSPECIFIED = "STATE_UNSPECIFIED"
|
||||
PROCESSING = "PROCESSING"
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ from fastapi import HTTPException, Request
|
|||
import litellm
|
||||
from litellm import Router
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.llms.base_llm.base_utils import BaseTokenCounter
|
||||
from litellm.llms.bedrock.common_utils import BedrockError
|
||||
from litellm.llms.bedrock.count_tokens.bedrock_token_counter import BedrockTokenCounter
|
||||
from litellm.llms.bedrock.count_tokens.handler import BedrockCountTokensHandler
|
||||
|
|
@ -29,7 +30,7 @@ from litellm.proxy._types import ProxyException, TokenCountRequest
|
|||
from litellm.proxy.anthropic_endpoints.endpoints import (
|
||||
count_tokens as anthropic_count_tokens,
|
||||
)
|
||||
from litellm.proxy.proxy_server import token_counter
|
||||
from litellm.proxy.proxy_server import _try_provider_token_count, token_counter
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
verbose_proxy_logger.setLevel(level=logging.DEBUG)
|
||||
|
|
@ -140,9 +141,7 @@ async def test_vLLM_token_counting():
|
|||
|
||||
print("response: ", response)
|
||||
|
||||
assert (
|
||||
response.tokenizer_type == "openai_tokenizer"
|
||||
) # SHOULD use the default tokenizer
|
||||
assert response.tokenizer_type == "openai_tokenizer" # SHOULD use the default tokenizer
|
||||
assert response.model_used == "wolfram/miquliz-120b-v2.0"
|
||||
|
||||
|
||||
|
|
@ -175,9 +174,7 @@ async def test_token_counting_model_not_in_model_list():
|
|||
|
||||
print("response: ", response)
|
||||
|
||||
assert (
|
||||
response.tokenizer_type == "openai_tokenizer"
|
||||
) # SHOULD use the OpenAI tokenizer
|
||||
assert response.tokenizer_type == "openai_tokenizer" # SHOULD use the OpenAI tokenizer
|
||||
assert response.model_used == "special-alias"
|
||||
|
||||
|
||||
|
|
@ -210,9 +207,7 @@ async def test_gpt_token_counting():
|
|||
|
||||
print("response: ", response)
|
||||
|
||||
assert (
|
||||
response.tokenizer_type == "openai_tokenizer"
|
||||
) # SHOULD use the OpenAI tokenizer
|
||||
assert response.tokenizer_type == "openai_tokenizer" # SHOULD use the OpenAI tokenizer
|
||||
assert response.request_model == "gpt-4"
|
||||
|
||||
|
||||
|
|
@ -249,9 +244,7 @@ async def test_anthropic_messages_count_tokens_endpoint():
|
|||
|
||||
# Mock the internal token_counter function to return a controlled response
|
||||
async def mock_token_counter(request, call_endpoint=False):
|
||||
assert (
|
||||
call_endpoint == True
|
||||
), "Should be called with call_endpoint=True for Anthropic endpoint"
|
||||
assert call_endpoint == True, "Should be called with call_endpoint=True for Anthropic endpoint"
|
||||
assert request.model == "claude-3-sonnet-20240229"
|
||||
assert request.messages == [{"role": "user", "content": "Hello Claude!"}]
|
||||
|
||||
|
|
@ -321,9 +314,7 @@ async def test_anthropic_messages_count_tokens_with_non_anthropic_model():
|
|||
|
||||
# Mock the internal token_counter function to return a controlled response
|
||||
async def mock_token_counter(request, call_endpoint=True):
|
||||
assert (
|
||||
call_endpoint == True
|
||||
), "Should be called with call_endpoint=True for Anthropic endpoint"
|
||||
assert call_endpoint == True, "Should be called with call_endpoint=True for Anthropic endpoint"
|
||||
assert request.model == "gpt-4"
|
||||
assert request.messages == [{"role": "user", "content": "Hello GPT!"}]
|
||||
|
||||
|
|
@ -480,9 +471,7 @@ async def test_factory_anthropic_endpoint_calls_anthropic_counter():
|
|||
|
||||
# Mock the global handler instance in token_counter module
|
||||
mock_handler = MagicMock()
|
||||
mock_handler.handle_count_tokens_request = AsyncMock(
|
||||
return_value={"input_tokens": 42}
|
||||
)
|
||||
mock_handler.handle_count_tokens_request = AsyncMock(return_value={"input_tokens": 42})
|
||||
|
||||
with patch(
|
||||
"litellm.llms.anthropic.count_tokens.token_counter.anthropic_count_tokens_handler",
|
||||
|
|
@ -537,9 +526,7 @@ async def test_factory_gpt4_endpoint_does_not_call_anthropic_counter():
|
|||
|
||||
# Mock the global handler instance in token_counter module
|
||||
mock_handler = MagicMock()
|
||||
mock_handler.handle_count_tokens_request = AsyncMock(
|
||||
return_value={"input_tokens": 42}
|
||||
)
|
||||
mock_handler.handle_count_tokens_request = AsyncMock(return_value={"input_tokens": 42})
|
||||
|
||||
with patch(
|
||||
"litellm.llms.anthropic.count_tokens.token_counter.anthropic_count_tokens_handler",
|
||||
|
|
@ -596,9 +583,7 @@ async def test_factory_normal_token_counter_endpoint_does_not_call_anthropic():
|
|||
|
||||
# Mock the global handler instance in token_counter module
|
||||
mock_handler = MagicMock()
|
||||
mock_handler.handle_count_tokens_request = AsyncMock(
|
||||
return_value={"input_tokens": 42}
|
||||
)
|
||||
mock_handler.handle_count_tokens_request = AsyncMock(return_value={"input_tokens": 42})
|
||||
|
||||
with patch(
|
||||
"litellm.llms.anthropic.count_tokens.token_counter.anthropic_count_tokens_handler",
|
||||
|
|
@ -613,9 +598,7 @@ async def test_factory_normal_token_counter_endpoint_does_not_call_anthropic():
|
|||
mock_router.model_list = [
|
||||
{
|
||||
"model_name": "claude-3-5-sonnet",
|
||||
"litellm_params": {
|
||||
"model": "anthropic/claude-3-5-sonnet-20241022"
|
||||
},
|
||||
"litellm_params": {"model": "anthropic/claude-3-5-sonnet-20241022"},
|
||||
"model_info": {},
|
||||
}
|
||||
]
|
||||
|
|
@ -624,9 +607,7 @@ async def test_factory_normal_token_counter_endpoint_does_not_call_anthropic():
|
|||
mock_router.async_get_available_deployment = AsyncMock(
|
||||
return_value={
|
||||
"model_name": "claude-3-5-sonnet",
|
||||
"litellm_params": {
|
||||
"model": "anthropic/claude-3-5-sonnet-20241022"
|
||||
},
|
||||
"litellm_params": {"model": "anthropic/claude-3-5-sonnet-20241022"},
|
||||
"model_info": {},
|
||||
}
|
||||
)
|
||||
|
|
@ -661,9 +642,7 @@ async def test_factory_registration():
|
|||
assert counter is not None
|
||||
|
||||
# Create test deployments
|
||||
anthropic_deployment = {
|
||||
"litellm_params": {"model": "anthropic/claude-3-5-sonnet-20241022"}
|
||||
}
|
||||
anthropic_deployment = {"litellm_params": {"model": "anthropic/claude-3-5-sonnet-20241022"}}
|
||||
|
||||
non_anthropic_deployment = {"litellm_params": {"model": "openai/gpt-4"}}
|
||||
|
||||
|
|
@ -678,9 +657,7 @@ async def test_factory_registration():
|
|||
assert not counter.should_use_token_counting_api(custom_llm_provider=None)
|
||||
|
||||
|
||||
@pytest.mark.skip(
|
||||
reason="Requires Google/Vertex AI credentials (GEMINI_API_KEY or VERTEX_AI_PRIVATE_KEY)."
|
||||
)
|
||||
@pytest.mark.skip(reason="Requires Google/Vertex AI credentials (GEMINI_API_KEY or VERTEX_AI_PRIVATE_KEY).")
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model_name", ["gemini-2.5-pro", "vertex-ai-gemini-2.5-pro"])
|
||||
async def test_vertex_ai_gemini_token_counting_with_contents(model_name):
|
||||
|
|
@ -711,9 +688,7 @@ async def test_vertex_ai_gemini_token_counting_with_contents(model_name):
|
|||
response = await token_counter(
|
||||
request=TokenCountRequest(
|
||||
model=model_name,
|
||||
contents=[
|
||||
{"parts": [{"text": "Hello world, how are you doing today? i am ij"}]}
|
||||
],
|
||||
contents=[{"parts": [{"text": "Hello world, how are you doing today? i am ij"}]}],
|
||||
),
|
||||
call_endpoint=True,
|
||||
)
|
||||
|
|
@ -750,9 +725,7 @@ async def test_bedrock_count_tokens_endpoint():
|
|||
model_list=[
|
||||
{
|
||||
"model_name": "claude-bedrock",
|
||||
"litellm_params": {
|
||||
"model": "bedrock/anthropic.claude-3-sonnet-20240229-v1:0"
|
||||
},
|
||||
"litellm_params": {"model": "bedrock/anthropic.claude-3-sonnet-20240229-v1:0"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
|
@ -766,9 +739,7 @@ async def test_bedrock_count_tokens_endpoint():
|
|||
}
|
||||
|
||||
# Test the mock handler directly to verify correct parameter extraction
|
||||
await mock_count_tokens_handler(
|
||||
request_data, {}, "anthropic.claude-3-sonnet-20240229-v1:0"
|
||||
)
|
||||
await mock_count_tokens_handler(request_data, {}, "anthropic.claude-3-sonnet-20240229-v1:0")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -830,10 +801,7 @@ async def test_vertex_ai_anthropic_token_counting():
|
|||
assert call_args is not None
|
||||
assert call_args.kwargs["model"] == "claude-3-5-sonnet-20241022"
|
||||
assert "messages" in call_args.kwargs["request_data"]
|
||||
assert (
|
||||
call_args.kwargs["request_data"]["messages"][0]["content"]
|
||||
== "Hello Claude on Vertex AI! How are you?"
|
||||
)
|
||||
assert call_args.kwargs["request_data"]["messages"][0]["content"] == "Hello Claude on Vertex AI! How are you?"
|
||||
|
||||
# Validate response structure
|
||||
assert response.model_used == "claude-3-5-sonnet-20241022"
|
||||
|
|
@ -866,9 +834,7 @@ def test_vertex_ai_partner_models_token_counting_endpoint(vertex_location):
|
|||
if vertex_location == "global":
|
||||
assert endpoint.startswith("https://aiplatform.googleapis.com")
|
||||
else:
|
||||
assert endpoint.startswith(
|
||||
f"https://{vertex_location}-aiplatform.googleapis.com"
|
||||
)
|
||||
assert endpoint.startswith(f"https://{vertex_location}-aiplatform.googleapis.com")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -880,13 +846,9 @@ async def test_bedrock_token_counter_error_propagation_bedrock_error():
|
|||
counter = BedrockTokenCounter()
|
||||
|
||||
# Mock the handler to raise BedrockError with specific status code
|
||||
with patch.object(
|
||||
counter, "count_tokens", wraps=counter.count_tokens
|
||||
) as mock_count:
|
||||
with patch.object(counter, "count_tokens", wraps=counter.count_tokens) as mock_count:
|
||||
# We need to patch at the handler level
|
||||
with patch(
|
||||
"litellm.llms.bedrock.count_tokens.bedrock_token_counter.BedrockCountTokensHandler"
|
||||
) as MockHandler:
|
||||
with patch("litellm.llms.bedrock.count_tokens.bedrock_token_counter.BedrockCountTokensHandler") as MockHandler:
|
||||
mock_handler_instance = MockHandler.return_value
|
||||
mock_handler_instance.handle_count_tokens_request = AsyncMock(
|
||||
side_effect=BedrockError(status_code=429, message="Rate limit exceeded")
|
||||
|
|
@ -915,13 +877,9 @@ async def test_bedrock_token_counter_error_propagation_generic_exception():
|
|||
"""
|
||||
counter = BedrockTokenCounter()
|
||||
|
||||
with patch(
|
||||
"litellm.llms.bedrock.count_tokens.bedrock_token_counter.BedrockCountTokensHandler"
|
||||
) as MockHandler:
|
||||
with patch("litellm.llms.bedrock.count_tokens.bedrock_token_counter.BedrockCountTokensHandler") as MockHandler:
|
||||
mock_handler_instance = MockHandler.return_value
|
||||
mock_handler_instance.handle_count_tokens_request = AsyncMock(
|
||||
side_effect=Exception("Unexpected error")
|
||||
)
|
||||
mock_handler_instance.handle_count_tokens_request = AsyncMock(side_effect=Exception("Unexpected error"))
|
||||
|
||||
result = await counter.count_tokens(
|
||||
model_to_use="anthropic.claude-3-sonnet",
|
||||
|
|
@ -958,20 +916,14 @@ async def test_bedrock_handler_httpx_error_status_code_propagation():
|
|||
|
||||
with patch.object(handler, "validate_count_tokens_request"):
|
||||
with patch.object(handler, "_get_aws_region_name", return_value="us-west-2"):
|
||||
with patch.object(
|
||||
handler, "transform_anthropic_to_bedrock_count_tokens", return_value={}
|
||||
):
|
||||
with patch.object(handler, "transform_anthropic_to_bedrock_count_tokens", return_value={}):
|
||||
with patch.object(
|
||||
handler,
|
||||
"get_bedrock_count_tokens_endpoint",
|
||||
return_value="https://example.com",
|
||||
):
|
||||
with patch.object(
|
||||
handler, "_sign_request", return_value=({}, "{}")
|
||||
):
|
||||
with patch(
|
||||
"litellm.llms.bedrock.count_tokens.handler.get_async_httpx_client"
|
||||
) as mock_client:
|
||||
with patch.object(handler, "_sign_request", return_value=({}, "{}")):
|
||||
with patch("litellm.llms.bedrock.count_tokens.handler.get_async_httpx_client") as mock_client:
|
||||
mock_async_client = AsyncMock()
|
||||
mock_async_client.post = AsyncMock(side_effect=http_error)
|
||||
mock_client.return_value = mock_async_client
|
||||
|
|
@ -980,9 +932,7 @@ async def test_bedrock_handler_httpx_error_status_code_propagation():
|
|||
await handler.handle_count_tokens_request(
|
||||
request_data={
|
||||
"model": "test",
|
||||
"messages": [
|
||||
{"role": "user", "content": "hello"}
|
||||
],
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
},
|
||||
litellm_params={},
|
||||
resolved_model="anthropic.claude-3-sonnet",
|
||||
|
|
@ -990,10 +940,7 @@ async def test_bedrock_handler_httpx_error_status_code_propagation():
|
|||
|
||||
assert exc_info.value.status_code == 403
|
||||
# Message should be the raw response text
|
||||
assert (
|
||||
exc_info.value.message
|
||||
== "Forbidden - Invalid credentials"
|
||||
)
|
||||
assert exc_info.value.message == "Forbidden - Invalid credentials"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1021,9 +968,7 @@ async def test_token_counter_httpx_status_error_raises_proxy_exception():
|
|||
mock_counter.count_tokens = AsyncMock(side_effect=http_error)
|
||||
|
||||
# Save originals
|
||||
original_get_provider_token_counter = (
|
||||
litellm.proxy.proxy_server._get_provider_token_counter
|
||||
)
|
||||
original_get_provider_token_counter = litellm.proxy.proxy_server._get_provider_token_counter
|
||||
original_router = litellm.proxy.proxy_server.llm_router
|
||||
|
||||
try:
|
||||
|
|
@ -1031,9 +976,7 @@ async def test_token_counter_httpx_status_error_raises_proxy_exception():
|
|||
def mock_get_provider_token_counter(deployment, model_to_use):
|
||||
return (mock_counter, "claude-4-6-sonnet", "vertex_ai")
|
||||
|
||||
litellm.proxy.proxy_server._get_provider_token_counter = (
|
||||
mock_get_provider_token_counter
|
||||
)
|
||||
litellm.proxy.proxy_server._get_provider_token_counter = mock_get_provider_token_counter
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.async_get_available_deployment = AsyncMock(
|
||||
|
|
@ -1061,12 +1004,74 @@ async def test_token_counter_httpx_status_error_raises_proxy_exception():
|
|||
assert exc_info.value.type == "token_counting_error"
|
||||
assert exc_info.value.param == "model"
|
||||
finally:
|
||||
litellm.proxy.proxy_server._get_provider_token_counter = (
|
||||
original_get_provider_token_counter
|
||||
)
|
||||
litellm.proxy.proxy_server._get_provider_token_counter = original_get_provider_token_counter
|
||||
litellm.proxy.proxy_server.llm_router = original_router
|
||||
|
||||
|
||||
class _RaisingCounter(BaseTokenCounter):
|
||||
def __init__(self, error: Exception) -> None:
|
||||
self._error = error
|
||||
|
||||
def should_use_token_counting_api(self, custom_llm_provider: str | None = None) -> bool:
|
||||
return True
|
||||
|
||||
async def count_tokens(
|
||||
self,
|
||||
model_to_use: str,
|
||||
messages: list | None,
|
||||
contents: list | None,
|
||||
deployment: dict | None = None,
|
||||
request_model: str = "",
|
||||
tools: list | None = None,
|
||||
system: object | None = None,
|
||||
) -> TokenCountResponse | None:
|
||||
raise self._error
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"provider_error, expected_code",
|
||||
[
|
||||
(
|
||||
litellm.APIError(
|
||||
status_code=400,
|
||||
message="contents is not specified",
|
||||
llm_provider="gemini",
|
||||
model="gemini-2.5-flash",
|
||||
),
|
||||
"400",
|
||||
),
|
||||
(
|
||||
litellm.APIConnectionError(message="connection refused", llm_provider="gemini", model="gemini-2.5-flash"),
|
||||
"500",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_provider_counter_raising_litellm_error_falls_back_or_surfaces_provider_status(
|
||||
provider_error, expected_code, monkeypatch
|
||||
):
|
||||
counter = _RaisingCounter(provider_error)
|
||||
call = dict(
|
||||
provider_counter=counter,
|
||||
custom_llm_provider="gemini",
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
contents=None,
|
||||
deployment={"litellm_params": {"model": "gemini/gemini-2.5-flash"}},
|
||||
request_model="gemini-flash",
|
||||
)
|
||||
|
||||
monkeypatch.setattr(litellm, "disable_token_counter", False)
|
||||
assert await _try_provider_token_count(**call) is None
|
||||
|
||||
monkeypatch.setattr(litellm, "disable_token_counter", True)
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _try_provider_token_count(**call)
|
||||
assert exc_info.value.code == expected_code
|
||||
assert provider_error.message in exc_info.value.message
|
||||
assert exc_info.value.type == "token_counting_error"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_token_counter_error_raises_exception_when_disabled():
|
||||
"""
|
||||
|
|
@ -1099,9 +1104,7 @@ async def test_proxy_token_counter_error_raises_exception_when_disabled():
|
|||
|
||||
# Save original value and function
|
||||
original_disable = litellm.disable_token_counter
|
||||
original_get_provider_token_counter = (
|
||||
litellm.proxy.proxy_server._get_provider_token_counter
|
||||
)
|
||||
original_get_provider_token_counter = litellm.proxy.proxy_server._get_provider_token_counter
|
||||
|
||||
try:
|
||||
litellm.disable_token_counter = True
|
||||
|
|
@ -1115,9 +1118,7 @@ async def test_proxy_token_counter_error_raises_exception_when_disabled():
|
|||
def mock_get_provider_token_counter(deployment, model_to_use):
|
||||
return (mock_counter, "anthropic.claude-3-sonnet", "bedrock")
|
||||
|
||||
litellm.proxy.proxy_server._get_provider_token_counter = (
|
||||
mock_get_provider_token_counter
|
||||
)
|
||||
litellm.proxy.proxy_server._get_provider_token_counter = mock_get_provider_token_counter
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await token_counter(
|
||||
|
|
@ -1132,9 +1133,7 @@ async def test_proxy_token_counter_error_raises_exception_when_disabled():
|
|||
assert "Rate limit exceeded" in exc_info.value.message
|
||||
finally:
|
||||
litellm.disable_token_counter = original_disable
|
||||
litellm.proxy.proxy_server._get_provider_token_counter = (
|
||||
original_get_provider_token_counter
|
||||
)
|
||||
litellm.proxy.proxy_server._get_provider_token_counter = original_get_provider_token_counter
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1169,9 +1168,7 @@ async def test_proxy_token_counter_error_falls_back_when_enabled():
|
|||
|
||||
# Save original value and function
|
||||
original_disable = litellm.disable_token_counter
|
||||
original_get_provider_token_counter = (
|
||||
litellm.proxy.proxy_server._get_provider_token_counter
|
||||
)
|
||||
original_get_provider_token_counter = litellm.proxy.proxy_server._get_provider_token_counter
|
||||
|
||||
try:
|
||||
litellm.disable_token_counter = False
|
||||
|
|
@ -1185,9 +1182,7 @@ async def test_proxy_token_counter_error_falls_back_when_enabled():
|
|||
def mock_get_provider_token_counter(deployment, model_to_use):
|
||||
return (mock_counter, "anthropic.claude-3-sonnet", "bedrock")
|
||||
|
||||
litellm.proxy.proxy_server._get_provider_token_counter = (
|
||||
mock_get_provider_token_counter
|
||||
)
|
||||
litellm.proxy.proxy_server._get_provider_token_counter = mock_get_provider_token_counter
|
||||
|
||||
# Should not raise, should fall back to local tokenizer
|
||||
result = await token_counter(
|
||||
|
|
@ -1204,9 +1199,7 @@ async def test_proxy_token_counter_error_falls_back_when_enabled():
|
|||
assert result.tokenizer_type != "bedrock_api"
|
||||
finally:
|
||||
litellm.disable_token_counter = original_disable
|
||||
litellm.proxy.proxy_server._get_provider_token_counter = (
|
||||
original_get_provider_token_counter
|
||||
)
|
||||
litellm.proxy.proxy_server._get_provider_token_counter = original_get_provider_token_counter
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ import json
|
|||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter
|
||||
|
||||
COUNT_TOKENS_URL = "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash:countTokens"
|
||||
|
|
@ -60,6 +61,81 @@ async def test_acount_tokens_keeps_contents_body_without_system_or_tools():
|
|||
assert body == {"contents": [{"role": "user", "parts": [{"text": "hi"}]}]}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acount_tokens_sends_generate_content_request_with_system_only():
|
||||
recorded: list[httpx.Request] = []
|
||||
|
||||
def _handler(request: httpx.Request) -> httpx.Response:
|
||||
recorded.append(request)
|
||||
return httpx.Response(200, json={"totalTokens": 9})
|
||||
|
||||
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",
|
||||
system_instruction={"parts": [{"text": "be terse"}]},
|
||||
client=client,
|
||||
)
|
||||
|
||||
body = json.loads(recorded[-1].content)
|
||||
assert body == {
|
||||
"generateContentRequest": {
|
||||
"model": "models/gemini-2.5-flash",
|
||||
"contents": [{"role": "user", "parts": [{"text": "hi"}]}],
|
||||
"systemInstruction": {"parts": [{"text": "be terse"}]},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acount_tokens_sends_generate_content_request_with_tools_only():
|
||||
recorded: list[httpx.Request] = []
|
||||
|
||||
def _handler(request: httpx.Request) -> httpx.Response:
|
||||
recorded.append(request)
|
||||
return httpx.Response(200, json={"totalTokens": 9})
|
||||
|
||||
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",
|
||||
tools=[{"function_declarations": [{"name": "get_weather"}]}],
|
||||
client=client,
|
||||
)
|
||||
|
||||
body = json.loads(recorded[-1].content)
|
||||
assert body == {
|
||||
"generateContentRequest": {
|
||||
"model": "models/gemini-2.5-flash",
|
||||
"contents": [{"role": "user", "parts": [{"text": "hi"}]}],
|
||||
"tools": [{"function_declarations": [{"name": "get_weather"}]}],
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acount_tokens_non_json_body_raises_api_error_with_response_status():
|
||||
def _handler(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, content=b"<html>proxy error page</html>")
|
||||
|
||||
client = httpx.AsyncClient(transport=httpx.MockTransport(_handler))
|
||||
|
||||
with pytest.raises(litellm.APIError) as exc_info:
|
||||
await GoogleAIStudioTokenCounter().acount_tokens(
|
||||
model="gemini-2.5-flash",
|
||||
contents=[{"role": "user", "parts": [{"text": "hi"}]}],
|
||||
api_key="test-key",
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 200
|
||||
assert "non-JSON" in exc_info.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acount_tokens_wraps_unexpected_error_in_api_error():
|
||||
import litellm
|
||||
|
|
|
|||
|
|
@ -297,6 +297,20 @@ def test_build_count_tokens_payload_folds_system_into_contents_for_models_withou
|
|||
assert [part.get("text") for part in payload.contents[0]["parts"]] == ["be nice", "hi"]
|
||||
|
||||
|
||||
def test_build_count_tokens_payload_returns_invalid_for_malformed_anthropic_input():
|
||||
from litellm.llms.gemini.count_tokens.transformation import InvalidCountTokensRequest
|
||||
|
||||
payload = build_count_tokens_payload(
|
||||
model="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
system={"not": "a valid system prompt"},
|
||||
tools=[{"name": "get_weather", "input_schema": {"type": "object"}}],
|
||||
)
|
||||
|
||||
assert isinstance(payload, InvalidCountTokensRequest)
|
||||
assert payload.message
|
||||
|
||||
|
||||
def test_normalize_count_tokens_tools_handles_each_tool_shape():
|
||||
from litellm.llms.gemini.count_tokens.transformation import normalize_count_tokens_tools
|
||||
|
||||
|
|
|
|||
|
|
@ -144,7 +144,13 @@ class TestGoogleAIStudioTokenCounter:
|
|||
assert result.original_response == mock_response
|
||||
|
||||
# Verify the mock was called correctly
|
||||
mock_acount_tokens.assert_called_once_with(model=model_to_use, contents=contents, client=None)
|
||||
mock_acount_tokens.assert_called_once_with(
|
||||
system_instruction=None,
|
||||
tools=None,
|
||||
client=None,
|
||||
model=model_to_use,
|
||||
contents=contents,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_translates_anthropic_messages_system_and_tools(self):
|
||||
|
|
@ -254,13 +260,13 @@ class TestGoogleAIStudioTokenCounter:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_translation_error_falls_back(self):
|
||||
"""A crash translating bad message shapes must surface as an error
|
||||
TokenCountResponse so the proxy falls back instead of 500ing."""
|
||||
"""Malformed message shapes surface as a 400 error TokenCountResponse so
|
||||
the proxy falls back instead of 500ing."""
|
||||
token_counter = GoogleAIStudioTokenCounter()
|
||||
|
||||
result = await token_counter.count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=[{"role": "tool", "content": "orphaned result", "tool_call_id": "missing-call"}],
|
||||
messages=[{"role": "user", "content": [{"type": "text", "text": 123}]}],
|
||||
contents=None,
|
||||
deployment={"litellm_params": {"api_key": "test-key"}},
|
||||
request_model="gemini/gemini-2.5-flash",
|
||||
|
|
@ -268,21 +274,56 @@ class TestGoogleAIStudioTokenCounter:
|
|||
|
||||
assert result is not None
|
||||
assert result.error is True
|
||||
assert result.status_code == 500
|
||||
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_unexpected_handler_error_returns_error_response(self):
|
||||
"""A non-litellm exception escaping the handler must still surface as an
|
||||
error TokenCountResponse so the proxy can fall back."""
|
||||
async def test_count_tokens_malformed_anthropic_input_returns_400_without_http_call(self):
|
||||
"""Input that fails Anthropic request validation returns a 400 error
|
||||
response before any request reaches the provider."""
|
||||
import httpx
|
||||
|
||||
recorded: list = []
|
||||
|
||||
def _handler(request):
|
||||
recorded.append(request)
|
||||
return httpx.Response(200, json={"totalTokens": 7})
|
||||
|
||||
token_counter = GoogleAIStudioTokenCounter()
|
||||
|
||||
result = await token_counter.count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
contents=None,
|
||||
deployment={"litellm_params": {"api_key": "test-key", "api_base": "https://gemini.example.test"}},
|
||||
request_model="gemini/gemini-2.5-flash",
|
||||
system={"not": "a valid system prompt"},
|
||||
tools=[{"name": "get_weather", "input_schema": {"type": "object"}}],
|
||||
client=httpx.AsyncClient(transport=httpx.MockTransport(_handler)),
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.error is True
|
||||
assert result.status_code == 400
|
||||
assert result.total_tokens == 0
|
||||
assert recorded == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_connection_error_returns_error_response(self):
|
||||
"""A provider APIConnectionError surfaces as an error TokenCountResponse
|
||||
so the proxy falls back 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 = RuntimeError("unexpected failure")
|
||||
mock_acount_tokens.side_effect = litellm.APIConnectionError(
|
||||
message="connection refused", llm_provider="gemini", model="gemini-2.5-flash"
|
||||
)
|
||||
|
||||
result = await token_counter.count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
|
|
@ -295,7 +336,38 @@ class TestGoogleAIStudioTokenCounter:
|
|||
assert result is not None
|
||||
assert result.error is True
|
||||
assert result.status_code == 500
|
||||
assert "unexpected failure" in (result.error_message or "")
|
||||
assert "connection refused" in (result.error_message or "")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_response_without_total_tokens_returns_502(self):
|
||||
"""A 200 body missing totalTokens is a malformed provider response, not
|
||||
a successful count, so it surfaces as a 502 error response."""
|
||||
import httpx
|
||||
|
||||
recorded: list = []
|
||||
|
||||
def _handler(request):
|
||||
recorded.append(request)
|
||||
return httpx.Response(200, json={"promptTokensDetails": []})
|
||||
|
||||
token_counter = GoogleAIStudioTokenCounter()
|
||||
|
||||
result = await token_counter.count_tokens(
|
||||
model_to_use="gemini-2.5-flash",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
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)),
|
||||
)
|
||||
|
||||
assert recorded != []
|
||||
assert result is not None
|
||||
assert result.error is True
|
||||
assert result.status_code == 502
|
||||
assert result.total_tokens == 0
|
||||
assert "totalTokens" in (result.error_message or "")
|
||||
assert result.original_response == {"promptTokensDetails": []}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_tokens_returns_none_without_contents_or_messages(self):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue